fedimint_gateway_server/
rate_limit.rs1use std::sync::Mutex;
2use std::time::SystemTime;
3
4use fedimint_core::time::now;
5
6#[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 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 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 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 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 let earlier = start - Duration::from_secs(60);
123 assert!(!limiter.try_acquire_at(earlier));
124 }
125}