1use std::collections::HashMap;
6use std::sync::Mutex;
7
8#[derive(Debug, Clone, Copy, PartialEq)]
10pub struct Rate {
11 pub burst: f64,
12 pub per_secs: f64,
13}
14
15impl Rate {
16 pub const fn new(burst: u32, per_secs: u32) -> Self {
17 Rate {
18 burst: burst as f64,
19 per_secs: per_secs as f64,
20 }
21 }
22}
23
24#[derive(Debug, Clone, Copy)]
25struct Bucket {
26 tokens: f64,
27 at: f64,
28}
29
30const MAX_KEYS: usize = 10_000;
32
33#[derive(Debug)]
34pub struct RateLimiter {
35 rate: Rate,
36 buckets: Mutex<HashMap<String, Bucket>>,
37}
38
39impl RateLimiter {
40 pub fn new(rate: Rate) -> Self {
41 RateLimiter {
42 rate,
43 buckets: Mutex::new(HashMap::new()),
44 }
45 }
46
47 pub fn take(&self, key: &str, now: f64) -> Result<(), u64> {
50 let mut m = self.buckets.lock().unwrap_or_else(|e| e.into_inner());
51 if m.len() >= MAX_KEYS && !m.contains_key(key) {
52 let r = self.rate;
53 m.retain(|_, b| refill(r, *b, now).tokens < r.burst);
54 if m.len() >= MAX_KEYS {
55 return Err(r.per_secs.ceil() as u64);
58 }
59 }
60 let b = m.entry(key.to_string()).or_insert(Bucket {
61 tokens: self.rate.burst,
62 at: now,
63 });
64 *b = refill(self.rate, *b, now);
65 if b.tokens >= 1.0 {
66 b.tokens -= 1.0;
67 Ok(())
68 } else {
69 Err(((1.0 - b.tokens) * self.rate.per_secs).ceil().max(1.0) as u64)
70 }
71 }
72}
73
74fn refill(r: Rate, b: Bucket, now: f64) -> Bucket {
75 let dt = (now - b.at).max(0.0);
76 Bucket {
77 tokens: (b.tokens + dt / r.per_secs).min(r.burst),
78 at: now,
79 }
80}
81
82#[cfg(test)]
83mod tests {
84 use super::*;
85
86 #[test]
87 fn bucket() {
88 let l = RateLimiter::new(Rate::new(3, 10));
89 for _ in 0..3 {
90 assert!(l.take("a", 100.0).is_ok());
91 }
92 assert_eq!(l.take("a", 100.0), Err(10));
93 assert!(l.take("b", 100.0).is_ok());
95 assert_eq!(l.take("a", 105.0), Err(5));
97 assert!(l.take("a", 110.0).is_ok());
98 assert!(l.take("a", 110.0).is_err());
99 for _ in 0..3 {
101 assert!(l.take("a", 10_000.0).is_ok());
102 }
103 assert!(l.take("a", 10_000.0).is_err());
104 }
105
106 #[test]
107 fn bounded() {
108 let l = RateLimiter::new(Rate::new(1, 60));
109 for i in 0..MAX_KEYS {
110 assert!(l.take(&format!("k{i}"), 0.0).is_ok());
111 }
112 assert!(l.take("new", 1.0).is_err());
114 assert!(l.take("new", 120.0).is_ok());
116 assert!(l.buckets.lock().unwrap().len() < 10);
117 }
118}