confium_coordinator/coordinator/
rate_limiter.rs1use std::collections::HashMap;
14use std::sync::Mutex;
15use std::time::Instant;
16
17pub trait RateLimiter: Send + Sync {
20 fn check(&self, key: &str) -> bool;
23
24 fn peek(&self, key: &str) -> u32;
27
28 fn reset(&self, key: &str);
30}
31
32#[derive(Debug, Clone)]
34pub struct TokenBucketConfig {
35 pub capacity: u32,
37 pub refill_per_second: f64,
39}
40
41pub 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 pub fn new(config: TokenBucketConfig) -> Self {
55 Self {
56 config,
57 buckets: Mutex::new(HashMap::new()),
58 }
59 }
60
61 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
115pub 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}