166 lines
5.0 KiB
Python
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)
|