Skip to main content

fedimint_gateway_server/
rate_limit.rs

1use std::sync::Mutex;
2use std::time::SystemTime;
3
4use fedimint_core::time::now;
5
6/// A global token bucket limiting the rate of unauthenticated requests that
7/// create state on the gateway or its Lightning node.
8///
9/// The bucket holds at most `burst` tokens and refills at `refill_per_second`;
10/// each request takes one token and is rejected if none is available. The
11/// limiter is transport-agnostic since it guards the request handler itself
12/// rather than the HTTP layer, so requests arriving over Iroh are limited as
13/// well.
14#[derive(Debug)]
15pub struct TokenBucketRateLimiter {
16    burst: f64,
17    refill_per_second: f64,
18    state: Mutex<TokenBucketState>,
19}
20
21#[derive(Debug)]
22struct TokenBucketState {
23    tokens: f64,
24    last_refill: SystemTime,
25}
26
27impl TokenBucketRateLimiter {
28    pub fn new(burst: u32, refill_per_second: u32) -> Self {
29        Self {
30            burst: f64::from(burst),
31            refill_per_second: f64::from(refill_per_second),
32            state: Mutex::new(TokenBucketState {
33                tokens: f64::from(burst),
34                last_refill: now(),
35            }),
36        }
37    }
38
39    /// Takes a token from the bucket if one is available, returning whether
40    /// the request may proceed.
41    pub fn try_acquire(&self) -> bool {
42        self.try_acquire_at(now())
43    }
44
45    fn try_acquire_at(&self, now: SystemTime) -> bool {
46        let mut state = self
47            .state
48            .lock()
49            .expect("No code holding the lock can panic");
50
51        // `SystemTime` is not monotonic; if the clock moved backwards, refill
52        // nothing and restart the refill measurement from the earlier time.
53        let elapsed = now
54            .duration_since(state.last_refill)
55            .unwrap_or_default()
56            .as_secs_f64();
57        state.tokens = (state.tokens + elapsed * self.refill_per_second).min(self.burst);
58        state.last_refill = now;
59
60        if state.tokens >= 1.0 {
61            state.tokens -= 1.0;
62            true
63        } else {
64            false
65        }
66    }
67}
68
69#[cfg(test)]
70mod tests {
71    use std::time::Duration;
72
73    use super::*;
74
75    #[test]
76    fn burst_is_granted_then_rejected() {
77        let limiter = TokenBucketRateLimiter::new(3, 1);
78        let start = now();
79
80        for _ in 0..3 {
81            assert!(limiter.try_acquire_at(start));
82        }
83        assert!(!limiter.try_acquire_at(start));
84    }
85
86    #[test]
87    fn tokens_refill_over_time() {
88        let limiter = TokenBucketRateLimiter::new(2, 5);
89        let start = now();
90
91        assert!(limiter.try_acquire_at(start));
92        assert!(limiter.try_acquire_at(start));
93        assert!(!limiter.try_acquire_at(start));
94
95        // 200ms at 5 tokens/sec refills exactly one token.
96        let later = start + Duration::from_millis(200);
97        assert!(limiter.try_acquire_at(later));
98        assert!(!limiter.try_acquire_at(later));
99    }
100
101    #[test]
102    fn refill_is_capped_at_burst() {
103        let limiter = TokenBucketRateLimiter::new(2, 5);
104        let start = now();
105
106        // After a long idle period only `burst` tokens are available.
107        let much_later = start + Duration::from_secs(60);
108        assert!(limiter.try_acquire_at(much_later));
109        assert!(limiter.try_acquire_at(much_later));
110        assert!(!limiter.try_acquire_at(much_later));
111    }
112
113    #[test]
114    fn backwards_clock_jump_refills_nothing() {
115        let limiter = TokenBucketRateLimiter::new(2, 5);
116        let start = now();
117
118        assert!(limiter.try_acquire_at(start));
119        assert!(limiter.try_acquire_at(start));
120
121        // A clock rollback must not grant tokens or panic.
122        let earlier = start - Duration::from_secs(60);
123        assert!(!limiter.try_acquire_at(earlier));
124    }
125}