Skip to main content

isb_server/auth/
limit.rs

1//! An in-memory token bucket per key (an email, an IP), to blunt password
2//! guessing. It resets when the daemon restarts, which is fine: the point is
3//! to make online guessing slow, and a restart is rare and visible.
4
5use std::collections::HashMap;
6use std::sync::Mutex;
7
8/// `burst` attempts at once, then one every `per_secs`.
9#[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
30/// Past this many keys, full buckets (idle keys) are dropped.
31const 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    /// Take one token for `key` at time `now` (seconds). `Err(secs)` says how
48    /// long until the next one.
49    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                // Still full of active keys: someone is spraying. Refuse new
56                // keys rather than grow without bound.
57                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        // Other keys are independent.
94        assert!(l.take("b", 100.0).is_ok());
95        // Refills at one per 10s.
96        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        // Never above the burst.
100        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        // Every bucket is empty (active): a new key is refused.
113        assert!(l.take("new", 1.0).is_err());
114        // Once they refill, idle keys are dropped and new ones fit.
115        assert!(l.take("new", 120.0).is_ok());
116        assert!(l.buckets.lock().unwrap().len() < 10);
117    }
118}