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

51 statements  

« prev     ^ index     » next       coverage.py v7.14.3, created at 2026-07-05 22:09 +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 dataclasses import dataclass, field 

25from typing import Callable, Optional 

26 

27from starlette.middleware.base import BaseHTTPMiddleware, RequestResponseEndpoint 

28from starlette.requests import Request 

29from starlette.responses import Response, JSONResponse 

30 

31 

32# ── Limiters ──────────────────────────────────────────────────────────────── 

33 

34class FixedWindowLimiter: 

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

36 

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

38 self.max_requests = max_requests 

39 self.window_seconds = window_seconds 

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

41 self._lock = threading.Lock() 

42 

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

44 now = int(time.time()) 

45 with self._lock: 

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

47 if now - window_start >= self.window_seconds: 

48 count, window_start = 0, now 

49 if count >= self.max_requests: 

50 reset_at = window_start + self.window_seconds 

51 return False, { 

52 "limit": self.max_requests, 

53 "remaining": 0, 

54 "reset": reset_at, 

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

56 } 

57 count += 1 

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

59 return True, { 

60 "limit": self.max_requests, 

61 "remaining": self.max_requests - count, 

62 "reset": window_start + self.window_seconds, 

63 "retry_after": 0, 

64 } 

65 

66 

67# ── Middleware ────────────────────────────────────────────────────────────── 

68 

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

70 

71 

72class RateLimitMiddleware(BaseHTTPMiddleware): 

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

74 

75 def __init__( 

76 self, 

77 app, 

78 limiter: FixedWindowLimiter, 

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

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

81 ): 

82 super().__init__(app) 

83 self.limiter = limiter 

84 self._key_func = key_func or self._default_key 

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

86 

87 @staticmethod 

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

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

90 if xff: 

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

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

93 

94 async def dispatch( 

95 self, request: Request, call_next: RequestResponseEndpoint 

96 ) -> Response: 

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

98 return await call_next(request) 

99 

100 key = self._key_func(request) 

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

102 

103 if not allowed: 

104 return JSONResponse( 

105 status_code=429, 

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

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

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"]