Skip to main content

tokio_prompt_orchestrator/
rate_limiter.rs

1//! Token-bucket and sliding-window rate limiter per model.
2//!
3//! Provides [`RateLimiterRegistry`] which holds a [`ModelRateLimiter`] for
4//! each registered model.  Each model limiter combines a [`TokenBucket`]
5//! (burst-capable refilling bucket) with a [`SlidingWindow`] (rolling
6//! request-count check) so that both instantaneous bursts and sustained
7//! throughput are controlled.
8
9use std::collections::{HashMap, VecDeque};
10use std::sync::atomic::{AtomicI64, AtomicU64, Ordering};
11use std::sync::{Mutex, RwLock};
12use std::time::Instant;
13
14use thiserror::Error;
15
16// ── RateLimitError ────────────────────────────────────────────────────────────
17
18/// Errors returned when a rate-limit check fails.
19#[derive(Debug, Error)]
20pub enum RateLimitError {
21    /// The token bucket has been exhausted; caller should retry after the
22    /// indicated delay.
23    #[error("token bucket exhausted; retry after {retry_after_ms} ms")]
24    TokenBucketExhausted {
25        /// Approximate milliseconds until enough tokens are available.
26        retry_after_ms: u64,
27    },
28    /// The sliding-window request count has been exceeded; caller should wait
29    /// until the window resets.
30    #[error("sliding window limit exceeded; window resets in {reset_at_ms} ms")]
31    WindowLimitExceeded {
32        /// Milliseconds until the oldest slot falls out of the window.
33        reset_at_ms: u64,
34    },
35    /// No limiter has been registered for the given model.
36    #[error("no rate limiter registered for model '{0}'")]
37    UnknownModel(String),
38}
39
40// ── RateLimiterConfig ─────────────────────────────────────────────────────────
41
42/// Configuration used when registering a model with the [`RateLimiterRegistry`].
43#[derive(Debug, Clone)]
44pub struct RateLimiterConfig {
45    /// Maximum requests allowed per minute in the sliding window.
46    pub requests_per_minute: u32,
47    /// Maximum tokens allowed per minute in the token bucket.
48    pub tokens_per_minute: u64,
49    /// Burst multiplier applied to `tokens_per_minute` to derive bucket
50    /// capacity (e.g. `2.0` allows a burst up to 2× the per-minute quota).
51    pub burst_multiplier: f64,
52}
53
54// ── TokenBucket ───────────────────────────────────────────────────────────────
55
56/// Refilling token bucket.
57///
58/// Tokens are stored as a signed 64-bit integer (scaled by 1000 to allow
59/// sub-token precision) and refilled lazily on every call to
60/// [`TokenBucket::try_consume`].
61pub struct TokenBucket {
62    /// Maximum tokens the bucket can hold (scaled ×1000).
63    capacity: u64,
64    /// Current token count (scaled ×1000); stored as `i64` for atomic CAS.
65    tokens: AtomicI64,
66    /// Tokens added per millisecond (scaled ×1000).
67    refill_rate: u64,
68    /// Wall-clock time of the last refill.
69    last_refill: Mutex<Instant>,
70}
71
72impl TokenBucket {
73    /// Create a new bucket with the given capacity (in whole tokens) and a
74    /// per-minute refill rate.  The bucket starts full.
75    pub fn new(capacity_tokens: u64, tokens_per_minute: u64) -> Self {
76        let capacity_scaled = capacity_tokens.saturating_mul(1_000);
77        // refill_rate in (tokens * 1000) per ms
78        let refill_rate = tokens_per_minute.saturating_mul(1_000) / 60_000;
79        Self {
80            capacity: capacity_scaled,
81            tokens: AtomicI64::new(capacity_scaled as i64),
82            refill_rate,
83            last_refill: Mutex::new(Instant::now()),
84        }
85    }
86
87    /// Attempt to consume `n` tokens.
88    ///
89    /// Refills the bucket based on elapsed time first, then deducts `n`.
90    /// Returns `true` if the tokens were available and consumed, `false`
91    /// otherwise (bucket remains unchanged on failure).
92    pub fn try_consume(&self, n: u64) -> bool {
93        // Refill
94        let elapsed_ms = {
95            let mut guard = self.last_refill.lock().unwrap_or_else(|e| e.into_inner());
96            let now = Instant::now();
97            let ms = now.duration_since(*guard).as_millis() as u64;
98            if ms > 0 {
99                *guard = now;
100            }
101            ms
102        };
103
104        let add = (elapsed_ms.saturating_mul(self.refill_rate)) as i64;
105        if add > 0 {
106            let cap = self.capacity as i64;
107            // Clamp to capacity
108            let _ = self.tokens.fetch_update(Ordering::AcqRel, Ordering::Acquire, |cur| {
109                Some((cur + add).min(cap))
110            });
111        }
112
113        // Try to consume
114        let needed = (n.saturating_mul(1_000)) as i64;
115        self.tokens
116            .fetch_update(Ordering::AcqRel, Ordering::Acquire, |cur| {
117                if cur >= needed {
118                    Some(cur - needed)
119                } else {
120                    None
121                }
122            })
123            .is_ok()
124    }
125
126    /// Estimated milliseconds until `n` tokens are available.
127    pub fn retry_after_ms(&self, n: u64) -> u64 {
128        let cur = self.tokens.load(Ordering::Acquire);
129        let needed = (n.saturating_mul(1_000)) as i64;
130        if cur >= needed {
131            return 0;
132        }
133        let deficit = (needed - cur) as u64;
134        if self.refill_rate == 0 {
135            return u64::MAX;
136        }
137        deficit / self.refill_rate + 1
138    }
139}
140
141// ── SlidingWindow ─────────────────────────────────────────────────────────────
142
143/// Rolling-window request counter.
144///
145/// Tracks `(timestamp, token_count)` slots and enforces a maximum request
146/// count over the configured window.
147pub struct SlidingWindow {
148    /// Window width in milliseconds.
149    window_ms: u64,
150    /// Slots of `(arrival_instant, token_count)`.
151    slots: VecDeque<(Instant, u64)>,
152    /// Maximum number of requests allowed within the window.
153    max_count: u32,
154}
155
156impl SlidingWindow {
157    /// Create a new sliding window.
158    pub fn new(window_ms: u64, max_count: u32) -> Self {
159        Self {
160            window_ms,
161            slots: VecDeque::new(),
162            max_count,
163        }
164    }
165
166    /// Evict slots that have fallen outside the current window.
167    fn evict(&mut self) {
168        let now = Instant::now();
169        let cutoff_ms = self.window_ms;
170        while let Some(&(ts, _)) = self.slots.front() {
171            if now.duration_since(ts).as_millis() as u64 >= cutoff_ms {
172                self.slots.pop_front();
173            } else {
174                break;
175            }
176        }
177    }
178
179    /// Record a new request with the given token count.
180    pub fn record(&mut self, tokens: u64) {
181        self.evict();
182        self.slots.push_back((Instant::now(), tokens));
183    }
184
185    /// Returns `true` if a new request can be admitted (count would not
186    /// exceed `max_count` after adding it).
187    pub fn check(&mut self) -> bool {
188        self.evict();
189        (self.slots.len() as u32) < self.max_count
190    }
191
192    /// Milliseconds until the oldest slot expires (0 if window is empty).
193    pub fn reset_at_ms(&mut self) -> u64 {
194        self.evict();
195        match self.slots.front() {
196            None => 0,
197            Some(&(ts, _)) => {
198                let elapsed = Instant::now().duration_since(ts).as_millis() as u64;
199                self.window_ms.saturating_sub(elapsed) + 1
200            }
201        }
202    }
203}
204
205// ── ModelRateLimiter ──────────────────────────────────────────────────────────
206
207/// Combined rate limiter for a single model.
208///
209/// Checks both the token bucket and the sliding window on every call to
210/// [`ModelRateLimiter::check_and_consume`].
211pub struct ModelRateLimiter {
212    pub(crate) token_bucket: TokenBucket,
213    pub(crate) sliding_window: Mutex<SlidingWindow>,
214    /// Model identifier.
215    pub model_id: String,
216    /// Cumulative count of requests that were throttled (either limiter).
217    pub total_throttled: AtomicU64,
218}
219
220impl ModelRateLimiter {
221    /// Create a limiter from a [`RateLimiterConfig`].
222    pub fn new(model_id: String, config: &RateLimiterConfig) -> Self {
223        let capacity = (config.tokens_per_minute as f64 * config.burst_multiplier) as u64;
224        Self {
225            token_bucket: TokenBucket::new(capacity, config.tokens_per_minute),
226            sliding_window: Mutex::new(SlidingWindow::new(60_000, config.requests_per_minute)),
227            model_id,
228            total_throttled: AtomicU64::new(0),
229        }
230    }
231
232    /// Check whether a request consuming `tokens` may proceed.
233    ///
234    /// On success, the tokens are deducted from the bucket and the request is
235    /// recorded in the sliding window.  On failure, `total_throttled` is
236    /// incremented and the appropriate [`RateLimitError`] is returned.
237    pub fn check_and_consume(&self, tokens: u64) -> Result<(), RateLimitError> {
238        // Check sliding window first (cheaper)
239        {
240            let mut win = self.sliding_window.lock().unwrap_or_else(|e| e.into_inner());
241            if !win.check() {
242                let reset = win.reset_at_ms();
243                self.total_throttled.fetch_add(1, Ordering::Relaxed);
244                return Err(RateLimitError::WindowLimitExceeded { reset_at_ms: reset });
245            }
246        }
247
248        // Check token bucket
249        if !self.token_bucket.try_consume(tokens) {
250            let retry = self.token_bucket.retry_after_ms(tokens);
251            self.total_throttled.fetch_add(1, Ordering::Relaxed);
252            return Err(RateLimitError::TokenBucketExhausted {
253                retry_after_ms: retry,
254            });
255        }
256
257        // Record in sliding window
258        {
259            let mut win = self.sliding_window.lock().unwrap_or_else(|e| e.into_inner());
260            win.record(tokens);
261        }
262
263        Ok(())
264    }
265}
266
267// ── RateLimiterRegistry ───────────────────────────────────────────────────────
268
269/// Registry that holds one [`ModelRateLimiter`] per model.
270#[derive(Default)]
271pub struct RateLimiterRegistry {
272    limiters: RwLock<HashMap<String, ModelRateLimiter>>,
273}
274
275impl RateLimiterRegistry {
276    /// Create an empty registry.
277    pub fn new() -> Self {
278        Self::default()
279    }
280
281    /// Register a model with the given configuration.
282    ///
283    /// Overwrites any previous limiter for the same `model_id`.
284    pub fn register(&self, model_id: String, config: RateLimiterConfig) {
285        let limiter = ModelRateLimiter::new(model_id.clone(), &config);
286        self.limiters.write().unwrap_or_else(|e| e.into_inner()).insert(model_id, limiter);
287    }
288
289    /// Check and consume `tokens` for the given model.
290    ///
291    /// Returns [`RateLimitError::UnknownModel`] if the model has not been
292    /// registered.
293    pub fn check(&self, model_id: &str, tokens: u64) -> Result<(), RateLimitError> {
294        let guard = self.limiters.read().unwrap_or_else(|e| e.into_inner());
295        match guard.get(model_id) {
296            Some(lim) => lim.check_and_consume(tokens),
297            None => Err(RateLimitError::UnknownModel(model_id.to_owned())),
298        }
299    }
300
301    /// Return `(model_id, total_throttled)` pairs for all registered models.
302    pub fn stats(&self) -> Vec<(String, u64)> {
303        self.limiters
304            .read()
305            .unwrap_or_else(|e| e.into_inner())
306            .iter()
307            .map(|(id, lim)| (id.clone(), lim.total_throttled.load(Ordering::Relaxed)))
308            .collect()
309    }
310}
311
312#[cfg(test)]
313mod tests {
314    use super::*;
315
316    #[test]
317    fn token_bucket_consume_and_exhaust() {
318        let bucket = TokenBucket::new(10, 600); // 10 tokens capacity, 600/min
319        assert!(bucket.try_consume(5));
320        assert!(bucket.try_consume(5));
321        assert!(!bucket.try_consume(1)); // exhausted
322    }
323
324    #[test]
325    fn sliding_window_count_limit() {
326        let mut win = SlidingWindow::new(60_000, 3);
327        assert!(win.check());
328        win.record(1);
329        win.record(1);
330        win.record(1);
331        assert!(!win.check()); // 3 requests recorded, limit reached
332    }
333
334    #[test]
335    fn registry_unknown_model() {
336        let reg = RateLimiterRegistry::new();
337        assert!(matches!(
338            reg.check("unknown", 1),
339            Err(RateLimitError::UnknownModel(_))
340        ));
341    }
342
343    #[test]
344    fn registry_allows_then_throttles() {
345        let reg = RateLimiterRegistry::new();
346        reg.register(
347            "gpt-4o".to_string(),
348            RateLimiterConfig {
349                requests_per_minute: 2,
350                tokens_per_minute: 1_000,
351                burst_multiplier: 1.0,
352            },
353        );
354        assert!(reg.check("gpt-4o", 1).is_ok());
355        assert!(reg.check("gpt-4o", 1).is_ok());
356        // Third request exceeds sliding window limit of 2
357        assert!(reg.check("gpt-4o", 1).is_err());
358    }
359}