Skip to main content

fraiseql_auth/
rate_limiting.rs

1//! Rate limiting for brute-force and abuse protection.
2//!
3//! Provides [`KeyedRateLimiter`] — a per-key sliding-window counter backed by
4//! a [`DashMap`] — and [`RateLimiters`], a pre-built set of limiters for
5//! each authentication endpoint.
6// # Threading Model
7//
8// Per-key updates are **atomic** with respect to concurrent access:
9// - check() holds a per-shard write reference through the entire read-current-time → load-record →
10//   update-counter sequence
11// - Different keys land on different shards and never contend
12// - This prevents race conditions where multiple threads simultaneously exceed the limit on the
13//   *same* key
14// - Periodic sweeps and capacity eviction are best-effort and run without holding any other shard's
15//   lock
16
17use std::{
18    sync::{
19        Arc,
20        atomic::{AtomicU64, AtomicUsize, Ordering},
21    },
22    time::{SystemTime, UNIX_EPOCH},
23};
24
25use dashmap::DashMap;
26
27use crate::error::{AuthError, Result};
28
29/// Abstraction over the wall-clock used by [`KeyedRateLimiter`].
30///
31/// Implementations must be `Send + Sync` because the limiter is shared across
32/// request-handler threads. The trait is generic over the rate-limiter type
33/// parameter, so production builds inline `SystemClock` without any virtual
34/// dispatch or heap allocation.
35///
36/// A blanket impl exists for `F: Fn() -> u64 + Send + Sync`, so test code can
37/// pass closures and `fn` pointers (such as `|| u64::MAX`) directly.
38///
39/// # Implementation note — production vs. test divergence
40///
41/// **Production code must use [`SystemClock`]** (the default type parameter on
42/// `KeyedRateLimiter`).  `SystemClock` reads `SystemTime::now()` and is
43/// monotonic-enough for sliding-window rate limiting in practice (it only goes
44/// backwards on explicit time-source failure, which the impl downgrades to a
45/// frozen `0` — see the `SystemClock::now_unix_secs` rustdoc).
46///
47/// The blanket `impl<F: Fn() -> u64 + Send + Sync> Clock for F` exists purely
48/// for **test ergonomics** so a test can write
49/// `KeyedRateLimiter::with_clock(|| 1_000)` without defining a new struct.
50/// Closures carry **none** of `SystemClock`'s implicit guarantees:
51///
52/// - A closure that returns a constant (`|| 0`) freezes the sliding window — counts never expire,
53///   requests stack up until they hit `max_requests`.
54/// - A closure that returns non-monotonic values (`|| rand::random()`) makes window expiry
55///   unpredictable; tests written against it are flaky.
56/// - A closure that returns `u64::MAX` (the canonical "broken clock" test input) deliberately
57///   exercises the saturating-arithmetic branches.
58///
59/// When writing a test that needs to advance time, prefer the
60/// `Arc<AtomicU64>`-backed `move || atomic.load(Ordering::Relaxed)` pattern
61/// used throughout `crates/fraiseql-auth/src/tests.rs` and
62/// `crates/fraiseql-auth/tests/rate_limiter_time_tests.rs`.  It is monotonic
63/// by construction (the test code only stores larger values) and reads like
64/// a `MockClock` without needing a named type.
65///
66/// In short: **closure-as-clock is a documented test seam, not a production
67/// extension point**.  Code review should reject `with_clock(|| ...)` outside
68/// `#[cfg(test)]` modules and integration-test files.
69pub trait Clock: Send + Sync {
70    /// Return the current time as a Unix timestamp (seconds since the epoch).
71    fn now_unix_secs(&self) -> u64;
72}
73
74impl<F> Clock for F
75where
76    F: Fn() -> u64 + Send + Sync,
77{
78    fn now_unix_secs(&self) -> u64 {
79        self()
80    }
81}
82
83/// Production wall-clock that reads `SystemTime::now()`.
84///
85/// On system time error, returns `0` (fail-closed): a timestamp of `0` is
86/// before any real `window_start`, so existing windows will not expire and
87/// rate limiting continues to be enforced with existing counters. New windows
88/// started while the clock is broken will have `window_start = 0`; when the
89/// clock recovers, those windows will immediately expire (since any real
90/// timestamp ≥ `0 + window_secs`) and reset naturally.
91#[derive(Debug, Default, Clone, Copy)]
92pub struct SystemClock;
93
94impl Clock for SystemClock {
95    fn now_unix_secs(&self) -> u64 {
96        match SystemTime::now().duration_since(UNIX_EPOCH) {
97            Ok(duration) => duration.as_secs(),
98            Err(e) => {
99                tracing::warn!(
100                    error = %e,
101                    "System time error in rate limiter — brute-force protection \
102                     continues using frozen timestamps. System clock may have moved \
103                     backward or time source is unavailable."
104                );
105                // Return 0 (not u64::MAX): existing windows will not expire,
106                // so rate limiting remains enforced during the clock failure.
107                0
108            },
109        }
110    }
111}
112
113/// Rate limit configuration for authentication endpoints (sliding-window algorithm).
114///
115/// Uses a per-key sliding-window counter for brute-force protection on
116/// authentication endpoints (login, token refresh, callback).
117///
118/// Distinct from `fraiseql_server::middleware::RateLimitConfig`, which uses
119/// a token-bucket algorithm for general request rate limiting.
120#[derive(Debug, Clone)]
121pub struct AuthRateLimitConfig {
122    /// Whether rate limiting is enabled for this endpoint
123    pub enabled:      bool,
124    /// Maximum number of requests allowed in the window
125    pub max_requests: u32,
126    /// Window duration in seconds
127    pub window_secs:  u64,
128}
129
130impl AuthRateLimitConfig {
131    /// IP-based rate limiting for public endpoints
132    /// 100 requests per 60 seconds (typical for auth/start, auth/callback)
133    #[must_use]
134    pub const fn per_ip_standard() -> Self {
135        Self {
136            enabled:      true,
137            max_requests: 100,
138            window_secs:  60,
139        }
140    }
141
142    /// Stricter IP-based rate limiting for sensitive endpoints
143    /// 50 requests per 60 seconds
144    #[must_use]
145    pub const fn per_ip_strict() -> Self {
146        Self {
147            enabled:      true,
148            max_requests: 50,
149            window_secs:  60,
150        }
151    }
152
153    /// User-based rate limiting for authenticated endpoints
154    /// 10 requests per 60 seconds
155    #[must_use]
156    pub const fn per_user_standard() -> Self {
157        Self {
158            enabled:      true,
159            max_requests: 10,
160            window_secs:  60,
161        }
162    }
163
164    /// Failed login attempt limiting
165    /// 5 failed attempts per 3600 seconds (1 hour)
166    #[must_use]
167    pub const fn failed_login_attempts() -> Self {
168        Self {
169            enabled:      true,
170            max_requests: 5,
171            window_secs:  3600,
172        }
173    }
174}
175
176/// Request record for tracking
177#[derive(Debug, Clone)]
178struct RequestRecord {
179    /// Number of requests in current window
180    count:        u32,
181    /// Unix timestamp of window start
182    window_start: u64,
183}
184
185/// How often (in number of `check()` calls) expired entries are purged from the map.
186///
187/// Stale entries accumulate when keys stop sending requests.  Every
188/// `PURGE_INTERVAL` calls the limiter performs a full sweep and removes entries
189/// whose window has elapsed, bounding the HashMap's memory footprint.
190const PURGE_INTERVAL: u64 = 1_000;
191
192/// Default maximum number of unique keys the limiter will track simultaneously.
193///
194/// When the cap is reached, new keys are denied immediately and a warning is logged.
195/// This prevents an attacker from exhausting memory by sending requests from millions
196/// of unique IP addresses. The cap is conservative: 100k entries × ~100 bytes ≈ 10 MB.
197const DEFAULT_MAX_ENTRIES: usize = 100_000;
198
199/// Per-key sliding-window rate limiter backed by a [`DashMap`].
200///
201/// Each unique key (IP address, user ID, etc.) gets its own independent counter.
202/// The check-and-update sequence for a given key is atomic: no TOCTOU race can
203/// allow more requests than `max_requests` in any single window, even under
204/// high concurrency.  Distinct keys live on different shards and never block
205/// each other on the update path.
206///
207/// The map is capped at `DEFAULT_MAX_ENTRIES` keys: when an insert would push
208/// `len()` past the cap the entry with the oldest `window_start` is evicted
209/// first.  The cap is enforced **strictly** — the check, eviction, and insert
210/// for new keys all run inside a single `insert_guard` critical section, so
211/// `len()` never exceeds `max_entries` at any observable instant.  Updates to
212/// already-present keys take the lock-free fast path and never contend on
213/// `insert_guard`.
214///
215/// # Deployment note
216///
217/// This rate limiter is **per-process**. In a multi-replica deployment, each
218/// replica enforces the limit independently — the effective limit across *N*
219/// replicas is *N × limit*. For true distributed enforcement, configure a
220/// Redis-backed rate limiter via the `redis-rate-limiting` Cargo feature (see
221/// the fraiseql-observers queue feature for the integration pattern). Call
222/// [`warn_if_single_node_rate_limiting`] during server startup to emit a
223/// reminder when no distributed backend is detected.
224///
225/// # Type parameter
226///
227/// `C: Clock` selects the time source. Production code uses the default
228/// [`SystemClock`] (a zero-sized type) so the clock is inlined and no virtual
229/// dispatch or heap allocation occurs. Tests can substitute any closure or
230/// custom clock via [`KeyedRateLimiter::with_clock`].
231///
232/// # Constructors
233///
234/// - [`KeyedRateLimiter::new`] — use the system wall clock (production).
235/// - [`KeyedRateLimiter::with_clock`] — inject a custom clock (testing).
236/// - [`KeyedRateLimiter::with_clock_and_max_entries`] — custom clock + cap (testing).
237pub struct KeyedRateLimiter<C: Clock = SystemClock> {
238    records:      Arc<DashMap<String, RequestRecord>>,
239    config:       AuthRateLimitConfig,
240    max_entries:  usize,
241    /// Monotonically increasing call counter for triggering periodic sweeps.
242    check_count:  AtomicU64,
243    /// Authoritative size counter for `records`, maintained by every code path
244    /// that adds or removes an entry while `insert_guard` is held.  DashMap's
245    /// own `len()` sums per-shard counters without a global lock and can
246    /// momentarily disagree with the actual entry count under concurrent
247    /// writes; this counter is the source of truth for the cap check.
248    record_count: Arc<AtomicUsize>,
249    /// Serialises the (cap-check → evict → insert) sequence for **new** keys.
250    /// Updates to existing keys never acquire this lock; it is held only on
251    /// the slow path that grows `records` so the `max_entries` cap is enforced
252    /// strictly under concurrent insertion.
253    insert_guard: Arc<parking_lot::Mutex<()>>,
254    /// Time source — defaults to [`SystemClock`].
255    clock:        C,
256}
257
258impl<C: Clock + Clone> Clone for KeyedRateLimiter<C> {
259    fn clone(&self) -> Self {
260        Self {
261            records:      Arc::clone(&self.records),
262            config:       self.config.clone(),
263            max_entries:  self.max_entries,
264            check_count:  AtomicU64::new(self.check_count.load(Ordering::Relaxed)),
265            record_count: Arc::clone(&self.record_count),
266            insert_guard: Arc::clone(&self.insert_guard),
267            clock:        self.clock.clone(),
268        }
269    }
270}
271
272impl KeyedRateLimiter<SystemClock> {
273    /// Create a new keyed rate limiter using wall-clock time.
274    #[must_use]
275    pub fn new(config: AuthRateLimitConfig) -> Self {
276        Self::with_parts(config, DEFAULT_MAX_ENTRIES, SystemClock)
277    }
278
279    /// Create a rate limiter with a custom entry cap.
280    ///
281    /// Use this when the deployment context calls for a tighter or looser bound
282    /// than `DEFAULT_MAX_ENTRIES`.  Setting `max_entries = 0` disables the cap
283    /// (unbounded — not recommended in production).
284    #[must_use]
285    pub fn with_max_entries(config: AuthRateLimitConfig, max_entries: usize) -> Self {
286        Self::with_parts(config, max_entries, SystemClock)
287    }
288}
289
290impl<C: Clock> KeyedRateLimiter<C> {
291    /// Create a rate limiter with an injectable clock (for testing).
292    ///
293    /// The `clock`'s `now_unix_secs` method is called on every `check()` to
294    /// obtain the current Unix timestamp. Pass `|| u64::MAX` to simulate a
295    /// broken system clock and verify fail-open behavior.
296    pub fn with_clock(config: AuthRateLimitConfig, clock: C) -> Self {
297        Self::with_parts(config, DEFAULT_MAX_ENTRIES, clock)
298    }
299
300    /// Create a rate limiter with both a custom clock and a custom entry cap (for testing).
301    ///
302    /// Combines the benefits of [`KeyedRateLimiter::with_clock`] and
303    /// [`KeyedRateLimiter::with_max_entries`] for deterministic eviction tests.
304    pub fn with_clock_and_max_entries(
305        config: AuthRateLimitConfig,
306        max_entries: usize,
307        clock: C,
308    ) -> Self {
309        Self::with_parts(config, max_entries, clock)
310    }
311
312    fn with_parts(config: AuthRateLimitConfig, max_entries: usize, clock: C) -> Self {
313        Self {
314            records: Arc::new(DashMap::new()),
315            config,
316            max_entries,
317            check_count: AtomicU64::new(0),
318            record_count: Arc::new(AtomicUsize::new(0)),
319            insert_guard: Arc::new(parking_lot::Mutex::new(())),
320            clock,
321        }
322    }
323
324    /// Check if a request should be allowed for the given key
325    ///
326    /// # Atomicity
327    ///
328    /// The check-and-update step for a given key is **atomic**: while
329    /// inspecting and mutating the `RequestRecord` for `key`, this function
330    /// holds the per-shard write reference for that key.  No concurrent thread
331    /// can observe a partial state for the same key, which prevents the
332    /// classic TOCTOU race where multiple threads simultaneously exceed the
333    /// rate limit.
334    ///
335    /// # Capacity cap
336    ///
337    /// When `max_entries > 0` the map's length is enforced **strictly**:
338    /// new-key inserts run under a serialising `insert_guard`, so the
339    /// cap-check, oldest-entry eviction, and insert all occur in a single
340    /// critical section.  `records.len() <= max_entries` therefore holds at
341    /// every observable instant, including under sustained concurrent burst.
342    /// The fast path (updates to keys already present) does not acquire the
343    /// guard and runs lock-free.
344    ///
345    /// The periodic expiry sweep is best-effort and runs outside the guard;
346    /// it only ever shrinks `records`, so it cannot push `len()` over the cap.
347    ///
348    /// # Returns
349    ///
350    /// `Ok(())` if the request is allowed and the counter has been incremented.
351    ///
352    /// # Errors
353    ///
354    /// Returns [`AuthError::RateLimited`] if the key has exceeded the configured
355    /// rate limit within the sliding window.
356    pub fn check(&self, key: &str) -> Result<()> {
357        // If rate limiting is disabled, always allow the request.
358        if !self.config.enabled {
359            return Ok(());
360        }
361
362        let now = self.clock.now_unix_secs();
363
364        // Periodic expiry sweep to bound DashMap growth.  Runs every
365        // PURGE_INTERVAL calls; overflow wraps silently which is fine.
366        // Held under `insert_guard` so `record_count` updates stay coherent.
367        let count = self.check_count.fetch_add(1, Ordering::Relaxed);
368        if count.is_multiple_of(PURGE_INTERVAL) {
369            let _sweep_guard = self.insert_guard.lock();
370            let mut removed: usize = 0;
371            self.records.retain(|_, r| {
372                let keep = now < r.window_start.saturating_add(self.config.window_secs);
373                if !keep {
374                    removed = removed.saturating_add(1);
375                }
376                keep
377            });
378            if removed > 0 {
379                self.record_count.fetch_sub(removed, Ordering::Relaxed);
380            }
381        }
382
383        // Fast path: if the key is already present, update its record under
384        // the per-shard write lock without touching `insert_guard`.  The fast
385        // path does not change `record_count` (we mutate an existing entry).
386        if let Some(mut record) = self.records.get_mut(key) {
387            return Self::tick_existing(&mut record, &self.config, now);
388        }
389
390        // Slow path: we need to insert.  Serialise (cap-check → evict → insert)
391        // under `insert_guard` so concurrent inserters cannot race past
392        // `max_entries`.  A concurrent thread may have inserted this key while
393        // we waited on the guard — re-check before evicting.
394        let _insert_guard = self.insert_guard.lock();
395
396        // Re-check the fast path under the guard.  This handles the race
397        // where another thread inserted `key` while we waited on the lock.
398        // The `get_mut` guard is dropped at the end of the if-let block, so
399        // we never hold a per-shard lock when calling `iter()` / `insert()`
400        // below — that would deadlock on the shard hosting `key`.
401        if let Some(mut record) = self.records.get_mut(key) {
402            return Self::tick_existing(&mut record, &self.config, now);
403        }
404
405        // Enforce the cap using the authoritative `record_count` counter.
406        // DashMap's own `len()` is non-atomic across shards and can briefly
407        // under-report under concurrent writes, which would silently let the
408        // cap drift upward by a small amount.  `record_count` is updated only
409        // under `insert_guard`, so it is exact at this point.
410        if self.max_entries > 0 && self.record_count.load(Ordering::Relaxed) >= self.max_entries {
411            if let Some(oldest_key) = self
412                .records
413                .iter()
414                .min_by_key(|r| r.value().window_start)
415                .map(|r| r.key().clone())
416            {
417                if self.records.remove(&oldest_key).is_some() {
418                    self.record_count.fetch_sub(1, Ordering::Relaxed);
419                    tracing::debug!(
420                        max_entries = self.max_entries,
421                        "Rate limiter at capacity — evicted oldest entry to make room for new key"
422                    );
423                }
424            }
425        }
426
427        // First request from this key — start a fresh window.  `insert`
428        // returns `None` for a previously-absent key, which we count toward
429        // `record_count`; a `Some` return means we replaced an existing entry
430        // (no count change).
431        if self
432            .records
433            .insert(
434                key.to_string(),
435                RequestRecord {
436                    count:        1,
437                    window_start: now,
438                },
439            )
440            .is_none()
441        {
442            self.record_count.fetch_add(1, Ordering::Relaxed);
443        }
444
445        Ok(())
446    }
447
448    /// Update an already-present record for the current window.
449    ///
450    /// Extracted so both the fast (`get_mut`) path and the slow (`entry()`
451    /// under guard) path share identical sliding-window semantics.
452    const fn tick_existing(
453        record: &mut RequestRecord,
454        config: &AuthRateLimitConfig,
455        now: u64,
456    ) -> Result<()> {
457        if now >= record.window_start.saturating_add(config.window_secs) {
458            // CASE 1: Window has expired - start a new window
459            record.count = 1;
460            record.window_start = now;
461            Ok(())
462        } else if record.count < config.max_requests {
463            // CASE 2: Window is active and we haven't exceeded the limit
464            record.count += 1;
465            Ok(())
466        } else {
467            // CASE 3: Window is active and we've reached the limit
468            Err(AuthError::RateLimited {
469                retry_after_secs: config.window_secs,
470            })
471        }
472    }
473
474    /// Get the number of active rate limiters (for monitoring).
475    ///
476    /// Returns the authoritative entry count maintained under `insert_guard`,
477    /// not DashMap's per-shard sum.  Reads are lock-free and reflect the
478    /// post-mutation count at every observable instant.
479    pub fn active_limiters(&self) -> usize {
480        self.record_count.load(Ordering::Relaxed)
481    }
482
483    /// Clear all rate limiters (for testing or reset).
484    pub fn clear(&self) {
485        let _guard = self.insert_guard.lock();
486        self.records.clear();
487        self.record_count.store(0, Ordering::Relaxed);
488    }
489
490    /// Create a copy for independent testing
491    pub fn clone_config(&self) -> AuthRateLimitConfig {
492        self.config.clone()
493    }
494}
495
496/// Emit a startup warning when no distributed rate-limiting backend is configured.
497///
498/// Call once during server startup. If the `FRAISEQL_RATE_LIMIT_WARN_SINGLE_NODE`
499/// environment variable is set to `true` or `1` (case-insensitive) and the
500/// `FRAISEQL_RATE_LIMIT_BACKEND` variable is unset, a `warn!` is emitted reminding
501/// operators that each replica enforces limits independently — the effective limit
502/// across *N* replicas is *N × limit*.
503///
504/// This is a documentation-only reminder; it does not change runtime behaviour.
505pub fn warn_if_single_node_rate_limiting() {
506    let should_warn = std::env::var("FRAISEQL_RATE_LIMIT_WARN_SINGLE_NODE")
507        .map(|v| v.eq_ignore_ascii_case("true") || v == "1")
508        .unwrap_or(false);
509    let has_backend = std::env::var("FRAISEQL_RATE_LIMIT_BACKEND").is_ok();
510    if should_warn && !has_backend {
511        tracing::warn!(
512            "Rate limiter is per-process; multi-replica deployments are not protected against \
513             distributed brute-force. Configure a Redis-backed rate limiter via the \
514             `redis-rate-limiting` feature for distributed enforcement."
515        );
516    }
517}
518
519/// Global rate limiters for different endpoints
520pub struct RateLimiters {
521    /// auth/start: per-IP, 100 req/min
522    pub auth_start:    KeyedRateLimiter,
523    /// auth/callback: per-IP, 50 req/min
524    pub auth_callback: KeyedRateLimiter,
525    /// auth/refresh: per-user, 10 req/min
526    pub auth_refresh:  KeyedRateLimiter,
527    /// auth/logout: per-user, 20 req/min
528    pub auth_logout:   KeyedRateLimiter,
529    /// Failed login tracking: per-user, 5 attempts/hour
530    pub failed_logins: KeyedRateLimiter,
531}
532
533impl RateLimiters {
534    /// Create default rate limiters for all endpoints
535    #[must_use]
536    pub fn new() -> Self {
537        Self {
538            auth_start:    KeyedRateLimiter::new(AuthRateLimitConfig::per_ip_standard()),
539            auth_callback: KeyedRateLimiter::new(AuthRateLimitConfig::per_ip_strict()),
540            auth_refresh:  KeyedRateLimiter::new(AuthRateLimitConfig::per_user_standard()),
541            auth_logout:   KeyedRateLimiter::new(AuthRateLimitConfig::per_user_standard()),
542            failed_logins: KeyedRateLimiter::new(AuthRateLimitConfig::failed_login_attempts()),
543        }
544    }
545
546    /// Create with custom configurations
547    #[must_use]
548    pub fn with_configs(
549        start_cfg: AuthRateLimitConfig,
550        callback_cfg: AuthRateLimitConfig,
551        refresh_cfg: AuthRateLimitConfig,
552        logout_cfg: AuthRateLimitConfig,
553        failed_cfg: AuthRateLimitConfig,
554    ) -> Self {
555        Self {
556            auth_start:    KeyedRateLimiter::new(start_cfg),
557            auth_callback: KeyedRateLimiter::new(callback_cfg),
558            auth_refresh:  KeyedRateLimiter::new(refresh_cfg),
559            auth_logout:   KeyedRateLimiter::new(logout_cfg),
560            failed_logins: KeyedRateLimiter::new(failed_cfg),
561        }
562    }
563}
564
565impl Default for RateLimiters {
566    fn default() -> Self {
567        Self::new()
568    }
569}