# 请求限流机制 - 简单内存限流器 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)