Coverage for agentos/api/rate_limiter.py: 0%
50 statements
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-06 23:17 +0800
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-06 23:17 +0800
1"""
2Pluggable rate limiting middleware for AgentOS.
4Strategies:
5- Fixed window (simple, memory-friendly)
6- Sliding window (burst-tolerant)
7- Token bucket (smooth rate control)
9Storage backends:
10- In-memory (default, single-process)
11- Redis (distributed)
13Usage:
14 from agentos.api.rate_limiter import RateLimitMiddleware, FixedWindowLimiter
16 limiter = FixedWindowLimiter(max_requests=100, window_seconds=60)
17 app.add_middleware(RateLimitMiddleware, limiter=limiter)
18"""
20from __future__ import annotations
22import threading
23import time
24from collections.abc import Callable
26from starlette.middleware.base import BaseHTTPMiddleware, RequestResponseEndpoint
27from starlette.requests import Request
28from starlette.responses import JSONResponse, Response
30# ── Limiters ────────────────────────────────────────────────────────────────
33class FixedWindowLimiter:
34 """Fixed-window counter. 100 req/min → resets every 60s."""
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()
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 }
66# ── Middleware ──────────────────────────────────────────────────────────────
68RATE_LIMIT_HEADERS = {"X-RateLimit-Limit", "X-RateLimit-Remaining", "X-RateLimit-Reset"}
71class RateLimitMiddleware(BaseHTTPMiddleware):
72 """Starlette middleware that rate-limits based on client IP or API key."""
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"])
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"
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)
97 key = self._key_func(request)
98 allowed, info = self.limiter.is_allowed(key)
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 )
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
120__all__ = ["RateLimitMiddleware", "FixedWindowLimiter"]