forked from xiaohei/taiji-AI-PAD
229 lines
7.6 KiB
Python
229 lines
7.6 KiB
Python
"""
|
||
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
|