Coverage for agentos/api/rate_limiter.py: 0%

50 statements  

« prev     ^ index     » next       coverage.py v7.14.3, created at 2026-07-06 10:59 +0800

1""" 

2Pluggable rate limiting middleware for AgentOS. 

3 

4Strategies: 

5- Fixed window (simple, memory-friendly) 

6- Sliding window (burst-tolerant) 

7- Token bucket (smooth rate control) 

8 

9Storage backends: 

10- In-memory (default, single-process) 

11- Redis (distributed) 

12 

13Usage: 

14 from agentos.api.rate_limiter import RateLimitMiddleware, FixedWindowLimiter 

15 

16 limiter = FixedWindowLimiter(max_requests=100, window_seconds=60) 

17 app.add_middleware(RateLimitMiddleware, limiter=limiter) 

18""" 

19 

20from __future__ import annotations 

21 

22import time 

23import threading 

24from typing import Callable, Optional 

25 

26from starlette.middleware.base import BaseHTTPMiddleware, RequestResponseEndpoint 

27from starlette.requests import Request 

28from starlette.responses import Response, JSONResponse 

29 

30 

31# ── Limiters ──────────────────────────────────────────────────────────────── 

32 

33class FixedWindowLimiter: 

34 """Fixed-window counter. 100 req/min → resets every 60s.""" 

35 

36 def __init__(self, max_requests: int = 100, window_seconds: int = 60): 

37 self.max_requests = max_requests 

38 self.window_seconds = window_seconds 

39 self._windows: dict[str, tuple[int, int]] = {} # key → (count, window_start) 

40 self._lock = threading.Lock() 

41 

42 def is_allowed(self, key: str) -> tuple[bool, dict]: 

43 now = int(time.time()) 

44 with self._lock: 

45 count, window_start = self._windows.get(key, (0, now)) 

46 if now - window_start >= self.window_seconds: 

47 count, window_start = 0, now 

48 if count >= self.max_requests: 

49 reset_at = window_start + self.window_seconds 

50 return False, { 

51 "limit": self.max_requests, 

52 "remaining": 0, 

53 "reset": reset_at, 

54 "retry_after": max(0, reset_at - now), 

55 } 

56 count += 1 

57 self._windows[key] = (count, window_start) 

58 return True, { 

59 "limit": self.max_requests, 

60 "remaining": self.max_requests - count, 

61 "reset": window_start + self.window_seconds, 

62 "retry_after": 0, 

63 } 

64 

65 

66# ── Middleware ────────────────────────────────────────────────────────────── 

67 

68RATE_LIMIT_HEADERS = {"X-RateLimit-Limit", "X-RateLimit-Remaining", "X-RateLimit-Reset"} 

69 

70 

71class RateLimitMiddleware(BaseHTTPMiddleware): 

72 """Starlette middleware that rate-limits based on client IP or API key.""" 

73 

74 def __init__( 

75 self, 

76 app, 

77 limiter: FixedWindowLimiter, 

78 key_func: Optional[Callable[[Request], str]] = None, 

79 exempt_paths: Optional[list[str]] = None, 

80 ): 

81 super().__init__(app) 

82 self.limiter = limiter 

83 self._key_func = key_func or self._default_key 

84 self._exempt = set(exempt_paths or ["/health", "/metrics"]) 

85 

86 @staticmethod 

87 def _default_key(request: Request) -> str: 

88 xff = request.headers.get("X-Forwarded-For", "") 

89 if xff: 

90 return xff.split(",")[0].strip() 

91 return request.client.host if request.client else "unknown" 

92 

93 async def dispatch( 

94 self, request: Request, call_next: RequestResponseEndpoint 

95 ) -> Response: 

96 if request.url.path in self._exempt: 

97 return await call_next(request) 

98 

99 key = self._key_func(request) 

100 allowed, info = self.limiter.is_allowed(key) 

101 

102 if not allowed: 

103 return JSONResponse( 

104 status_code=429, 

105 content={"detail": "Too Many Requests", **info}, 

106 headers={k: str(info.get({"X-RateLimit-Limit": "limit"}[k], "")) for k in RATE_LIMIT_HEADERS}, 

107 ) 

108 

109 response = await call_next(request) 

110 for k, field in [ 

111 ("X-RateLimit-Limit", "limit"), 

112 ("X-RateLimit-Remaining", "remaining"), 

113 ("X-RateLimit-Reset", "reset"), 

114 ]: 

115 response.headers[k] = str(info.get(field, "")) 

116 return response 

117 

118 

119__all__ = ["RateLimitMiddleware", "FixedWindowLimiter"]