1use std::collections::HashMap;
7use std::fmt;
8use std::time::{Duration, Instant};
9
10use serde::{Deserialize, Serialize};
11
12#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
24pub struct RateLimitRule {
25 resource_type: String,
26 max_requests: u32,
27 #[serde(with = "duration_serde")]
28 window: Duration,
29}
30
31impl RateLimitRule {
32 #[must_use]
34 pub fn new(resource_type: impl Into<String>, max_requests: u32, window: Duration) -> Self {
35 Self { resource_type: resource_type.into(), max_requests, window }
36 }
37
38 #[must_use]
40 pub fn resource_type(&self) -> &str {
41 &self.resource_type
42 }
43
44 #[must_use]
46 pub const fn max_requests(&self) -> u32 {
47 self.max_requests
48 }
49
50 #[must_use]
52 pub const fn window(&self) -> Duration {
53 self.window
54 }
55}
56
57#[derive(Debug, Clone, PartialEq, Eq)]
70#[non_exhaustive]
71pub enum RateLimitDecision {
72 Allowed {
74 remaining: u32,
76 },
77 Exceeded {
79 retry_after: Duration,
81 },
82}
83
84impl RateLimitDecision {
85 #[must_use]
87 pub const fn is_allowed(&self) -> bool {
88 matches!(self, Self::Allowed { .. })
89 }
90}
91
92impl fmt::Display for RateLimitDecision {
93 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
94 match self {
95 Self::Allowed { remaining } => write!(f, "allowed ({remaining} remaining)"),
96 Self::Exceeded { retry_after } => {
97 write!(f, "exceeded (retry after {}ms)", retry_after.as_millis())
98 }
99 }
100 }
101}
102
103#[derive(Debug, Clone)]
105struct RateLimitState {
106 requests: Vec<Instant>,
107}
108
109impl RateLimitState {
110 const fn new() -> Self {
111 Self { requests: Vec::new() }
112 }
113
114 fn cleanup_and_count(&mut self, window: Duration, now: Instant) -> usize {
116 let cutoff = now.checked_sub(window).unwrap_or(now);
117 self.requests.retain(|&t| t > cutoff);
118 self.requests.len()
119 }
120
121 fn record(&mut self, now: Instant) {
122 self.requests.push(now);
123 }
124
125 fn oldest(&self) -> Option<Instant> {
127 self.requests.first().copied()
128 }
129}
130
131#[derive(Debug, Clone, PartialEq, Eq, Hash)]
133struct StateKey {
134 actor_id: String,
135 resource_type: String,
136}
137
138impl StateKey {
139 fn new(actor_id: &str, resource_type: &str) -> Self {
140 Self { actor_id: actor_id.to_owned(), resource_type: resource_type.to_owned() }
141 }
142}
143
144fn state_key(actor_id: &str, resource_type: &str) -> StateKey {
145 StateKey::new(actor_id, resource_type)
146}
147
148#[derive(Debug)]
167pub struct RateLimiter {
168 rules: HashMap<String, RateLimitRule>,
169 state: HashMap<StateKey, RateLimitState>,
170 ops_since_cleanup: u16,
171}
172
173impl RateLimiter {
174 const AUTO_CLEANUP_INTERVAL_OPS: u16 = 1024;
176
177 #[must_use]
179 pub fn new() -> Self {
180 Self { rules: HashMap::new(), state: HashMap::new(), ops_since_cleanup: 0 }
181 }
182
183 pub fn add_rule(&mut self, rule: RateLimitRule) {
186 self.rules.insert(rule.resource_type.clone(), rule);
187 }
188
189 #[must_use]
193 pub fn check(&mut self, actor_id: &str, resource_type: &str) -> RateLimitDecision {
194 self.maybe_cleanup();
195 self.check_at(actor_id, resource_type, Instant::now())
196 }
197
198 pub fn check_and_record(&mut self, actor_id: &str, resource_type: &str) -> RateLimitDecision {
200 self.maybe_cleanup();
201 self.check_and_record_at(actor_id, resource_type, Instant::now())
202 }
203
204 pub fn record(&mut self, actor_id: &str, resource_type: &str) {
207 self.maybe_cleanup();
208 self.record_at(actor_id, resource_type, Instant::now());
209 }
210
211 pub fn cleanup(&mut self) {
213 let now = Instant::now();
214 self.state.retain(|key, state| {
215 if let Some(rule) = self.rules.get(key.resource_type.as_str()) {
217 state.cleanup_and_count(rule.window, now);
218 !state.requests.is_empty()
219 } else {
220 false
221 }
222 });
223 }
224
225 #[must_use]
227 pub fn rule_count(&self) -> usize {
228 self.rules.len()
229 }
230
231 fn maybe_cleanup(&mut self) {
234 self.ops_since_cleanup = self.ops_since_cleanup.saturating_add(1);
235 if self.ops_since_cleanup >= Self::AUTO_CLEANUP_INTERVAL_OPS {
236 self.cleanup();
237 self.ops_since_cleanup = 0;
238 }
239 }
240
241 fn check_at(&mut self, actor_id: &str, resource_type: &str, now: Instant) -> RateLimitDecision {
242 let Some(rule) = self.rules.get(resource_type) else {
243 return RateLimitDecision::Allowed { remaining: u32::MAX };
245 };
246
247 let key = state_key(actor_id, resource_type);
248 let state = self.state.entry(key).or_insert_with(RateLimitState::new);
249 let count = state.cleanup_and_count(rule.window, now) as u32;
250
251 if count >= rule.max_requests {
252 let retry_after = state.oldest().map_or(rule.window, |oldest| {
253 let window_end = oldest + rule.window;
254 window_end.saturating_duration_since(now)
255 });
256
257 RateLimitDecision::Exceeded { retry_after }
258 } else {
259 RateLimitDecision::Allowed { remaining: rule.max_requests - count }
260 }
261 }
262
263 fn check_and_record_at(
264 &mut self,
265 actor_id: &str,
266 resource_type: &str,
267 now: Instant,
268 ) -> RateLimitDecision {
269 let decision = self.check_at(actor_id, resource_type, now);
270 if decision.is_allowed() {
271 self.record_at(actor_id, resource_type, now);
272 if let RateLimitDecision::Allowed { remaining } = decision {
274 return RateLimitDecision::Allowed { remaining: remaining.saturating_sub(1) };
275 }
276 }
277 decision
278 }
279
280 fn record_at(&mut self, actor_id: &str, resource_type: &str, now: Instant) {
281 let key = state_key(actor_id, resource_type);
282 let state = self.state.entry(key).or_insert_with(RateLimitState::new);
283 state.record(now);
284 }
285}
286
287impl Default for RateLimiter {
288 fn default() -> Self {
289 Self::new()
290 }
291}
292
293mod duration_serde {
295 use std::time::Duration;
296
297 use serde::{self, Deserialize, Deserializer, Serializer};
298
299 pub(super) fn serialize<S>(duration: &Duration, serializer: S) -> Result<S::Ok, S::Error>
300 where
301 S: Serializer,
302 {
303 serializer.serialize_u64(duration.as_millis() as u64)
304 }
305
306 pub(super) fn deserialize<'de, D>(deserializer: D) -> Result<Duration, D::Error>
307 where
308 D: Deserializer<'de>,
309 {
310 let ms = u64::deserialize(deserializer)?;
311 Ok(Duration::from_millis(ms))
312 }
313}
314
315#[cfg(test)]
316mod tests {
317 use super::*;
318
319 fn rule_2_per_60s() -> RateLimitRule {
320 RateLimitRule::new("orders", 2, Duration::from_secs(60))
321 }
322
323 #[test]
324 fn no_rule_means_no_limit() {
325 let mut limiter = RateLimiter::new();
326 let d = limiter.check("actor-1", "orders");
327 assert!(d.is_allowed());
328 if let RateLimitDecision::Allowed { remaining } = d {
329 assert_eq!(remaining, u32::MAX);
330 }
331 }
332
333 #[test]
334 fn under_limit() {
335 let mut limiter = RateLimiter::new();
336 limiter.add_rule(rule_2_per_60s());
337
338 let d = limiter.check_and_record("actor-1", "orders");
339 assert!(d.is_allowed());
340 if let RateLimitDecision::Allowed { remaining } = d {
341 assert_eq!(remaining, 1);
342 }
343 }
344
345 #[test]
346 fn at_limit() {
347 let mut limiter = RateLimiter::new();
348 limiter.add_rule(rule_2_per_60s());
349
350 let now = Instant::now();
351 limiter.record_at("a", "orders", now);
352 limiter.record_at("a", "orders", now);
353
354 let d = limiter.check_at("a", "orders", now);
355 assert!(!d.is_allowed());
356 }
357
358 #[test]
359 fn over_limit_shows_retry_after() {
360 let mut limiter = RateLimiter::new();
361 limiter.add_rule(rule_2_per_60s());
362
363 let now = Instant::now();
364 limiter.record_at("a", "orders", now);
365 limiter.record_at("a", "orders", now);
366
367 let d = limiter.check_at("a", "orders", now);
368 if let RateLimitDecision::Exceeded { retry_after } = d {
369 assert!(retry_after.as_secs() <= 60);
371 assert!(retry_after.as_secs() >= 59);
372 } else {
373 panic!("expected exceeded");
374 }
375 }
376
377 #[test]
378 fn window_expiry() {
379 let mut limiter = RateLimiter::new();
380 limiter.add_rule(rule_2_per_60s());
381
382 let start = Instant::now();
383 limiter.record_at("a", "orders", start);
384 limiter.record_at("a", "orders", start);
385
386 let after_window = start + Duration::from_secs(61);
388 let d = limiter.check_at("a", "orders", after_window);
389 assert!(d.is_allowed());
390 }
391
392 #[test]
393 fn multiple_actors_independent() {
394 let mut limiter = RateLimiter::new();
395 limiter.add_rule(rule_2_per_60s());
396
397 let now = Instant::now();
398 limiter.record_at("alice", "orders", now);
399 limiter.record_at("alice", "orders", now);
400
401 assert!(!limiter.check_at("alice", "orders", now).is_allowed());
403
404 assert!(limiter.check_at("bob", "orders", now).is_allowed());
406 }
407
408 #[test]
409 fn multiple_resources_independent() {
410 let mut limiter = RateLimiter::new();
411 limiter.add_rule(RateLimitRule::new("orders", 1, Duration::from_secs(60)));
412 limiter.add_rule(RateLimitRule::new("customers", 1, Duration::from_secs(60)));
413
414 let now = Instant::now();
415 limiter.record_at("a", "orders", now);
416
417 assert!(!limiter.check_at("a", "orders", now).is_allowed());
419
420 assert!(limiter.check_at("a", "customers", now).is_allowed());
422 }
423
424 #[test]
425 fn check_and_record_blocks_after_limit() {
426 let mut limiter = RateLimiter::new();
427 limiter.add_rule(RateLimitRule::new("orders", 3, Duration::from_secs(60)));
428
429 let now = Instant::now();
430 assert!(limiter.check_and_record_at("a", "orders", now).is_allowed());
431 assert!(limiter.check_and_record_at("a", "orders", now).is_allowed());
432 assert!(limiter.check_and_record_at("a", "orders", now).is_allowed());
433 assert!(!limiter.check_and_record_at("a", "orders", now).is_allowed());
434 }
435
436 #[test]
437 fn check_and_record_does_not_record_on_exceed() {
438 let mut limiter = RateLimiter::new();
439 limiter.add_rule(RateLimitRule::new("orders", 1, Duration::from_secs(60)));
440
441 let now = Instant::now();
442 assert!(limiter.check_and_record_at("a", "orders", now).is_allowed());
443 assert!(!limiter.check_and_record_at("a", "orders", now).is_allowed());
445
446 let later = now + Duration::from_secs(61);
448 let d = limiter.check_at("a", "orders", later);
449 assert!(d.is_allowed());
450 if let RateLimitDecision::Allowed { remaining } = d {
451 assert_eq!(remaining, 1);
452 }
453 }
454
455 #[test]
456 fn cleanup_removes_expired() {
457 let mut limiter = RateLimiter::new();
458 limiter.add_rule(RateLimitRule::new("orders", 10, Duration::from_secs(1)));
459
460 let old = Instant::now();
462 limiter.record_at("a", "orders", old);
463
464 limiter.cleanup();
467 }
468
469 #[test]
470 fn cleanup_handles_colons_in_actor_id() {
471 let mut limiter = RateLimiter::new();
472 limiter.add_rule(RateLimitRule::new("orders", 1, Duration::from_secs(60)));
473
474 let now = Instant::now();
475 limiter.record_at("tenant:alice", "orders", now);
476 limiter.cleanup();
477
478 assert!(!limiter.check_at("tenant:alice", "orders", now).is_allowed());
480 }
481
482 #[test]
483 fn actor_and_resource_with_colons_use_distinct_buckets() {
484 let mut limiter = RateLimiter::new();
485 limiter.add_rule(RateLimitRule::new("c", 1, Duration::from_secs(60)));
486 limiter.add_rule(RateLimitRule::new("b:c", 1, Duration::from_secs(60)));
487
488 let now = Instant::now();
489 assert!(limiter.check_and_record_at("a:b", "c", now).is_allowed());
490 assert!(limiter.check_and_record_at("a", "b:c", now).is_allowed());
491
492 assert!(!limiter.check_and_record_at("a:b", "c", now).is_allowed());
494 assert!(!limiter.check_and_record_at("a", "b:c", now).is_allowed());
495 }
496
497 #[test]
498 fn rule_count() {
499 let mut limiter = RateLimiter::new();
500 assert_eq!(limiter.rule_count(), 0);
501 limiter.add_rule(rule_2_per_60s());
502 assert_eq!(limiter.rule_count(), 1);
503 }
504
505 #[test]
506 fn rule_replacement() {
507 let mut limiter = RateLimiter::new();
508 limiter.add_rule(RateLimitRule::new("orders", 5, Duration::from_secs(60)));
509 limiter.add_rule(RateLimitRule::new("orders", 10, Duration::from_secs(120)));
510 assert_eq!(limiter.rule_count(), 1);
511 }
512
513 #[test]
514 fn display_allowed() {
515 let d = RateLimitDecision::Allowed { remaining: 5 };
516 assert_eq!(d.to_string(), "allowed (5 remaining)");
517 }
518
519 #[test]
520 fn display_exceeded() {
521 let d = RateLimitDecision::Exceeded { retry_after: Duration::from_secs(30) };
522 assert_eq!(d.to_string(), "exceeded (retry after 30000ms)");
523 }
524
525 #[test]
526 fn rule_serde_roundtrip() {
527 let rule = RateLimitRule::new("orders", 100, Duration::from_secs(60));
528 let json = serde_json::to_string(&rule).unwrap();
529 let parsed: RateLimitRule = serde_json::from_str(&json).unwrap();
530 assert_eq!(parsed, rule);
531 }
532
533 #[test]
534 fn rule_accessors() {
535 let rule = RateLimitRule::new("test", 42, Duration::from_millis(500));
536 assert_eq!(rule.resource_type(), "test");
537 assert_eq!(rule.max_requests(), 42);
538 assert_eq!(rule.window(), Duration::from_millis(500));
539 }
540
541 #[test]
542 fn default_impl() {
543 let limiter = RateLimiter::default();
544 assert_eq!(limiter.rule_count(), 0);
545 }
546
547 #[test]
548 fn auto_cleanup_runs_on_operation_threshold() {
549 let mut limiter = RateLimiter::new();
550 limiter.add_rule(rule_2_per_60s());
551
552 let base = Instant::now();
553 limiter.record_at("stale", "orders", base - Duration::from_secs(120));
554 limiter.record_at("fresh", "orders", base);
555 assert_eq!(limiter.state.len(), 2);
556
557 limiter.ops_since_cleanup = RateLimiter::AUTO_CLEANUP_INTERVAL_OPS - 1;
558 let _ = limiter.check("fresh", "orders");
559
560 assert!(limiter.state.contains_key(&state_key("fresh", "orders")));
561 assert!(!limiter.state.contains_key(&state_key("stale", "orders")));
562 assert_eq!(limiter.ops_since_cleanup, 0);
563 }
564}