Skip to main content

codoseo_web/agent/
limiter.rs

1//! A small in-memory limiter for the no-key MCP tools: how many calls one client address may
2//! make in a sliding window. It guards the database from a client that polls or probes in a
3//! tight loop; the audit and email limits that cost something live in the store. One web
4//! container is enough for that, so the counts live here and a restart forgets them.
5
6use std::collections::{HashMap, VecDeque};
7use std::sync::Mutex;
8use std::time::{Duration, Instant};
9
10/// Calls a direct (non-shared) no-key client may make per window.
11pub const ANON_CALLS_PER_WINDOW: usize = 30;
12pub const ANON_WINDOW: Duration = Duration::from_secs(60);
13
14/// Past this many tracked clients, idle ones are dropped on the next call.
15const 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    /// The limit for the no-key tools: 30 calls a minute.
33    pub fn for_anon_tools() -> CallLimiter {
34        CallLimiter::new(ANON_CALLS_PER_WINDOW, ANON_WINDOW)
35    }
36
37    /// Counts one call by `key` at `now`. `Err` holds how long until the oldest counted call
38    /// leaves the window (nothing is counted then).
39    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        // A refusal counts nothing, and another client is unaffected.
80        assert!(limiter.check(b"b", start).is_ok());
81        // The window slides: the first three age out together.
82        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}