79 lines
2.5 KiB
Python
79 lines
2.5 KiB
Python
import time
|
|
|
|
from fastapi import Request, Response
|
|
from starlette.middleware.base import BaseHTTPMiddleware, RequestResponseEndpoint
|
|
from starlette.responses import JSONResponse
|
|
|
|
from app.config import settings
|
|
from app.utils.redis import get_redis
|
|
|
|
|
|
RATE_LIMITS = {
|
|
"POST:/api/auth/login": (10, 60),
|
|
"POST:/api/auth/change-password": (5, 60),
|
|
"POST:/api/generation-records/optimize": (20, 60),
|
|
"POST:/api/generation-records/*/generate": (10, 60),
|
|
"POST:/api/generation-records/*/retry": (10, 60),
|
|
}
|
|
|
|
|
|
def _match_rate_limit(path: str, method: str) -> tuple[int, int] | None:
|
|
key = f"{method}:{path}"
|
|
if key in RATE_LIMITS:
|
|
return RATE_LIMITS[key]
|
|
# Check wildcard patterns
|
|
for pattern, limit in RATE_LIMITS.items():
|
|
pattern_method, pattern_path = pattern.split(":", 1)
|
|
if method != pattern_method:
|
|
continue
|
|
pattern_parts = pattern_path.split("/")
|
|
key_parts = path.split("/")
|
|
if len(pattern_parts) != len(key_parts):
|
|
continue
|
|
match = True
|
|
for pp, kp in zip(pattern_parts, key_parts):
|
|
if pp != "*" and pp != kp:
|
|
match = False
|
|
break
|
|
if match:
|
|
return limit
|
|
return None
|
|
|
|
|
|
class RateLimitMiddleware(BaseHTTPMiddleware):
|
|
async def dispatch(
|
|
self, request: Request, call_next: RequestResponseEndpoint
|
|
) -> Response:
|
|
if not settings.RATE_LIMIT_ENABLED:
|
|
return await call_next(request)
|
|
|
|
limit_info = _match_rate_limit(request.url.path, request.method)
|
|
if not limit_info:
|
|
return await call_next(request)
|
|
|
|
max_requests, window = limit_info
|
|
# Use user_id from token or IP as identifier
|
|
identifier = request.client.host if request.client else "unknown"
|
|
|
|
redis = get_redis()
|
|
if not redis:
|
|
return await call_next(request)
|
|
|
|
redis_key = f"ratelimit:{request.method}:{request.url.path}:{identifier}"
|
|
|
|
try:
|
|
current = await redis.incr(redis_key)
|
|
if current == 1:
|
|
await redis.expire(redis_key, window)
|
|
if current > max_requests:
|
|
ttl = await redis.ttl(redis_key)
|
|
return JSONResponse(
|
|
status_code=429,
|
|
content={"detail": "请求过于频繁,请稍后重试"},
|
|
headers={"Retry-After": str(max(ttl, 1))},
|
|
)
|
|
except Exception:
|
|
pass # Redis unavailable, skip rate limiting
|
|
|
|
return await call_next(request)
|