1use std::collections::HashMap;
14use std::sync::{Arc, RwLock};
15use std::time::{Duration, Instant};
16
17pub const DEFAULT_MAX_KEYS: usize = 10_000;
22
23pub trait RateLimiter: Send + Sync {
24 fn acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError>;
25 fn try_acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError>;
26 fn reset(&self, key: &str) -> Result<(), RateLimitError>;
27}
28
29#[derive(Debug, Clone)]
30pub struct RateLimitResult {
31 pub allowed: bool,
32 pub remaining: u64,
33 pub reset_at: i64,
34}
35
36impl RateLimitResult {
37 pub fn allowed(remaining: u64, reset_at: i64) -> Self {
38 Self {
39 allowed: true,
40 remaining,
41 reset_at,
42 }
43 }
44
45 pub fn rejected(remaining: u64, reset_at: i64) -> Self {
46 Self {
47 allowed: false,
48 remaining,
49 reset_at,
50 }
51 }
52}
53
54pub struct SlidingWindowRateLimiter {
55 max_requests: u64,
56 window_size: Duration,
57 entries: Arc<RwLock<HashMap<String, SlidingWindowEntry>>>,
58 max_keys: usize,
60}
61
62#[derive(Clone)]
63struct SlidingWindowEntry {
64 requests: Vec<Instant>,
65}
66
67impl SlidingWindowRateLimiter {
68 pub fn new(max_requests: u64, window_size: Duration) -> Self {
69 Self {
70 max_requests,
71 window_size,
72 entries: Arc::new(RwLock::new(HashMap::new())),
73 max_keys: DEFAULT_MAX_KEYS,
74 }
75 }
76
77 pub fn with_max_keys(mut self, max_keys: usize) -> Self {
82 self.max_keys = max_keys;
83 self
84 }
85
86 fn cleanup_old_requests(&self, entry: &mut SlidingWindowEntry) {
87 let now = Instant::now();
88 entry
89 .requests
90 .retain(|&time| now.duration_since(time) < self.window_size);
91 }
92
93 fn enforce_max_keys(&self, entries: &mut HashMap<String, SlidingWindowEntry>) {
98 while entries.len() > self.max_keys {
99 let now = Instant::now();
101 let oldest_key = entries
102 .iter()
103 .min_by_key(|(_, e)| e.requests.first().copied().unwrap_or(now))
104 .map(|(k, _)| k.clone());
105 match oldest_key {
106 Some(k) => {
107 entries.remove(&k);
108 }
109 None => break,
110 }
111 }
112 }
113}
114
115impl RateLimiter for SlidingWindowRateLimiter {
116 fn acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
117 let mut entries = self
118 .entries
119 .write()
120 .map_err(|e| RateLimitError::Internal(e.to_string()))?;
121
122 if entries.len() >= self.max_keys && !entries.contains_key(key) {
124 self.enforce_max_keys(&mut entries);
125 }
126
127 let entry = entries
128 .entry(key.to_string())
129 .or_insert_with(|| SlidingWindowEntry {
130 requests: Vec::new(),
131 });
132
133 self.cleanup_old_requests(entry);
134
135 if entry.requests.len() < self.max_requests as usize {
136 entry.requests.push(Instant::now());
137 let remaining = self.max_requests - entry.requests.len() as u64;
138 let reset_at = now_timestamp() + self.window_size.as_millis() as i64;
139 Ok(RateLimitResult::allowed(remaining, reset_at))
140 } else {
141 let oldest = entry
142 .requests
143 .first()
144 .map(|t| {
145 let elapsed = t.elapsed().as_millis() as i64;
146 let window_ms = self.window_size.as_millis() as i64;
147 now_timestamp() + (window_ms - elapsed)
148 })
149 .unwrap_or(now_timestamp());
150
151 let remaining = 0;
152 Ok(RateLimitResult::rejected(remaining, oldest))
153 }
154 }
155
156 fn try_acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
157 self.acquire(key)
158 }
159
160 fn reset(&self, key: &str) -> Result<(), RateLimitError> {
161 let mut entries = self
162 .entries
163 .write()
164 .map_err(|e| RateLimitError::Internal(e.to_string()))?;
165 entries.remove(key);
166 Ok(())
167 }
168}
169
170pub struct TokenBucketRateLimiter {
171 capacity: f64,
172 refill_rate: f64,
173 entries: Arc<RwLock<HashMap<String, TokenBucketEntry>>>,
174 max_keys: usize,
176}
177
178#[derive(Clone)]
179struct TokenBucketEntry {
180 tokens: f64,
181 last_refill: Instant,
182}
183
184impl TokenBucketRateLimiter {
185 pub fn new(capacity: u64, refill_per_second: f64) -> Self {
186 Self {
187 capacity: capacity as f64,
188 refill_rate: refill_per_second,
189 entries: Arc::new(RwLock::new(HashMap::new())),
190 max_keys: DEFAULT_MAX_KEYS,
191 }
192 }
193
194 pub fn with_max_keys(mut self, max_keys: usize) -> Self {
196 self.max_keys = max_keys;
197 self
198 }
199
200 fn refill(&self, entry: &mut TokenBucketEntry) {
201 let now = Instant::now();
202 let elapsed = now.duration_since(entry.last_refill).as_secs_f64();
203 let tokens_to_add = if self.refill_rate > 0.0 {
206 elapsed * self.refill_rate
207 } else {
208 0.0
209 };
210
211 entry.tokens = (entry.tokens + tokens_to_add).min(self.capacity);
212 entry.last_refill = now;
213 }
214
215 fn enforce_max_keys(&self, entries: &mut HashMap<String, TokenBucketEntry>) {
220 while entries.len() > self.max_keys {
221 let oldest_key = entries
222 .iter()
223 .min_by_key(|(_, e)| e.last_refill)
224 .map(|(k, _)| k.clone());
225 match oldest_key {
226 Some(k) => {
227 entries.remove(&k);
228 }
229 None => break,
230 }
231 }
232 }
233}
234
235impl RateLimiter for TokenBucketRateLimiter {
236 fn acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
237 let mut entries = self
238 .entries
239 .write()
240 .map_err(|e| RateLimitError::Internal(e.to_string()))?;
241
242 if entries.len() >= self.max_keys && !entries.contains_key(key) {
244 self.enforce_max_keys(&mut entries);
245 }
246
247 let entry = entries
248 .entry(key.to_string())
249 .or_insert_with(|| TokenBucketEntry {
250 tokens: self.capacity,
251 last_refill: Instant::now(),
252 });
253
254 self.refill(entry);
255
256 if entry.tokens >= 1.0 {
257 entry.tokens -= 1.0;
258 let remaining = entry.tokens.floor() as u64;
259 let reset_at = if self.refill_rate > 0.0 {
262 now_timestamp() + (1000.0 / self.refill_rate) as i64
263 } else {
264 i64::MAX
266 };
267 Ok(RateLimitResult::allowed(remaining, reset_at))
268 } else {
269 let reset_at = if self.refill_rate > 0.0 {
271 let wait_time = ((1.0 - entry.tokens) / self.refill_rate * 1000.0) as i64;
272 now_timestamp() + wait_time
273 } else {
274 i64::MAX
276 };
277 Ok(RateLimitResult::rejected(0, reset_at))
278 }
279 }
280
281 fn try_acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
282 self.acquire(key)
283 }
284
285 fn reset(&self, key: &str) -> Result<(), RateLimitError> {
286 let mut entries = self
287 .entries
288 .write()
289 .map_err(|e| RateLimitError::Internal(e.to_string()))?;
290 entries.remove(key);
291 Ok(())
292 }
293}
294
295#[derive(Debug)]
296pub enum RateLimitError {
297 KeyNotFound(String),
298 Internal(String),
299}
300
301impl std::fmt::Display for RateLimitError {
302 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
303 match self {
304 RateLimitError::KeyNotFound(key) => write!(f, "Key not found: {}", key),
305 RateLimitError::Internal(msg) => write!(f, "Internal error: {}", msg),
306 }
307 }
308}
309
310impl std::error::Error for RateLimitError {}
311
312impl serde::Serialize for RateLimitError {
313 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
314 where
315 S: serde::Serializer,
316 {
317 serializer.serialize_str(&self.to_string())
318 }
319}
320
321fn now_timestamp() -> i64 {
322 use std::time::{SystemTime, UNIX_EPOCH};
323 SystemTime::now()
324 .duration_since(UNIX_EPOCH)
325 .unwrap_or_default()
326 .as_millis() as i64
327}
328
329#[cfg(test)]
330mod tests {
331 use super::*;
332
333 #[test]
334 fn test_rate_limit_result_allowed() {
335 let result = RateLimitResult::allowed(5, 1000);
336 assert!(result.allowed);
337 assert_eq!(result.remaining, 5);
338 assert_eq!(result.reset_at, 1000);
339 }
340
341 #[test]
342 fn test_rate_limit_result_rejected() {
343 let result = RateLimitResult::rejected(0, 2000);
344 assert!(!result.allowed);
345 assert_eq!(result.remaining, 0);
346 }
347
348 #[test]
349 fn test_sliding_window_limiter_new() {
350 let limiter = SlidingWindowRateLimiter::new(10, Duration::from_secs(60));
351 let result = limiter.acquire("test-key");
352 assert!(result.is_ok());
353 assert!(result.unwrap().allowed);
354 }
355
356 #[test]
357 fn test_sliding_window_limiter_full() {
358 let limiter = SlidingWindowRateLimiter::new(2, Duration::from_secs(60));
359
360 let r1 = limiter.acquire("key1").unwrap();
361 assert!(r1.allowed);
362
363 let r2 = limiter.acquire("key1").unwrap();
364 assert!(r2.allowed);
365
366 let r3 = limiter.acquire("key1").unwrap();
367 assert!(!r3.allowed);
368 }
369
370 #[test]
371 fn test_sliding_window_different_keys() {
372 let limiter = SlidingWindowRateLimiter::new(1, Duration::from_secs(60));
373
374 let r1 = limiter.acquire("key-a").unwrap();
375 assert!(r1.allowed);
376
377 let r2 = limiter.acquire("key-b").unwrap();
378 assert!(r2.allowed);
379 }
380
381 #[test]
382 fn test_sliding_window_reset() {
383 let limiter = SlidingWindowRateLimiter::new(1, Duration::from_secs(60));
384
385 limiter.acquire("reset-key").unwrap();
386 limiter.acquire("reset-key").unwrap();
387
388 limiter.reset("reset-key").unwrap();
389
390 let result = limiter.acquire("reset-key").unwrap();
391 assert!(result.allowed);
392 }
393
394 #[test]
395 fn test_token_bucket_limiter_new() {
396 let limiter = TokenBucketRateLimiter::new(10, 1.0);
397 let result = limiter.acquire("test-key");
398 assert!(result.is_ok());
399 assert!(result.unwrap().allowed);
400 }
401
402 #[test]
403 fn test_token_bucket_limiter_depletes() {
404 let limiter = TokenBucketRateLimiter::new(2, 1.0);
405
406 let r1 = limiter.acquire("key1").unwrap();
407 assert!(r1.allowed);
408 assert_eq!(r1.remaining, 1);
409
410 let r2 = limiter.acquire("key1").unwrap();
411 assert!(r2.allowed);
412 assert_eq!(r2.remaining, 0);
413
414 let r3 = limiter.acquire("key1").unwrap();
415 assert!(!r3.allowed);
416 }
417
418 #[test]
419 fn test_token_bucket_different_keys() {
420 let limiter = TokenBucketRateLimiter::new(1, 1.0);
421
422 let r1 = limiter.acquire("key-a").unwrap();
423 assert!(r1.allowed);
424
425 let r2 = limiter.acquire("key-b").unwrap();
426 assert!(r2.allowed);
427 }
428
429 #[test]
430 fn test_token_bucket_reset() {
431 let limiter = TokenBucketRateLimiter::new(1, 1.0);
432
433 limiter.acquire("reset-key").unwrap();
434 limiter.acquire("reset-key").unwrap();
435
436 limiter.reset("reset-key").unwrap();
437
438 let result = limiter.acquire("reset-key").unwrap();
439 assert!(result.allowed);
440 }
441
442 #[test]
443 fn test_limiter_try_acquire() {
444 let limiter = SlidingWindowRateLimiter::new(1, Duration::from_secs(60));
445
446 let r1 = limiter.try_acquire("key").unwrap();
447 assert!(r1.allowed);
448
449 let r2 = limiter.try_acquire("key").unwrap();
450 assert!(!r2.allowed);
451 }
452
453 #[test]
456 fn test_token_bucket_zero_refill_rate_does_not_panic() {
457 let limiter = TokenBucketRateLimiter::new(1, 0.0);
460
461 let r1 = limiter.acquire("zero-refill").unwrap();
462 assert!(r1.allowed, "first acquire should be allowed");
463
464 let r2 = limiter.acquire("zero-refill").unwrap();
466 assert!(!r2.allowed, "second acquire should be rejected");
467 assert!(
469 r2.reset_at > 0,
470 "reset_at should be a valid timestamp, got: {}",
471 r2.reset_at
472 );
473 }
474
475 #[test]
476 fn test_token_bucket_negative_refill_rate_does_not_panic() {
477 let limiter = TokenBucketRateLimiter::new(1, -1.0);
479
480 let r1 = limiter.acquire("neg-refill").unwrap();
481 assert!(r1.allowed, "first acquire should be allowed");
482
483 let r2 = limiter.acquire("neg-refill").unwrap();
484 assert!(!r2.allowed, "second acquire should be rejected");
485 assert!(
486 r2.reset_at > 0,
487 "reset_at should be a valid timestamp, got: {}",
488 r2.reset_at
489 );
490 }
491}