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
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-05 22:09 +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 time
23import threading
24from dataclasses import dataclass, field
25from typing import Callable, Optional
27from starlette.middleware.base import BaseHTTPMiddleware, RequestResponseEndpoint
28from starlette.requests import Request
29from starlette.responses import Response, JSONResponse
32# ── Limiters ────────────────────────────────────────────────────────────────
34class FixedWindowLimiter:
35 """Fixed-window counter. 100 req/min → resets every 60s."""
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()
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 }
67# ── Middleware ──────────────────────────────────────────────────────────────
69RATE_LIMIT_HEADERS = {"X-RateLimit-Limit", "X-RateLimit-Remaining", "X-RateLimit-Reset"}
72class RateLimitMiddleware(BaseHTTPMiddleware):
73 """Starlette middleware that rate-limits based on client IP or API key."""
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"])
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"
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)
100 key = self._key_func(request)
101 allowed, info = self.limiter.is_allowed(key)
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 )
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"]