1use std::collections::HashMap;
72use std::future::Future;
73use std::pin::Pin;
74use std::sync::atomic::{AtomicU64, Ordering};
75use std::time::{SystemTime, UNIX_EPOCH};
76
77use a2a_protocol_types::error::{A2aError, A2aResult};
78use tokio::sync::RwLock;
79
80use crate::call_context::CallContext;
81use crate::error::{ServerError, ServerResult};
82use crate::interceptor::ServerInterceptor;
83
84#[derive(Debug, Clone)]
86pub struct RateLimitConfig {
87 pub requests_per_window: u64,
91
92 pub window_secs: u64,
96
97 pub trusted_proxy_hops: usize,
112
113 pub max_buckets: usize,
120}
121
122pub const DEFAULT_MAX_BUCKETS: usize = 10_000;
124
125impl Default for RateLimitConfig {
126 fn default() -> Self {
127 Self {
128 requests_per_window: 100,
129 window_secs: 60,
130 trusted_proxy_hops: 0,
131 max_buckets: DEFAULT_MAX_BUCKETS,
132 }
133 }
134}
135
136struct CallerBucket {
138 window_start: AtomicU64,
140 count: AtomicU64,
142}
143
144pub struct RateLimitInterceptor {
155 config: RateLimitConfig,
156 buckets: RwLock<HashMap<String, CallerBucket>>,
157 check_count: AtomicU64,
159}
160
161const CLEANUP_INTERVAL: u64 = 256;
163
164impl std::fmt::Debug for RateLimitInterceptor {
165 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
166 f.debug_struct("RateLimitInterceptor")
167 .field("config", &self.config)
168 .finish_non_exhaustive()
169 }
170}
171
172impl RateLimitInterceptor {
173 pub fn new(config: RateLimitConfig) -> ServerResult<Self> {
182 if config.requests_per_window == 0 {
183 return Err(ServerError::InvalidParams(
184 "rate limit requests_per_window must be greater than zero".into(),
185 ));
186 }
187 if config.window_secs == 0 {
188 return Err(ServerError::InvalidParams(
189 "rate limit window_secs must be greater than zero".into(),
190 ));
191 }
192 if config.max_buckets == 0 {
193 return Err(ServerError::InvalidParams(
194 "rate limit max_buckets must be greater than zero".into(),
195 ));
196 }
197 Ok(Self {
198 config,
199 buckets: RwLock::new(HashMap::new()),
200 check_count: AtomicU64::new(0),
201 })
202 }
203
204 fn caller_key(&self, ctx: &CallContext) -> String {
209 if let Some(identity) = ctx.caller_identity() {
210 return identity.to_owned();
211 }
212 let hops = self.config.trusted_proxy_hops;
213 if hops > 0 {
214 if let Some(xff) = ctx.http_headers().get("x-forwarded-for") {
215 let entries: Vec<&str> = xff
216 .split(',')
217 .map(str::trim)
218 .filter(|e| !e.is_empty())
219 .collect();
220 if entries.len() >= hops {
224 return canonicalize_caller_ip(entries[entries.len() - hops]);
225 }
226 }
230 }
231 "anonymous".to_string()
232 }
233
234 const fn window_number(&self, now_secs: u64) -> u64 {
236 now_secs / self.config.window_secs
237 }
238
239 fn evict_stale(buckets: &mut HashMap<String, CallerBucket>, current_window: u64) {
241 buckets.retain(|_, bucket| {
242 bucket.window_start.load(Ordering::Relaxed) >= current_window.saturating_sub(1)
243 });
244 }
245
246 async fn cleanup_stale_buckets(&self) {
251 let now_secs = SystemTime::now()
252 .duration_since(UNIX_EPOCH)
253 .unwrap_or_default()
254 .as_secs();
255 let current_window = self.window_number(now_secs);
256
257 let mut buckets = self.buckets.write().await;
258 Self::evict_stale(&mut buckets, current_window);
259 }
260
261 #[allow(clippy::too_many_lines)]
263 async fn check(&self, key: &str) -> A2aResult<()> {
264 let now_secs = SystemTime::now()
265 .duration_since(UNIX_EPOCH)
266 .unwrap_or_default()
267 .as_secs();
268 let current_window = self.window_number(now_secs);
269
270 let count = self.check_count.fetch_add(1, Ordering::Relaxed);
272 if count > 0 && count.is_multiple_of(CLEANUP_INTERVAL) {
273 self.cleanup_stale_buckets().await;
274 }
275
276 {
278 let buckets = self.buckets.read().await;
279 if let Some(bucket) = buckets.get(key) {
280 loop {
284 let bucket_window = bucket.window_start.load(Ordering::Acquire);
285 if bucket_window == current_window {
286 let count = bucket.count.fetch_add(1, Ordering::Relaxed) + 1;
287 if count > self.config.requests_per_window {
288 return Err(A2aError::internal(format!(
289 "rate limit exceeded: {} requests per {} seconds",
290 self.config.requests_per_window, self.config.window_secs
291 )));
292 }
293 return Ok(());
294 }
295 if bucket
299 .window_start
300 .compare_exchange(
301 bucket_window,
302 current_window,
303 Ordering::AcqRel,
304 Ordering::Acquire,
305 )
306 .is_ok()
307 {
308 bucket.count.store(1, Ordering::Release);
309 return Ok(());
310 }
311 }
313 }
314 }
315
316 let mut buckets = self.buckets.write().await;
318 if let Some(bucket) = buckets.get(key) {
320 let bucket_window = bucket.window_start.load(Ordering::Acquire);
321 if bucket_window == current_window {
322 let count = bucket.count.fetch_add(1, Ordering::Relaxed) + 1;
323 if count > self.config.requests_per_window {
324 return Err(A2aError::internal(format!(
325 "rate limit exceeded: {} requests per {} seconds",
326 self.config.requests_per_window, self.config.window_secs
327 )));
328 }
329 } else {
330 bucket.window_start.store(current_window, Ordering::Release);
331 bucket.count.store(1, Ordering::Release);
332 }
333 return Ok(());
334 }
335 if buckets.len() >= self.config.max_buckets {
336 Self::evict_stale(&mut buckets, current_window);
338 if buckets.len() >= self.config.max_buckets {
339 return Err(A2aError::internal(format!(
340 "rate limiter caller capacity exhausted ({} buckets); request rejected",
341 self.config.max_buckets
342 )));
343 }
344 }
345 buckets.insert(
346 key.to_string(),
347 CallerBucket {
348 window_start: AtomicU64::new(current_window),
349 count: AtomicU64::new(1),
350 },
351 );
352 drop(buckets);
353 Ok(())
354 }
355}
356
357impl ServerInterceptor for RateLimitInterceptor {
358 fn before<'a>(
359 &'a self,
360 ctx: &'a CallContext,
361 ) -> Pin<Box<dyn Future<Output = A2aResult<()>> + Send + 'a>> {
362 Box::pin(async move {
363 let key = self.caller_key(ctx);
364 self.check(&key).await
365 })
366 }
367
368 fn after<'a>(
369 &'a self,
370 _ctx: &'a CallContext,
371 ) -> Pin<Box<dyn Future<Output = A2aResult<()>> + Send + 'a>> {
372 Box::pin(async { Ok(()) })
373 }
374}
375
376fn canonicalize_caller_ip(entry: &str) -> String {
386 use std::net::IpAddr;
387 let trimmed = entry.trim().trim_start_matches('[').trim_end_matches(']');
388 match trimmed.parse::<IpAddr>() {
389 Ok(IpAddr::V6(v6)) => v6
390 .to_ipv4_mapped()
391 .map_or_else(|| IpAddr::V6(v6).to_string(), |v4| v4.to_string()),
392 Ok(ip) => ip.to_string(),
393 Err(_) => trimmed.to_string(),
394 }
395}
396
397#[cfg(test)]
398mod tests {
399 use super::*;
400 use std::collections::HashMap;
401
402 #[test]
403 fn caller_ip_canonicalization_collapses_equivalent_forms() {
404 assert_eq!(canonicalize_caller_ip("::ffff:203.0.113.7"), "203.0.113.7");
406 assert_eq!(canonicalize_caller_ip("203.0.113.7"), "203.0.113.7");
407 assert_eq!(
409 canonicalize_caller_ip("[2001:db8::1]"),
410 canonicalize_caller_ip("2001:0db8:0000:0000:0000:0000:0000:0001")
411 );
412 assert_eq!(canonicalize_caller_ip(" not-an-ip "), "not-an-ip");
414 }
415
416 fn make_ctx(identity: Option<&str>) -> CallContext {
417 let mut ctx = CallContext::new("message/send");
418 if let Some(id) = identity {
419 ctx = ctx.with_caller_identity(id.to_owned());
420 }
421 ctx
422 }
423
424 #[tokio::test]
425 async fn allows_requests_within_limit() {
426 let limiter = RateLimitInterceptor::new(RateLimitConfig {
427 requests_per_window: 5,
428 window_secs: 60,
429 ..RateLimitConfig::default()
430 })
431 .expect("valid config");
432 let ctx = make_ctx(Some("user-1"));
433 for _ in 0..5 {
434 assert!(limiter.before(&ctx).await.is_ok());
435 }
436 }
437
438 #[tokio::test]
439 async fn rejects_requests_over_limit() {
440 let limiter = RateLimitInterceptor::new(RateLimitConfig {
441 requests_per_window: 3,
442 window_secs: 60,
443 ..RateLimitConfig::default()
444 })
445 .expect("valid config");
446 let ctx = make_ctx(Some("user-2"));
447 for _ in 0..3 {
448 assert!(limiter.before(&ctx).await.is_ok());
449 }
450 let result = limiter.before(&ctx).await;
451 assert!(result.is_err());
452 }
453
454 #[tokio::test]
455 async fn different_callers_have_separate_limits() {
456 let limiter = RateLimitInterceptor::new(RateLimitConfig {
457 requests_per_window: 2,
458 window_secs: 60,
459 ..RateLimitConfig::default()
460 })
461 .expect("valid config");
462 let ctx_a = make_ctx(Some("alice"));
463 let ctx_b = make_ctx(Some("bob"));
464
465 assert!(limiter.before(&ctx_a).await.is_ok());
466 assert!(limiter.before(&ctx_a).await.is_ok());
467 assert!(limiter.before(&ctx_a).await.is_err()); assert!(limiter.before(&ctx_b).await.is_ok());
471 assert!(limiter.before(&ctx_b).await.is_ok());
472 }
473
474 #[tokio::test]
475 async fn anonymous_fallback_when_no_identity() {
476 let limiter = RateLimitInterceptor::new(RateLimitConfig {
477 requests_per_window: 1,
478 window_secs: 60,
479 ..RateLimitConfig::default()
480 })
481 .expect("valid config");
482 let ctx = make_ctx(None);
483 assert!(limiter.before(&ctx).await.is_ok());
484 assert!(limiter.before(&ctx).await.is_err());
485 }
486
487 #[tokio::test]
491 async fn default_config_ignores_forged_x_forwarded_for() {
492 let limiter = RateLimitInterceptor::new(RateLimitConfig {
493 requests_per_window: 1,
494 window_secs: 60,
495 ..RateLimitConfig::default()
496 })
497 .expect("valid config");
498 let ctx1 = CallContext::new("message/send").with_http_header("x-forwarded-for", "10.0.0.1");
501 let ctx2 = CallContext::new("message/send").with_http_header("x-forwarded-for", "10.0.0.2");
502 assert!(limiter.before(&ctx1).await.is_ok());
503 assert!(
504 limiter.before(&ctx2).await.is_err(),
505 "forged x-forwarded-for must not evade the limit"
506 );
507 assert_eq!(limiter.buckets.read().await.len(), 1);
509 }
510
511 #[tokio::test]
515 async fn trusted_hop_uses_rightmost_entry_and_resists_spoofing() {
516 let limiter = RateLimitInterceptor::new(RateLimitConfig {
517 requests_per_window: 1,
518 window_secs: 60,
519 trusted_proxy_hops: 1,
520 ..RateLimitConfig::default()
521 })
522 .expect("valid config");
523 let ctx1 = CallContext::new("message/send")
525 .with_http_header("x-forwarded-for", "6.6.6.1, 203.0.113.7");
526 let ctx2 = CallContext::new("message/send")
527 .with_http_header("x-forwarded-for", "6.6.6.2, 203.0.113.7");
528 assert!(limiter.before(&ctx1).await.is_ok());
529 assert!(
530 limiter.before(&ctx2).await.is_err(),
531 "spoofed left-hand entries must map to the same real client"
532 );
533 let ctx3 =
535 CallContext::new("message/send").with_http_header("x-forwarded-for", "203.0.113.8");
536 assert!(limiter.before(&ctx3).await.is_ok());
537 }
538
539 #[tokio::test]
541 async fn trusted_hops_two_takes_second_from_right() {
542 let limiter = RateLimitInterceptor::new(RateLimitConfig {
543 requests_per_window: 1,
544 window_secs: 60,
545 trusted_proxy_hops: 2,
546 ..RateLimitConfig::default()
547 })
548 .expect("valid config");
549 let ctx1 = CallContext::new("message/send")
551 .with_http_header("x-forwarded-for", "6.6.6.1, 198.51.100.9, 10.0.0.5");
552 let ctx2 = CallContext::new("message/send")
553 .with_http_header("x-forwarded-for", "6.6.6.2, 198.51.100.9, 10.0.0.5");
554 assert!(limiter.before(&ctx1).await.is_ok());
555 assert!(
556 limiter.before(&ctx2).await.is_err(),
557 "same client, same bucket"
558 );
559 }
560
561 #[tokio::test]
564 async fn short_xff_chain_falls_back_to_anonymous() {
565 let limiter = RateLimitInterceptor::new(RateLimitConfig {
566 requests_per_window: 1,
567 window_secs: 60,
568 trusted_proxy_hops: 3,
569 ..RateLimitConfig::default()
570 })
571 .expect("valid config");
572 let ctx1 = CallContext::new("message/send").with_http_header("x-forwarded-for", "1.2.3.4");
573 let ctx2 = CallContext::new("message/send").with_http_header("x-forwarded-for", "5.6.7.8");
574 assert!(limiter.before(&ctx1).await.is_ok());
575 assert!(
576 limiter.before(&ctx2).await.is_err(),
577 "short chains must share the anonymous bucket, not be trusted"
578 );
579 }
580
581 #[test]
586 fn new_rejects_zero_window_secs() {
587 let err = RateLimitInterceptor::new(RateLimitConfig {
588 window_secs: 0,
589 ..RateLimitConfig::default()
590 })
591 .expect_err("zero window_secs must be rejected");
592 assert!(err.to_string().contains("window_secs"), "got: {err}");
593 }
594
595 #[test]
596 fn new_rejects_zero_requests_per_window() {
597 let err = RateLimitInterceptor::new(RateLimitConfig {
598 requests_per_window: 0,
599 ..RateLimitConfig::default()
600 })
601 .expect_err("zero requests_per_window must be rejected");
602 assert!(
603 err.to_string().contains("requests_per_window"),
604 "got: {err}"
605 );
606 }
607
608 #[test]
609 fn new_rejects_zero_max_buckets() {
610 let err = RateLimitInterceptor::new(RateLimitConfig {
611 max_buckets: 0,
612 ..RateLimitConfig::default()
613 })
614 .expect_err("zero max_buckets must be rejected");
615 assert!(err.to_string().contains("max_buckets"), "got: {err}");
616 }
617
618 #[tokio::test]
623 async fn bucket_map_is_bounded() {
624 let limiter = RateLimitInterceptor::new(RateLimitConfig {
625 requests_per_window: 10,
626 window_secs: 60,
627 max_buckets: 2,
628 ..RateLimitConfig::default()
629 })
630 .expect("valid config");
631 assert!(limiter.before(&make_ctx(Some("a"))).await.is_ok());
632 assert!(limiter.before(&make_ctx(Some("b"))).await.is_ok());
633 let err = limiter
634 .before(&make_ctx(Some("c")))
635 .await
636 .expect_err("third caller must be rejected at capacity");
637 assert!(err.to_string().contains("capacity"), "got: {err}");
638 assert_eq!(limiter.buckets.read().await.len(), 2);
639 assert!(limiter.before(&make_ctx(Some("a"))).await.is_ok());
641 }
642
643 #[tokio::test]
646 async fn full_map_evicts_stale_buckets_before_rejecting() {
647 let limiter = RateLimitInterceptor::new(RateLimitConfig {
648 requests_per_window: 10,
649 window_secs: 60,
650 max_buckets: 2,
651 ..RateLimitConfig::default()
652 })
653 .expect("valid config");
654 assert!(limiter.before(&make_ctx(Some("live"))).await.is_ok());
656 {
657 let mut buckets = limiter.buckets.write().await;
658 buckets.insert(
659 "ancient".to_string(),
660 CallerBucket {
661 window_start: AtomicU64::new(0),
662 count: AtomicU64::new(1),
663 },
664 );
665 }
666 assert!(
668 limiter.before(&make_ctx(Some("newcomer"))).await.is_ok(),
669 "stale bucket should be evicted to admit the new caller"
670 );
671 let buckets = limiter.buckets.read().await;
672 assert!(!buckets.contains_key("ancient"));
673 assert!(buckets.contains_key("live"));
674 assert!(buckets.contains_key("newcomer"));
675 drop(buckets);
676 }
677
678 #[tokio::test]
681 async fn concurrent_distinct_callers_respect_bucket_cap() {
682 use std::sync::Arc;
683
684 let limiter = RateLimitInterceptor::new(RateLimitConfig {
685 requests_per_window: 10,
686 window_secs: 60,
687 max_buckets: 10,
688 ..RateLimitConfig::default()
689 })
690 .expect("valid config");
691 let limiter = Arc::new(limiter);
692
693 let mut handles = Vec::new();
694 for i in 0..50 {
695 let lim = Arc::clone(&limiter);
696 handles.push(tokio::spawn(async move {
697 let ctx =
698 CallContext::new("message/send").with_caller_identity(format!("user-{i}"));
699 lim.before(&ctx).await
700 }));
701 }
702
703 let mut ok_count = 0;
704 let mut err_count = 0;
705 for handle in handles {
706 match handle.await.unwrap() {
707 Ok(()) => ok_count += 1,
708 Err(_) => err_count += 1,
709 }
710 }
711 assert_eq!(ok_count, 10, "exactly max_buckets callers admitted");
712 assert_eq!(err_count, 40);
713 assert_eq!(limiter.buckets.read().await.len(), 10);
714 }
715
716 #[tokio::test]
717 async fn concurrent_rate_limit_checks() {
718 use std::sync::Arc;
719
720 let limiter = Arc::new(
721 RateLimitInterceptor::new(RateLimitConfig {
722 requests_per_window: 100,
723 window_secs: 60,
724 ..RateLimitConfig::default()
725 })
726 .expect("valid config"),
727 );
728
729 let mut handles = Vec::new();
731 for _ in 0..200 {
732 let lim = Arc::clone(&limiter);
733 handles.push(tokio::spawn(async move {
734 let ctx =
735 CallContext::new("message/send").with_caller_identity("concurrent-user".into());
736 lim.before(&ctx).await
737 }));
738 }
739
740 let mut ok_count = 0;
741 let mut err_count = 0;
742 for handle in handles {
743 match handle.await.unwrap() {
744 Ok(()) => ok_count += 1,
745 Err(_) => err_count += 1,
746 }
747 }
748
749 assert_eq!(ok_count, 100, "expected 100 allowed, got {ok_count}");
751 assert_eq!(err_count, 100, "expected 100 rejected, got {err_count}");
752 }
753
754 #[tokio::test]
755 async fn stale_bucket_cleanup() {
756 let limiter = RateLimitInterceptor::new(RateLimitConfig {
757 requests_per_window: 10,
758 window_secs: 60,
759 ..RateLimitConfig::default()
760 })
761 .expect("valid config");
762
763 let ctx_a = make_ctx(Some("stale-a"));
765 let ctx_b = make_ctx(Some("stale-b"));
766 assert!(limiter.before(&ctx_a).await.is_ok());
767 assert!(limiter.before(&ctx_b).await.is_ok());
768
769 assert_eq!(limiter.buckets.read().await.len(), 2);
770
771 limiter.cleanup_stale_buckets().await;
773 assert_eq!(
774 limiter.buckets.read().await.len(),
775 2,
776 "current-window buckets should not be evicted"
777 );
778 }
779
780 #[test]
781 fn debug_format_includes_config() {
782 let limiter = RateLimitInterceptor::new(RateLimitConfig {
783 requests_per_window: 42,
784 window_secs: 10,
785 ..RateLimitConfig::default()
786 })
787 .expect("valid config");
788 let debug = format!("{limiter:?}");
789 assert!(
790 debug.contains("RateLimitInterceptor"),
791 "Debug output should contain struct name"
792 );
793 assert!(
794 debug.contains("config"),
795 "Debug output should contain config field"
796 );
797 }
798
799 #[test]
801 fn default_config_values() {
802 let config = RateLimitConfig::default();
803 assert_eq!(config.requests_per_window, 100);
804 assert_eq!(config.window_secs, 60);
805 }
806
807 #[tokio::test]
809 async fn after_hook_is_noop() {
810 let limiter = RateLimitInterceptor::new(RateLimitConfig::default()).expect("valid config");
811 let ctx = make_ctx(Some("user"));
812 let result = limiter.after(&ctx).await;
813 assert_eq!(result.unwrap(), (), "after hook should return Ok(())");
814 }
815
816 #[test]
817 fn window_number_correctness() {
818 let limiter = RateLimitInterceptor::new(RateLimitConfig {
819 requests_per_window: 10,
820 window_secs: 60,
821 ..RateLimitConfig::default()
822 })
823 .expect("valid config");
824
825 assert_eq!(limiter.window_number(0), 0);
827 assert_eq!(limiter.window_number(59), 0);
829 assert_eq!(limiter.window_number(60), 1);
831 assert_eq!(limiter.window_number(120), 2);
833 assert_eq!(limiter.window_number(61), 1);
835 }
836
837 #[tokio::test]
838 async fn cleanup_stale_buckets_removes_old_entries() {
839 let limiter = RateLimitInterceptor::new(RateLimitConfig {
840 requests_per_window: 100,
841 window_secs: 60,
842 ..RateLimitConfig::default()
843 })
844 .expect("valid config");
845
846 {
848 let mut buckets = limiter.buckets.write().await;
849 buckets.insert(
850 "ancient-user".to_string(),
851 CallerBucket {
852 window_start: AtomicU64::new(0), count: AtomicU64::new(5),
854 },
855 );
856 }
857 assert_eq!(limiter.buckets.read().await.len(), 1);
858
859 limiter.cleanup_stale_buckets().await;
861 assert_eq!(
862 limiter.buckets.read().await.len(),
863 0,
864 "ancient bucket should be evicted"
865 );
866 }
867
868 #[tokio::test]
869 async fn check_triggers_cleanup_at_interval() {
870 let limiter = RateLimitInterceptor::new(RateLimitConfig {
871 requests_per_window: 10000,
872 window_secs: 60,
873 ..RateLimitConfig::default()
874 })
875 .expect("valid config");
876
877 {
879 let mut buckets = limiter.buckets.write().await;
880 buckets.insert(
881 "stale-for-cleanup".to_string(),
882 CallerBucket {
883 window_start: AtomicU64::new(0),
884 count: AtomicU64::new(1),
885 },
886 );
887 }
888
889 limiter
892 .check_count
893 .store(CLEANUP_INTERVAL, Ordering::Relaxed);
894
895 let ctx = make_ctx(Some("cleanup-trigger-user"));
896 assert!(limiter.before(&ctx).await.is_ok());
898
899 let buckets = limiter.buckets.read().await;
901 let has_stale = buckets.contains_key("stale-for-cleanup");
902 drop(buckets);
903 assert!(
904 !has_stale,
905 "stale bucket should be cleaned up after CLEANUP_INTERVAL checks"
906 );
907 }
908
909 #[tokio::test]
910 async fn slow_path_double_check_same_window() {
911 let limiter = RateLimitInterceptor::new(RateLimitConfig {
915 requests_per_window: 2,
916 window_secs: 60,
917 ..RateLimitConfig::default()
918 })
919 .expect("valid config");
920
921 let ctx = make_ctx(Some("race-user"));
922 assert!(limiter.before(&ctx).await.is_ok());
924 assert!(limiter.before(&ctx).await.is_ok());
926 assert!(limiter.before(&ctx).await.is_err());
928 }
929
930 #[tokio::test]
933 async fn slow_path_double_check_stale_window() {
934 let limiter = RateLimitInterceptor::new(RateLimitConfig {
935 requests_per_window: 10,
936 window_secs: 60,
937 ..RateLimitConfig::default()
938 })
939 .expect("valid config");
940
941 let key = "slow-path-stale";
944 {
945 let mut buckets = limiter.buckets.write().await;
946 buckets.insert(
947 key.to_string(),
948 CallerBucket {
949 window_start: AtomicU64::new(1), count: AtomicU64::new(5),
951 },
952 );
953 }
954
955 let result = limiter.check(key).await;
959 assert!(
960 result.is_ok(),
961 "slow-path stale-window reset should succeed"
962 );
963
964 assert_eq!(
966 limiter
967 .buckets
968 .read()
969 .await
970 .get(key)
971 .expect("bucket should exist")
972 .count
973 .load(Ordering::Relaxed),
974 1,
975 "count should be reset to 1 after window advance"
976 );
977 }
978
979 #[tokio::test]
982 async fn slow_path_rate_limit_exceeded() {
983 let limiter = RateLimitInterceptor::new(RateLimitConfig {
984 requests_per_window: 1,
985 window_secs: 60,
986 ..RateLimitConfig::default()
987 })
988 .expect("valid config");
989
990 let now_secs = SystemTime::now()
991 .duration_since(UNIX_EPOCH)
992 .unwrap()
993 .as_secs();
994 let current_window = limiter.window_number(now_secs);
995
996 let key = "slow-path-exceeded";
998 {
999 let mut buckets = limiter.buckets.write().await;
1000 buckets.insert(
1001 key.to_string(),
1002 CallerBucket {
1003 window_start: AtomicU64::new(current_window),
1004 count: AtomicU64::new(1), },
1006 );
1007 }
1008
1009 let result = limiter.check(key).await;
1012 assert!(
1013 result.is_err(),
1014 "slow-path should reject when count exceeds limit"
1015 );
1016 }
1017
1018 #[tokio::test]
1020 async fn fast_path_rate_limit_exceeded() {
1021 let limiter = RateLimitInterceptor::new(RateLimitConfig {
1022 requests_per_window: 2,
1023 window_secs: 60,
1024 ..RateLimitConfig::default()
1025 })
1026 .expect("valid config");
1027
1028 let ctx = make_ctx(Some("fast-path-user"));
1030 assert!(limiter.before(&ctx).await.is_ok());
1031 assert!(limiter.before(&ctx).await.is_ok());
1032 let result = limiter.before(&ctx).await;
1034 assert!(
1035 result.is_err(),
1036 "fast-path should reject when count exceeds limit"
1037 );
1038 let err = result.unwrap_err();
1039 assert!(
1040 err.to_string().contains("rate limit exceeded"),
1041 "error message should mention rate limit exceeded, got: {err}"
1042 );
1043 }
1044
1045 #[tokio::test]
1048 async fn fast_path_window_advancement_resets_count() {
1049 let limiter = RateLimitInterceptor::new(RateLimitConfig {
1050 requests_per_window: 1,
1051 window_secs: 60,
1052 ..RateLimitConfig::default()
1053 })
1054 .expect("valid config");
1055
1056 let key = "fast-path-window-advance";
1057 {
1059 let mut buckets = limiter.buckets.write().await;
1060 buckets.insert(
1061 key.to_string(),
1062 CallerBucket {
1063 window_start: AtomicU64::new(1), count: AtomicU64::new(999),
1065 },
1066 );
1067 }
1068
1069 let result = limiter.check(key).await;
1072 assert_eq!(
1073 result.unwrap(),
1074 (),
1075 "fast-path window advance should return Ok(())"
1076 );
1077
1078 assert_eq!(
1079 limiter
1080 .buckets
1081 .read()
1082 .await
1083 .get(key)
1084 .expect("bucket should exist")
1085 .count
1086 .load(Ordering::Relaxed),
1087 1,
1088 "count should be reset to 1 after window advance"
1089 );
1090 }
1091
1092 #[tokio::test]
1098 async fn cleanup_does_not_run_on_first_call() {
1099 let limiter = RateLimitInterceptor::new(RateLimitConfig {
1100 requests_per_window: 10000,
1101 window_secs: 60,
1102 ..RateLimitConfig::default()
1103 })
1104 .expect("valid config");
1105
1106 {
1108 let mut buckets = limiter.buckets.write().await;
1109 buckets.insert(
1110 "stale-first-call".to_string(),
1111 CallerBucket {
1112 window_start: AtomicU64::new(0),
1113 count: AtomicU64::new(1),
1114 },
1115 );
1116 }
1117
1118 let ctx = make_ctx(Some("first-caller"));
1121 assert!(limiter.before(&ctx).await.is_ok());
1122
1123 assert!(
1125 limiter
1126 .buckets
1127 .read()
1128 .await
1129 .contains_key("stale-first-call"),
1130 "stale bucket should not be cleaned up on the very first call"
1131 );
1132 }
1133
1134 #[tokio::test]
1137 async fn x_forwarded_for_single_ip_with_trusted_hop() {
1138 let limiter = RateLimitInterceptor::new(RateLimitConfig {
1139 requests_per_window: 1,
1140 window_secs: 60,
1141 trusted_proxy_hops: 1,
1142 ..RateLimitConfig::default()
1143 })
1144 .expect("valid config");
1145 let mut headers = HashMap::new();
1146 headers.insert("x-forwarded-for".to_string(), "192.168.1.1".to_string());
1147 let ctx = CallContext::new("message/send").with_http_headers(headers);
1148 assert!(limiter.before(&ctx).await.is_ok());
1149 assert!(limiter.before(&ctx).await.is_err());
1151 }
1152}