Skip to main content

stateset_authz/
rate_limit.rs

1//! Window-based rate limiting.
2//!
3//! Provides a simple, IO-free rate limiter that tracks per-actor, per-resource
4//! request counts within configurable time windows.
5
6use std::collections::HashMap;
7use std::fmt;
8use std::time::{Duration, Instant};
9
10use serde::{Deserialize, Serialize};
11
12/// Configuration for a single rate limit rule.
13///
14/// ```rust
15/// use stateset_authz::RateLimitRule;
16/// use std::time::Duration;
17///
18/// let rule = RateLimitRule::new("orders", 100, Duration::from_secs(60));
19/// assert_eq!(rule.resource_type(), "orders");
20/// assert_eq!(rule.max_requests(), 100);
21/// assert_eq!(rule.window(), Duration::from_secs(60));
22/// ```
23#[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    /// Creates a new rate limit rule.
33    #[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    /// Returns the resource type this rule applies to.
39    #[must_use]
40    pub fn resource_type(&self) -> &str {
41        &self.resource_type
42    }
43
44    /// Returns the maximum number of requests allowed per window.
45    #[must_use]
46    pub const fn max_requests(&self) -> u32 {
47        self.max_requests
48    }
49
50    /// Returns the time window duration.
51    #[must_use]
52    pub const fn window(&self) -> Duration {
53        self.window
54    }
55}
56
57/// The result of a rate limit check.
58///
59/// ```rust
60/// use stateset_authz::RateLimitDecision;
61/// use std::time::Duration;
62///
63/// let allowed = RateLimitDecision::Allowed { remaining: 5 };
64/// assert!(allowed.is_allowed());
65///
66/// let exceeded = RateLimitDecision::Exceeded { retry_after: Duration::from_secs(30) };
67/// assert!(!exceeded.is_allowed());
68/// ```
69#[derive(Debug, Clone, PartialEq, Eq)]
70#[non_exhaustive]
71pub enum RateLimitDecision {
72    /// The request is within limits.
73    Allowed {
74        /// How many requests remain in the current window.
75        remaining: u32,
76    },
77    /// The rate limit has been exceeded.
78    Exceeded {
79        /// How long until the next request would be allowed.
80        retry_after: Duration,
81    },
82}
83
84impl RateLimitDecision {
85    /// Returns `true` if the request is allowed.
86    #[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/// Tracks request timestamps for a single actor+resource bucket.
104#[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    /// Removes timestamps outside the window and returns the count within the window.
115    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    /// Returns the oldest timestamp in the window, if any.
126    fn oldest(&self) -> Option<Instant> {
127        self.requests.first().copied()
128    }
129}
130
131/// Composite key for per-actor, per-resource state lookups.
132#[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/// A window-based rate limiter.
149///
150/// ```rust
151/// use stateset_authz::{RateLimiter, RateLimitRule};
152/// use std::time::Duration;
153///
154/// let mut limiter = RateLimiter::new();
155/// limiter.add_rule(RateLimitRule::new("orders", 2, Duration::from_secs(60)));
156///
157/// let d1 = limiter.check_and_record("actor-1", "orders");
158/// assert!(d1.is_allowed());
159///
160/// let d2 = limiter.check_and_record("actor-1", "orders");
161/// assert!(d2.is_allowed());
162///
163/// let d3 = limiter.check_and_record("actor-1", "orders");
164/// assert!(!d3.is_allowed());
165/// ```
166#[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    /// Run global stale-entry cleanup every N operations.
175    const AUTO_CLEANUP_INTERVAL_OPS: u16 = 1024;
176
177    /// Creates an empty rate limiter with no rules.
178    #[must_use]
179    pub fn new() -> Self {
180        Self { rules: HashMap::new(), state: HashMap::new(), ops_since_cleanup: 0 }
181    }
182
183    /// Adds a rate limit rule. If a rule for the same resource type already exists,
184    /// it is replaced.
185    pub fn add_rule(&mut self, rule: RateLimitRule) {
186        self.rules.insert(rule.resource_type.clone(), rule);
187    }
188
189    /// Checks whether a request from `actor_id` for `resource_type` is within limits,
190    /// **without** recording the request. Use [`check_and_record`](Self::check_and_record)
191    /// to atomically check and record.
192    #[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    /// Checks and records a request in one step.
199    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    /// Records a request without checking. Useful when the decision has already
205    /// been made externally.
206    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    /// Removes expired entries from state. Call periodically for long-lived limiters.
212    pub fn cleanup(&mut self) {
213        let now = Instant::now();
214        self.state.retain(|key, state| {
215            // Find the applicable window; if no rule, drop the entry.
216            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    /// Returns the number of rules configured.
226    #[must_use]
227    pub fn rule_count(&self) -> usize {
228        self.rules.len()
229    }
230
231    // -- Internal helpers with injectable `now` for testing --
232
233    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            // No rule means no limit
244            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            // Adjust remaining to reflect the state *after* recording
273            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
293/// Serde helpers for `Duration` (as milliseconds).
294mod 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            // oldest was `now`, window is 60s, so retry_after ≈ 60s
370            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        // After the window expires, should be allowed again
387        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        // Alice is at limit
402        assert!(!limiter.check_at("alice", "orders", now).is_allowed());
403
404        // Bob still has room
405        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        // Orders at limit
418        assert!(!limiter.check_at("a", "orders", now).is_allowed());
419
420        // Customers still fine
421        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        // Second check exceeds — should NOT record
444        assert!(!limiter.check_and_record_at("a", "orders", now).is_allowed());
445
446        // After window, should be allowed (only 1 recorded, not 2)
447        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        // Record some entries that will be old
461        let old = Instant::now();
462        limiter.record_at("a", "orders", old);
463
464        // Simulate time passing — cleanup after window
465        // (In real code, cleanup() uses Instant::now(), but we can at least test it runs.)
466        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        // If cleanup parses the state key incorrectly, it drops the entry and this would become allowed.
479        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        // Each tuple should be independently limited to 1.
493        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}