codoseo_web/agent/
limiter.rs1use std::collections::{HashMap, VecDeque};
7use std::sync::Mutex;
8use std::time::{Duration, Instant};
9
10pub const ANON_CALLS_PER_WINDOW: usize = 30;
12pub const ANON_WINDOW: Duration = Duration::from_secs(60);
13
14const PRUNE_ABOVE: usize = 4096;
16
17pub struct CallLimiter {
18 max: usize,
19 window: Duration,
20 calls: Mutex<HashMap<Vec<u8>, VecDeque<Instant>>>,
21}
22
23impl CallLimiter {
24 pub fn new(max: usize, window: Duration) -> CallLimiter {
25 CallLimiter {
26 max,
27 window,
28 calls: Mutex::new(HashMap::new()),
29 }
30 }
31
32 pub fn for_anon_tools() -> CallLimiter {
34 CallLimiter::new(ANON_CALLS_PER_WINDOW, ANON_WINDOW)
35 }
36
37 pub fn check(&self, key: &[u8], now: Instant) -> Result<(), Duration> {
40 let mut calls = self.calls.lock().unwrap_or_else(|e| e.into_inner());
41 if calls.len() > PRUNE_ABOVE {
42 calls.retain(|_, times| {
43 times
44 .back()
45 .is_some_and(|last| now.saturating_duration_since(*last) < self.window)
46 });
47 }
48 let times = calls.entry(key.to_vec()).or_default();
49 while times
50 .front()
51 .is_some_and(|first| now.saturating_duration_since(*first) >= self.window)
52 {
53 times.pop_front();
54 }
55 if times.len() >= self.max {
56 let oldest = times.front().copied().unwrap_or(now);
57 return Err((oldest + self.window).saturating_duration_since(now));
58 }
59 times.push_back(now);
60 Ok(())
61 }
62}
63
64#[cfg(test)]
65mod tests {
66 use super::*;
67
68 #[test]
69 fn it_allows_the_limit_per_window_per_client_and_forgets_old_calls() {
70 let limiter = CallLimiter::new(3, Duration::from_secs(60));
71 let start = Instant::now();
72 for _ in 0..3 {
73 assert!(limiter.check(b"a", start).is_ok());
74 }
75 let wait = limiter
76 .check(b"a", start + Duration::from_secs(10))
77 .unwrap_err();
78 assert_eq!(wait, Duration::from_secs(50));
79 assert!(limiter.check(b"b", start).is_ok());
81 assert!(limiter.check(b"a", start + Duration::from_secs(60)).is_ok());
83 }
84
85 #[test]
86 fn idle_clients_are_dropped_when_many_are_tracked() {
87 let limiter = CallLimiter::new(1, Duration::from_secs(60));
88 let start = Instant::now();
89 for n in 0..=PRUNE_ABOVE {
90 limiter.check(&n.to_be_bytes(), start).unwrap();
91 }
92 limiter
93 .check(b"late", start + Duration::from_secs(120))
94 .unwrap();
95 assert_eq!(limiter.calls.lock().unwrap().len(), 1);
96 }
97}