""" API Key 限流中间件 基于 API Key 的请求限流,支持: - 每分钟请求数限制(默认 60) - 每日请求数限制(默认 10000) """ from __future__ import annotations import time from typing import Dict, Tuple from datetime import datetime, timedelta from collections import defaultdict from threading import Lock from fastapi import Request, HTTPException, status from starlette.middleware.base import BaseHTTPMiddleware from starlette.responses import JSONResponse import structlog logger = structlog.get_logger(__name__) class InMemoryRateLimiter: """ 基于内存的限流器 使用滑动窗口算法实现限流 """ def __init__(self): # 存储每个 API Key 的请求时间戳 # key_id -> [(timestamp, request_count), ...] self.minute_requests: Dict[str, list] = defaultdict(list) self.daily_requests: Dict[str, int] = defaultdict(int) self.daily_reset: Dict[str, datetime] = {} self.lock = Lock() def check_rate_limit( self, key_id: str, rpm_limit: int = 60, daily_limit: int = 10000, ) -> Tuple[bool, Dict[str, any]]: """ 检查是否超出限流 Args: key_id: API Key ID rpm_limit: 每分钟请求数限制 daily_limit: 每日请求数限制 Returns: (是否允许, 限流信息字典) """ with self.lock: now = time.time() current_date = datetime.utcnow().date() # 1. 清理过期的分钟级请求记录(保留最近1分钟) one_minute_ago = now - 60 self.minute_requests[key_id] = [ ts for ts in self.minute_requests[key_id] if ts > one_minute_ago ] # 2. 检查每分钟限制 minute_count = len(self.minute_requests[key_id]) if minute_count >= rpm_limit: return False, { "limit_type": "rpm", "limit": rpm_limit, "current": minute_count, "reset_at": int(one_minute_ago + 60), "retry_after": int(60 - (now - min(self.minute_requests[key_id]))), } # 3. 检查每日限制 # 如果日期变更,重置计数 if key_id not in self.daily_reset or self.daily_reset[key_id] < current_date: self.daily_requests[key_id] = 0 self.daily_reset[key_id] = current_date daily_count = self.daily_requests[key_id] if daily_count >= daily_limit: # 计算距离明天0点的秒数 tomorrow = datetime.utcnow().replace( hour=0, minute=0, second=0, microsecond=0 ) + timedelta(days=1) retry_after = int((tomorrow - datetime.utcnow()).total_seconds()) return False, { "limit_type": "daily", "limit": daily_limit, "current": daily_count, "reset_at": int(tomorrow.timestamp()), "retry_after": retry_after, } # 4. 记录本次请求 self.minute_requests[key_id].append(now) self.daily_requests[key_id] += 1 # 5. 返回限流信息 return True, { "rpm_limit": rpm_limit, "rpm_remaining": rpm_limit - minute_count - 1, "daily_limit": daily_limit, "daily_remaining": daily_limit - daily_count - 1, } def get_stats(self, key_id: str) -> Dict[str, int]: """获取某个 API Key 的统计信息""" with self.lock: now = time.time() one_minute_ago = now - 60 # 清理过期记录 self.minute_requests[key_id] = [ ts for ts in self.minute_requests[key_id] if ts > one_minute_ago ] return { "minute_requests": len(self.minute_requests[key_id]), "daily_requests": self.daily_requests.get(key_id, 0), } # 全局限流器实例 _rate_limiter = InMemoryRateLimiter() class RateLimitMiddleware(BaseHTTPMiddleware): """ 限流中间件 对使用 API Key 认证的请求进行限流 JWT Token 认证的请求不受影响 """ async def dispatch(self, request: Request, call_next): # 跳过不需要限流的路径 path = request.url.path skip_paths = { "/health", "/metrics", "/docs", "/redoc", "/openapi.json", } if path in skip_paths or not path.startswith("/api"): return await call_next(request) # 只对 API Key 认证的请求进行限流 principal = getattr(request.state, "principal", None) if not principal or principal.get("type") != "api_key": # JWT Token 认证或未认证的请求,不限流 return await call_next(request) api_key_id = principal.get("api_key_id") if not api_key_id: return await call_next(request) # 获取限流配置(可以从数据库读取,这里使用默认值) rpm_limit = 60 # 每分钟60次 daily_limit = 10000 # 每日10000次 # 检查限流 allowed, info = _rate_limiter.check_rate_limit( key_id=api_key_id, rpm_limit=rpm_limit, daily_limit=daily_limit, ) if not allowed: # 超出限流 limit_type = info.get("limit_type") error_message = ( f"超出每分钟请求限制 ({info['limit']})" if limit_type == "rpm" else f"超出每日请求限制 ({info['limit']})" ) logger.warning( "rate_limit_exceeded", api_key_id=api_key_id, limit_type=limit_type, limit=info["limit"], current=info["current"], ) return JSONResponse( status_code=status.HTTP_429_TOO_MANY_REQUESTS, content={ "success": False, "error": error_message, "limit_type": limit_type, "limit": info["limit"], "current": info["current"], "retry_after": info["retry_after"], }, headers={ "X-RateLimit-Limit": str(info["limit"]), "X-RateLimit-Remaining": "0", "X-RateLimit-Reset": str(info["reset_at"]), "Retry-After": str(info["retry_after"]), } ) # 继续处理请求,添加限流信息到响应头 response = await call_next(request) # 添加限流响应头 response.headers["X-RateLimit-Limit-Minute"] = str(info["rpm_limit"]) response.headers["X-RateLimit-Remaining-Minute"] = str(info["rpm_remaining"]) response.headers["X-RateLimit-Limit-Daily"] = str(info["daily_limit"]) response.headers["X-RateLimit-Remaining-Daily"] = str(info["daily_remaining"]) return response def get_rate_limiter() -> InMemoryRateLimiter: """获取全局限流器实例""" return _rate_limiter