tokio_prompt_orchestrator/
rate_limiter.rs1use 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#[derive(Debug, Error)]
20pub enum RateLimitError {
21 #[error("token bucket exhausted; retry after {retry_after_ms} ms")]
24 TokenBucketExhausted {
25 retry_after_ms: u64,
27 },
28 #[error("sliding window limit exceeded; window resets in {reset_at_ms} ms")]
31 WindowLimitExceeded {
32 reset_at_ms: u64,
34 },
35 #[error("no rate limiter registered for model '{0}'")]
37 UnknownModel(String),
38}
39
40#[derive(Debug, Clone)]
44pub struct RateLimiterConfig {
45 pub requests_per_minute: u32,
47 pub tokens_per_minute: u64,
49 pub burst_multiplier: f64,
52}
53
54pub struct TokenBucket {
62 capacity: u64,
64 tokens: AtomicI64,
66 refill_rate: u64,
68 last_refill: Mutex<Instant>,
70}
71
72impl TokenBucket {
73 pub fn new(capacity_tokens: u64, tokens_per_minute: u64) -> Self {
76 let capacity_scaled = capacity_tokens.saturating_mul(1_000);
77 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 pub fn try_consume(&self, n: u64) -> bool {
93 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 let _ = self.tokens.fetch_update(Ordering::AcqRel, Ordering::Acquire, |cur| {
109 Some((cur + add).min(cap))
110 });
111 }
112
113 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 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
141pub struct SlidingWindow {
148 window_ms: u64,
150 slots: VecDeque<(Instant, u64)>,
152 max_count: u32,
154}
155
156impl SlidingWindow {
157 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 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 pub fn record(&mut self, tokens: u64) {
181 self.evict();
182 self.slots.push_back((Instant::now(), tokens));
183 }
184
185 pub fn check(&mut self) -> bool {
188 self.evict();
189 (self.slots.len() as u32) < self.max_count
190 }
191
192 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
205pub struct ModelRateLimiter {
212 pub(crate) token_bucket: TokenBucket,
213 pub(crate) sliding_window: Mutex<SlidingWindow>,
214 pub model_id: String,
216 pub total_throttled: AtomicU64,
218}
219
220impl ModelRateLimiter {
221 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 pub fn check_and_consume(&self, tokens: u64) -> Result<(), RateLimitError> {
238 {
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 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 {
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#[derive(Default)]
271pub struct RateLimiterRegistry {
272 limiters: RwLock<HashMap<String, ModelRateLimiter>>,
273}
274
275impl RateLimiterRegistry {
276 pub fn new() -> Self {
278 Self::default()
279 }
280
281 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 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 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); assert!(bucket.try_consume(5));
320 assert!(bucket.try_consume(5));
321 assert!(!bucket.try_consume(1)); }
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()); }
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 assert!(reg.check("gpt-4o", 1).is_err());
358 }
359}