Skip to main content

confium_coordinator/coordinator/
rate_limiter.rs

1//! Rate limiter — token bucket implementation for DoS protection.
2//!
3//! Limits request rates per client (or per any arbitrary key). Uses
4//! the token bucket algorithm: tokens refill at a fixed rate up to a
5//! capacity ceiling. Each request consumes one token.
6//!
7//! ## OCP design
8//!
9//! New rate-limiting algorithms (sliding window, leaky bucket) are
10//! added by implementing the [`RateLimiter`] trait — no existing code
11//! is modified.
12
13use std::collections::HashMap;
14use std::sync::Mutex;
15use std::time::Instant;
16
17/// Trait for rate limiters. Implementations decide whether a request
18/// from `key` (typically a client ID or IP) should be allowed.
19pub trait RateLimiter: Send + Sync {
20    /// Check if a request is allowed. Consumes one token if allowed.
21    /// Returns `true` if allowed, `false` if rate-limited.
22    fn check(&self, key: &str) -> bool;
23
24    /// Peek at remaining tokens without consuming. Returns the
25    /// approximate number of available tokens (0 means limited).
26    fn peek(&self, key: &str) -> u32;
27
28    /// Reset the limiter state for a key (clears the bucket).
29    fn reset(&self, key: &str);
30}
31
32/// Configuration for a token bucket rate limiter.
33#[derive(Debug, Clone)]
34pub struct TokenBucketConfig {
35    /// Maximum tokens a bucket can hold.
36    pub capacity: u32,
37    /// Tokens added per second.
38    pub refill_per_second: f64,
39}
40
41/// Token bucket rate limiter with per-key buckets.
42pub struct TokenBucketRateLimiter {
43    config: TokenBucketConfig,
44    buckets: Mutex<HashMap<String, Bucket>>,
45}
46
47struct Bucket {
48    tokens: f64,
49    last_refill: Instant,
50}
51
52impl TokenBucketRateLimiter {
53    /// Create a new rate limiter with the given configuration.
54    pub fn new(config: TokenBucketConfig) -> Self {
55        Self {
56            config,
57            buckets: Mutex::new(HashMap::new()),
58        }
59    }
60
61    /// Create a rate limiter with `capacity` tokens, refilling at
62    /// `refill_per_second` tokens/second.
63    pub fn with_rate(capacity: u32, refill_per_second: f64) -> Self {
64        Self::new(TokenBucketConfig {
65            capacity,
66            refill_per_second,
67        })
68    }
69
70    fn refill_bucket(&self, bucket: &mut Bucket) {
71        let now = Instant::now();
72        let elapsed = now.duration_since(bucket.last_refill).as_secs_f64();
73        let refilled = elapsed * self.config.refill_per_second;
74        bucket.tokens = (bucket.tokens + refilled).min(self.config.capacity as f64);
75        bucket.last_refill = now;
76    }
77
78    fn get_or_create_bucket(&self, _key: &str) -> Bucket {
79        Bucket {
80            tokens: self.config.capacity as f64,
81            last_refill: Instant::now(),
82        }
83    }
84}
85
86impl RateLimiter for TokenBucketRateLimiter {
87    fn check(&self, key: &str) -> bool {
88        let mut buckets = self.buckets.lock().unwrap();
89        let bucket = buckets
90            .entry(key.to_string())
91            .or_insert_with(|| self.get_or_create_bucket(key));
92        self.refill_bucket(bucket);
93        if bucket.tokens >= 1.0 {
94            bucket.tokens -= 1.0;
95            true
96        } else {
97            false
98        }
99    }
100
101    fn peek(&self, key: &str) -> u32 {
102        let mut buckets = self.buckets.lock().unwrap();
103        let bucket = buckets
104            .entry(key.to_string())
105            .or_insert_with(|| self.get_or_create_bucket(key));
106        self.refill_bucket(bucket);
107        bucket.tokens as u32
108    }
109
110    fn reset(&self, key: &str) {
111        self.buckets.lock().unwrap().remove(key);
112    }
113}
114
115/// A rate limiter that always allows. Useful for testing and
116/// development environments where rate limiting is disabled.
117pub struct NoopRateLimiter;
118
119impl RateLimiter for NoopRateLimiter {
120    fn check(&self, _key: &str) -> bool {
121        true
122    }
123    fn peek(&self, _key: &str) -> u32 {
124        u32::MAX
125    }
126    fn reset(&self, _key: &str) {}
127}
128
129#[cfg(test)]
130mod tests {
131    use super::*;
132    use std::time::Duration;
133
134    #[test]
135    fn allows_until_capacity() {
136        let limiter = TokenBucketRateLimiter::with_rate(5, 0.0);
137        for _ in 0..5 {
138            assert!(limiter.check("client-1"));
139        }
140        assert!(!limiter.check("client-1"));
141    }
142
143    #[test]
144    fn separate_keys_have_separate_buckets() {
145        let limiter = TokenBucketRateLimiter::with_rate(2, 0.0);
146        assert!(limiter.check("a"));
147        assert!(limiter.check("a"));
148        assert!(limiter.check("b"));
149        assert!(limiter.check("b"));
150        assert!(!limiter.check("a"));
151        assert!(!limiter.check("b"));
152    }
153
154    #[test]
155    fn noop_always_allows() {
156        let limiter = NoopRateLimiter;
157        for _ in 0..100 {
158            assert!(limiter.check("anyone"));
159        }
160    }
161
162    #[test]
163    fn peek_does_not_consume() {
164        let limiter = TokenBucketRateLimiter::with_rate(3, 0.0);
165        assert_eq!(limiter.peek("k"), 3);
166        assert_eq!(limiter.peek("k"), 3);
167        limiter.check("k");
168        assert_eq!(limiter.peek("k"), 2);
169    }
170
171    #[test]
172    fn reset_clears_bucket() {
173        let limiter = TokenBucketRateLimiter::with_rate(1, 0.0);
174        assert!(limiter.check("k"));
175        assert!(!limiter.check("k"));
176        limiter.reset("k");
177        assert!(limiter.check("k"));
178    }
179
180    #[test]
181    fn refill_restores_tokens_over_time() {
182        let limiter = TokenBucketRateLimiter::with_rate(1, 1000.0);
183        assert!(limiter.check("k"));
184        assert!(!limiter.check("k"));
185        std::thread::sleep(Duration::from_millis(10));
186        assert!(limiter.check("k"));
187    }
188
189    #[test]
190    fn capacity_ceiling_prevents_overfill() {
191        let limiter = TokenBucketRateLimiter::with_rate(3, 1000.0);
192        std::thread::sleep(Duration::from_millis(50));
193        assert_eq!(limiter.peek("k"), 3);
194    }
195
196    #[test]
197    fn new_key_starts_at_full_capacity() {
198        let limiter = TokenBucketRateLimiter::with_rate(7, 1.0);
199        assert_eq!(limiter.peek("fresh"), 7);
200    }
201
202    #[test]
203    fn unknown_key_peek_creates_full_bucket() {
204        let limiter = TokenBucketRateLimiter::with_rate(5, 0.0);
205        assert_eq!(limiter.peek("unknown"), 5);
206    }
207
208    #[test]
209    fn config_values_preserved() {
210        let config = TokenBucketConfig {
211            capacity: 10,
212            refill_per_second: 2.5,
213        };
214        let limiter = TokenBucketRateLimiter::new(config);
215        assert_eq!(limiter.peek("k"), 10);
216    }
217}