jiachenlong/backend/app/core/rate_limit.py

166 lines
5.0 KiB
Python

# 请求限流机制 - 简单内存限流器
import time
from collections import defaultdict
from functools import wraps
from typing import Callable, Optional
from fastapi import HTTPException, Request
from app.core.config import settings
class RateLimiter:
"""简单内存限流器"""
def __init__(self):
self._requests = defaultdict(list)
self._enabled = settings.RATE_LIMIT_ENABLED
def _cleanup(self, key: str, window: int):
"""清理过期的请求记录"""
now = time.time()
self._requests[key] = [
ts for ts in self._requests[key]
if now - ts < window
]
def is_allowed(self, key: str, max_requests: int, window: int = 60) -> tuple[bool, int]:
"""
检查是否允许请求
Args:
key: 限流键 (如 IP、用户ID、手机号等)
max_requests: 时间窗口内最大请求数
window: 时间窗口秒数
Returns:
(是否允许, 剩余请求数)
"""
if not self._enabled:
return True, max_requests
now = time.time()
self._cleanup(key, window)
current_count = len(self._requests[key])
remaining = max(0, max_requests - current_count)
if current_count >= max_requests:
return False, 0
self._requests[key].append(now)
return True, remaining - 1
def get_retry_after(self, key: str, window: int = 60) -> int:
"""获取重试前需要等待的秒数"""
if not self._requests[key]:
return 0
oldest = min(self._requests[key])
now = time.time()
elapsed = now - oldest
remaining = window - elapsed
return max(0, int(remaining))
# 全局限流器实例
_rate_limiter = RateLimiter()
def get_rate_limiter() -> RateLimiter:
"""获取限流器实例"""
return _rate_limiter
def rate_limit(key_func: Callable[[Request], str], max_requests: int, window: int = 60):
"""
限流装饰器
Args:
key_func: 从请求中提取限流键的函数
max_requests: 最大请求数
window: 时间窗口(秒)
Example:
@rate_limit(lambda r: r.client.host, 10, 60)
async def my_endpoint():
...
"""
def decorator(func):
@wraps(func)
async def wrapper(request: Request, *args, **kwargs):
limiter = get_rate_limiter()
key = key_func(request)
allowed, remaining = limiter.is_allowed(key, max_requests, window)
if not allowed:
retry_after = limiter.get_retry_after(key, window)
raise HTTPException(
status_code=429,
detail=f"请求过于频繁,请 {retry_after} 秒后重试",
headers={"Retry-After": str(retry_after)}
)
response = await func(request, *args, **kwargs)
# 如果返回的是 Response 对象,添加限流头
if hasattr(response, 'headers'):
response.headers['X-RateLimit-Remaining'] = str(remaining)
response.headers['X-RateLimit-Limit'] = str(max_requests)
return response
# 对于非 async 函数
if not hasattr(wrapper, '__wrapped__'):
@wraps(func)
def sync_wrapper(*args, **kwargs):
return func(*args, **kwargs)
return sync_wrapper
return wrapper
return decorator
def rate_limit_by_ip(max_requests: int = 60, window: int = 60):
"""按IP限流的装饰器"""
return rate_limit(lambda r: r.client.host if r.client else "unknown", max_requests, window)
def rate_limit_by_phone(phone: str, max_requests: int, window: int = 60) -> bool:
"""
按手机号限流(用于短信发送等场景)
Returns:
是否允许发送
"""
limiter = get_rate_limiter()
allowed, _ = limiter.is_allowed(f"phone:{phone}", max_requests, window)
return allowed
def rate_limit_sms():
"""短信限流装饰器工厂"""
def key_func(request: Request) -> str:
# 尝试从body获取手机号
import json
try:
body = json.loads(request.body.decode())
phone = body.get("phone", "")
except:
phone = ""
return f"phone:{phone}" if phone else request.client.host
return rate_limit(key_func, settings.RATE_LIMIT_SMS_PER_MINUTE, 60)
def rate_limit_ocr():
"""OCR接口限流装饰器"""
return rate_limit(lambda r: r.client.host if r.client else "unknown",
settings.RATE_LIMIT_OCR_PER_MINUTE, 60)
def rate_limit_batch():
"""批量解析接口限流装饰器"""
return rate_limit(lambda r: r.client.host if r.client else "unknown",
settings.RATE_LIMIT_BATCH_PER_MINUTE, 60)