1use 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
36pub const RATE_LIMIT_ERROR_CODE: i32 = -32005;
40
41const MAX_CLIENT_ID_BYTES: usize = 4096;
46
47const MAX_CLIENT_PARTITIONS: usize = 4096;
50
51const MAX_NAMED_CLIENT_PARTITIONS: usize = MAX_CLIENT_PARTITIONS - 1;
54
55const 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
69const 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 (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#[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#[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#[derive(Debug)]
292pub struct TokenBucketRateLimiter {
293 capacity: usize,
295 refill_rate: f64,
297 tokens: Mutex<f64>,
299 last_refill: Mutex<Instant>,
301}
302
303impl TokenBucketRateLimiter {
304 #[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 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 *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 #[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 *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#[derive(Debug)]
397pub struct SlidingWindowRateLimiter {
398 max_requests: usize,
400 window_seconds: u64,
402 requests: Mutex<VecDeque<Instant>>,
404}
405
406impl SlidingWindowRateLimiter {
407 #[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 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 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 #[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 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
505pub type ClientIdExtractor =
507 Box<dyn Fn(&McpContext, &JsonRpcRequest) -> Option<String> + Send + Sync>;
508
509pub struct RateLimitingMiddleware {
524 max_requests_per_second: f64,
526 burst_capacity: usize,
528 get_client_id: Option<ClientIdExtractor>,
530 global_limit: bool,
532 limiters: Mutex<PartitionStore<TokenBucketRateLimiter>>,
534 partition_idle_ttl: Duration,
536 overflow_limiter: TokenBucketRateLimiter,
538 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 #[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 #[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 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 #[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 #[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 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 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
732pub struct SlidingWindowRateLimitingMiddleware {
746 max_requests: usize,
748 window_seconds: u64,
750 get_client_id: Option<ClientIdExtractor>,
752 limiters: Mutex<PartitionStore<SlidingWindowRateLimiter>>,
754 partition_idle_ttl: Duration,
756 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 #[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 #[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 #[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 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 #[test]
943 fn test_token_bucket_allows_burst() {
944 let limiter = TokenBucketRateLimiter::new(5, 1.0);
945
946 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 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); assert!(limiter.try_consume(1));
963 assert!(limiter.try_consume(1));
964 assert!(!limiter.try_consume(1));
965
966 std::thread::sleep(std::time::Duration::from_millis(15));
968
969 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 #[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()); }
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 #[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 assert!(middleware.on_request(&ctx, &request).is_ok());
1029 assert!(middleware.on_request(&ctx, &request).is_ok());
1030
1031 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 assert!(middleware.on_request(&ctx, &request1).is_ok());
1051 assert!(middleware.on_request(&ctx, &request2).is_ok());
1052
1053 assert!(middleware.on_request(&ctx, &request1).is_err());
1055 assert!(middleware.on_request(&ctx, &request2).is_err());
1056 }
1057
1058 #[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 for _ in 0..100 {
1085 assert!(middleware.on_request(&ctx, &request).is_ok());
1086 }
1087
1088 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 #[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 #[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 assert!(limiter.try_consume(5));
1132 assert!(limiter.try_consume(5));
1134 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 assert!(!limiter.try_consume(6));
1143 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); 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 #[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()); assert_eq!(limiter.current_requests(), 2);
1188 }
1189
1190 #[test]
1195 fn rate_limiting_middleware_default_burst_capacity() {
1196 let m = RateLimitingMiddleware::new(10.0);
1197 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")); }
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 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 assert!(m.global_limiter.is_some());
1261
1262 let ctx = test_context();
1263 let req = test_request("test");
1264 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 #[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 #[test]
1322 fn rate_limiting_middleware_without_extractor_is_method_scoped() {
1323 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 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 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 #[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); }
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 #[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 #[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 assert!(m.on_request(&ctx, &req_a).is_ok());
1512 assert!(m.on_request(&ctx, &req_b).is_ok());
1513
1514 assert!(m.on_request(&ctx, &req_a).is_err());
1516 assert!(m.on_request(&ctx, &req_b).is_err());
1517 }
1518
1519 #[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 #[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 assert!(m.get_or_create_limiter(partition));
1570 assert!(m.get_or_create_limiter(partition));
1572 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 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); assert!(limiter.is_allowed());
1595 assert!(limiter.is_allowed());
1596 assert!(!limiter.is_allowed()); std::thread::sleep(std::time::Duration::from_millis(1100));
1600
1601 assert!(limiter.is_allowed());
1603 }
1604
1605 #[test]
1606 fn sliding_window_current_requests_resets_after_window() {
1607 let limiter = SlidingWindowRateLimiter::new(5, 1); 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 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 limiter.try_consume(3);
1634 assert!(!limiter.try_consume(1)); 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); assert!(!limiter.try_consume(2));
1644 assert!(!limiter.try_consume(1));
1645
1646 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}