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

50 statements  

« prev     ^ index     » next       coverage.py v7.14.3, created at 2026-07-08 01:44 +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 threading 

23import time 

24from collections.abc import Callable 

25 

26from starlette.middleware.base import BaseHTTPMiddleware, RequestResponseEndpoint 

27from starlette.requests import Request 

28from starlette.responses import JSONResponse, Response 

29 

30# ── Limiters ──────────────────────────────────────────────────────────────── 

31 

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: Callable[[Request], str] | None = None, 

79 exempt_paths: list[str] | None = 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(self, request: Request, call_next: RequestResponseEndpoint) -> Response: 

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

95 return await call_next(request) 

96 

97 key = self._key_func(request) 

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

99 

100 if not allowed: 

101 return JSONResponse( 

102 status_code=429, 

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

104 headers={ 

105 k: str(info.get({"X-RateLimit-Limit": "limit"}[k], "")) 

106 for k in RATE_LIMIT_HEADERS 

107 }, 

108 ) 

109 

110 response = await call_next(request) 

111 for k, field in [ 

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

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

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

115 ]: 

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

117 return response 

118 

119 

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