Skip to main content

fastmcp_server/
rate_limiting.rs

1//! Rate limiting middleware for protecting FastMCP servers from abuse.
2//!
3//! This module provides two rate limiting strategies:
4//!
5//! - [`RateLimitingMiddleware`]: Token bucket algorithm for burst-friendly limits
6//! - [`SlidingWindowRateLimitingMiddleware`]: Sliding window for precise tracking
7//!
8//! # Example
9//!
10//! ```ignore
11//! use fastmcp_rust::prelude::*;
12//! use fastmcp_rust::rate_limiting::RateLimitingMiddleware;
13//!
14//! // Allow 10 requests per second with bursts up to 20
15//! let rate_limiter = RateLimitingMiddleware::new(10.0)
16//!     .burst_capacity(20);
17//!
18//! Server::new("my-server", "1.0.0")
19//!     .middleware(rate_limiter)
20//!     .build()
21//!     .run_stdio();
22//! ```
23
24use std::collections::{BTreeSet, HashMap, VecDeque};
25use std::sync::Mutex;
26use std::time::{Duration, Instant};
27
28use fastmcp_core::{
29    McpContext, McpError, McpErrorCode, McpResult, SHA256_DIGEST_BYTES, Sha256Digest,
30    sha256_bounded,
31};
32use fastmcp_protocol::JsonRpcRequest;
33
34use crate::{Middleware, MiddlewareDecision};
35
36/// Error code for rate limit exceeded (-32005).
37///
38/// This is in the MCP server error range (-32000 to -32099).
39pub const RATE_LIMIT_ERROR_CODE: i32 = -32005;
40
41/// Maximum accepted byte length for a custom client identifier.
42///
43/// Identifiers are hashed immediately after this bound is enforced. The raw
44/// value is never retained by the middleware or included in diagnostics.
45const MAX_CLIENT_ID_BYTES: usize = 4096;
46
47/// Maximum number of limiter partitions, including the shared overflow
48/// partition.
49const MAX_CLIENT_PARTITIONS: usize = 4096;
50
51/// One partition is permanently reserved for identifiers first observed after
52/// the named-partition cap has been reached.
53const MAX_NAMED_CLIENT_PARTITIONS: usize = MAX_CLIENT_PARTITIONS - 1;
54
55/// Minimum inactivity period before a dedicated partition may be reclaimed.
56/// Reclamation additionally requires the limiter to have naturally returned
57/// to its initial state, so eviction cannot reset an active limit.
58const CLIENT_PARTITION_IDLE_TTL: Duration = Duration::from_secs(60);
59
60const DEFAULT_CLIENT_ID: &[u8] = b"fastmcp-default-rate-limit-partition";
61const RATE_LIMIT_EXCEEDED_MESSAGE: &str = "Rate limit exceeded";
62const RATE_LIMIT_METHOD_PARTITION_DOMAIN: &[u8] = b"fastmcp-rate-limit-method-partition-v1\0";
63const MAX_RATE_LIMIT_METHOD_BYTES: usize = 512;
64const MAX_RATE_LIMIT_PARTITION_INPUT_BYTES: usize = RATE_LIMIT_METHOD_PARTITION_DOMAIN.len()
65    + SHA256_DIGEST_BYTES
66    + std::mem::size_of::<u64>()
67    + MAX_RATE_LIMIT_METHOD_BYTES;
68
69// Token counts are stored as `f64`. Above 2^53 - 1, subtracting one can round
70// back to the original value and silently turn a configured bucket into an
71// effectively unlimited one. On 32-bit targets this cast naturally yields
72// `usize::MAX`, where every `usize` remains exactly representable for the
73// integer operations used here.
74const MAX_EXACT_TOKEN_CAPACITY: usize = ((1_u64 << 53) - 1) as usize;
75
76#[derive(Debug, Clone, Copy, PartialEq, Eq)]
77enum RateLimitAdmission {
78    Allowed,
79    Rejected { retry_after_ms: u64 },
80}
81
82impl RateLimitAdmission {
83    const fn is_allowed(self) -> bool {
84        matches!(self, Self::Allowed)
85    }
86}
87
88fn sanitized_rate(rate: f64) -> f64 {
89    if rate.is_finite() && rate > 0.0 {
90        rate
91    } else {
92        0.0
93    }
94}
95
96fn default_burst_capacity(rate: f64) -> usize {
97    if rate <= 0.0 {
98        return 0;
99    }
100
101    let doubled = rate * 2.0;
102    if !doubled.is_finite() || doubled > MAX_EXACT_TOKEN_CAPACITY as f64 {
103        0
104    } else {
105        // Every accepted positive rate must admit at least one initial
106        // request. Truncating a fractional default below 0.5 requests/second
107        // to zero otherwise creates a limiter that can never recover.
108        (doubled as usize).max(1)
109    }
110}
111
112fn default_client_partition() -> McpResult<Sha256Digest> {
113    sha256_bounded(DEFAULT_CLIENT_ID, MAX_CLIENT_ID_BYTES)
114        .map_err(|_| rate_limit_error(RATE_LIMIT_EXCEEDED_MESSAGE))
115}
116
117fn rate_limit_method_partition(
118    client_partition: Sha256Digest,
119    method: &str,
120) -> McpResult<Sha256Digest> {
121    if method.len() > MAX_RATE_LIMIT_METHOD_BYTES {
122        return Err(rate_limit_error(RATE_LIMIT_EXCEEDED_MESSAGE));
123    }
124
125    let method_len =
126        u64::try_from(method.len()).map_err(|_| rate_limit_error(RATE_LIMIT_EXCEEDED_MESSAGE))?;
127    let mut input = Vec::with_capacity(MAX_RATE_LIMIT_PARTITION_INPUT_BYTES);
128    input.extend_from_slice(RATE_LIMIT_METHOD_PARTITION_DOMAIN);
129    input.extend_from_slice(client_partition.as_bytes());
130    input.extend_from_slice(&method_len.to_be_bytes());
131    input.extend_from_slice(method.as_bytes());
132    sha256_bounded(&input, MAX_RATE_LIMIT_PARTITION_INPUT_BYTES)
133        .map_err(|_| rate_limit_error(RATE_LIMIT_EXCEEDED_MESSAGE))
134}
135
136fn retry_after_millis(deficit: f64, refill_rate: f64) -> u64 {
137    if !deficit.is_finite() || !refill_rate.is_finite() || deficit <= 0.0 || refill_rate <= 0.0 {
138        return u64::MAX;
139    }
140
141    let millis = (deficit / refill_rate) * 1_000.0;
142    if !millis.is_finite() || millis >= u64::MAX as f64 {
143        u64::MAX
144    } else {
145        (millis.ceil() as u64).max(1)
146    }
147}
148
149fn rate_limit_retry_error(request: &JsonRpcRequest, retry_after_ms: u64) -> McpError {
150    McpError::with_data(
151        McpErrorCode::Custom(RATE_LIMIT_ERROR_CODE),
152        RATE_LIMIT_EXCEEDED_MESSAGE,
153        serde_json::json!({
154            "method": request.method.clone(),
155            "requestId": request.id.clone(),
156            "retryAfterMs": retry_after_ms,
157        }),
158    )
159}
160
161#[derive(Debug)]
162struct ClientPartition<L> {
163    limiter: L,
164    last_seen: Instant,
165}
166
167/// Bounded partition storage with an ordered recency index.
168///
169/// The index makes the full-and-live churn path logarithmic instead of
170/// scanning every partition for each unseen identifier. Digest bytes provide
171/// a deterministic tie-breaker when the monotonic clock returns equal values.
172#[derive(Debug)]
173struct PartitionStore<L> {
174    entries: HashMap<Sha256Digest, ClientPartition<L>>,
175    recency: BTreeSet<(Instant, [u8; SHA256_DIGEST_BYTES])>,
176}
177
178impl<L> PartitionStore<L> {
179    fn new() -> Self {
180        Self {
181            entries: HashMap::new(),
182            recency: BTreeSet::new(),
183        }
184    }
185
186    fn use_existing<T, F>(&mut self, key: Sha256Digest, operation: F) -> Option<T>
187    where
188        F: FnOnce(&L) -> T,
189    {
190        let key_bytes = key.into_bytes();
191        let (old_last_seen, new_last_seen, result) = {
192            let entry = self.entries.get_mut(&key)?;
193            let old_last_seen = entry.last_seen;
194            let result = operation(&entry.limiter);
195            let new_last_seen = Instant::now();
196            entry.last_seen = new_last_seen;
197            (old_last_seen, new_last_seen, result)
198        };
199
200        let removed = self.recency.remove(&(old_last_seen, key_bytes));
201        let inserted = self.recency.insert((new_last_seen, key_bytes));
202        debug_assert!(removed, "existing partition must have a recency entry");
203        debug_assert!(inserted, "updated recency entry must be unique");
204        Some(result)
205    }
206
207    fn insert(&mut self, key: Sha256Digest, limiter: L) {
208        let key_bytes = key.into_bytes();
209        let last_seen = Instant::now();
210        let previous = self
211            .entries
212            .insert(key, ClientPartition { limiter, last_seen });
213        debug_assert!(previous.is_none(), "partition insertion must be unique");
214        let inserted = self.recency.insert((last_seen, key_bytes));
215        debug_assert!(inserted, "new recency entry must be unique");
216    }
217
218    fn reclaim_oldest_if<F>(&mut self, now: Instant, idle_ttl: Duration, mut is_reset: F) -> bool
219    where
220        F: FnMut(&L) -> bool,
221    {
222        let mut candidate = None;
223        for &(last_seen, key_bytes) in &self.recency {
224            let Some(idle_for) = now.checked_duration_since(last_seen) else {
225                break;
226            };
227            if idle_for < idle_ttl {
228                break;
229            }
230
231            let key = Sha256Digest::from_bytes(key_bytes);
232            let Some(entry) = self.entries.get(&key) else {
233                continue;
234            };
235            if is_reset(&entry.limiter) {
236                candidate = Some((last_seen, key_bytes, key));
237                break;
238            }
239        }
240
241        let Some((last_seen, key_bytes, key)) = candidate else {
242            return false;
243        };
244        let removed_recency = self.recency.remove(&(last_seen, key_bytes));
245        let removed_entry = self.entries.remove(&key);
246        debug_assert!(removed_recency, "recency entry must exist");
247        debug_assert!(removed_entry.is_some(), "partition entry must exist");
248        true
249    }
250
251    fn len(&self) -> usize {
252        self.entries.len()
253    }
254
255    #[cfg(test)]
256    fn is_empty(&self) -> bool {
257        self.entries.is_empty()
258    }
259
260    #[cfg(test)]
261    fn contains_key(&self, key: &Sha256Digest) -> bool {
262        self.entries.contains_key(key)
263    }
264
265    #[cfg(test)]
266    fn set_last_seen(&mut self, key: Sha256Digest, last_seen: Instant) -> bool {
267        let Some(entry) = self.entries.get_mut(&key) else {
268            return false;
269        };
270        let key_bytes = key.into_bytes();
271        let removed = self.recency.remove(&(entry.last_seen, key_bytes));
272        entry.last_seen = last_seen;
273        let inserted = self.recency.insert((last_seen, key_bytes));
274        debug_assert!(removed, "test partition must have a recency entry");
275        debug_assert!(inserted, "test recency entry must be unique");
276        true
277    }
278}
279
280/// Creates a rate limit exceeded error.
281#[must_use]
282pub fn rate_limit_error(message: impl Into<String>) -> McpError {
283    McpError::new(McpErrorCode::Custom(RATE_LIMIT_ERROR_CODE), message)
284}
285
286/// Token bucket implementation for rate limiting.
287///
288/// The token bucket algorithm allows for burst traffic while maintaining
289/// a sustainable long-term rate. Tokens are added at a constant rate and
290/// consumed when requests arrive.
291#[derive(Debug)]
292pub struct TokenBucketRateLimiter {
293    /// Maximum number of tokens in the bucket.
294    capacity: usize,
295    /// Tokens added per second.
296    refill_rate: f64,
297    /// Current number of tokens (as f64 for fractional tokens).
298    tokens: Mutex<f64>,
299    /// Last time tokens were refilled.
300    last_refill: Mutex<Instant>,
301}
302
303impl TokenBucketRateLimiter {
304    /// Creates a new token bucket rate limiter.
305    ///
306    /// # Arguments
307    ///
308    /// * `capacity` - Maximum number of tokens (burst capacity)
309    /// * `refill_rate` - Tokens added per second (sustained rate)
310    #[must_use]
311    pub fn new(capacity: usize, refill_rate: f64) -> Self {
312        let refill_rate = sanitized_rate(refill_rate);
313        let capacity = if refill_rate > 0.0 && capacity <= MAX_EXACT_TOKEN_CAPACITY {
314            capacity
315        } else {
316            0
317        };
318        Self {
319            capacity,
320            refill_rate,
321            tokens: Mutex::new(capacity as f64),
322            last_refill: Mutex::new(Instant::now()),
323        }
324    }
325
326    /// Tries to consume tokens from the bucket.
327    ///
328    /// Returns `true` if tokens were available and consumed, `false` otherwise.
329    pub fn try_consume(&self, tokens: usize) -> bool {
330        self.try_consume_with_retry(tokens).is_allowed()
331    }
332
333    fn try_consume_with_retry(&self, tokens: usize) -> RateLimitAdmission {
334        let mut current_tokens = self
335            .tokens
336            .lock()
337            .unwrap_or_else(std::sync::PoisonError::into_inner);
338        let mut last_refill = self
339            .last_refill
340            .lock()
341            .unwrap_or_else(std::sync::PoisonError::into_inner);
342
343        let now = Instant::now();
344        let elapsed = now.duration_since(*last_refill).as_secs_f64();
345
346        // Add tokens based on elapsed time
347        *current_tokens = (*current_tokens + elapsed * self.refill_rate).min(self.capacity as f64);
348        *last_refill = now;
349
350        let tokens_needed = tokens as f64;
351        if *current_tokens >= tokens_needed {
352            *current_tokens -= tokens_needed;
353            RateLimitAdmission::Allowed
354        } else {
355            RateLimitAdmission::Rejected {
356                retry_after_ms: retry_after_millis(
357                    tokens_needed - *current_tokens,
358                    self.refill_rate,
359                ),
360            }
361        }
362    }
363
364    /// Returns the current number of available tokens.
365    #[must_use]
366    pub fn available_tokens(&self) -> f64 {
367        let mut current_tokens = self
368            .tokens
369            .lock()
370            .unwrap_or_else(std::sync::PoisonError::into_inner);
371        let mut last_refill = self
372            .last_refill
373            .lock()
374            .unwrap_or_else(std::sync::PoisonError::into_inner);
375
376        let now = Instant::now();
377        let elapsed = now.duration_since(*last_refill).as_secs_f64();
378
379        // Update tokens without consuming
380        *current_tokens = (*current_tokens + elapsed * self.refill_rate).min(self.capacity as f64);
381        *last_refill = now;
382
383        *current_tokens
384    }
385
386    fn is_fully_refilled(&self) -> bool {
387        self.available_tokens() >= self.capacity as f64
388    }
389}
390
391/// Sliding window rate limiter implementation.
392///
393/// Tracks individual request timestamps within a time window for precise
394/// rate limiting. More memory-intensive than token bucket but provides
395/// exact request counting.
396#[derive(Debug)]
397pub struct SlidingWindowRateLimiter {
398    /// Maximum requests allowed in the time window.
399    max_requests: usize,
400    /// Time window in seconds.
401    window_seconds: u64,
402    /// Request timestamps (as durations from a fixed start time).
403    requests: Mutex<VecDeque<Instant>>,
404}
405
406impl SlidingWindowRateLimiter {
407    /// Creates a new sliding window rate limiter.
408    ///
409    /// # Arguments
410    ///
411    /// * `max_requests` - Maximum requests allowed in the time window
412    /// * `window_seconds` - Time window duration in seconds
413    #[must_use]
414    pub fn new(max_requests: usize, window_seconds: u64) -> Self {
415        Self {
416            max_requests,
417            window_seconds,
418            requests: Mutex::new(VecDeque::new()),
419        }
420    }
421
422    /// Checks if a request is allowed under the rate limit.
423    ///
424    /// If allowed, records the request timestamp and returns `true`.
425    /// Otherwise returns `false`.
426    pub fn is_allowed(&self) -> bool {
427        self.is_allowed_with_retry().is_allowed()
428    }
429
430    fn is_allowed_with_retry(&self) -> RateLimitAdmission {
431        if self.window_seconds == 0 {
432            return RateLimitAdmission::Rejected {
433                retry_after_ms: u64::MAX,
434            };
435        }
436
437        let mut requests = self
438            .requests
439            .lock()
440            .unwrap_or_else(std::sync::PoisonError::into_inner);
441
442        let now = Instant::now();
443        let cutoff = now.checked_sub(std::time::Duration::from_secs(self.window_seconds));
444
445        // Remove old requests outside the window
446        if let Some(cutoff) = cutoff {
447            while let Some(&oldest) = requests.front() {
448                if oldest < cutoff {
449                    requests.pop_front();
450                } else {
451                    break;
452                }
453            }
454        }
455
456        if requests.len() < self.max_requests {
457            requests.push_back(now);
458            RateLimitAdmission::Allowed
459        } else {
460            let retry_after_ms = requests.front().map_or(u64::MAX, |oldest| {
461                let elapsed = now.saturating_duration_since(*oldest);
462                let window = Duration::from_secs(self.window_seconds);
463                if elapsed >= window {
464                    1
465                } else {
466                    u64::try_from((window - elapsed).as_millis())
467                        .unwrap_or(u64::MAX)
468                        .saturating_add(1)
469                }
470            });
471            RateLimitAdmission::Rejected { retry_after_ms }
472        }
473    }
474
475    /// Returns the current number of requests in the window.
476    #[must_use]
477    pub fn current_requests(&self) -> usize {
478        if self.window_seconds == 0 {
479            return 0;
480        }
481
482        let mut requests = self
483            .requests
484            .lock()
485            .unwrap_or_else(std::sync::PoisonError::into_inner);
486
487        let now = Instant::now();
488        let cutoff = now.checked_sub(std::time::Duration::from_secs(self.window_seconds));
489
490        // Remove old requests outside the window
491        if let Some(cutoff) = cutoff {
492            while let Some(&oldest) = requests.front() {
493                if oldest < cutoff {
494                    requests.pop_front();
495                } else {
496                    break;
497                }
498            }
499        }
500
501        requests.len()
502    }
503}
504
505/// Function type for extracting client ID from request context.
506pub type ClientIdExtractor =
507    Box<dyn Fn(&McpContext, &JsonRpcRequest) -> Option<String> + Send + Sync>;
508
509/// Rate limiting middleware using token bucket algorithm.
510///
511/// Uses a token bucket algorithm by default, allowing for burst traffic
512/// while maintaining a sustainable long-term rate.
513///
514/// # Example
515///
516/// ```ignore
517/// use fastmcp_server::rate_limiting::RateLimitingMiddleware;
518///
519/// // Allow 10 requests per second with bursts up to 20
520/// let rate_limiter = RateLimitingMiddleware::new(10.0)
521///     .burst_capacity(20);
522/// ```
523pub struct RateLimitingMiddleware {
524    /// Sustained requests per second allowed.
525    max_requests_per_second: f64,
526    /// Maximum burst capacity.
527    burst_capacity: usize,
528    /// Function to extract client ID for method-scoped client partitions.
529    get_client_id: Option<ClientIdExtractor>,
530    /// If true, apply limit globally; if false, per-client.
531    global_limit: bool,
532    /// Storage for a bounded number of fixed-width client partitions.
533    limiters: Mutex<PartitionStore<TokenBucketRateLimiter>>,
534    /// Minimum inactivity required before safe least-recently-used eviction.
535    partition_idle_ttl: Duration,
536    /// Shared partition for new identifiers after the named-partition cap.
537    overflow_limiter: TokenBucketRateLimiter,
538    /// Global rate limiter (used when global_limit is true).
539    global_limiter: Option<TokenBucketRateLimiter>,
540}
541
542impl std::fmt::Debug for RateLimitingMiddleware {
543    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
544        f.debug_struct("RateLimitingMiddleware")
545            .field("max_requests_per_second", &self.max_requests_per_second)
546            .field("burst_capacity", &self.burst_capacity)
547            .field("global_limit", &self.global_limit)
548            .finish()
549    }
550}
551
552impl RateLimitingMiddleware {
553    /// Creates a new rate limiting middleware with the specified rate.
554    ///
555    /// # Arguments
556    ///
557    /// * `max_requests_per_second` - Sustained requests per second allowed
558    ///
559    /// Burst capacity defaults to 2x the sustained rate.
560    #[must_use]
561    pub fn new(max_requests_per_second: f64) -> Self {
562        let max_requests_per_second = sanitized_rate(max_requests_per_second);
563        let burst_capacity = default_burst_capacity(max_requests_per_second);
564        Self {
565            max_requests_per_second,
566            burst_capacity,
567            get_client_id: None,
568            global_limit: false,
569            limiters: Mutex::new(PartitionStore::new()),
570            partition_idle_ttl: CLIENT_PARTITION_IDLE_TTL,
571            overflow_limiter: TokenBucketRateLimiter::new(burst_capacity, max_requests_per_second),
572            global_limiter: None,
573        }
574    }
575
576    /// Sets the burst capacity (maximum tokens in the bucket).
577    #[must_use]
578    pub fn burst_capacity(mut self, capacity: usize) -> Self {
579        let capacity = if capacity <= MAX_EXACT_TOKEN_CAPACITY {
580            capacity
581        } else {
582            0
583        };
584        self.burst_capacity = capacity;
585        self.overflow_limiter = TokenBucketRateLimiter::new(capacity, self.max_requests_per_second);
586        // Re-create global limiter if it exists
587        if self.global_limit {
588            self.global_limiter = Some(TokenBucketRateLimiter::new(
589                capacity,
590                self.max_requests_per_second,
591            ));
592        }
593        self
594    }
595
596    /// Sets a custom function to extract client ID from the request context.
597    ///
598    /// Identifiers longer than 4096 bytes are rejected. Accepted identifiers
599    /// are immediately reduced to fixed-width SHA-256 partition keys; their raw
600    /// values are neither retained nor included in rate-limit errors or debug
601    /// output. At most 4095 distinct keys receive dedicated partitions. At the
602    /// cap, a least-recently-used partition idle for at least 60 seconds is
603    /// reclaimed only after its limiter has naturally reset; otherwise new
604    /// keys share one overflow partition.
605    ///
606    /// If not set, all clients share the default identity while each method
607    /// retains an independent rate-limit partition.
608    #[must_use]
609    pub fn client_id_extractor<F>(mut self, extractor: F) -> Self
610    where
611        F: Fn(&McpContext, &JsonRpcRequest) -> Option<String> + Send + Sync + 'static,
612    {
613        self.get_client_id = Some(Box::new(extractor));
614        self
615    }
616
617    #[cfg(test)]
618    fn with_partition_idle_ttl(mut self, idle_ttl: Duration) -> Self {
619        self.partition_idle_ttl = idle_ttl;
620        self
621    }
622
623    /// Enables global rate limiting (all clients share one limit).
624    ///
625    /// When enabled, all requests count against a single rate limit
626    /// regardless of client identity.
627    #[must_use]
628    pub fn global(mut self) -> Self {
629        self.global_limit = true;
630        self.global_limiter = Some(TokenBucketRateLimiter::new(
631            self.burst_capacity,
632            self.max_requests_per_second,
633        ));
634        self
635    }
636
637    fn client_partition_key(
638        &self,
639        ctx: &McpContext,
640        request: &JsonRpcRequest,
641    ) -> McpResult<Sha256Digest> {
642        if let Some(ref extractor) = self.get_client_id {
643            if let Some(id) = extractor(ctx, request) {
644                return sha256_bounded(id.as_bytes(), MAX_CLIENT_ID_BYTES)
645                    .map_err(|_| rate_limit_error(RATE_LIMIT_EXCEEDED_MESSAGE));
646            }
647        }
648        default_client_partition()
649    }
650
651    fn request_partition_key(
652        &self,
653        ctx: &McpContext,
654        request: &JsonRpcRequest,
655    ) -> McpResult<Sha256Digest> {
656        rate_limit_method_partition(self.client_partition_key(ctx, request)?, &request.method)
657    }
658
659    fn get_or_create_limiter_with_retry(&self, partition: Sha256Digest) -> RateLimitAdmission {
660        let mut limiters = self
661            .limiters
662            .lock()
663            .unwrap_or_else(std::sync::PoisonError::into_inner);
664
665        if let Some(admission) =
666            limiters.use_existing(partition, |limiter| limiter.try_consume_with_retry(1))
667        {
668            return admission;
669        }
670
671        if limiters.len() >= MAX_NAMED_CLIENT_PARTITIONS
672            && !limiters.reclaim_oldest_if(
673                Instant::now(),
674                self.partition_idle_ttl,
675                TokenBucketRateLimiter::is_fully_refilled,
676            )
677        {
678            return self.overflow_limiter.try_consume_with_retry(1);
679        }
680
681        let limiter =
682            TokenBucketRateLimiter::new(self.burst_capacity, self.max_requests_per_second);
683        let admission = limiter.try_consume_with_retry(1);
684        limiters.insert(partition, limiter);
685        admission
686    }
687
688    fn get_or_create_limiter(&self, partition: Sha256Digest) -> bool {
689        self.get_or_create_limiter_with_retry(partition)
690            .is_allowed()
691    }
692}
693
694impl Middleware for RateLimitingMiddleware {
695    fn on_request(
696        &self,
697        ctx: &McpContext,
698        request: &JsonRpcRequest,
699    ) -> McpResult<MiddlewareDecision> {
700        ctx.ensure_live().map_err(McpError::from)?;
701        if self.max_requests_per_second <= 0.0 || self.burst_capacity == 0 {
702            return Err(rate_limit_error(RATE_LIMIT_EXCEEDED_MESSAGE));
703        }
704
705        let admission = if self.global_limit {
706            // Global rate limiting
707            if let Some(ref limiter) = self.global_limiter {
708                limiter.try_consume_with_retry(1)
709            } else {
710                RateLimitAdmission::Rejected {
711                    retry_after_ms: u64::MAX,
712                }
713            }
714        } else {
715            // The request ID remains correlation-only. Including it in the
716            // bucket key would let a retry with a fresh JSON-RPC ID evade the
717            // method's configured admission budget.
718            let partition = self.request_partition_key(ctx, request)?;
719            self.get_or_create_limiter_with_retry(partition)
720        };
721
722        ctx.ensure_live().map_err(McpError::from)?;
723        match admission {
724            RateLimitAdmission::Allowed => Ok(MiddlewareDecision::Continue),
725            RateLimitAdmission::Rejected { retry_after_ms } => {
726                Err(rate_limit_retry_error(request, retry_after_ms))
727            }
728        }
729    }
730}
731
732/// Rate limiting middleware using sliding window algorithm.
733///
734/// Uses a sliding window approach which provides more precise rate limiting
735/// but uses more memory to track individual request timestamps.
736///
737/// # Example
738///
739/// ```ignore
740/// use fastmcp_server::rate_limiting::SlidingWindowRateLimitingMiddleware;
741///
742/// // Allow 100 requests per minute
743/// let rate_limiter = SlidingWindowRateLimitingMiddleware::new(100, 60);
744/// ```
745pub struct SlidingWindowRateLimitingMiddleware {
746    /// Maximum requests allowed in the time window.
747    max_requests: usize,
748    /// Time window in seconds.
749    window_seconds: u64,
750    /// Function to extract client ID for method-scoped client partitions.
751    get_client_id: Option<ClientIdExtractor>,
752    /// Storage for a bounded number of fixed-width client partitions.
753    limiters: Mutex<PartitionStore<SlidingWindowRateLimiter>>,
754    /// Minimum inactivity required before safe least-recently-used eviction.
755    partition_idle_ttl: Duration,
756    /// Shared partition for new identifiers after the named-partition cap.
757    overflow_limiter: SlidingWindowRateLimiter,
758}
759
760impl std::fmt::Debug for SlidingWindowRateLimitingMiddleware {
761    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
762        f.debug_struct("SlidingWindowRateLimitingMiddleware")
763            .field("max_requests", &self.max_requests)
764            .field("window_seconds", &self.window_seconds)
765            .finish()
766    }
767}
768
769impl SlidingWindowRateLimitingMiddleware {
770    /// Creates a new sliding window rate limiting middleware.
771    ///
772    /// # Arguments
773    ///
774    /// * `max_requests` - Maximum requests allowed in the time window
775    /// * `window_seconds` - Time window duration in seconds
776    #[must_use]
777    pub fn new(max_requests: usize, window_seconds: u64) -> Self {
778        Self {
779            max_requests,
780            window_seconds,
781            get_client_id: None,
782            limiters: Mutex::new(PartitionStore::new()),
783            partition_idle_ttl: CLIENT_PARTITION_IDLE_TTL,
784            overflow_limiter: SlidingWindowRateLimiter::new(max_requests, window_seconds),
785        }
786    }
787
788    /// Creates a sliding window rate limiter with minutes-based window.
789    ///
790    /// # Arguments
791    ///
792    /// * `max_requests` - Maximum requests allowed in the time window
793    /// * `window_minutes` - Time window duration in minutes
794    #[must_use]
795    pub fn per_minute(max_requests: usize, window_minutes: u64) -> Self {
796        Self::new(max_requests, window_minutes.checked_mul(60).unwrap_or(0))
797    }
798
799    /// Sets a custom function to extract client ID from the request context.
800    ///
801    /// Identifiers longer than 4096 bytes are rejected. Accepted identifiers
802    /// are immediately reduced to fixed-width SHA-256 partition keys; their raw
803    /// values are neither retained nor included in rate-limit errors or debug
804    /// output. At most 4095 distinct keys receive dedicated partitions. At the
805    /// cap, a least-recently-used partition idle for at least 60 seconds is
806    /// reclaimed only after its limiter has naturally reset; otherwise new
807    /// keys share one overflow partition.
808    #[must_use]
809    pub fn client_id_extractor<F>(mut self, extractor: F) -> Self
810    where
811        F: Fn(&McpContext, &JsonRpcRequest) -> Option<String> + Send + Sync + 'static,
812    {
813        self.get_client_id = Some(Box::new(extractor));
814        self
815    }
816
817    #[cfg(test)]
818    fn with_partition_idle_ttl(mut self, idle_ttl: Duration) -> Self {
819        self.partition_idle_ttl = idle_ttl;
820        self
821    }
822
823    fn client_partition_key(
824        &self,
825        ctx: &McpContext,
826        request: &JsonRpcRequest,
827    ) -> McpResult<Sha256Digest> {
828        if let Some(ref extractor) = self.get_client_id {
829            if let Some(id) = extractor(ctx, request) {
830                return sha256_bounded(id.as_bytes(), MAX_CLIENT_ID_BYTES)
831                    .map_err(|_| rate_limit_error(RATE_LIMIT_EXCEEDED_MESSAGE));
832            }
833        }
834        default_client_partition()
835    }
836
837    fn request_partition_key(
838        &self,
839        ctx: &McpContext,
840        request: &JsonRpcRequest,
841    ) -> McpResult<Sha256Digest> {
842        rate_limit_method_partition(self.client_partition_key(ctx, request)?, &request.method)
843    }
844
845    fn is_request_allowed_with_retry(&self, partition: Sha256Digest) -> RateLimitAdmission {
846        let mut limiters = self
847            .limiters
848            .lock()
849            .unwrap_or_else(std::sync::PoisonError::into_inner);
850
851        if let Some(admission) =
852            limiters.use_existing(partition, |limiter| limiter.is_allowed_with_retry())
853        {
854            return admission;
855        }
856
857        if limiters.len() >= MAX_NAMED_CLIENT_PARTITIONS
858            && !limiters.reclaim_oldest_if(Instant::now(), self.partition_idle_ttl, |limiter| {
859                limiter.current_requests() == 0
860            })
861        {
862            return self.overflow_limiter.is_allowed_with_retry();
863        }
864
865        let limiter = SlidingWindowRateLimiter::new(self.max_requests, self.window_seconds);
866        let admission = limiter.is_allowed_with_retry();
867        limiters.insert(partition, limiter);
868        admission
869    }
870
871    fn is_request_allowed(&self, partition: Sha256Digest) -> bool {
872        self.is_request_allowed_with_retry(partition).is_allowed()
873    }
874}
875
876impl Middleware for SlidingWindowRateLimitingMiddleware {
877    fn on_request(
878        &self,
879        ctx: &McpContext,
880        request: &JsonRpcRequest,
881    ) -> McpResult<MiddlewareDecision> {
882        ctx.ensure_live().map_err(McpError::from)?;
883        if self.max_requests == 0 || self.window_seconds == 0 {
884            return Err(rate_limit_error(RATE_LIMIT_EXCEEDED_MESSAGE));
885        }
886
887        let partition = self.request_partition_key(ctx, request)?;
888        let admission = self.is_request_allowed_with_retry(partition);
889
890        ctx.ensure_live().map_err(McpError::from)?;
891        match admission {
892            RateLimitAdmission::Allowed => Ok(MiddlewareDecision::Continue),
893            RateLimitAdmission::Rejected { retry_after_ms } => {
894                Err(rate_limit_retry_error(request, retry_after_ms))
895            }
896        }
897    }
898}
899
900#[cfg(test)]
901mod tests {
902    use super::*;
903    use asupersync::Cx;
904
905    fn test_context() -> McpContext {
906        let cx = Cx::for_testing();
907        McpContext::new(cx, 1)
908    }
909
910    fn test_request(method: &str) -> JsonRpcRequest {
911        JsonRpcRequest {
912            jsonrpc: std::borrow::Cow::Borrowed(fastmcp_protocol::JSONRPC_VERSION),
913            method: method.to_string(),
914            params: None,
915            id: Some(fastmcp_protocol::RequestId::Number(1)),
916        }
917    }
918
919    /// Live partitions are scoped by client AND method (aced504); tests that
920    /// address stored partitions directly must compose the same key. These
921    /// suites extract the client id from the request method, so both digest
922    /// stages consume the same string.
923    fn method_scoped_test_key(id: &str) -> Sha256Digest {
924        let client = sha256_bounded(id.as_bytes(), MAX_CLIENT_ID_BYTES)
925            .expect("test identifier is within the bound");
926        rate_limit_method_partition(client, id)
927            .expect("test method partition input is within the bound")
928    }
929
930    fn modern_test_request(method: &str, id: &str) -> JsonRpcRequest {
931        JsonRpcRequest::new(
932            method,
933            None,
934            fastmcp_protocol::RequestId::String(id.to_string()),
935        )
936    }
937
938    // ========================================
939    // TokenBucketRateLimiter tests
940    // ========================================
941
942    #[test]
943    fn test_token_bucket_allows_burst() {
944        let limiter = TokenBucketRateLimiter::new(5, 1.0);
945
946        // Should allow burst up to capacity
947        assert!(limiter.try_consume(1));
948        assert!(limiter.try_consume(1));
949        assert!(limiter.try_consume(1));
950        assert!(limiter.try_consume(1));
951        assert!(limiter.try_consume(1));
952
953        // Should deny once capacity exhausted
954        assert!(!limiter.try_consume(1));
955    }
956
957    #[test]
958    fn test_token_bucket_refills_over_time() {
959        let limiter = TokenBucketRateLimiter::new(2, 100.0); // 100 tokens per second
960
961        // Exhaust tokens
962        assert!(limiter.try_consume(1));
963        assert!(limiter.try_consume(1));
964        assert!(!limiter.try_consume(1));
965
966        // Wait for refill (10ms should add ~1 token at 100 t/s)
967        std::thread::sleep(std::time::Duration::from_millis(15));
968
969        // Should have refilled
970        assert!(limiter.try_consume(1));
971    }
972
973    #[test]
974    fn test_token_bucket_available_tokens() {
975        let limiter = TokenBucketRateLimiter::new(10, 1.0);
976        assert!((limiter.available_tokens() - 10.0).abs() < 0.1);
977
978        limiter.try_consume(5);
979        assert!((limiter.available_tokens() - 5.0).abs() < 0.1);
980    }
981
982    // ========================================
983    // SlidingWindowRateLimiter tests
984    // ========================================
985
986    #[test]
987    fn test_sliding_window_allows_up_to_limit() {
988        let limiter = SlidingWindowRateLimiter::new(3, 60);
989
990        assert!(limiter.is_allowed());
991        assert!(limiter.is_allowed());
992        assert!(limiter.is_allowed());
993        assert!(!limiter.is_allowed()); // Fourth request denied
994    }
995
996    #[test]
997    fn test_sliding_window_current_requests() {
998        let limiter = SlidingWindowRateLimiter::new(10, 60);
999
1000        assert_eq!(limiter.current_requests(), 0);
1001        limiter.is_allowed();
1002        assert_eq!(limiter.current_requests(), 1);
1003        limiter.is_allowed();
1004        assert_eq!(limiter.current_requests(), 2);
1005    }
1006
1007    // ========================================
1008    // RateLimitingMiddleware tests
1009    // ========================================
1010
1011    #[test]
1012    fn test_rate_limiting_middleware_allows_initial_requests() {
1013        let middleware = RateLimitingMiddleware::new(10.0).global();
1014        let ctx = test_context();
1015        let request = test_request("tools/call");
1016
1017        let result = middleware.on_request(&ctx, &request);
1018        assert!(matches!(result, Ok(MiddlewareDecision::Continue)));
1019    }
1020
1021    #[test]
1022    fn test_rate_limiting_middleware_denies_after_burst() {
1023        let middleware = RateLimitingMiddleware::new(10.0).burst_capacity(2).global();
1024        let ctx = test_context();
1025        let request = test_request("tools/call");
1026
1027        // First two should succeed (burst capacity = 2)
1028        assert!(middleware.on_request(&ctx, &request).is_ok());
1029        assert!(middleware.on_request(&ctx, &request).is_ok());
1030
1031        // Third should fail
1032        let result = middleware.on_request(&ctx, &request);
1033        assert!(result.is_err());
1034        let err = result.unwrap_err();
1035        assert_eq!(i32::from(err.code), RATE_LIMIT_ERROR_CODE);
1036        assert_eq!(err.message, RATE_LIMIT_EXCEEDED_MESSAGE);
1037    }
1038
1039    #[test]
1040    fn test_rate_limiting_middleware_per_client() {
1041        let middleware = RateLimitingMiddleware::new(10.0)
1042            .burst_capacity(1)
1043            .client_id_extractor(|_ctx, req| Some(req.method.clone()));
1044        let ctx = test_context();
1045
1046        let request1 = test_request("method_a");
1047        let request2 = test_request("method_b");
1048
1049        // Each "client" (method) gets their own bucket
1050        assert!(middleware.on_request(&ctx, &request1).is_ok());
1051        assert!(middleware.on_request(&ctx, &request2).is_ok());
1052
1053        // Now both are exhausted
1054        assert!(middleware.on_request(&ctx, &request1).is_err());
1055        assert!(middleware.on_request(&ctx, &request2).is_err());
1056    }
1057
1058    // ========================================
1059    // SlidingWindowRateLimitingMiddleware tests
1060    // ========================================
1061
1062    #[test]
1063    fn test_sliding_window_middleware_allows_up_to_limit() {
1064        let middleware = SlidingWindowRateLimitingMiddleware::new(2, 60);
1065        let ctx = test_context();
1066        let request = test_request("tools/call");
1067
1068        assert!(middleware.on_request(&ctx, &request).is_ok());
1069        assert!(middleware.on_request(&ctx, &request).is_ok());
1070
1071        let result = middleware.on_request(&ctx, &request);
1072        assert!(result.is_err());
1073        let err = result.unwrap_err();
1074        assert_eq!(i32::from(err.code), RATE_LIMIT_ERROR_CODE);
1075    }
1076
1077    #[test]
1078    fn test_sliding_window_middleware_per_minute() {
1079        let middleware = SlidingWindowRateLimitingMiddleware::per_minute(100, 1);
1080        let ctx = test_context();
1081        let request = test_request("tools/call");
1082
1083        // Should allow many requests
1084        for _ in 0..100 {
1085            assert!(middleware.on_request(&ctx, &request).is_ok());
1086        }
1087
1088        // 101st should fail
1089        assert!(middleware.on_request(&ctx, &request).is_err());
1090    }
1091
1092    #[test]
1093    fn test_rate_limit_error_code() {
1094        let err = rate_limit_error("test");
1095        assert_eq!(i32::from(err.code), RATE_LIMIT_ERROR_CODE);
1096        assert_eq!(err.message, "test");
1097    }
1098
1099    // ========================================
1100    // rate_limit_error / RATE_LIMIT_ERROR_CODE
1101    // ========================================
1102
1103    #[test]
1104    fn rate_limit_error_code_value() {
1105        assert_eq!(RATE_LIMIT_ERROR_CODE, -32005);
1106    }
1107
1108    #[test]
1109    fn rate_limit_error_from_string() {
1110        let err = rate_limit_error(String::from("custom message"));
1111        assert_eq!(err.message, "custom message");
1112        assert_eq!(i32::from(err.code), RATE_LIMIT_ERROR_CODE);
1113    }
1114
1115    // ========================================
1116    // TokenBucketRateLimiter — additional
1117    // ========================================
1118
1119    #[test]
1120    fn token_bucket_debug() {
1121        let limiter = TokenBucketRateLimiter::new(10, 5.0);
1122        let debug = format!("{:?}", limiter);
1123        assert!(debug.contains("TokenBucketRateLimiter"));
1124        assert!(debug.contains("10"));
1125    }
1126
1127    #[test]
1128    fn token_bucket_consume_multiple_at_once() {
1129        let limiter = TokenBucketRateLimiter::new(10, 1.0);
1130        // Consume 5 at once — should succeed
1131        assert!(limiter.try_consume(5));
1132        // Consume another 5 — should succeed (exactly 10 tokens)
1133        assert!(limiter.try_consume(5));
1134        // No tokens left
1135        assert!(!limiter.try_consume(1));
1136    }
1137
1138    #[test]
1139    fn token_bucket_consume_more_than_capacity() {
1140        let limiter = TokenBucketRateLimiter::new(5, 1.0);
1141        // Request more than capacity — should fail immediately
1142        assert!(!limiter.try_consume(6));
1143        // Bucket still has tokens (nothing was consumed on failure)
1144        assert!(limiter.try_consume(5));
1145    }
1146
1147    #[test]
1148    fn token_bucket_available_tokens_caps_at_capacity() {
1149        let limiter = TokenBucketRateLimiter::new(5, 1000.0); // Very high refill
1150        // Even with high refill rate, wait a bit — should not exceed capacity
1151        std::thread::sleep(std::time::Duration::from_millis(10));
1152        assert!(limiter.available_tokens() <= 5.0 + 0.1);
1153    }
1154
1155    #[test]
1156    fn token_bucket_available_tokens_after_full_drain() {
1157        let limiter = TokenBucketRateLimiter::new(3, 1.0);
1158        limiter.try_consume(3);
1159        assert!(limiter.available_tokens() < 1.0);
1160    }
1161
1162    // ========================================
1163    // SlidingWindowRateLimiter — additional
1164    // ========================================
1165
1166    #[test]
1167    fn sliding_window_debug() {
1168        let limiter = SlidingWindowRateLimiter::new(100, 60);
1169        let debug = format!("{:?}", limiter);
1170        assert!(debug.contains("SlidingWindowRateLimiter"));
1171        assert!(debug.contains("100"));
1172    }
1173
1174    #[test]
1175    fn sliding_window_current_requests_starts_at_zero() {
1176        let limiter = SlidingWindowRateLimiter::new(10, 60);
1177        assert_eq!(limiter.current_requests(), 0);
1178    }
1179
1180    #[test]
1181    fn sliding_window_denied_request_not_counted() {
1182        let limiter = SlidingWindowRateLimiter::new(2, 60);
1183        assert!(limiter.is_allowed());
1184        assert!(limiter.is_allowed());
1185        assert!(!limiter.is_allowed()); // denied
1186        // Only 2 requests counted (not the denied one)
1187        assert_eq!(limiter.current_requests(), 2);
1188    }
1189
1190    // ========================================
1191    // RateLimitingMiddleware — construction/Debug
1192    // ========================================
1193
1194    #[test]
1195    fn rate_limiting_middleware_default_burst_capacity() {
1196        let m = RateLimitingMiddleware::new(10.0);
1197        // Default burst capacity is 2x rate = 20
1198        assert_eq!(m.burst_capacity, 20);
1199        assert!(!m.global_limit);
1200        assert!(m.global_limiter.is_none());
1201        assert!(m.get_client_id.is_none());
1202    }
1203
1204    #[test]
1205    fn fractional_positive_rates_have_a_nonzero_default_burst() {
1206        for rate in [0.1, 0.49] {
1207            let middleware = RateLimitingMiddleware::new(rate).global();
1208            assert_eq!(middleware.burst_capacity, 1);
1209
1210            let ctx = test_context();
1211            let request = test_request("tools/call");
1212            assert!(middleware.on_request(&ctx, &request).is_ok());
1213            assert!(middleware.on_request(&ctx, &request).is_err());
1214        }
1215    }
1216
1217    #[test]
1218    fn invalid_rates_remain_fail_closed() {
1219        for rate in [0.0, -1.0, f64::NAN, f64::INFINITY, f64::NEG_INFINITY] {
1220            let middleware = RateLimitingMiddleware::new(rate).global();
1221            assert_eq!(middleware.burst_capacity, 0);
1222
1223            let ctx = test_context();
1224            let request = test_request("tools/call");
1225            assert!(middleware.on_request(&ctx, &request).is_err());
1226        }
1227    }
1228
1229    #[test]
1230    fn rate_limiting_middleware_debug() {
1231        let m = RateLimitingMiddleware::new(10.0)
1232            .burst_capacity(30)
1233            .global();
1234        let debug = format!("{:?}", m);
1235        assert!(debug.contains("RateLimitingMiddleware"));
1236        assert!(debug.contains("30"));
1237        assert!(debug.contains("true")); // global_limit
1238    }
1239
1240    #[test]
1241    fn rate_limiting_middleware_global_creates_limiter() {
1242        let m = RateLimitingMiddleware::new(5.0).global();
1243        assert!(m.global_limit);
1244        assert!(m.global_limiter.is_some());
1245    }
1246
1247    #[test]
1248    fn rate_limiting_middleware_burst_capacity_without_global() {
1249        let m = RateLimitingMiddleware::new(10.0).burst_capacity(50);
1250        // No global limiter created when not in global mode
1251        assert!(m.global_limiter.is_none());
1252        assert_eq!(m.burst_capacity, 50);
1253    }
1254
1255    #[test]
1256    fn rate_limiting_middleware_burst_capacity_with_global_recreates_limiter() {
1257        let m = RateLimitingMiddleware::new(10.0).global().burst_capacity(3);
1258        assert_eq!(m.burst_capacity, 3);
1259        // Global limiter should exist with new capacity
1260        assert!(m.global_limiter.is_some());
1261
1262        let ctx = test_context();
1263        let req = test_request("test");
1264        // Should allow exactly 3 requests (burst capacity)
1265        assert!(m.on_request(&ctx, &req).is_ok());
1266        assert!(m.on_request(&ctx, &req).is_ok());
1267        assert!(m.on_request(&ctx, &req).is_ok());
1268        assert!(m.on_request(&ctx, &req).is_err());
1269    }
1270
1271    // ========================================
1272    // RateLimitingMiddleware — client ID extraction
1273    // ========================================
1274
1275    #[test]
1276    fn rate_limiting_middleware_no_extractor_uses_global_key() {
1277        let m = RateLimitingMiddleware::new(10.0);
1278        let ctx = test_context();
1279        let req = test_request("tools/call");
1280        let partition = m
1281            .client_partition_key(&ctx, &req)
1282            .expect("default partition must be valid");
1283        assert_eq!(
1284            partition,
1285            default_client_partition().expect("default partition must be valid")
1286        );
1287    }
1288
1289    #[test]
1290    fn rate_limiting_middleware_extractor_returning_none_uses_global() {
1291        let m = RateLimitingMiddleware::new(10.0).client_id_extractor(|_ctx, _req| None);
1292        let ctx = test_context();
1293        let req = test_request("tools/call");
1294        let partition = m
1295            .client_partition_key(&ctx, &req)
1296            .expect("default partition must be valid");
1297        assert_eq!(
1298            partition,
1299            default_client_partition().expect("default partition must be valid")
1300        );
1301    }
1302
1303    #[test]
1304    fn rate_limiting_middleware_extractor_returning_some() {
1305        let m = RateLimitingMiddleware::new(10.0)
1306            .client_id_extractor(|_ctx, _req| Some("user-42".to_string()));
1307        let ctx = test_context();
1308        let req = test_request("tools/call");
1309        let partition = m
1310            .client_partition_key(&ctx, &req)
1311            .expect("bounded custom partition must be valid");
1312        let expected = sha256_bounded(b"user-42", MAX_CLIENT_ID_BYTES)
1313            .expect("test identifier is within the bound");
1314        assert_eq!(partition, expected);
1315    }
1316
1317    // ========================================
1318    // RateLimitingMiddleware — per-client without extractor
1319    // ========================================
1320
1321    #[test]
1322    fn rate_limiting_middleware_without_extractor_is_method_scoped() {
1323        // Without an extractor, all clients share the default identity, but
1324        // each validated method retains its own admission budget.
1325        let m = RateLimitingMiddleware::new(10.0).burst_capacity(2);
1326        let ctx = test_context();
1327        let req_a = test_request("method_a");
1328        let req_b = test_request("method_b");
1329
1330        // Distinct methods use distinct partitions.
1331        assert!(m.on_request(&ctx, &req_a).is_ok());
1332        assert!(m.on_request(&ctx, &req_b).is_ok());
1333        assert!(m.on_request(&ctx, &req_a).is_ok());
1334        assert!(m.on_request(&ctx, &req_b).is_ok());
1335        // Each method exhausts only its own bucket.
1336        assert!(m.on_request(&ctx, &req_a).is_err());
1337        assert!(m.on_request(&ctx, &req_b).is_err());
1338    }
1339
1340    #[test]
1341    fn modern_rate_limit_allows_a_distinct_method_for_the_same_client() {
1342        let middleware = RateLimitingMiddleware::new(1.0e-300)
1343            .burst_capacity(1)
1344            .client_id_extractor(|_ctx, _request| Some("modern-tenant".to_string()));
1345        let first_ctx = McpContext::new(Cx::for_testing(), 41);
1346        let retry_ctx = McpContext::new(Cx::for_testing(), 42);
1347        let first = modern_test_request("tools/call", "first-id");
1348        let distinct_method = modern_test_request("resources/read", "retry-id");
1349
1350        assert!(middleware.on_request(&first_ctx, &first).is_ok());
1351        assert!(middleware.on_request(&retry_ctx, &distinct_method).is_ok());
1352    }
1353
1354    #[test]
1355    fn modern_rate_limit_rejects_a_new_id_retry_for_the_same_method() {
1356        let middleware = RateLimitingMiddleware::new(1.0e-300)
1357            .burst_capacity(1)
1358            .client_id_extractor(|_ctx, _request| Some("modern-tenant".to_string()));
1359        let first_ctx = McpContext::new(Cx::for_testing(), 41);
1360        let retry_ctx = McpContext::new(Cx::for_testing(), 42);
1361        let first = modern_test_request("tools/call", "first-id");
1362        let retry = modern_test_request("tools/call", "retry-id");
1363
1364        assert!(middleware.on_request(&first_ctx, &first).is_ok());
1365        let error = middleware
1366            .on_request(&retry_ctx, &retry)
1367            .expect_err("RH-5 planted negative: a fresh request ID must not reset a method limit");
1368        assert_eq!(error.code, McpErrorCode::Custom(RATE_LIMIT_ERROR_CODE));
1369        assert_eq!(error.message, RATE_LIMIT_EXCEEDED_MESSAGE);
1370        assert_eq!(
1371            error.data,
1372            Some(serde_json::json!({
1373                "method": "tools/call",
1374                "requestId": "retry-id",
1375                "retryAfterMs": u64::MAX,
1376            }))
1377        );
1378    }
1379
1380    #[test]
1381    fn cancelled_modern_request_does_not_consume_a_method_limit() {
1382        let middleware = RateLimitingMiddleware::new(1.0e-300).burst_capacity(1);
1383        let cancelled_cx = Cx::for_testing();
1384        cancelled_cx.set_cancel_requested(true);
1385        let cancelled_ctx = McpContext::new(cancelled_cx, 41);
1386        let cancelled = modern_test_request("tools/call", "cancelled-id");
1387
1388        let error = middleware
1389            .on_request(&cancelled_ctx, &cancelled)
1390            .expect_err("cancelled requests must not receive an admission token");
1391        assert_eq!(error.code, McpErrorCode::RequestCancelled);
1392
1393        let live_ctx = McpContext::new(Cx::for_testing(), 42);
1394        let live = modern_test_request("tools/call", "live-id");
1395        assert!(middleware.on_request(&live_ctx, &live).is_ok());
1396    }
1397
1398    #[test]
1399    fn rate_limiting_middleware_error_is_generic_per_client() {
1400        let m = RateLimitingMiddleware::new(10.0)
1401            .burst_capacity(1)
1402            .client_id_extractor(|_ctx, _req| Some("alice".to_string()));
1403        let ctx = test_context();
1404        let req = test_request("tools/call");
1405
1406        m.on_request(&ctx, &req).unwrap();
1407        let err = m.on_request(&ctx, &req).unwrap_err();
1408        assert_eq!(err.message, RATE_LIMIT_EXCEEDED_MESSAGE);
1409        assert!(!err.message.contains("alice"));
1410    }
1411
1412    #[test]
1413    fn rate_limiting_middleware_error_msg_global() {
1414        let m = RateLimitingMiddleware::new(10.0).burst_capacity(1).global();
1415        let ctx = test_context();
1416        let req = test_request("tools/call");
1417
1418        m.on_request(&ctx, &req).unwrap();
1419        let err = m.on_request(&ctx, &req).unwrap_err();
1420        assert_eq!(err.message, RATE_LIMIT_EXCEEDED_MESSAGE);
1421    }
1422
1423    // ========================================
1424    // SlidingWindowRateLimitingMiddleware — construction/Debug
1425    // ========================================
1426
1427    #[test]
1428    fn sliding_window_middleware_new_fields() {
1429        let m = SlidingWindowRateLimitingMiddleware::new(50, 120);
1430        assert_eq!(m.max_requests, 50);
1431        assert_eq!(m.window_seconds, 120);
1432        assert!(m.get_client_id.is_none());
1433    }
1434
1435    #[test]
1436    fn sliding_window_middleware_per_minute_converts() {
1437        let m = SlidingWindowRateLimitingMiddleware::per_minute(100, 5);
1438        assert_eq!(m.max_requests, 100);
1439        assert_eq!(m.window_seconds, 300); // 5 * 60
1440    }
1441
1442    #[test]
1443    fn sliding_window_middleware_debug() {
1444        let m = SlidingWindowRateLimitingMiddleware::new(50, 120);
1445        let debug = format!("{:?}", m);
1446        assert!(debug.contains("SlidingWindowRateLimitingMiddleware"));
1447        assert!(debug.contains("50"));
1448        assert!(debug.contains("120"));
1449    }
1450
1451    // ========================================
1452    // SlidingWindowRateLimitingMiddleware — client ID
1453    // ========================================
1454
1455    #[test]
1456    fn sliding_window_middleware_no_extractor_uses_global() {
1457        let m = SlidingWindowRateLimitingMiddleware::new(10, 60);
1458        let ctx = test_context();
1459        let req = test_request("tools/call");
1460        let partition = m
1461            .client_partition_key(&ctx, &req)
1462            .expect("default partition must be valid");
1463        assert_eq!(
1464            partition,
1465            default_client_partition().expect("default partition must be valid")
1466        );
1467    }
1468
1469    #[test]
1470    fn sliding_window_middleware_extractor_returning_none_uses_global() {
1471        let m =
1472            SlidingWindowRateLimitingMiddleware::new(10, 60).client_id_extractor(|_ctx, _req| None);
1473        let ctx = test_context();
1474        let req = test_request("tools/call");
1475        let partition = m
1476            .client_partition_key(&ctx, &req)
1477            .expect("default partition must be valid");
1478        assert_eq!(
1479            partition,
1480            default_client_partition().expect("default partition must be valid")
1481        );
1482    }
1483
1484    #[test]
1485    fn sliding_window_middleware_extractor_returning_some() {
1486        let m = SlidingWindowRateLimitingMiddleware::new(10, 60)
1487            .client_id_extractor(|_ctx, _req| Some("bob".to_string()));
1488        let ctx = test_context();
1489        let req = test_request("tools/call");
1490        let partition = m
1491            .client_partition_key(&ctx, &req)
1492            .expect("bounded custom partition must be valid");
1493        let expected = sha256_bounded(b"bob", MAX_CLIENT_ID_BYTES)
1494            .expect("test identifier is within the bound");
1495        assert_eq!(partition, expected);
1496    }
1497
1498    // ========================================
1499    // SlidingWindowRateLimitingMiddleware — per-client
1500    // ========================================
1501
1502    #[test]
1503    fn sliding_window_middleware_per_client() {
1504        let m = SlidingWindowRateLimitingMiddleware::new(1, 60)
1505            .client_id_extractor(|_ctx, req| Some(req.method.clone()));
1506        let ctx = test_context();
1507        let req_a = test_request("method_a");
1508        let req_b = test_request("method_b");
1509
1510        // Each client gets their own window
1511        assert!(m.on_request(&ctx, &req_a).is_ok());
1512        assert!(m.on_request(&ctx, &req_b).is_ok());
1513
1514        // Both exhausted
1515        assert!(m.on_request(&ctx, &req_a).is_err());
1516        assert!(m.on_request(&ctx, &req_b).is_err());
1517    }
1518
1519    // ========================================
1520    // SlidingWindowRateLimitingMiddleware — error messages
1521    // ========================================
1522
1523    #[test]
1524    fn sliding_window_middleware_error_is_generic_for_seconds_window() {
1525        let m = SlidingWindowRateLimitingMiddleware::new(1, 30);
1526        let ctx = test_context();
1527        let req = test_request("tools/call");
1528
1529        m.on_request(&ctx, &req).unwrap();
1530        let err = m.on_request(&ctx, &req).unwrap_err();
1531        assert_eq!(err.message, RATE_LIMIT_EXCEEDED_MESSAGE);
1532    }
1533
1534    #[test]
1535    fn sliding_window_middleware_error_is_generic_for_minutes_window() {
1536        let m = SlidingWindowRateLimitingMiddleware::new(1, 120);
1537        let ctx = test_context();
1538        let req = test_request("tools/call");
1539
1540        m.on_request(&ctx, &req).unwrap();
1541        let err = m.on_request(&ctx, &req).unwrap_err();
1542        assert_eq!(err.message, RATE_LIMIT_EXCEEDED_MESSAGE);
1543    }
1544
1545    #[test]
1546    fn sliding_window_middleware_error_omits_client_id() {
1547        let m = SlidingWindowRateLimitingMiddleware::new(1, 60)
1548            .client_id_extractor(|_ctx, _req| Some("alice".to_string()));
1549        let ctx = test_context();
1550        let req = test_request("tools/call");
1551
1552        m.on_request(&ctx, &req).unwrap();
1553        let err = m.on_request(&ctx, &req).unwrap_err();
1554        assert_eq!(err.message, RATE_LIMIT_EXCEEDED_MESSAGE);
1555        assert!(!err.message.contains("alice"));
1556        assert_eq!(i32::from(err.code), RATE_LIMIT_ERROR_CODE);
1557    }
1558
1559    // ========================================
1560    // Edge cases
1561    // ========================================
1562
1563    #[test]
1564    fn rate_limiting_middleware_get_or_create_limiter_creates_new() {
1565        let m = RateLimitingMiddleware::new(10.0).burst_capacity(2);
1566        let partition = sha256_bounded(b"new-client", MAX_CLIENT_ID_BYTES)
1567            .expect("test identifier is within the bound");
1568        // First call for a new client creates a limiter
1569        assert!(m.get_or_create_limiter(partition));
1570        // Second call reuses the same limiter
1571        assert!(m.get_or_create_limiter(partition));
1572        // Third call exhausts it
1573        assert!(!m.get_or_create_limiter(partition));
1574    }
1575
1576    #[test]
1577    fn sliding_window_middleware_is_request_allowed_creates_new() {
1578        let m = SlidingWindowRateLimitingMiddleware::new(2, 60);
1579        let c1 = sha256_bounded(b"c1", MAX_CLIENT_ID_BYTES)
1580            .expect("test identifier is within the bound");
1581        let c2 = sha256_bounded(b"c2", MAX_CLIENT_ID_BYTES)
1582            .expect("test identifier is within the bound");
1583        assert!(m.is_request_allowed(c1));
1584        assert!(m.is_request_allowed(c1));
1585        assert!(!m.is_request_allowed(c1));
1586
1587        // Different client gets its own limiter
1588        assert!(m.is_request_allowed(c2));
1589    }
1590
1591    #[test]
1592    fn sliding_window_requests_expire_after_window() {
1593        let limiter = SlidingWindowRateLimiter::new(2, 1); // 2 requests per 1 second
1594        assert!(limiter.is_allowed());
1595        assert!(limiter.is_allowed());
1596        assert!(!limiter.is_allowed()); // exhausted
1597
1598        // Wait for window to expire
1599        std::thread::sleep(std::time::Duration::from_millis(1100));
1600
1601        // Requests should be allowed again
1602        assert!(limiter.is_allowed());
1603    }
1604
1605    #[test]
1606    fn sliding_window_current_requests_resets_after_window() {
1607        let limiter = SlidingWindowRateLimiter::new(5, 1); // 1 second window
1608        limiter.is_allowed();
1609        limiter.is_allowed();
1610        assert_eq!(limiter.current_requests(), 2);
1611
1612        std::thread::sleep(std::time::Duration::from_millis(1100));
1613
1614        // Old requests should have expired
1615        assert_eq!(limiter.current_requests(), 0);
1616    }
1617
1618    #[test]
1619    fn sliding_window_error_exactly_60_seconds_is_generic() {
1620        let m = SlidingWindowRateLimitingMiddleware::new(1, 60);
1621        let ctx = test_context();
1622        let req = test_request("tools/call");
1623
1624        m.on_request(&ctx, &req).unwrap();
1625        let err = m.on_request(&ctx, &req).unwrap_err();
1626        assert_eq!(err.message, RATE_LIMIT_EXCEEDED_MESSAGE);
1627    }
1628
1629    #[test]
1630    fn token_bucket_try_consume_zero_always_succeeds() {
1631        let limiter = TokenBucketRateLimiter::new(3, 1.0);
1632        // Drain all tokens
1633        limiter.try_consume(3);
1634        assert!(!limiter.try_consume(1)); // exhausted
1635
1636        // Consuming zero should still succeed
1637        assert!(limiter.try_consume(0));
1638    }
1639
1640    #[test]
1641    fn token_bucket_refill_rate_zero_fails_closed() {
1642        let limiter = TokenBucketRateLimiter::new(2, 0.0); // zero refill rate
1643        assert!(!limiter.try_consume(2));
1644        assert!(!limiter.try_consume(1));
1645
1646        // Even after waiting, no refill
1647        std::thread::sleep(std::time::Duration::from_millis(50));
1648        assert!(!limiter.try_consume(1));
1649    }
1650
1651    #[test]
1652    fn token_bucket_invalid_rates_fail_closed() {
1653        for rate in [f64::NAN, f64::INFINITY, f64::NEG_INFINITY, -1.0] {
1654            let limiter = TokenBucketRateLimiter::new(10, rate);
1655            assert!(
1656                !limiter.try_consume(1),
1657                "invalid rate {rate:?} admitted traffic"
1658            );
1659            assert!(limiter.available_tokens().abs() <= f64::EPSILON);
1660
1661            let middleware = RateLimitingMiddleware::new(rate)
1662                .burst_capacity(10)
1663                .global();
1664            let result = middleware.on_request(&test_context(), &test_request("tools/call"));
1665            assert!(
1666                result.is_err(),
1667                "invalid rate {rate:?} admitted middleware traffic"
1668            );
1669        }
1670    }
1671
1672    #[test]
1673    fn token_bucket_exact_integer_capacity_still_decrements() {
1674        let limiter = TokenBucketRateLimiter::new(MAX_EXACT_TOKEN_CAPACITY, f64::MIN_POSITIVE);
1675
1676        assert!(limiter.try_consume(1));
1677        let expected_tokens = (MAX_EXACT_TOKEN_CAPACITY - 1) as f64;
1678        assert_eq!(
1679            limiter.available_tokens().total_cmp(&expected_tokens),
1680            std::cmp::Ordering::Equal
1681        );
1682    }
1683
1684    #[cfg(target_pointer_width = "64")]
1685    #[test]
1686    fn token_bucket_inexact_integer_capacity_fails_closed() {
1687        let inexact_capacity = MAX_EXACT_TOKEN_CAPACITY + 1;
1688        let limiter = TokenBucketRateLimiter::new(inexact_capacity, 1.0);
1689        assert!(!limiter.try_consume(1));
1690        assert!(limiter.available_tokens().abs() <= f64::EPSILON);
1691
1692        let middleware = RateLimitingMiddleware::new(1.0)
1693            .burst_capacity(inexact_capacity)
1694            .global();
1695        assert_eq!(middleware.burst_capacity, 0);
1696        assert!(
1697            middleware
1698                .on_request(&test_context(), &test_request("tools/call"))
1699                .is_err()
1700        );
1701    }
1702
1703    #[test]
1704    fn reclamation_skips_oldest_penalized_partition_for_later_reset_candidate() {
1705        let oldest_key = sha256_bounded(b"oldest-penalized", MAX_CLIENT_ID_BYTES)
1706            .expect("test identifier is within the bound");
1707        let reset_key = sha256_bounded(b"later-reset", MAX_CLIENT_ID_BYTES)
1708            .expect("test identifier is within the bound");
1709        let oldest_limiter = TokenBucketRateLimiter::new(1, 1.0e-300);
1710        assert!(oldest_limiter.try_consume(1));
1711        let reset_limiter = TokenBucketRateLimiter::new(1, 1.0e-300);
1712
1713        let mut partitions = PartitionStore::new();
1714        partitions.insert(oldest_key, oldest_limiter);
1715        partitions.insert(reset_key, reset_limiter);
1716
1717        let idle_ttl = Duration::from_millis(1);
1718        let now = Instant::now();
1719        let oldest_last_seen = now
1720            .checked_sub(Duration::from_millis(3))
1721            .expect("the test idle interval must fit in Instant");
1722        let reset_last_seen = now
1723            .checked_sub(Duration::from_millis(2))
1724            .expect("the test idle interval must fit in Instant");
1725        assert!(partitions.set_last_seen(oldest_key, oldest_last_seen));
1726        assert!(partitions.set_last_seen(reset_key, reset_last_seen));
1727
1728        assert!(partitions.reclaim_oldest_if(
1729            now,
1730            idle_ttl,
1731            TokenBucketRateLimiter::is_fully_refilled,
1732        ));
1733        assert!(partitions.contains_key(&oldest_key));
1734        assert!(!partitions.contains_key(&reset_key));
1735        assert_eq!(partitions.len(), 1);
1736        assert_eq!(partitions.recency.len(), 1);
1737    }
1738
1739    #[test]
1740    fn token_bucket_partitions_are_bounded_and_overflow_is_shared() {
1741        let middleware = RateLimitingMiddleware::new(1.0e-300)
1742            .burst_capacity(1)
1743            .client_id_extractor(|_ctx, request| Some(request.method.clone()));
1744        let ctx = test_context();
1745
1746        for index in 0..MAX_NAMED_CLIENT_PARTITIONS {
1747            let request = test_request(&format!("named-client-{index}"));
1748            assert!(middleware.on_request(&ctx, &request).is_ok());
1749        }
1750        assert_eq!(
1751            middleware
1752                .limiters
1753                .lock()
1754                .unwrap_or_else(std::sync::PoisonError::into_inner)
1755                .len(),
1756            MAX_NAMED_CLIENT_PARTITIONS
1757        );
1758
1759        assert!(
1760            middleware
1761                .on_request(&ctx, &test_request("overflow-client-a"))
1762                .is_ok()
1763        );
1764        assert!(
1765            middleware
1766                .on_request(&ctx, &test_request("overflow-client-b"))
1767                .is_err(),
1768            "a fresh identifier must not reset the shared overflow limit"
1769        );
1770        assert_eq!(
1771            middleware
1772                .limiters
1773                .lock()
1774                .unwrap_or_else(std::sync::PoisonError::into_inner)
1775                .len(),
1776            MAX_NAMED_CLIENT_PARTITIONS
1777        );
1778    }
1779
1780    #[test]
1781    fn sliding_window_partitions_are_bounded_and_overflow_is_shared() {
1782        let middleware = SlidingWindowRateLimitingMiddleware::new(1, u64::MAX)
1783            .client_id_extractor(|_ctx, request| Some(request.method.clone()));
1784        let ctx = test_context();
1785
1786        for index in 0..MAX_NAMED_CLIENT_PARTITIONS {
1787            let request = test_request(&format!("named-client-{index}"));
1788            assert!(middleware.on_request(&ctx, &request).is_ok());
1789        }
1790        assert_eq!(
1791            middleware
1792                .limiters
1793                .lock()
1794                .unwrap_or_else(std::sync::PoisonError::into_inner)
1795                .len(),
1796            MAX_NAMED_CLIENT_PARTITIONS
1797        );
1798
1799        assert!(
1800            middleware
1801                .on_request(&ctx, &test_request("overflow-client-a"))
1802                .is_ok()
1803        );
1804        assert!(
1805            middleware
1806                .on_request(&ctx, &test_request("overflow-client-b"))
1807                .is_err(),
1808            "a fresh identifier must not reset the shared overflow limit"
1809        );
1810        assert_eq!(
1811            middleware
1812                .limiters
1813                .lock()
1814                .unwrap_or_else(std::sync::PoisonError::into_inner)
1815                .len(),
1816            MAX_NAMED_CLIENT_PARTITIONS
1817        );
1818    }
1819
1820    #[test]
1821    fn token_bucket_reclaims_reset_stale_partition_and_preserves_recent_client() {
1822        let idle_ttl = Duration::from_millis(1);
1823        let middleware = RateLimitingMiddleware::new(1.0e-300)
1824            .burst_capacity(1)
1825            .client_id_extractor(|_ctx, request| Some(request.method.clone()))
1826            .with_partition_idle_ttl(idle_ttl);
1827        let ctx = test_context();
1828        let legitimate_id = "recent-legitimate-token-client";
1829
1830        assert!(
1831            middleware
1832                .on_request(&ctx, &test_request(legitimate_id))
1833                .is_ok()
1834        );
1835        for index in 0..(MAX_NAMED_CLIENT_PARTITIONS - 1) {
1836            assert!(
1837                middleware
1838                    .on_request(
1839                        &ctx,
1840                        &test_request(&format!("stale-token-attacker-{index}"))
1841                    )
1842                    .is_ok()
1843            );
1844        }
1845        assert!(
1846            middleware
1847                .on_request(&ctx, &test_request(legitimate_id))
1848                .is_err(),
1849            "touching an exhausted legitimate partition must preserve its limit"
1850        );
1851
1852        let stale_key = method_scoped_test_key("stale-token-attacker-0");
1853        let legitimate_key = method_scoped_test_key(legitimate_id);
1854        let stale_last_seen = Instant::now()
1855            .checked_sub(idle_ttl + idle_ttl)
1856            .expect("the test idle interval must fit in Instant");
1857        {
1858            let mut partitions = middleware
1859                .limiters
1860                .lock()
1861                .unwrap_or_else(std::sync::PoisonError::into_inner);
1862            assert!(partitions.set_last_seen(stale_key, stale_last_seen));
1863            assert!(partitions.set_last_seen(legitimate_key, Instant::now()));
1864        }
1865
1866        let active_state_probe = "token-active-state-overflow-probe";
1867        assert!(
1868            middleware
1869                .on_request(&ctx, &test_request(active_state_probe))
1870                .is_ok(),
1871            "stale but non-reset state must use the shared overflow limiter"
1872        );
1873        let active_state_probe_key = method_scoped_test_key(active_state_probe);
1874        {
1875            let partitions = middleware
1876                .limiters
1877                .lock()
1878                .unwrap_or_else(std::sync::PoisonError::into_inner);
1879            assert!(!partitions.contains_key(&active_state_probe_key));
1880        }
1881        {
1882            let partitions = middleware
1883                .limiters
1884                .lock()
1885                .unwrap_or_else(std::sync::PoisonError::into_inner);
1886            let stale = partitions
1887                .entries
1888                .get(&stale_key)
1889                .expect("attacker partition must exist");
1890            let mut tokens = stale
1891                .limiter
1892                .tokens
1893                .lock()
1894                .unwrap_or_else(std::sync::PoisonError::into_inner);
1895            *tokens = stale.limiter.capacity as f64;
1896        }
1897
1898        let newcomer_id = "new-legitimate-token-client";
1899        assert!(
1900            middleware
1901                .on_request(&ctx, &test_request(newcomer_id))
1902                .is_ok(),
1903            "a safely reset stale attacker partition should be reclaimed"
1904        );
1905        let newcomer_key = method_scoped_test_key(newcomer_id);
1906        let partitions = middleware
1907            .limiters
1908            .lock()
1909            .unwrap_or_else(std::sync::PoisonError::into_inner);
1910        assert!(!partitions.contains_key(&stale_key));
1911        assert!(partitions.contains_key(&legitimate_key));
1912        assert!(partitions.contains_key(&newcomer_key));
1913        assert_eq!(partitions.len(), MAX_NAMED_CLIENT_PARTITIONS);
1914        assert_eq!(partitions.recency.len(), MAX_NAMED_CLIENT_PARTITIONS);
1915        drop(partitions);
1916
1917        assert!(
1918            middleware
1919                .on_request(&ctx, &test_request(legitimate_id))
1920                .is_err(),
1921            "the recent legitimate client's exhausted limiter must not be reset"
1922        );
1923    }
1924
1925    #[test]
1926    fn sliding_window_reclaims_empty_stale_partition_and_preserves_recent_client() {
1927        let idle_ttl = Duration::from_millis(1);
1928        let middleware = SlidingWindowRateLimitingMiddleware::new(1, 60)
1929            .client_id_extractor(|_ctx, request| Some(request.method.clone()))
1930            .with_partition_idle_ttl(idle_ttl);
1931        let ctx = test_context();
1932        let legitimate_id = "recent-legitimate-window-client";
1933
1934        assert!(
1935            middleware
1936                .on_request(&ctx, &test_request(legitimate_id))
1937                .is_ok()
1938        );
1939        for index in 0..(MAX_NAMED_CLIENT_PARTITIONS - 1) {
1940            assert!(
1941                middleware
1942                    .on_request(
1943                        &ctx,
1944                        &test_request(&format!("stale-window-attacker-{index}"))
1945                    )
1946                    .is_ok()
1947            );
1948        }
1949        assert!(
1950            middleware
1951                .on_request(&ctx, &test_request(legitimate_id))
1952                .is_err(),
1953            "touching a limited legitimate partition must preserve its window"
1954        );
1955
1956        let stale_key = method_scoped_test_key("stale-window-attacker-0");
1957        let legitimate_key = method_scoped_test_key(legitimate_id);
1958        let stale_last_seen = Instant::now()
1959            .checked_sub(idle_ttl + idle_ttl)
1960            .expect("the test idle interval must fit in Instant");
1961        {
1962            let mut partitions = middleware
1963                .limiters
1964                .lock()
1965                .unwrap_or_else(std::sync::PoisonError::into_inner);
1966            assert!(partitions.set_last_seen(stale_key, stale_last_seen));
1967            assert!(partitions.set_last_seen(legitimate_key, Instant::now()));
1968        }
1969
1970        let active_state_probe = "window-active-state-overflow-probe";
1971        assert!(
1972            middleware
1973                .on_request(&ctx, &test_request(active_state_probe))
1974                .is_ok(),
1975            "stale but active window state must use the shared overflow limiter"
1976        );
1977        let active_state_probe_key = method_scoped_test_key(active_state_probe);
1978        {
1979            let partitions = middleware
1980                .limiters
1981                .lock()
1982                .unwrap_or_else(std::sync::PoisonError::into_inner);
1983            assert!(!partitions.contains_key(&active_state_probe_key));
1984        }
1985        {
1986            let partitions = middleware
1987                .limiters
1988                .lock()
1989                .unwrap_or_else(std::sync::PoisonError::into_inner);
1990            partitions
1991                .entries
1992                .get(&stale_key)
1993                .expect("attacker partition must exist")
1994                .limiter
1995                .requests
1996                .lock()
1997                .unwrap_or_else(std::sync::PoisonError::into_inner)
1998                .clear();
1999        }
2000
2001        let newcomer_id = "new-legitimate-window-client";
2002        assert!(
2003            middleware
2004                .on_request(&ctx, &test_request(newcomer_id))
2005                .is_ok(),
2006            "an empty stale attacker partition should be reclaimed"
2007        );
2008        let newcomer_key = method_scoped_test_key(newcomer_id);
2009        let partitions = middleware
2010            .limiters
2011            .lock()
2012            .unwrap_or_else(std::sync::PoisonError::into_inner);
2013        assert!(!partitions.contains_key(&stale_key));
2014        assert!(partitions.contains_key(&legitimate_key));
2015        assert!(partitions.contains_key(&newcomer_key));
2016        assert_eq!(partitions.len(), MAX_NAMED_CLIENT_PARTITIONS);
2017        assert_eq!(partitions.recency.len(), MAX_NAMED_CLIENT_PARTITIONS);
2018        drop(partitions);
2019
2020        assert!(
2021            middleware
2022                .on_request(&ctx, &test_request(legitimate_id))
2023                .is_err(),
2024            "the recent legitimate client's active window must not be reset"
2025        );
2026    }
2027
2028    #[test]
2029    fn custom_identifier_canary_is_absent_from_errors_and_debug() {
2030        const CANARY: &str = "secret-client-canary-71d8f0";
2031        let token = RateLimitingMiddleware::new(1.0e-300)
2032            .burst_capacity(1)
2033            .client_id_extractor(|_ctx, _request| Some(CANARY.to_string()));
2034        let sliding = SlidingWindowRateLimitingMiddleware::new(1, 60)
2035            .client_id_extractor(|_ctx, _request| Some(CANARY.to_string()));
2036        let ctx = test_context();
2037        let request = test_request("tools/call");
2038
2039        assert!(token.on_request(&ctx, &request).is_ok());
2040        let token_error = token.on_request(&ctx, &request).unwrap_err();
2041        assert!(!token_error.message.contains(CANARY));
2042        assert!(!format!("{token:?}").contains(CANARY));
2043
2044        assert!(sliding.on_request(&ctx, &request).is_ok());
2045        let sliding_error = sliding.on_request(&ctx, &request).unwrap_err();
2046        assert!(!sliding_error.message.contains(CANARY));
2047        assert!(!format!("{sliding:?}").contains(CANARY));
2048    }
2049
2050    #[test]
2051    fn oversized_custom_identifiers_fail_closed_without_partition_growth() {
2052        let oversized = "x".repeat(MAX_CLIENT_ID_BYTES + 1);
2053        let token_oversized = oversized.clone();
2054        let token = RateLimitingMiddleware::new(10.0)
2055            .client_id_extractor(move |_ctx, _request| Some(token_oversized.clone()));
2056        let sliding = SlidingWindowRateLimitingMiddleware::new(10, 60)
2057            .client_id_extractor(move |_ctx, _request| Some(oversized.clone()));
2058        let ctx = test_context();
2059        let request = test_request("tools/call");
2060
2061        let token_error = token.on_request(&ctx, &request).unwrap_err();
2062        assert_eq!(token_error.message, RATE_LIMIT_EXCEEDED_MESSAGE);
2063        assert!(
2064            token
2065                .limiters
2066                .lock()
2067                .unwrap_or_else(std::sync::PoisonError::into_inner)
2068                .is_empty()
2069        );
2070
2071        let sliding_error = sliding.on_request(&ctx, &request).unwrap_err();
2072        assert_eq!(sliding_error.message, RATE_LIMIT_EXCEEDED_MESSAGE);
2073        assert!(
2074            sliding
2075                .limiters
2076                .lock()
2077                .unwrap_or_else(std::sync::PoisonError::into_inner)
2078                .is_empty()
2079        );
2080    }
2081
2082    #[test]
2083    fn zero_and_overflowing_windows_fail_closed() {
2084        let zero_window = SlidingWindowRateLimiter::new(10, 0);
2085        assert!(!zero_window.is_allowed());
2086        assert_eq!(zero_window.current_requests(), 0);
2087
2088        let zero_window_middleware = SlidingWindowRateLimitingMiddleware::new(10, 0);
2089        assert!(
2090            zero_window_middleware
2091                .on_request(&test_context(), &test_request("tools/call"))
2092                .is_err()
2093        );
2094
2095        let overflowing_minutes = SlidingWindowRateLimitingMiddleware::per_minute(10, u64::MAX);
2096        assert_eq!(overflowing_minutes.window_seconds, 0);
2097        assert!(
2098            overflowing_minutes
2099                .on_request(&test_context(), &test_request("tools/call"))
2100                .is_err()
2101        );
2102
2103        let maximum_seconds = SlidingWindowRateLimiter::new(1, u64::MAX);
2104        assert!(maximum_seconds.is_allowed());
2105        assert!(!maximum_seconds.is_allowed());
2106    }
2107}