2025-12-10 20:52:44 +08:00
|
|
|
|
"""
|
2026-01-10 18:43:53 +08:00
|
|
|
|
RPM 限制管理器 - 支持 Redis 或内存的 Key 级别 RPM 限制
|
2025-12-10 20:52:44 +08:00
|
|
|
|
|
|
|
|
|
|
功能:
|
2026-01-10 18:43:53 +08:00
|
|
|
|
1. ProviderAPIKey 级别的 RPM 限制(按分钟窗口计数)
|
|
|
|
|
|
2. 分布式环境下优先使用 Redis,多实例共享
|
|
|
|
|
|
3. 在开发/单实例场景下自动降级为内存计数
|
|
|
|
|
|
4. 支持缓存用户优先级(预留槽位机制)
|
2025-12-10 20:52:44 +08:00
|
|
|
|
"""
|
|
|
|
|
|
|
2026-01-30 03:10:21 +08:00
|
|
|
|
from __future__ import annotations
|
|
|
|
|
|
|
2025-12-10 20:52:44 +08:00
|
|
|
|
import asyncio
|
|
|
|
|
|
import math
|
2026-01-10 18:43:53 +08:00
|
|
|
|
import os
|
|
|
|
|
|
import time
|
2025-12-10 20:52:44 +08:00
|
|
|
|
from contextlib import asynccontextmanager
|
2026-02-01 17:28:00 +08:00
|
|
|
|
from typing import Any
|
2025-12-10 20:52:44 +08:00
|
|
|
|
|
|
|
|
|
|
import redis.asyncio as aioredis
|
2026-01-10 18:43:53 +08:00
|
|
|
|
|
|
|
|
|
|
from src.config.constants import RPMDefaults
|
2025-12-10 20:52:44 +08:00
|
|
|
|
from src.core.logger import logger
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
class ConcurrencyManager:
|
2026-01-10 18:43:53 +08:00
|
|
|
|
"""Key RPM 限制管理器"""
|
2025-12-10 20:52:44 +08:00
|
|
|
|
|
2026-01-30 03:10:21 +08:00
|
|
|
|
_instance: ConcurrencyManager | None = None
|
|
|
|
|
|
_redis: aioredis.Redis | None = None
|
2025-12-10 20:52:44 +08:00
|
|
|
|
|
2026-01-30 14:30:57 +08:00
|
|
|
|
def __new__(cls) -> "ConcurrencyManager":
|
2025-12-10 20:52:44 +08:00
|
|
|
|
"""单例模式"""
|
|
|
|
|
|
if cls._instance is None:
|
|
|
|
|
|
cls._instance = super().__new__(cls)
|
|
|
|
|
|
return cls._instance
|
|
|
|
|
|
|
2026-01-30 14:30:57 +08:00
|
|
|
|
def __init__(self) -> None:
|
2025-12-10 20:52:44 +08:00
|
|
|
|
"""初始化内存后端结构(只执行一次)"""
|
|
|
|
|
|
if hasattr(self, "_memory_initialized"):
|
|
|
|
|
|
return
|
|
|
|
|
|
|
2026-02-16 11:00:48 +08:00
|
|
|
|
from src.config.settings import config
|
|
|
|
|
|
|
|
|
|
|
|
self._key_rpm_bucket_seconds: int = config.rpm_bucket_seconds
|
|
|
|
|
|
self._key_rpm_key_ttl_seconds: int = config.rpm_key_ttl_seconds
|
|
|
|
|
|
|
2025-12-10 20:52:44 +08:00
|
|
|
|
self._memory_lock: asyncio.Lock = asyncio.Lock()
|
2026-01-10 18:43:53 +08:00
|
|
|
|
# Key RPM 计数器:{key_id: (bucket, count)},bucket = floor(now / 60)
|
|
|
|
|
|
self._memory_key_rpm_counts: dict[str, tuple[int, int]] = {}
|
2025-12-18 00:35:46 +08:00
|
|
|
|
self._owns_redis: bool = False
|
2026-01-10 18:43:53 +08:00
|
|
|
|
self._last_cleanup_bucket: int = 0 # 上次清理时的 bucket,用于定期清理过期数据
|
|
|
|
|
|
self._last_cleanup_time: float = 0 # 上次清理的时间戳,用于强制定期清理
|
2026-02-16 11:00:48 +08:00
|
|
|
|
self._cleanup_interval_seconds: int = config.rpm_cleanup_interval_seconds
|
2026-01-30 03:10:21 +08:00
|
|
|
|
self._cleanup_task: asyncio.Task | None = None # 后台清理任务
|
2026-01-10 18:43:53 +08:00
|
|
|
|
|
|
|
|
|
|
# 内存模式下的最大条目限制,防止内存泄漏(支持环境变量覆盖)
|
|
|
|
|
|
self._max_memory_rpm_entries: int = int(
|
|
|
|
|
|
os.getenv("RPM_MAX_MEMORY_ENTRIES", str(RPMDefaults.MAX_MEMORY_RPM_ENTRIES))
|
|
|
|
|
|
)
|
|
|
|
|
|
# 早期告警阈值(达到此比例时记录警告)
|
|
|
|
|
|
self._memory_warning_threshold: float = float(
|
|
|
|
|
|
os.getenv("RPM_MEMORY_WARNING_THRESHOLD", str(RPMDefaults.MEMORY_WARNING_THRESHOLD))
|
|
|
|
|
|
)
|
2025-12-10 20:52:44 +08:00
|
|
|
|
self._memory_initialized = True
|
|
|
|
|
|
|
|
|
|
|
|
async def initialize(self) -> None:
|
|
|
|
|
|
"""初始化 Redis 连接"""
|
|
|
|
|
|
if self._redis is not None:
|
|
|
|
|
|
return
|
|
|
|
|
|
|
2025-12-18 00:35:46 +08:00
|
|
|
|
try:
|
|
|
|
|
|
# 复用全局 Redis 客户端(带熔断/降级),避免重复创建连接池
|
|
|
|
|
|
from src.clients.redis_client import get_redis_client
|
2025-12-10 20:52:44 +08:00
|
|
|
|
|
2025-12-18 00:35:46 +08:00
|
|
|
|
self._redis = await get_redis_client(require_redis=False)
|
|
|
|
|
|
self._owns_redis = False
|
|
|
|
|
|
if self._redis:
|
|
|
|
|
|
logger.info("[OK] ConcurrencyManager 已复用全局 Redis 客户端")
|
2025-12-10 20:52:44 +08:00
|
|
|
|
else:
|
2026-02-01 17:28:00 +08:00
|
|
|
|
logger.warning(
|
|
|
|
|
|
"[WARN] Redis 不可用,RPM 限制降级为内存模式(仅在单实例环境下安全)"
|
|
|
|
|
|
)
|
2026-01-10 18:43:53 +08:00
|
|
|
|
# 内存模式下启动后台清理任务
|
|
|
|
|
|
self._start_background_cleanup()
|
2025-12-10 20:52:44 +08:00
|
|
|
|
except Exception as e:
|
2026-02-15 16:32:23 +08:00
|
|
|
|
logger.error("[ERROR] 获取全局 Redis 客户端失败: {}", e)
|
2026-01-10 18:43:53 +08:00
|
|
|
|
logger.warning("[WARN] RPM 限制将降级为内存模式(仅在单实例环境下安全)")
|
2025-12-10 20:52:44 +08:00
|
|
|
|
self._redis = None
|
2025-12-18 00:35:46 +08:00
|
|
|
|
self._owns_redis = False
|
2026-01-10 18:43:53 +08:00
|
|
|
|
# 内存模式下启动后台清理任务
|
|
|
|
|
|
self._start_background_cleanup()
|
|
|
|
|
|
|
|
|
|
|
|
def _start_background_cleanup(self) -> None:
|
|
|
|
|
|
"""启动后台定期清理任务(仅内存模式需要)"""
|
|
|
|
|
|
if self._cleanup_task is not None:
|
|
|
|
|
|
return # 已经启动
|
|
|
|
|
|
|
2026-01-30 14:30:57 +08:00
|
|
|
|
async def cleanup_loop() -> None:
|
2026-01-10 18:43:53 +08:00
|
|
|
|
"""后台清理循环"""
|
|
|
|
|
|
while True:
|
|
|
|
|
|
try:
|
|
|
|
|
|
await asyncio.sleep(60) # 每分钟检查一次
|
|
|
|
|
|
async with self._memory_lock:
|
|
|
|
|
|
current_bucket = self._get_rpm_bucket()
|
|
|
|
|
|
self._cleanup_expired_memory_rpm_counts(current_bucket, force=False)
|
|
|
|
|
|
except asyncio.CancelledError:
|
|
|
|
|
|
break
|
|
|
|
|
|
except Exception as e:
|
2026-02-15 16:32:23 +08:00
|
|
|
|
logger.debug("后台清理任务异常: {}", e)
|
2026-01-10 18:43:53 +08:00
|
|
|
|
|
|
|
|
|
|
try:
|
|
|
|
|
|
self._cleanup_task = asyncio.create_task(cleanup_loop())
|
|
|
|
|
|
logger.debug("[OK] 内存模式后台清理任务已启动")
|
|
|
|
|
|
except RuntimeError:
|
|
|
|
|
|
# 没有事件循环时忽略
|
|
|
|
|
|
pass
|
2025-12-10 20:52:44 +08:00
|
|
|
|
|
|
|
|
|
|
async def close(self) -> None:
|
|
|
|
|
|
"""关闭 Redis 连接"""
|
2026-01-10 18:43:53 +08:00
|
|
|
|
# 停止后台清理任务
|
|
|
|
|
|
if self._cleanup_task is not None:
|
|
|
|
|
|
self._cleanup_task.cancel()
|
|
|
|
|
|
try:
|
|
|
|
|
|
await self._cleanup_task
|
|
|
|
|
|
except asyncio.CancelledError:
|
|
|
|
|
|
pass
|
|
|
|
|
|
self._cleanup_task = None
|
|
|
|
|
|
|
2025-12-18 00:35:46 +08:00
|
|
|
|
if self._redis and self._owns_redis:
|
2025-12-10 20:52:44 +08:00
|
|
|
|
await self._redis.close()
|
2025-12-18 00:35:46 +08:00
|
|
|
|
logger.info("ConcurrencyManager Redis 连接已关闭")
|
|
|
|
|
|
self._redis = None
|
|
|
|
|
|
self._owns_redis = False
|
2025-12-10 20:52:44 +08:00
|
|
|
|
|
2026-02-16 11:00:48 +08:00
|
|
|
|
def _get_rpm_bucket(self, now_ts: float | None = None) -> int:
|
2026-01-10 18:43:53 +08:00
|
|
|
|
"""获取当前 RPM 计数桶(按分钟)"""
|
|
|
|
|
|
ts = now_ts if now_ts is not None else time.time()
|
2026-02-16 11:00:48 +08:00
|
|
|
|
return int(ts // self._key_rpm_bucket_seconds)
|
2026-01-10 18:43:53 +08:00
|
|
|
|
|
2026-02-16 11:00:48 +08:00
|
|
|
|
def _get_key_key(self, key_id: str, bucket: int | None = None) -> str:
|
2026-01-10 18:43:53 +08:00
|
|
|
|
"""获取 ProviderAPIKey RPM 计数的 Redis Key(按分钟桶)"""
|
2026-02-16 11:00:48 +08:00
|
|
|
|
b = bucket if bucket is not None else self._get_rpm_bucket()
|
2026-01-10 18:43:53 +08:00
|
|
|
|
return f"rpm:key:{key_id}:{b}"
|
|
|
|
|
|
|
|
|
|
|
|
def _get_memory_key_rpm_count(self, key_id: str, bucket: int) -> int:
|
|
|
|
|
|
"""获取内存模式下 Key 在指定 bucket 的 RPM 计数"""
|
|
|
|
|
|
stored = self._memory_key_rpm_counts.get(key_id)
|
|
|
|
|
|
if not stored:
|
|
|
|
|
|
return 0
|
|
|
|
|
|
stored_bucket, count = stored
|
|
|
|
|
|
if stored_bucket != bucket:
|
|
|
|
|
|
# 旧桶数据已过期,删除以防止内存泄漏
|
|
|
|
|
|
del self._memory_key_rpm_counts[key_id]
|
|
|
|
|
|
return 0
|
|
|
|
|
|
return count
|
|
|
|
|
|
|
|
|
|
|
|
def _set_memory_key_rpm_count(self, key_id: str, bucket: int, count: int) -> None:
|
|
|
|
|
|
"""设置内存模式下 Key 在指定 bucket 的 RPM 计数"""
|
|
|
|
|
|
current_size = len(self._memory_key_rpm_counts)
|
|
|
|
|
|
warning_threshold = int(self._max_memory_rpm_entries * self._memory_warning_threshold)
|
|
|
|
|
|
high_threshold = int(self._max_memory_rpm_entries * 0.8)
|
|
|
|
|
|
critical_threshold = int(self._max_memory_rpm_entries * 0.95)
|
|
|
|
|
|
|
|
|
|
|
|
# 分级告警:根据使用率记录不同级别的日志
|
|
|
|
|
|
if current_size >= critical_threshold and key_id not in self._memory_key_rpm_counts:
|
|
|
|
|
|
logger.critical(
|
2026-02-15 16:32:23 +08:00
|
|
|
|
"[CRITICAL] 内存 RPM 计数器接近上限 ({}/{}),"
|
|
|
|
|
|
"强烈建议启用 Redis!继续增长可能导致 RPM 限制失效",
|
|
|
|
|
|
current_size,
|
|
|
|
|
|
self._max_memory_rpm_entries,
|
2026-01-10 18:43:53 +08:00
|
|
|
|
)
|
|
|
|
|
|
elif current_size >= high_threshold and key_id not in self._memory_key_rpm_counts:
|
|
|
|
|
|
# 每 100 个条目告警一次,避免日志过多
|
|
|
|
|
|
if current_size % 100 == 0:
|
|
|
|
|
|
logger.error(
|
2026-02-15 16:32:23 +08:00
|
|
|
|
"[ERROR] 内存 RPM 计数器使用率过高 ({}/{}),建议启用 Redis",
|
|
|
|
|
|
current_size,
|
|
|
|
|
|
self._max_memory_rpm_entries,
|
2026-01-10 18:43:53 +08:00
|
|
|
|
)
|
|
|
|
|
|
elif current_size >= warning_threshold and key_id not in self._memory_key_rpm_counts:
|
|
|
|
|
|
if current_size == warning_threshold:
|
|
|
|
|
|
logger.warning(
|
2026-02-15 16:32:23 +08:00
|
|
|
|
"[WARN] 内存 RPM 计数器达到 {:.0%} 阈值 ({}/{}),建议启用 Redis",
|
|
|
|
|
|
self._memory_warning_threshold,
|
|
|
|
|
|
current_size,
|
|
|
|
|
|
self._max_memory_rpm_entries,
|
2026-01-10 18:43:53 +08:00
|
|
|
|
)
|
2025-12-10 20:52:44 +08:00
|
|
|
|
|
2026-01-10 18:43:53 +08:00
|
|
|
|
# 检查是否超过最大条目限制
|
|
|
|
|
|
if (
|
|
|
|
|
|
key_id not in self._memory_key_rpm_counts
|
|
|
|
|
|
and current_size >= self._max_memory_rpm_entries
|
|
|
|
|
|
):
|
|
|
|
|
|
# 触发强制清理
|
|
|
|
|
|
self._cleanup_expired_memory_rpm_counts(bucket, force=True)
|
|
|
|
|
|
# 如果清理后仍然超过限制,执行 LRU 淘汰(删除最旧的 20%)
|
|
|
|
|
|
if len(self._memory_key_rpm_counts) >= self._max_memory_rpm_entries:
|
|
|
|
|
|
evict_count = max(1, self._max_memory_rpm_entries // 5)
|
|
|
|
|
|
# 按 bucket(时间)排序,删除最旧的
|
|
|
|
|
|
sorted_keys = sorted(
|
2026-02-01 17:28:00 +08:00
|
|
|
|
self._memory_key_rpm_counts.items(), key=lambda x: x[1][0] # 按 bucket 排序
|
2026-01-10 18:43:53 +08:00
|
|
|
|
)
|
|
|
|
|
|
for k, _ in sorted_keys[:evict_count]:
|
|
|
|
|
|
del self._memory_key_rpm_counts[k]
|
2026-02-15 16:32:23 +08:00
|
|
|
|
logger.warning("[WARN] 内存 RPM 计数器达到上限,已淘汰 {} 个最旧条目", evict_count)
|
2026-01-10 18:43:53 +08:00
|
|
|
|
self._memory_key_rpm_counts[key_id] = (bucket, count)
|
|
|
|
|
|
|
|
|
|
|
|
def _cleanup_expired_memory_rpm_counts(self, current_bucket: int, force: bool = False) -> None:
|
|
|
|
|
|
"""
|
|
|
|
|
|
清理内存中过期的 RPM 计数(必须在持有 _memory_lock 时调用)
|
2025-12-10 20:52:44 +08:00
|
|
|
|
|
2026-01-10 18:43:53 +08:00
|
|
|
|
清理策略:
|
|
|
|
|
|
- 常规清理:每分钟最多执行一次(当 bucket 变化时)
|
|
|
|
|
|
- 强制清理:每 5 分钟执行一次(防止长时间无请求导致内存泄漏)
|
2025-12-10 20:52:44 +08:00
|
|
|
|
"""
|
2026-01-10 18:43:53 +08:00
|
|
|
|
now = time.time()
|
|
|
|
|
|
|
|
|
|
|
|
# 检查是否需要清理
|
|
|
|
|
|
should_cleanup = (
|
|
|
|
|
|
current_bucket != self._last_cleanup_bucket # 分钟切换
|
|
|
|
|
|
or force # 强制清理
|
|
|
|
|
|
or (now - self._last_cleanup_time > self._cleanup_interval_seconds) # 超时清理
|
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
if not should_cleanup:
|
|
|
|
|
|
return
|
|
|
|
|
|
|
|
|
|
|
|
self._last_cleanup_bucket = current_bucket
|
|
|
|
|
|
self._last_cleanup_time = now
|
|
|
|
|
|
expired_keys = []
|
|
|
|
|
|
|
|
|
|
|
|
for key_id, (stored_bucket, _count) in self._memory_key_rpm_counts.items():
|
|
|
|
|
|
if stored_bucket < current_bucket:
|
|
|
|
|
|
expired_keys.append(key_id)
|
2025-12-10 20:52:44 +08:00
|
|
|
|
|
2026-01-10 18:43:53 +08:00
|
|
|
|
for key_id in expired_keys:
|
|
|
|
|
|
del self._memory_key_rpm_counts[key_id]
|
|
|
|
|
|
|
|
|
|
|
|
if expired_keys:
|
2026-02-15 16:32:23 +08:00
|
|
|
|
logger.debug("[CLEANUP] 清理了 {} 个过期的内存 RPM 计数", len(expired_keys))
|
2026-01-10 18:43:53 +08:00
|
|
|
|
|
|
|
|
|
|
async def get_key_rpm_count(self, key_id: str) -> int:
|
|
|
|
|
|
"""
|
|
|
|
|
|
获取 Key 当前 RPM 计数
|
2026-01-08 13:34:59 +08:00
|
|
|
|
|
2025-12-10 20:52:44 +08:00
|
|
|
|
Args:
|
2026-01-10 18:43:53 +08:00
|
|
|
|
key_id: ProviderAPIKey ID
|
2025-12-10 20:52:44 +08:00
|
|
|
|
|
|
|
|
|
|
Returns:
|
2026-01-10 18:43:53 +08:00
|
|
|
|
当前分钟窗口内的请求数
|
2025-12-10 20:52:44 +08:00
|
|
|
|
"""
|
|
|
|
|
|
if self._redis is None:
|
|
|
|
|
|
async with self._memory_lock:
|
2026-01-10 18:43:53 +08:00
|
|
|
|
bucket = self._get_rpm_bucket()
|
|
|
|
|
|
# 定期清理过期数据,避免内存泄漏
|
|
|
|
|
|
self._cleanup_expired_memory_rpm_counts(bucket)
|
|
|
|
|
|
return self._get_memory_key_rpm_count(key_id, bucket)
|
2025-12-10 20:52:44 +08:00
|
|
|
|
|
|
|
|
|
|
try:
|
2026-01-10 18:43:53 +08:00
|
|
|
|
key_key = self._get_key_key(key_id)
|
|
|
|
|
|
result = await self._redis.get(key_key)
|
|
|
|
|
|
return int(result) if result else 0
|
2025-12-10 20:52:44 +08:00
|
|
|
|
except Exception as e:
|
2026-02-15 16:32:23 +08:00
|
|
|
|
logger.error("获取 RPM 计数失败: {}", e)
|
2026-01-10 18:43:53 +08:00
|
|
|
|
return 0
|
2025-12-10 20:52:44 +08:00
|
|
|
|
|
2026-01-10 18:43:53 +08:00
|
|
|
|
async def check_rpm_available(
|
2025-12-10 20:52:44 +08:00
|
|
|
|
self,
|
|
|
|
|
|
key_id: str,
|
2026-01-30 03:10:21 +08:00
|
|
|
|
key_rpm_limit: int | None,
|
2026-01-10 18:43:53 +08:00
|
|
|
|
is_cached_user: bool = False,
|
2026-01-30 03:10:21 +08:00
|
|
|
|
cache_reservation_ratio: float | None = None,
|
2025-12-10 20:52:44 +08:00
|
|
|
|
) -> bool:
|
|
|
|
|
|
"""
|
2026-01-10 18:43:53 +08:00
|
|
|
|
检查是否可以通过 RPM 限制(不实际增加计数)
|
2025-12-10 20:52:44 +08:00
|
|
|
|
|
|
|
|
|
|
Args:
|
|
|
|
|
|
key_id: ProviderAPIKey ID
|
2026-01-10 18:43:53 +08:00
|
|
|
|
key_rpm_limit: Key RPM 限制(每分钟最大请求数,None 表示不限制)
|
|
|
|
|
|
is_cached_user: 是否是缓存用户
|
|
|
|
|
|
cache_reservation_ratio: 缓存预留比例
|
2025-12-10 20:52:44 +08:00
|
|
|
|
|
|
|
|
|
|
Returns:
|
|
|
|
|
|
是否可用(True/False)
|
|
|
|
|
|
"""
|
2026-01-10 18:43:53 +08:00
|
|
|
|
if key_rpm_limit is None:
|
|
|
|
|
|
return True
|
2025-12-10 20:52:44 +08:00
|
|
|
|
|
2026-01-10 18:43:53 +08:00
|
|
|
|
# 从配置读取默认值
|
|
|
|
|
|
from src.config.settings import config
|
2025-12-10 20:52:44 +08:00
|
|
|
|
|
2026-01-10 18:43:53 +08:00
|
|
|
|
if cache_reservation_ratio is None:
|
|
|
|
|
|
cache_reservation_ratio = config.cache_reservation_ratio
|
2025-12-10 20:52:44 +08:00
|
|
|
|
|
2026-01-10 18:43:53 +08:00
|
|
|
|
key_count = await self.get_key_rpm_count(key_id)
|
2025-12-10 20:52:44 +08:00
|
|
|
|
|
2026-01-10 18:43:53 +08:00
|
|
|
|
if is_cached_user:
|
|
|
|
|
|
return key_count < key_rpm_limit
|
|
|
|
|
|
else:
|
|
|
|
|
|
# 新用户只能使用 (1 - cache_reservation_ratio) 的槽位
|
|
|
|
|
|
available_for_new = max(1, math.floor(key_rpm_limit * (1 - cache_reservation_ratio)))
|
|
|
|
|
|
return key_count < available_for_new
|
2025-12-10 20:52:44 +08:00
|
|
|
|
|
2026-01-10 18:43:53 +08:00
|
|
|
|
async def acquire_rpm_slot(
|
2025-12-10 20:52:44 +08:00
|
|
|
|
self,
|
|
|
|
|
|
key_id: str,
|
2026-01-30 03:10:21 +08:00
|
|
|
|
key_rpm_limit: int | None,
|
2026-01-10 18:43:53 +08:00
|
|
|
|
is_cached_user: bool = False,
|
2026-01-30 03:10:21 +08:00
|
|
|
|
cache_reservation_ratio: float | None = None,
|
2025-12-10 20:52:44 +08:00
|
|
|
|
) -> bool:
|
|
|
|
|
|
"""
|
2026-01-10 18:43:53 +08:00
|
|
|
|
尝试获取 RPM 槽位(支持缓存用户优先级)
|
2025-12-10 20:52:44 +08:00
|
|
|
|
|
|
|
|
|
|
Args:
|
|
|
|
|
|
key_id: ProviderAPIKey ID
|
2026-01-10 18:43:53 +08:00
|
|
|
|
key_rpm_limit: Key RPM 限制(每分钟最大请求数,None 表示不限制)
|
2025-12-10 20:52:44 +08:00
|
|
|
|
is_cached_user: 是否是缓存用户(缓存用户可使用全部槽位)
|
2025-12-12 15:42:45 +08:00
|
|
|
|
cache_reservation_ratio: 缓存预留比例,None 时从配置读取
|
2025-12-10 20:52:44 +08:00
|
|
|
|
|
|
|
|
|
|
Returns:
|
|
|
|
|
|
是否成功获取(True/False)
|
|
|
|
|
|
|
|
|
|
|
|
缓存预留机制说明:
|
2026-01-10 18:43:53 +08:00
|
|
|
|
- 假设 key_rpm_limit = 100, cache_reservation_ratio = 0.3
|
|
|
|
|
|
- 新用户最多使用: 70 RPM (100 * (1 - 0.3))
|
|
|
|
|
|
- 缓存用户最多使用: 100 RPM(全部)
|
|
|
|
|
|
- 预留的 30 RPM 专门给缓存用户,保证他们的请求优先
|
2025-12-10 20:52:44 +08:00
|
|
|
|
"""
|
2025-12-12 15:42:45 +08:00
|
|
|
|
# 从配置读取默认值
|
|
|
|
|
|
from src.config.settings import config
|
|
|
|
|
|
|
|
|
|
|
|
if cache_reservation_ratio is None:
|
|
|
|
|
|
cache_reservation_ratio = config.cache_reservation_ratio
|
|
|
|
|
|
|
2025-12-10 20:52:44 +08:00
|
|
|
|
if self._redis is None:
|
|
|
|
|
|
async with self._memory_lock:
|
2026-01-10 18:43:53 +08:00
|
|
|
|
bucket = self._get_rpm_bucket()
|
|
|
|
|
|
# 定期清理过期数据,避免内存泄漏
|
|
|
|
|
|
self._cleanup_expired_memory_rpm_counts(bucket)
|
2025-12-10 20:52:44 +08:00
|
|
|
|
|
2026-01-10 18:43:53 +08:00
|
|
|
|
key_count = self._get_memory_key_rpm_count(key_id, bucket)
|
|
|
|
|
|
|
|
|
|
|
|
# Key RPM 限制,包含缓存预留
|
|
|
|
|
|
if key_rpm_limit is not None:
|
2025-12-10 20:52:44 +08:00
|
|
|
|
if is_cached_user:
|
2026-01-10 18:43:53 +08:00
|
|
|
|
if key_count >= key_rpm_limit:
|
2025-12-10 20:52:44 +08:00
|
|
|
|
return False
|
|
|
|
|
|
else:
|
2026-01-10 18:43:53 +08:00
|
|
|
|
# 新用户只能使用 (1 - cache_reservation_ratio) 的槽位
|
2025-12-10 20:52:44 +08:00
|
|
|
|
available_for_new = max(
|
2026-01-10 18:43:53 +08:00
|
|
|
|
1, math.floor(key_rpm_limit * (1 - cache_reservation_ratio))
|
2025-12-10 20:52:44 +08:00
|
|
|
|
)
|
|
|
|
|
|
if key_count >= available_for_new:
|
|
|
|
|
|
return False
|
|
|
|
|
|
|
|
|
|
|
|
# 通过限制,更新计数
|
2026-01-10 18:43:53 +08:00
|
|
|
|
self._set_memory_key_rpm_count(key_id, bucket, key_count + 1)
|
2025-12-10 20:52:44 +08:00
|
|
|
|
return True
|
|
|
|
|
|
|
2026-01-10 18:43:53 +08:00
|
|
|
|
bucket = self._get_rpm_bucket()
|
|
|
|
|
|
key_key = self._get_key_key(key_id, bucket=bucket)
|
2025-12-10 20:52:44 +08:00
|
|
|
|
|
|
|
|
|
|
try:
|
2026-01-10 18:43:53 +08:00
|
|
|
|
# 使用 Lua 脚本保证原子性(支持缓存预留逻辑)
|
2025-12-10 20:52:44 +08:00
|
|
|
|
lua_script = """
|
2026-01-10 18:43:53 +08:00
|
|
|
|
local key_key = KEYS[1]
|
|
|
|
|
|
local key_max = tonumber(ARGV[1])
|
|
|
|
|
|
local key_ttl = tonumber(ARGV[2])
|
|
|
|
|
|
local is_cached = tonumber(ARGV[3]) -- 0=新用户, 1=缓存用户
|
|
|
|
|
|
local cache_ratio = tonumber(ARGV[4]) -- 缓存预留比例
|
2025-12-10 20:52:44 +08:00
|
|
|
|
|
|
|
|
|
|
-- 获取当前值
|
|
|
|
|
|
local key_count = tonumber(redis.call('GET', key_key) or '0')
|
|
|
|
|
|
|
|
|
|
|
|
-- 检查 key 限制(支持缓存预留)
|
|
|
|
|
|
if key_max >= 0 then
|
|
|
|
|
|
if is_cached == 0 then
|
|
|
|
|
|
-- 新用户:只能使用 (1 - cache_ratio) 的槽位
|
2026-01-10 18:43:53 +08:00
|
|
|
|
local available_for_new = math.max(1, math.floor(key_max * (1 - cache_ratio)))
|
2025-12-10 20:52:44 +08:00
|
|
|
|
if key_count >= available_for_new then
|
|
|
|
|
|
return 0 -- 失败:新用户配额已满
|
|
|
|
|
|
end
|
|
|
|
|
|
else
|
|
|
|
|
|
-- 缓存用户:可以使用全部槽位
|
|
|
|
|
|
if key_count >= key_max then
|
|
|
|
|
|
return 0 -- 失败:总配额已满
|
|
|
|
|
|
end
|
|
|
|
|
|
end
|
|
|
|
|
|
end
|
|
|
|
|
|
|
|
|
|
|
|
-- 增加计数
|
|
|
|
|
|
redis.call('INCR', key_key)
|
2026-01-10 18:43:53 +08:00
|
|
|
|
redis.call('EXPIRE', key_key, key_ttl)
|
2025-12-10 20:52:44 +08:00
|
|
|
|
|
|
|
|
|
|
return 1 -- 成功
|
|
|
|
|
|
"""
|
|
|
|
|
|
|
|
|
|
|
|
# 执行脚本
|
|
|
|
|
|
result = await self._redis.eval(
|
|
|
|
|
|
lua_script,
|
2026-01-10 18:43:53 +08:00
|
|
|
|
1, # 1 个 KEY
|
2025-12-10 20:52:44 +08:00
|
|
|
|
key_key,
|
2026-01-10 18:43:53 +08:00
|
|
|
|
key_rpm_limit if key_rpm_limit is not None else -1,
|
|
|
|
|
|
self._key_rpm_key_ttl_seconds,
|
2025-12-10 20:52:44 +08:00
|
|
|
|
1 if is_cached_user else 0, # 缓存用户标志
|
|
|
|
|
|
cache_reservation_ratio, # 预留比例
|
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
success = result == 1
|
|
|
|
|
|
|
|
|
|
|
|
if success:
|
|
|
|
|
|
user_type = "缓存用户" if is_cached_user else "新用户"
|
2026-02-15 16:32:23 +08:00
|
|
|
|
logger.debug("[OK] 获取 RPM 槽位成功: key={}, 类型={}", key_id, user_type)
|
2025-12-10 20:52:44 +08:00
|
|
|
|
else:
|
2026-01-10 18:43:53 +08:00
|
|
|
|
key_count = await self.get_key_rpm_count(key_id)
|
2025-12-10 20:52:44 +08:00
|
|
|
|
|
2026-01-10 18:43:53 +08:00
|
|
|
|
# 计算新用户可用 RPM
|
|
|
|
|
|
if key_rpm_limit and not is_cached_user:
|
|
|
|
|
|
available_for_new = int(key_rpm_limit * (1 - cache_reservation_ratio))
|
2025-12-10 20:52:44 +08:00
|
|
|
|
user_info = f"新用户配额={available_for_new}, 当前={key_count}"
|
|
|
|
|
|
else:
|
2026-01-10 18:43:53 +08:00
|
|
|
|
user_info = f"缓存用户, 当前={key_count}/{key_rpm_limit}"
|
2025-12-10 20:52:44 +08:00
|
|
|
|
|
2026-02-15 16:32:23 +08:00
|
|
|
|
logger.warning("[WARN] RPM 限制已达上限: key={}({})", key_id, user_info)
|
2025-12-10 20:52:44 +08:00
|
|
|
|
|
|
|
|
|
|
return success
|
|
|
|
|
|
|
|
|
|
|
|
except Exception as e:
|
2026-02-15 16:32:23 +08:00
|
|
|
|
logger.error("获取 RPM 槽位失败,降级到内存模式: {}", e)
|
2025-12-10 20:52:44 +08:00
|
|
|
|
# Redis 异常时降级到内存模式进行保守限流
|
|
|
|
|
|
async with self._memory_lock:
|
2026-01-10 18:43:53 +08:00
|
|
|
|
bucket = self._get_rpm_bucket()
|
|
|
|
|
|
self._cleanup_expired_memory_rpm_counts(bucket)
|
|
|
|
|
|
|
|
|
|
|
|
key_count = self._get_memory_key_rpm_count(key_id, bucket)
|
2025-12-10 20:52:44 +08:00
|
|
|
|
|
|
|
|
|
|
# 降级模式下使用更保守的限制(50%)
|
2026-01-10 18:43:53 +08:00
|
|
|
|
fallback_rpm_limit = (
|
|
|
|
|
|
max(1, key_rpm_limit // 2) if key_rpm_limit is not None else None
|
2025-12-10 20:52:44 +08:00
|
|
|
|
)
|
|
|
|
|
|
|
2026-01-10 18:43:53 +08:00
|
|
|
|
if fallback_rpm_limit is not None and key_count >= fallback_rpm_limit:
|
2025-12-10 20:52:44 +08:00
|
|
|
|
logger.warning(
|
2026-02-15 16:32:23 +08:00
|
|
|
|
"[FALLBACK] Key RPM 达到降级限制: {}/{}",
|
|
|
|
|
|
key_count,
|
|
|
|
|
|
fallback_rpm_limit,
|
2025-12-10 20:52:44 +08:00
|
|
|
|
)
|
|
|
|
|
|
return False
|
|
|
|
|
|
|
|
|
|
|
|
# 更新内存计数
|
2026-01-10 18:43:53 +08:00
|
|
|
|
self._set_memory_key_rpm_count(key_id, bucket, key_count + 1)
|
2026-02-15 16:32:23 +08:00
|
|
|
|
logger.debug("[FALLBACK] 使用内存模式获取 RPM 槽位: key={}", key_id)
|
2025-12-10 20:52:44 +08:00
|
|
|
|
return True
|
|
|
|
|
|
|
|
|
|
|
|
@asynccontextmanager
|
2026-01-10 18:43:53 +08:00
|
|
|
|
async def rpm_guard(
|
2025-12-10 20:52:44 +08:00
|
|
|
|
self,
|
|
|
|
|
|
key_id: str,
|
2026-01-30 03:10:21 +08:00
|
|
|
|
key_rpm_limit: int | None,
|
2026-01-10 18:43:53 +08:00
|
|
|
|
is_cached_user: bool = False,
|
2026-01-30 03:10:21 +08:00
|
|
|
|
cache_reservation_ratio: float | None = None,
|
2026-01-30 14:30:57 +08:00
|
|
|
|
) -> Any:
|
2025-12-10 20:52:44 +08:00
|
|
|
|
"""
|
2026-01-10 18:43:53 +08:00
|
|
|
|
RPM 限制上下文管理器(支持缓存用户优先级)
|
2025-12-10 20:52:44 +08:00
|
|
|
|
|
|
|
|
|
|
用法:
|
2026-01-10 18:43:53 +08:00
|
|
|
|
async with manager.rpm_guard(
|
|
|
|
|
|
key_id, key_rpm_limit,
|
2025-12-10 20:52:44 +08:00
|
|
|
|
is_cached_user=True # 缓存用户
|
|
|
|
|
|
):
|
|
|
|
|
|
# 执行请求
|
|
|
|
|
|
response = await send_request(...)
|
|
|
|
|
|
|
|
|
|
|
|
如果获取失败,会抛出 ConcurrencyLimitError 异常
|
2026-01-10 18:43:53 +08:00
|
|
|
|
|
|
|
|
|
|
注意:RPM 是按分钟窗口计数,不需要在请求结束后释放
|
2025-12-10 20:52:44 +08:00
|
|
|
|
"""
|
2025-12-12 15:42:45 +08:00
|
|
|
|
# 从配置读取默认值
|
|
|
|
|
|
from src.config.settings import config
|
|
|
|
|
|
|
|
|
|
|
|
if cache_reservation_ratio is None:
|
|
|
|
|
|
cache_reservation_ratio = config.cache_reservation_ratio
|
|
|
|
|
|
|
2025-12-10 20:52:44 +08:00
|
|
|
|
# 尝试获取槽位(传递缓存用户参数)
|
2026-01-10 18:43:53 +08:00
|
|
|
|
acquired = await self.acquire_rpm_slot(
|
2025-12-10 20:52:44 +08:00
|
|
|
|
key_id,
|
2026-01-10 18:43:53 +08:00
|
|
|
|
key_rpm_limit,
|
2025-12-10 20:52:44 +08:00
|
|
|
|
is_cached_user,
|
|
|
|
|
|
cache_reservation_ratio,
|
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
if not acquired:
|
|
|
|
|
|
from src.core.exceptions import ConcurrencyLimitError
|
|
|
|
|
|
|
2026-02-15 16:32:23 +08:00
|
|
|
|
# Keep the client-facing message generic; do not leak internal IDs.
|
|
|
|
|
|
raise ConcurrencyLimitError("服务暂时繁忙,请稍后重试")
|
2025-12-10 20:52:44 +08:00
|
|
|
|
|
|
|
|
|
|
# 记录开始时间和状态
|
|
|
|
|
|
import time
|
|
|
|
|
|
|
|
|
|
|
|
slot_acquired_at = time.time()
|
|
|
|
|
|
exception_occurred = False
|
|
|
|
|
|
|
|
|
|
|
|
try:
|
|
|
|
|
|
yield # 执行请求
|
2026-01-10 18:43:53 +08:00
|
|
|
|
except Exception:
|
2025-12-10 20:52:44 +08:00
|
|
|
|
# 记录异常
|
|
|
|
|
|
exception_occurred = True
|
|
|
|
|
|
raise
|
|
|
|
|
|
finally:
|
|
|
|
|
|
# 计算槽位占用时长
|
|
|
|
|
|
slot_duration = time.time() - slot_acquired_at
|
|
|
|
|
|
|
|
|
|
|
|
# 记录 Prometheus 指标
|
|
|
|
|
|
try:
|
|
|
|
|
|
from src.core.metrics import (
|
|
|
|
|
|
concurrency_slot_duration_seconds,
|
|
|
|
|
|
concurrency_slot_release_total,
|
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
# 记录槽位占用时长分布
|
|
|
|
|
|
concurrency_slot_duration_seconds.labels(
|
|
|
|
|
|
exception=str(exception_occurred),
|
|
|
|
|
|
).observe(slot_duration)
|
|
|
|
|
|
|
|
|
|
|
|
# 记录槽位释放计数
|
|
|
|
|
|
concurrency_slot_release_total.labels(
|
2026-03-10 09:52:21 +08:00
|
|
|
|
exception=str(exception_occurred),
|
2025-12-10 20:52:44 +08:00
|
|
|
|
).inc()
|
|
|
|
|
|
|
|
|
|
|
|
# 告警:槽位占用时间过长(超过 60 秒)
|
|
|
|
|
|
if slot_duration > 60:
|
|
|
|
|
|
logger.warning(
|
2026-02-15 16:32:23 +08:00
|
|
|
|
"[WARN] 请求耗时过长: key_id={}..., duration={:.1f}s, exception={}",
|
|
|
|
|
|
key_id[:8] if key_id else "unknown",
|
|
|
|
|
|
slot_duration,
|
|
|
|
|
|
exception_occurred,
|
2025-12-10 20:52:44 +08:00
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
except Exception as metric_error:
|
|
|
|
|
|
# 指标记录失败不应影响业务逻辑
|
2026-02-15 16:32:23 +08:00
|
|
|
|
logger.debug("记录指标失败: {}", metric_error)
|
2025-12-10 20:52:44 +08:00
|
|
|
|
|
2026-01-10 18:43:53 +08:00
|
|
|
|
# 注意:RPM 计数不需要在请求结束后释放,它会在分钟窗口过期后自动重置
|
2025-12-10 20:52:44 +08:00
|
|
|
|
|
2026-01-10 18:43:53 +08:00
|
|
|
|
async def reset_key_rpm(self, key_id: str) -> None:
|
2025-12-10 20:52:44 +08:00
|
|
|
|
"""
|
2026-01-10 18:43:53 +08:00
|
|
|
|
重置 Key RPM 计数(管理功能,慎用)
|
2025-12-10 20:52:44 +08:00
|
|
|
|
|
|
|
|
|
|
Args:
|
2026-01-10 18:43:53 +08:00
|
|
|
|
key_id: ProviderAPIKey ID
|
2025-12-10 20:52:44 +08:00
|
|
|
|
"""
|
|
|
|
|
|
if self._redis is None:
|
|
|
|
|
|
async with self._memory_lock:
|
2026-01-10 18:43:53 +08:00
|
|
|
|
self._memory_key_rpm_counts.pop(key_id, None)
|
2026-02-15 16:32:23 +08:00
|
|
|
|
logger.info("[RESET] 重置 Key RPM 计数(内存): {}", key_id)
|
2025-12-10 20:52:44 +08:00
|
|
|
|
return
|
|
|
|
|
|
|
|
|
|
|
|
try:
|
2026-01-10 18:43:53 +08:00
|
|
|
|
deleted_count = await self._scan_and_delete(f"rpm:key:{key_id}:*")
|
2026-02-15 16:32:23 +08:00
|
|
|
|
logger.info("[RESET] 重置 Key RPM 计数: {}, 删除 {} 个键", key_id, deleted_count)
|
2026-01-10 18:43:53 +08:00
|
|
|
|
except Exception as e:
|
2026-02-15 16:32:23 +08:00
|
|
|
|
logger.error("重置 Key RPM 计数失败: {}", e)
|
2026-01-10 18:43:53 +08:00
|
|
|
|
|
|
|
|
|
|
async def reset_all_rpm(self) -> None:
|
|
|
|
|
|
"""重置所有 Key RPM 计数(管理功能,慎用)"""
|
|
|
|
|
|
if self._redis is None:
|
|
|
|
|
|
async with self._memory_lock:
|
|
|
|
|
|
count = len(self._memory_key_rpm_counts)
|
|
|
|
|
|
self._memory_key_rpm_counts.clear()
|
|
|
|
|
|
if count:
|
2026-02-15 16:32:23 +08:00
|
|
|
|
logger.info("[RESET] 重置所有 Key RPM 计数(内存): {} 个", count)
|
2026-01-10 18:43:53 +08:00
|
|
|
|
return
|
2025-12-10 20:52:44 +08:00
|
|
|
|
|
2026-01-10 18:43:53 +08:00
|
|
|
|
try:
|
|
|
|
|
|
deleted_count = await self._scan_and_delete("rpm:key:*")
|
|
|
|
|
|
if deleted_count:
|
2026-02-15 16:32:23 +08:00
|
|
|
|
logger.info("[RESET] 重置所有 Key RPM 计数: {} 个", deleted_count)
|
2025-12-10 20:52:44 +08:00
|
|
|
|
except Exception as e:
|
2026-02-15 16:32:23 +08:00
|
|
|
|
logger.error("重置所有 Key RPM 计数失败: {}", e)
|
2026-01-10 18:43:53 +08:00
|
|
|
|
|
|
|
|
|
|
async def _scan_and_delete(self, pattern: str, batch_size: int = 100) -> int:
|
|
|
|
|
|
"""使用 SCAN 遍历并分批删除匹配的键,避免阻塞 Redis"""
|
|
|
|
|
|
if self._redis is None:
|
|
|
|
|
|
return 0
|
|
|
|
|
|
|
|
|
|
|
|
deleted_count = 0
|
|
|
|
|
|
cursor = 0
|
|
|
|
|
|
while True:
|
|
|
|
|
|
cursor, keys = await self._redis.scan(cursor, match=pattern, count=batch_size)
|
|
|
|
|
|
if keys:
|
|
|
|
|
|
# 分批删除,每批最多 batch_size 个
|
|
|
|
|
|
for i in range(0, len(keys), batch_size):
|
|
|
|
|
|
batch = keys[i : i + batch_size]
|
|
|
|
|
|
await self._redis.delete(*batch)
|
|
|
|
|
|
deleted_count += len(batch)
|
|
|
|
|
|
if cursor == 0:
|
|
|
|
|
|
break
|
|
|
|
|
|
return deleted_count
|
2025-12-10 20:52:44 +08:00
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
# 全局单例
|
2026-01-30 03:10:21 +08:00
|
|
|
|
_concurrency_manager: ConcurrencyManager | None = None
|
2025-12-10 20:52:44 +08:00
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
async def get_concurrency_manager() -> ConcurrencyManager:
|
|
|
|
|
|
"""获取全局 ConcurrencyManager 实例"""
|
|
|
|
|
|
global _concurrency_manager
|
|
|
|
|
|
|
|
|
|
|
|
if _concurrency_manager is None:
|
|
|
|
|
|
_concurrency_manager = ConcurrencyManager()
|
|
|
|
|
|
await _concurrency_manager.initialize()
|
|
|
|
|
|
|
|
|
|
|
|
return _concurrency_manager
|