mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 09:20:22 +08:00
refactor: 移除 Python 后端源码,全面迁移至 Rust gateway 架构
- 删除全部 Python 源码 (src/) 及 Alembic 迁移脚本,归档至 _deprecated_py_src/ - 重构 Rust gateway ai_pipeline: 拆分 planner/finalize 模块,新增 contracts/adaptation 层 - 重组 handlers 模块为 admin/public/proxy/internal/shared 子模块结构 - 新增 executor 模块,引入 Rust 原生数据库迁移 (aether-data/migrations) - 简化 CI/Docker 构建流程,移除 base image 二级构建,统一为单一 app image - 移除 Python 相关基础设施文件 (entrypoint.sh, gunicorn_conf.py, Dockerfile.base)
This commit is contained in:
26
_deprecated_py_src/services/rate_limit/__init__.py
Normal file
26
_deprecated_py_src/services/rate_limit/__init__.py
Normal file
@@ -0,0 +1,26 @@
|
||||
"""
|
||||
限流服务模块
|
||||
|
||||
包含自适应 RPM 控制、并发管理、IP限流等功能。
|
||||
"""
|
||||
|
||||
from src.services.rate_limit.adaptive_rpm import AdaptiveConcurrencyManager # 向后兼容别名
|
||||
from src.services.rate_limit.adaptive_rpm import (
|
||||
AdaptiveRPMManager,
|
||||
get_adaptive_rpm_manager,
|
||||
)
|
||||
from src.services.rate_limit.concurrency_manager import ConcurrencyManager
|
||||
from src.services.rate_limit.detector import RateLimitDetector
|
||||
from src.services.rate_limit.ip_limiter import IPRateLimiter
|
||||
from src.services.rate_limit.user_rpm_limiter import UserRpmLimiter, get_user_rpm_limiter
|
||||
|
||||
__all__ = [
|
||||
"AdaptiveConcurrencyManager", # 向后兼容
|
||||
"AdaptiveRPMManager",
|
||||
"ConcurrencyManager",
|
||||
"IPRateLimiter",
|
||||
"RateLimitDetector",
|
||||
"UserRpmLimiter",
|
||||
"get_adaptive_rpm_manager",
|
||||
"get_user_rpm_limiter",
|
||||
]
|
||||
337
_deprecated_py_src/services/rate_limit/adaptive_reservation.py
Normal file
337
_deprecated_py_src/services/rate_limit/adaptive_reservation.py
Normal file
@@ -0,0 +1,337 @@
|
||||
"""
|
||||
自适应预留比例管理器
|
||||
|
||||
根据学习置信度和当前负载动态计算缓存用户预留比例,
|
||||
解决固定 30% 预留在学习初期和负载变化时的不适应问题。
|
||||
|
||||
核心思路:
|
||||
1. 探测阶段:使用低预留,让系统快速学习真实并发限制
|
||||
2. 稳定阶段:根据置信度和负载动态调整预留比例
|
||||
3. 置信度计算:综合考虑连续成功次数、429冷却时间、调整历史稳定性
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import statistics
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime, timezone
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from src.config.constants import AdaptiveReservationDefaults
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from src.models.database import ProviderAPIKey
|
||||
|
||||
|
||||
@dataclass
|
||||
class ReservationConfig:
|
||||
"""预留比例配置(使用统一常量作为默认值)"""
|
||||
|
||||
# 探测阶段配置
|
||||
probe_phase_requests: int = field(
|
||||
default_factory=lambda: AdaptiveReservationDefaults.PROBE_PHASE_REQUESTS
|
||||
)
|
||||
probe_reservation: float = field(
|
||||
default_factory=lambda: AdaptiveReservationDefaults.PROBE_RESERVATION
|
||||
)
|
||||
|
||||
# 稳定阶段配置
|
||||
stable_min_reservation: float = field(
|
||||
default_factory=lambda: AdaptiveReservationDefaults.STABLE_MIN_RESERVATION
|
||||
)
|
||||
stable_max_reservation: float = field(
|
||||
default_factory=lambda: AdaptiveReservationDefaults.STABLE_MAX_RESERVATION
|
||||
)
|
||||
|
||||
# 置信度计算参数
|
||||
success_count_for_full_confidence: int = field(
|
||||
default_factory=lambda: AdaptiveReservationDefaults.SUCCESS_COUNT_FOR_FULL_CONFIDENCE
|
||||
)
|
||||
cooldown_hours_for_full_confidence: int = field(
|
||||
default_factory=lambda: AdaptiveReservationDefaults.COOLDOWN_HOURS_FOR_FULL_CONFIDENCE
|
||||
)
|
||||
|
||||
# 负载阈值
|
||||
low_load_threshold: float = field(
|
||||
default_factory=lambda: AdaptiveReservationDefaults.LOW_LOAD_THRESHOLD
|
||||
)
|
||||
high_load_threshold: float = field(
|
||||
default_factory=lambda: AdaptiveReservationDefaults.HIGH_LOAD_THRESHOLD
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ReservationResult:
|
||||
"""预留比例计算结果"""
|
||||
|
||||
ratio: float # 最终预留比例
|
||||
phase: str # 当前阶段: "probe" | "stable"
|
||||
confidence: float # 置信度 (0-1)
|
||||
load_factor: float # 负载因子 (0-1)
|
||||
details: dict[str, Any] # 详细信息
|
||||
|
||||
|
||||
class AdaptiveReservationManager:
|
||||
"""
|
||||
自适应预留比例管理器
|
||||
|
||||
工作原理:
|
||||
1. 探测阶段(请求数 < 阈值):
|
||||
- 使用低预留比例(10%),不浪费资源
|
||||
- 让系统快速探测真实并发限制
|
||||
|
||||
2. 稳定阶段(请求数 >= 阈值):
|
||||
- 根据置信度和负载动态计算预留比例
|
||||
- 置信度高 + 负载高 = 高预留(保护缓存用户)
|
||||
- 置信度低或负载低 = 低预留(避免浪费)
|
||||
|
||||
置信度因素:
|
||||
- 连续成功次数:越多说明当前限制越准确
|
||||
- 429冷却时间:距离上次429越久越稳定
|
||||
- 调整历史稳定性:最近调整的方差越小越稳定
|
||||
"""
|
||||
|
||||
def __init__(self, config: ReservationConfig | None = None):
|
||||
self.config = config or ReservationConfig()
|
||||
|
||||
def calculate_reservation(
|
||||
self,
|
||||
key: ProviderAPIKey,
|
||||
current_usage: int = 0,
|
||||
effective_limit: int | None = None,
|
||||
) -> ReservationResult:
|
||||
"""
|
||||
计算当前应使用的预留比例
|
||||
|
||||
Args:
|
||||
key: ProviderAPIKey 对象
|
||||
current_usage: 当前使用量(RPM 计数)
|
||||
effective_limit: 有效限制(学习值或配置值)
|
||||
|
||||
Returns:
|
||||
ReservationResult 包含预留比例和详细信息
|
||||
"""
|
||||
# 计算总请求数(用于判断阶段)
|
||||
total_requests = self._get_total_requests(key)
|
||||
|
||||
# 计算负载率
|
||||
load_ratio = self._calculate_load_ratio(current_usage, effective_limit)
|
||||
|
||||
# 阶段1: 探测阶段
|
||||
if total_requests < self.config.probe_phase_requests:
|
||||
return ReservationResult(
|
||||
ratio=self.config.probe_reservation,
|
||||
phase="probe",
|
||||
confidence=0.0,
|
||||
load_factor=load_ratio,
|
||||
details={
|
||||
"total_requests": total_requests,
|
||||
"probe_threshold": self.config.probe_phase_requests,
|
||||
"reason": "探测阶段,使用低预留让系统学习真实限制",
|
||||
},
|
||||
)
|
||||
|
||||
# 阶段2: 稳定阶段
|
||||
confidence = self._calculate_confidence(key)
|
||||
ratio = self._calculate_stable_ratio(confidence, load_ratio)
|
||||
|
||||
return ReservationResult(
|
||||
ratio=ratio,
|
||||
phase="stable",
|
||||
confidence=confidence,
|
||||
load_factor=load_ratio,
|
||||
details={
|
||||
"total_requests": total_requests,
|
||||
"confidence_factors": self._get_confidence_breakdown(key),
|
||||
"reason": self._get_ratio_reason(confidence, load_ratio),
|
||||
},
|
||||
)
|
||||
|
||||
def _get_total_requests(self, key: ProviderAPIKey) -> int:
|
||||
"""获取总请求数(用于判断是否过了探测阶段)"""
|
||||
# 使用总请求计数作为基准
|
||||
request_count = key.request_count or 0
|
||||
|
||||
# 如果 request_count 为 0,使用 429 计数 + 成功计数作为近似值
|
||||
if request_count == 0:
|
||||
concurrent_429 = key.concurrent_429_count or 0
|
||||
rpm_429 = key.rpm_429_count or 0
|
||||
success_count = key.success_count or 0
|
||||
# 调整历史中的记录数也可以参考
|
||||
history_count = len(key.adjustment_history or []) * 10
|
||||
return concurrent_429 + rpm_429 + success_count + history_count
|
||||
|
||||
return request_count
|
||||
|
||||
def _calculate_load_ratio(self, current_usage: int, effective_limit: int | None) -> float:
|
||||
"""计算当前负载率"""
|
||||
if not effective_limit or effective_limit <= 0:
|
||||
return 0.0
|
||||
return min(current_usage / effective_limit, 1.0)
|
||||
|
||||
def _calculate_confidence(self, key: ProviderAPIKey) -> float:
|
||||
"""
|
||||
计算学习值的置信度 (0-1)
|
||||
|
||||
三个因素各占一定权重:
|
||||
- 成功率:40%(基于总成功数/总请求数)
|
||||
- 429冷却时间:30%
|
||||
- 调整历史稳定性:30%
|
||||
"""
|
||||
scores = self._get_confidence_breakdown(key)
|
||||
return min(
|
||||
scores["success_score"] + scores["cooldown_score"] + scores["stability_score"], 1.0
|
||||
)
|
||||
|
||||
def _get_confidence_breakdown(self, key: ProviderAPIKey) -> dict[str, float]:
|
||||
"""获取置信度各因素的详细分数"""
|
||||
# 因素1: 成功率(权重 40%)
|
||||
# 使用成功率而非连续成功次数,更准确反映 Key 的稳定性
|
||||
request_count = key.request_count or 0
|
||||
success_count = key.success_count or 0
|
||||
|
||||
if request_count >= self.config.success_count_for_full_confidence:
|
||||
# 请求数足够时,根据成功率计算
|
||||
success_rate = success_count / request_count if request_count > 0 else 0
|
||||
success_score = success_rate * 0.4
|
||||
elif request_count > 0:
|
||||
# 请求数不足时,按比例折算
|
||||
progress_ratio = request_count / self.config.success_count_for_full_confidence
|
||||
success_rate = success_count / request_count
|
||||
success_score = success_rate * progress_ratio * 0.4
|
||||
else:
|
||||
success_score = 0.0
|
||||
|
||||
# 因素2: 429冷却时间(权重 30%)
|
||||
if key.last_429_at:
|
||||
now = datetime.now(timezone.utc)
|
||||
# 确保 last_429_at 有时区信息
|
||||
last_429 = key.last_429_at
|
||||
if last_429.tzinfo is None:
|
||||
last_429 = last_429.replace(tzinfo=timezone.utc)
|
||||
hours_since_429 = (now - last_429).total_seconds() / 3600
|
||||
cooldown_ratio = min(
|
||||
hours_since_429 / self.config.cooldown_hours_for_full_confidence, 1.0
|
||||
)
|
||||
cooldown_score = cooldown_ratio * 0.3
|
||||
else:
|
||||
# 从未触发 429,给满分
|
||||
cooldown_score = 0.3
|
||||
|
||||
# 因素3: 调整历史稳定性(权重 30%)
|
||||
history = key.adjustment_history or []
|
||||
if len(history) >= 3:
|
||||
# 取最近的调整记录
|
||||
recent = history[-5:] if len(history) >= 5 else history
|
||||
limits = [h.get("new_limit", 0) for h in recent if h.get("new_limit")]
|
||||
|
||||
if len(limits) >= 2:
|
||||
try:
|
||||
variance = statistics.variance(limits)
|
||||
# 方差越小越稳定,方差为10时分数接近0
|
||||
stability_ratio = max(0, 1 - variance / 10)
|
||||
stability_score = stability_ratio * 0.3
|
||||
except statistics.StatisticsError:
|
||||
stability_score = 0.15
|
||||
else:
|
||||
stability_score = 0.15
|
||||
else:
|
||||
# 历史数据不足,给一半分
|
||||
stability_score = 0.15
|
||||
|
||||
# 计算成功率用于返回
|
||||
success_rate_pct = (success_count / request_count * 100) if request_count > 0 else None
|
||||
|
||||
return {
|
||||
"success_score": round(success_score, 3),
|
||||
"cooldown_score": round(cooldown_score, 3),
|
||||
"stability_score": round(stability_score, 3),
|
||||
"request_count": request_count,
|
||||
"success_count": success_count,
|
||||
"success_rate": round(success_rate_pct, 1) if success_rate_pct is not None else None,
|
||||
"hours_since_429": (
|
||||
round(
|
||||
(
|
||||
datetime.now(timezone.utc) - key.last_429_at.replace(tzinfo=timezone.utc)
|
||||
).total_seconds()
|
||||
/ 3600,
|
||||
1,
|
||||
)
|
||||
if key.last_429_at
|
||||
else None
|
||||
),
|
||||
"history_count": len(history),
|
||||
}
|
||||
|
||||
def _calculate_stable_ratio(self, confidence: float, load_ratio: float) -> float:
|
||||
"""
|
||||
计算稳定阶段的预留比例
|
||||
|
||||
策略:
|
||||
- 低负载(<50%):使用最小预留,槽位充足无需过多预留
|
||||
- 中等负载(50-80%):根据置信度线性增加预留
|
||||
- 高负载(>80%):根据置信度使用较高预留保护缓存用户
|
||||
"""
|
||||
min_r = self.config.stable_min_reservation
|
||||
max_r = self.config.stable_max_reservation
|
||||
|
||||
if load_ratio < self.config.low_load_threshold:
|
||||
# 低负载:使用最小预留
|
||||
return min_r
|
||||
|
||||
if load_ratio < self.config.high_load_threshold:
|
||||
# 中等负载:根据置信度和负载线性插值
|
||||
# 负载越高、置信度越高,预留越多
|
||||
load_factor = (load_ratio - self.config.low_load_threshold) / (
|
||||
self.config.high_load_threshold - self.config.low_load_threshold
|
||||
)
|
||||
return min_r + confidence * load_factor * (max_r - min_r)
|
||||
|
||||
# 高负载:根据置信度决定预留比例
|
||||
# 置信度高 → 接近最大预留
|
||||
# 置信度低 → 保守预留(避免基于不准确的学习值过度预留)
|
||||
return min_r + confidence * (max_r - min_r)
|
||||
|
||||
def _get_ratio_reason(self, confidence: float, load_ratio: float) -> str:
|
||||
"""生成预留比例的解释"""
|
||||
if load_ratio < self.config.low_load_threshold:
|
||||
return f"低负载({load_ratio:.0%}),使用最小预留"
|
||||
|
||||
if confidence < 0.3:
|
||||
return f"置信度低({confidence:.0%}),保守预留避免浪费"
|
||||
|
||||
if confidence > 0.7 and load_ratio > self.config.high_load_threshold:
|
||||
return f"高置信度({confidence:.0%})+高负载({load_ratio:.0%}),使用较高预留保护缓存用户"
|
||||
|
||||
return f"置信度{confidence:.0%},负载{load_ratio:.0%},动态计算预留"
|
||||
|
||||
def get_stats(self) -> dict[str, Any]:
|
||||
"""获取管理器统计信息"""
|
||||
return {
|
||||
"config": {
|
||||
"probe_phase_requests": self.config.probe_phase_requests,
|
||||
"probe_reservation": self.config.probe_reservation,
|
||||
"stable_min_reservation": self.config.stable_min_reservation,
|
||||
"stable_max_reservation": self.config.stable_max_reservation,
|
||||
"low_load_threshold": self.config.low_load_threshold,
|
||||
"high_load_threshold": self.config.high_load_threshold,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
# 全局单例
|
||||
_reservation_manager: AdaptiveReservationManager | None = None
|
||||
|
||||
|
||||
def get_adaptive_reservation_manager() -> AdaptiveReservationManager:
|
||||
"""获取全局自适应预留管理器单例"""
|
||||
global _reservation_manager
|
||||
if _reservation_manager is None:
|
||||
_reservation_manager = AdaptiveReservationManager()
|
||||
return _reservation_manager
|
||||
|
||||
|
||||
def reset_adaptive_reservation_manager() -> None:
|
||||
"""重置全局单例(用于测试)"""
|
||||
global _reservation_manager
|
||||
_reservation_manager = None
|
||||
755
_deprecated_py_src/services/rate_limit/adaptive_rpm.py
Normal file
755
_deprecated_py_src/services/rate_limit/adaptive_rpm.py
Normal file
@@ -0,0 +1,755 @@
|
||||
"""
|
||||
自适应 RPM 调整器 - 基于置信度衰减的 RPM 限制学习
|
||||
|
||||
核心算法:多次观察确认 + 置信度衰减
|
||||
- 收到 429 时记录观察(本地 RPM + 上游 header 限制值)
|
||||
- 多次一致的观察才确认限制(header 需 2 次,无 header 需 3 次)
|
||||
- confidence 随时间自然衰减,限制永远不会固化
|
||||
- confidence 低于阈值时停止本地 RPM 限制执行,让上游 429 透传
|
||||
|
||||
设计原则:
|
||||
1. 限制永远不固化 -- confidence 需要持续的 429 观察来维持
|
||||
2. 即使有上游 header 也要多次确认
|
||||
3. 学习期间 429 直接透传给客户端
|
||||
4. 优先使用上游 header 声明的限制值,而非本地 RPM 计数
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from statistics import median
|
||||
from typing import Any, cast
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.config.constants import RPMDefaults
|
||||
from src.core.batch_committer import get_batch_committer
|
||||
from src.core.logger import logger
|
||||
from src.models.database import ProviderAPIKey
|
||||
from src.services.rate_limit.detector import RateLimitInfo, RateLimitType
|
||||
|
||||
|
||||
class AdaptiveStrategy:
|
||||
"""自适应策略类型"""
|
||||
|
||||
AIMD = "aimd" # 加性增-乘性减 (Additive Increase Multiplicative Decrease)
|
||||
CONSERVATIVE = "conservative" # 保守策略(只减不增)
|
||||
AGGRESSIVE = "aggressive" # 激进策略(快速探测)
|
||||
|
||||
|
||||
class AdaptiveRPMManager:
|
||||
"""
|
||||
自适应 RPM 管理器
|
||||
|
||||
核心算法:多次观察确认 + 置信度衰减
|
||||
- 收到 429 时记录观察(本地 RPM 计数 + 上游 header 限制值)
|
||||
- 有 header 的观察需 MIN_HEADER_CONFIRMATIONS 次一致才确认
|
||||
- 无 header 的观察需 MIN_CONSISTENT_OBSERVATIONS 次一致才确认
|
||||
- confidence 随时间自然衰减(CONFIDENCE_DECAY_PER_MINUTE)
|
||||
- confidence < ENFORCEMENT_CONFIDENCE_THRESHOLD 时停止本地限制执行
|
||||
|
||||
扩容条件(满足任一即可):
|
||||
1. 利用率扩容:窗口内高利用率比例 >= 60%,且当前限制 < 边界
|
||||
2. 探测性扩容:距上次 429 超过 30 分钟,可以尝试突破边界
|
||||
"""
|
||||
|
||||
# 默认配置
|
||||
DEFAULT_INITIAL_LIMIT = RPMDefaults.INITIAL_LIMIT
|
||||
MIN_RPM_LIMIT = RPMDefaults.MIN_RPM_LIMIT
|
||||
MAX_RPM_LIMIT = RPMDefaults.MAX_RPM_LIMIT
|
||||
|
||||
# AIMD 参数
|
||||
INCREASE_STEP = RPMDefaults.INCREASE_STEP
|
||||
|
||||
# 滑动窗口参数
|
||||
UTILIZATION_WINDOW_SIZE = RPMDefaults.UTILIZATION_WINDOW_SIZE
|
||||
UTILIZATION_WINDOW_SECONDS = RPMDefaults.UTILIZATION_WINDOW_SECONDS
|
||||
UTILIZATION_THRESHOLD = RPMDefaults.UTILIZATION_THRESHOLD
|
||||
HIGH_UTILIZATION_RATIO = RPMDefaults.HIGH_UTILIZATION_RATIO
|
||||
MIN_SAMPLES_FOR_DECISION = RPMDefaults.MIN_SAMPLES_FOR_DECISION
|
||||
|
||||
# 探测性扩容参数
|
||||
PROBE_INCREASE_INTERVAL_MINUTES = RPMDefaults.PROBE_INCREASE_INTERVAL_MINUTES
|
||||
PROBE_INCREASE_MIN_REQUESTS = RPMDefaults.PROBE_INCREASE_MIN_REQUESTS
|
||||
|
||||
# 记录历史数量
|
||||
MAX_HISTORY_RECORDS = 20
|
||||
|
||||
# 置信度学习参数
|
||||
MIN_CONSISTENT_OBSERVATIONS = RPMDefaults.MIN_CONSISTENT_OBSERVATIONS
|
||||
MIN_HEADER_CONFIRMATIONS = RPMDefaults.MIN_HEADER_CONFIRMATIONS
|
||||
OBSERVATION_CONSISTENCY_THRESHOLD = RPMDefaults.OBSERVATION_CONSISTENCY_THRESHOLD
|
||||
HEADER_LIMIT_SAFETY_MARGIN = RPMDefaults.HEADER_LIMIT_SAFETY_MARGIN
|
||||
OBSERVATION_LIMIT_SAFETY_MARGIN = RPMDefaults.OBSERVATION_LIMIT_SAFETY_MARGIN
|
||||
ENFORCEMENT_CONFIDENCE_THRESHOLD = RPMDefaults.ENFORCEMENT_CONFIDENCE_THRESHOLD
|
||||
CONFIDENCE_DECAY_PER_MINUTE = RPMDefaults.CONFIDENCE_DECAY_PER_MINUTE
|
||||
|
||||
def __init__(self, strategy: str = AdaptiveStrategy.AIMD):
|
||||
"""
|
||||
初始化自适应 RPM 管理器
|
||||
|
||||
Args:
|
||||
strategy: 调整策略
|
||||
"""
|
||||
self.strategy = strategy
|
||||
|
||||
@staticmethod
|
||||
def _persist_metadata_update(db: Session) -> None:
|
||||
"""
|
||||
延后持久化自适应学习元数据,避免在请求热路径上同步 flush。
|
||||
|
||||
这些更新只修改既有 ProviderAPIKey 行上的统计/学习字段,不依赖数据库生成值。
|
||||
请求级事务会在中间件统一 commit,后台会话则由 BatchCommitter 负责后续提交。
|
||||
"""
|
||||
get_batch_committer().mark_dirty(db)
|
||||
|
||||
# ==================== 429 处理 ====================
|
||||
|
||||
def handle_429_error(
|
||||
self,
|
||||
db: Session,
|
||||
key: ProviderAPIKey,
|
||||
rate_limit_info: RateLimitInfo,
|
||||
current_rpm: int | None = None,
|
||||
) -> int | None:
|
||||
"""
|
||||
处理 429 错误,记录观察并基于一致性评估是否设置限制
|
||||
|
||||
不再单次 429 就设限,而是:
|
||||
1. 记录 429 观察(本地 RPM + 上游 header 限制值)
|
||||
2. 评估历史观察的一致性
|
||||
3. 一致性达标时设置 learned_rpm_limit 并赋予 confidence
|
||||
4. 一致性不够时保持学习期(429 透传给客户端)
|
||||
|
||||
Returns:
|
||||
调整后的 RPM 限制,或 None(学习期间)
|
||||
"""
|
||||
is_adaptive = key.rpm_limit is None
|
||||
|
||||
if not is_adaptive:
|
||||
logger.debug(f"Key {key.id} 设置了固定 RPM 限制 ({key.rpm_limit}),跳过自适应调整")
|
||||
return int(key.rpm_limit) # type: ignore[arg-type]
|
||||
|
||||
# 更新 429 统计
|
||||
key.last_429_at = datetime.now(timezone.utc) # type: ignore[assignment]
|
||||
key.last_429_type = rate_limit_info.limit_type # type: ignore[assignment]
|
||||
|
||||
# 清空利用率采样窗口
|
||||
key.utilization_samples = [] # type: ignore[assignment]
|
||||
|
||||
if rate_limit_info.limit_type == RateLimitType.RPM:
|
||||
key.rpm_429_count = int(key.rpm_429_count or 0) + 1 # type: ignore[assignment]
|
||||
|
||||
upstream_limit = rate_limit_info.limit_value
|
||||
|
||||
# 记录 429 观察
|
||||
self._record_429_observation(key, current_rpm, upstream_limit)
|
||||
|
||||
# 评估观察一致性,决定是否设置/更新限制
|
||||
evaluated_limit, confidence = self._evaluate_observations(key)
|
||||
|
||||
old_limit = key.learned_rpm_limit
|
||||
|
||||
if evaluated_limit is not None and confidence >= self.ENFORCEMENT_CONFIDENCE_THRESHOLD:
|
||||
# 一致性达标,设置限制
|
||||
self._record_adjustment(
|
||||
key,
|
||||
old_limit=old_limit or 0,
|
||||
new_limit=evaluated_limit,
|
||||
reason="rpm_429",
|
||||
current_rpm=current_rpm,
|
||||
upstream_limit=upstream_limit,
|
||||
confidence=round(confidence, 3),
|
||||
learning_source="header" if upstream_limit else "observation",
|
||||
)
|
||||
key.learned_rpm_limit = evaluated_limit # type: ignore[assignment]
|
||||
|
||||
# 更新 last_rpm_peak:优先使用 upstream header
|
||||
if upstream_limit and upstream_limit > 0:
|
||||
key.last_rpm_peak = upstream_limit # type: ignore[assignment]
|
||||
elif current_rpm and current_rpm > 0:
|
||||
key.last_rpm_peak = current_rpm # type: ignore[assignment]
|
||||
|
||||
logger.warning(
|
||||
f"[RPM] 限制已确认: Key {key.id[:8]}... | "
|
||||
f"当前 RPM: {current_rpm} | "
|
||||
f"上游 header: {upstream_limit} | "
|
||||
f"调整: {old_limit} -> {evaluated_limit} | "
|
||||
f"confidence: {confidence:.2f}"
|
||||
)
|
||||
else:
|
||||
# 一致性不够,保持学习期
|
||||
logger.info(
|
||||
f"[RPM] 学习中: Key {key.id[:8]}... | "
|
||||
f"当前 RPM: {current_rpm} | "
|
||||
f"上游 header: {upstream_limit} | "
|
||||
f"观察已记录,暂不设限"
|
||||
)
|
||||
|
||||
elif rate_limit_info.limit_type == RateLimitType.CONCURRENT:
|
||||
key.concurrent_429_count = int(key.concurrent_429_count or 0) + 1 # type: ignore[assignment]
|
||||
logger.info(
|
||||
f"[CONCURRENT] 并发限制触发: Key {key.id[:8]}... | "
|
||||
f"不调整 RPM 限制(这是并发问题,非 RPM 问题)"
|
||||
)
|
||||
|
||||
else:
|
||||
# 未知类型:保守处理(仅在已有学习值时减少)
|
||||
old_limit = key.learned_rpm_limit
|
||||
if old_limit is not None:
|
||||
logger.warning(
|
||||
f"[UNKNOWN] 未知429类型: Key {key.id[:8]}... | "
|
||||
f"当前 RPM: {current_rpm} | "
|
||||
f"保守减少 RPM: {old_limit} -> {max(int(old_limit * 0.95), self.MIN_RPM_LIMIT)}"
|
||||
)
|
||||
else:
|
||||
logger.info(
|
||||
f"[UNKNOWN] 未知429类型: Key {key.id[:8]}... | "
|
||||
f"当前 RPM: {current_rpm} | "
|
||||
f"无学习值,跳过调整"
|
||||
)
|
||||
if old_limit is not None:
|
||||
new_limit = max(int(old_limit * 0.95), self.MIN_RPM_LIMIT)
|
||||
self._record_adjustment(
|
||||
key,
|
||||
old_limit=int(old_limit),
|
||||
new_limit=new_limit,
|
||||
reason="unknown_429",
|
||||
current_rpm=current_rpm,
|
||||
)
|
||||
key.learned_rpm_limit = new_limit # type: ignore[assignment]
|
||||
|
||||
self._persist_metadata_update(db)
|
||||
|
||||
return key.learned_rpm_limit if key.learned_rpm_limit is not None else None
|
||||
|
||||
# ==================== 观察记录与评估 ====================
|
||||
|
||||
def _record_429_observation(
|
||||
self,
|
||||
key: ProviderAPIKey,
|
||||
current_rpm: int | None,
|
||||
upstream_limit: int | None,
|
||||
) -> None:
|
||||
"""在 adjustment_history 中记录一次 429 观察"""
|
||||
history: list[dict[str, Any]] = list(key.adjustment_history or [])
|
||||
|
||||
observation: dict[str, Any] = {
|
||||
"type": "429_observation",
|
||||
"timestamp": datetime.now(timezone.utc).isoformat(),
|
||||
"current_rpm": current_rpm,
|
||||
"upstream_limit": upstream_limit,
|
||||
}
|
||||
history.append(observation)
|
||||
|
||||
key.adjustment_history = self._trim_history(history) # type: ignore[assignment]
|
||||
|
||||
def _evaluate_observations(self, key: ProviderAPIKey) -> tuple[int | None, float]:
|
||||
"""
|
||||
评估历史 429 观察的一致性,决定是否确认限制
|
||||
|
||||
优先使用有 header 的观察(upstream_limit),其次使用纯本地观察(current_rpm)。
|
||||
|
||||
Returns:
|
||||
(limit, confidence):
|
||||
- limit: 新确认的限制值,或 None(一致性不够,不设/不更新限制)
|
||||
- confidence: 置信度分数 0.0~1.0
|
||||
"""
|
||||
history: list[dict[str, Any]] = list(key.adjustment_history or [])
|
||||
observations = [h for h in history if h.get("type") == "429_observation"]
|
||||
|
||||
if not observations:
|
||||
return None, 0.0
|
||||
|
||||
# 优先评估有 header 的观察
|
||||
header_obs = [
|
||||
o
|
||||
for o in observations
|
||||
if o.get("upstream_limit") is not None and o["upstream_limit"] > 0
|
||||
]
|
||||
if len(header_obs) >= self.MIN_HEADER_CONFIRMATIONS:
|
||||
recent = header_obs[-self.MIN_HEADER_CONFIRMATIONS * 2 :]
|
||||
values = [o["upstream_limit"] for o in recent]
|
||||
last_n = values[-self.MIN_HEADER_CONFIRMATIONS :]
|
||||
|
||||
if self._check_consistency(last_n):
|
||||
limit_val = int(median(last_n) * self.HEADER_LIMIT_SAFETY_MARGIN)
|
||||
limit_val = max(limit_val, self.MIN_RPM_LIMIT)
|
||||
limit_val = min(limit_val, self.MAX_RPM_LIMIT)
|
||||
return limit_val, 0.8
|
||||
|
||||
# 其次评估纯本地观察(无 header)
|
||||
local_obs = [
|
||||
o for o in observations if o.get("current_rpm") is not None and o["current_rpm"] > 0
|
||||
]
|
||||
if len(local_obs) >= self.MIN_CONSISTENT_OBSERVATIONS:
|
||||
recent = local_obs[-self.MIN_CONSISTENT_OBSERVATIONS * 2 :]
|
||||
values = [o["current_rpm"] for o in recent]
|
||||
last_n = values[-self.MIN_CONSISTENT_OBSERVATIONS :]
|
||||
|
||||
if self._check_consistency(last_n):
|
||||
limit_val = int(median(last_n) * self.OBSERVATION_LIMIT_SAFETY_MARGIN)
|
||||
limit_val = max(limit_val, self.MIN_RPM_LIMIT)
|
||||
limit_val = min(limit_val, self.MAX_RPM_LIMIT)
|
||||
return limit_val, 0.6
|
||||
|
||||
# 一致性不够,不设/不更新限制(已有的 learned_rpm_limit 不在此处处理)
|
||||
return None, 0.0
|
||||
|
||||
def _check_consistency(self, values: list[int]) -> bool:
|
||||
"""检查一组数值是否在 OBSERVATION_CONSISTENCY_THRESHOLD 偏差范围内"""
|
||||
if not values:
|
||||
return False
|
||||
med = median(values)
|
||||
if med <= 0:
|
||||
return False
|
||||
return all(abs(v - med) / med <= self.OBSERVATION_CONSISTENCY_THRESHOLD for v in values)
|
||||
|
||||
# ==================== 置信度计算 ====================
|
||||
|
||||
def get_confidence(self, key: ProviderAPIKey) -> float:
|
||||
"""
|
||||
计算当前 confidence 分数(0.0~1.0),包含时间衰减
|
||||
|
||||
confidence 基于最后一次 429 评估的基础值,随时间自然衰减。
|
||||
确保限制永远不会固化:长时间没有新 429 观察 → confidence 降至 0。
|
||||
|
||||
Returns:
|
||||
当前 confidence(0.0~1.0)
|
||||
"""
|
||||
if key.learned_rpm_limit is None:
|
||||
return 0.0
|
||||
|
||||
# 从历史中获取基础 confidence
|
||||
base_confidence = self._get_base_confidence(key)
|
||||
|
||||
if base_confidence <= 0:
|
||||
return 0.0
|
||||
|
||||
# 时间衰减
|
||||
if key.last_429_at is not None:
|
||||
last_429_at = cast(datetime, key.last_429_at)
|
||||
minutes_since = max(
|
||||
0.0, (datetime.now(timezone.utc) - last_429_at).total_seconds() / 60.0
|
||||
)
|
||||
time_decay = minutes_since * self.CONFIDENCE_DECAY_PER_MINUTE
|
||||
else:
|
||||
time_decay = 1.0 # 没有 429 记录,直接衰减到 0
|
||||
|
||||
final = max(0.0, base_confidence - time_decay)
|
||||
return min(final, 1.0)
|
||||
|
||||
def is_enforcement_active(self, key: ProviderAPIKey) -> bool:
|
||||
"""confidence 是否达到执行阈值,达标才执行本地 RPM 限制"""
|
||||
return self.get_confidence(key) >= self.ENFORCEMENT_CONFIDENCE_THRESHOLD
|
||||
|
||||
def get_effective_limit(self, key: ProviderAPIKey) -> int | None:
|
||||
"""
|
||||
获取 key 当前有效的 RPM 限制(统一入口)
|
||||
|
||||
- rpm_limit=NULL(自适应):learned_rpm_limit + confidence 达标才返回
|
||||
- rpm_limit=数字(固定):直接返回固定值
|
||||
- 其余情况返回 None(不限制)
|
||||
"""
|
||||
if key.rpm_limit is not None:
|
||||
return int(key.rpm_limit)
|
||||
# 自适应模式
|
||||
if key.learned_rpm_limit is not None and self.is_enforcement_active(key):
|
||||
return int(key.learned_rpm_limit)
|
||||
return None
|
||||
|
||||
def _get_base_confidence(self, key: ProviderAPIKey) -> float:
|
||||
"""从最近的 adjustment 记录中获取基础 confidence"""
|
||||
history: list[dict[str, Any]] = list(key.adjustment_history or [])
|
||||
|
||||
# 从最新的 adjustment 记录中查找 confidence
|
||||
for record in reversed(history):
|
||||
if record.get("type") != "429_observation" and "confidence" in record:
|
||||
return float(record["confidence"])
|
||||
|
||||
# 没有 confidence 记录(旧数据迁移):尝试从观察中重新评估
|
||||
evaluated_limit, confidence = self._evaluate_observations(key)
|
||||
if confidence > 0:
|
||||
return confidence
|
||||
|
||||
# 有 learned_rpm_limit 但无法从观察中确认(旧数据),给予低基线置信度
|
||||
if key.learned_rpm_limit is not None:
|
||||
return 0.3
|
||||
|
||||
return 0.0
|
||||
|
||||
# ==================== 成功处理 ====================
|
||||
|
||||
def handle_success(
|
||||
self,
|
||||
db: Session,
|
||||
key: ProviderAPIKey,
|
||||
current_rpm: int,
|
||||
) -> int | None:
|
||||
"""
|
||||
处理成功请求,基于滑动窗口利用率考虑增加 RPM 限制
|
||||
|
||||
Returns:
|
||||
调整后的 RPM 限制(如果有调整),否则返回 None
|
||||
"""
|
||||
is_adaptive = key.rpm_limit is None
|
||||
|
||||
if not is_adaptive:
|
||||
return None
|
||||
|
||||
# 未碰壁学习前,不主动设置限制
|
||||
if key.learned_rpm_limit is None:
|
||||
return None
|
||||
|
||||
# confidence 太低时不做扩容逻辑(系统已在自由运行模式)
|
||||
confidence = self.get_confidence(key)
|
||||
if confidence < self.ENFORCEMENT_CONFIDENCE_THRESHOLD:
|
||||
return None
|
||||
|
||||
current_limit = int(key.learned_rpm_limit)
|
||||
|
||||
# 获取已知边界(上次触发 429 时的 RPM)
|
||||
known_boundary = key.last_rpm_peak
|
||||
|
||||
# 计算当前利用率
|
||||
utilization = float(current_rpm / current_limit) if current_limit > 0 else 0.0
|
||||
|
||||
now = datetime.now(timezone.utc)
|
||||
now_ts = now.timestamp()
|
||||
|
||||
# 更新滑动窗口
|
||||
samples = self._update_utilization_window(key, now_ts, utilization)
|
||||
|
||||
# 检查是否满足扩容条件
|
||||
increase_reason = self._check_increase_conditions(key, samples, now, known_boundary)
|
||||
|
||||
if increase_reason and current_limit < self.MAX_RPM_LIMIT:
|
||||
old_limit = current_limit
|
||||
is_probe = increase_reason == "probe_increase"
|
||||
new_limit = self._increase_limit(current_limit, known_boundary, is_probe)
|
||||
|
||||
# 如果没有实际增长(已达边界),跳过
|
||||
if new_limit <= old_limit:
|
||||
return None
|
||||
|
||||
# 计算窗口统计用于日志
|
||||
avg_util = sum(s["util"] for s in samples) / len(samples) if samples else 0
|
||||
high_util_count = sum(1 for s in samples if s["util"] >= self.UTILIZATION_THRESHOLD)
|
||||
high_util_ratio = high_util_count / len(samples) if samples else 0
|
||||
|
||||
boundary_info = f"边界: {known_boundary}" if known_boundary else "无边界"
|
||||
logger.info(
|
||||
f"[INCREASE] {increase_reason}: Key {key.id[:8]}... | "
|
||||
f"窗口采样: {len(samples)} | "
|
||||
f"平均利用率: {avg_util:.1%} | "
|
||||
f"高利用率比例: {high_util_ratio:.1%} | "
|
||||
f"{boundary_info} | "
|
||||
f"调整: {old_limit} -> {new_limit}"
|
||||
)
|
||||
|
||||
# 记录调整历史
|
||||
self._record_adjustment(
|
||||
key,
|
||||
old_limit=old_limit,
|
||||
new_limit=new_limit,
|
||||
reason=increase_reason,
|
||||
avg_utilization=round(avg_util, 2),
|
||||
high_util_ratio=round(high_util_ratio, 2),
|
||||
sample_count=len(samples),
|
||||
current_rpm=current_rpm,
|
||||
known_boundary=known_boundary,
|
||||
confidence=round(confidence, 3),
|
||||
)
|
||||
|
||||
# 更新限制
|
||||
key.learned_rpm_limit = new_limit # type: ignore[assignment]
|
||||
|
||||
# 如果是探测性扩容,更新探测时间
|
||||
if is_probe:
|
||||
key.last_probe_increase_at = now # type: ignore[assignment]
|
||||
|
||||
# 扩容后清空采样窗口,重新开始收集
|
||||
key.utilization_samples = [] # type: ignore[assignment]
|
||||
|
||||
self._persist_metadata_update(db)
|
||||
|
||||
return new_limit
|
||||
|
||||
# 定期持久化采样数据(每5个采样保存一次)
|
||||
if len(samples) % 5 == 0:
|
||||
self._persist_metadata_update(db)
|
||||
|
||||
return None
|
||||
|
||||
# ==================== 滑动窗口 ====================
|
||||
|
||||
def _update_utilization_window(
|
||||
self, key: ProviderAPIKey, now_ts: float, utilization: float
|
||||
) -> list[dict[str, Any]]:
|
||||
"""更新利用率滑动窗口"""
|
||||
samples: list[dict[str, Any]] = list(key.utilization_samples or [])
|
||||
|
||||
samples.append({"ts": now_ts, "util": round(utilization, 3)})
|
||||
|
||||
cutoff_ts = now_ts - self.UTILIZATION_WINDOW_SECONDS
|
||||
samples = [s for s in samples if s["ts"] > cutoff_ts]
|
||||
|
||||
if len(samples) > self.UTILIZATION_WINDOW_SIZE:
|
||||
samples = samples[-self.UTILIZATION_WINDOW_SIZE :]
|
||||
|
||||
key.utilization_samples = samples # type: ignore[assignment]
|
||||
return samples
|
||||
|
||||
# ==================== 扩容条件 ====================
|
||||
|
||||
def _check_increase_conditions(
|
||||
self,
|
||||
key: ProviderAPIKey,
|
||||
samples: list[dict[str, Any]],
|
||||
now: datetime,
|
||||
known_boundary: int | None = None,
|
||||
) -> str | None:
|
||||
"""检查是否满足扩容条件"""
|
||||
if self._is_in_cooldown(key):
|
||||
return None
|
||||
|
||||
current_limit = int(key.learned_rpm_limit or self.DEFAULT_INITIAL_LIMIT)
|
||||
|
||||
# 条件1:滑动窗口扩容(不超过边界)
|
||||
if len(samples) >= self.MIN_SAMPLES_FOR_DECISION:
|
||||
high_util_count = sum(1 for s in samples if s["util"] >= self.UTILIZATION_THRESHOLD)
|
||||
high_util_ratio = high_util_count / len(samples)
|
||||
|
||||
if high_util_ratio >= self.HIGH_UTILIZATION_RATIO:
|
||||
if known_boundary:
|
||||
if current_limit < known_boundary:
|
||||
return "high_utilization"
|
||||
else:
|
||||
return "high_utilization"
|
||||
|
||||
# 条件2:探测性扩容
|
||||
if self._should_probe_increase(key, samples, now):
|
||||
return "probe_increase"
|
||||
|
||||
return None
|
||||
|
||||
def _should_probe_increase(
|
||||
self, key: ProviderAPIKey, samples: list[dict[str, Any]], now: datetime
|
||||
) -> bool:
|
||||
"""检查是否应该进行探测性扩容"""
|
||||
probe_interval_seconds = self.PROBE_INCREASE_INTERVAL_MINUTES * 60
|
||||
|
||||
if key.last_429_at:
|
||||
last_429_at = cast(datetime, key.last_429_at)
|
||||
time_since_429 = (now - last_429_at).total_seconds()
|
||||
if time_since_429 < probe_interval_seconds:
|
||||
return False
|
||||
|
||||
if key.last_probe_increase_at:
|
||||
last_probe = cast(datetime, key.last_probe_increase_at)
|
||||
time_since_probe = (now - last_probe).total_seconds()
|
||||
if time_since_probe < probe_interval_seconds:
|
||||
return False
|
||||
|
||||
if len(samples) < self.PROBE_INCREASE_MIN_REQUESTS:
|
||||
return False
|
||||
|
||||
avg_util = sum(s["util"] for s in samples) / len(samples)
|
||||
if avg_util < 0.3:
|
||||
return False
|
||||
|
||||
return True
|
||||
|
||||
def _is_in_cooldown(self, key: ProviderAPIKey) -> bool:
|
||||
"""检查是否在 429 错误后的冷却期内"""
|
||||
if key.last_429_at is None:
|
||||
return False
|
||||
|
||||
last_429_at = cast(datetime, key.last_429_at)
|
||||
time_since_429 = (datetime.now(timezone.utc) - last_429_at).total_seconds()
|
||||
cooldown_seconds = RPMDefaults.COOLDOWN_AFTER_429_MINUTES * 60
|
||||
|
||||
return bool(time_since_429 < cooldown_seconds)
|
||||
|
||||
# ==================== 限制调整 ====================
|
||||
|
||||
def _increase_limit(
|
||||
self,
|
||||
current_limit: int,
|
||||
known_boundary: int | None = None,
|
||||
is_probe: bool = False,
|
||||
) -> int:
|
||||
"""增加 RPM 限制(考虑边界保护)"""
|
||||
if is_probe:
|
||||
new_limit = current_limit + 1
|
||||
else:
|
||||
new_limit = current_limit + self.INCREASE_STEP
|
||||
if known_boundary:
|
||||
if new_limit > known_boundary:
|
||||
new_limit = known_boundary
|
||||
|
||||
new_limit = min(new_limit, self.MAX_RPM_LIMIT)
|
||||
|
||||
if new_limit <= current_limit:
|
||||
return current_limit
|
||||
|
||||
return new_limit
|
||||
|
||||
# ==================== 历史记录 ====================
|
||||
|
||||
def _record_adjustment(
|
||||
self,
|
||||
key: ProviderAPIKey,
|
||||
old_limit: int,
|
||||
new_limit: int,
|
||||
reason: str,
|
||||
**extra_data: Any,
|
||||
) -> None:
|
||||
"""记录 RPM 调整历史"""
|
||||
history: list[dict[str, Any]] = list(key.adjustment_history or [])
|
||||
|
||||
record = {
|
||||
"timestamp": datetime.now(timezone.utc).isoformat(),
|
||||
"old_limit": old_limit,
|
||||
"new_limit": new_limit,
|
||||
"reason": reason,
|
||||
**extra_data,
|
||||
}
|
||||
history.append(record)
|
||||
|
||||
key.adjustment_history = self._trim_history(history) # type: ignore[assignment]
|
||||
|
||||
def _trim_history(self, history: list[dict[str, Any]]) -> list[dict[str, Any]]:
|
||||
"""
|
||||
截断历史记录,优先保留 429_observation(学习数据源)
|
||||
|
||||
策略:超出 MAX_HISTORY_RECORDS 时,先淘汰最旧的非观察记录,
|
||||
仍超限则淘汰最旧的观察记录。
|
||||
"""
|
||||
if len(history) <= self.MAX_HISTORY_RECORDS:
|
||||
return history
|
||||
|
||||
observations = [h for h in history if h.get("type") == "429_observation"]
|
||||
adjustments = [h for h in history if h.get("type") != "429_observation"]
|
||||
|
||||
# 按时间戳排序(最新在后)
|
||||
observations.sort(key=lambda h: h.get("timestamp", ""))
|
||||
adjustments.sort(key=lambda h: h.get("timestamp", ""))
|
||||
|
||||
# 保留尽可能多的观察记录:先缩减 adjustment,再缩减 observation
|
||||
overflow = len(history) - self.MAX_HISTORY_RECORDS
|
||||
trim_adj = min(overflow, len(adjustments))
|
||||
adjustments = adjustments[trim_adj:]
|
||||
overflow -= trim_adj
|
||||
|
||||
if overflow > 0:
|
||||
observations = observations[overflow:]
|
||||
|
||||
merged = observations + adjustments
|
||||
merged.sort(key=lambda h: h.get("timestamp", ""))
|
||||
return merged
|
||||
|
||||
# ==================== 统计与管理 ====================
|
||||
|
||||
def get_adjustment_stats(self, key: ProviderAPIKey) -> dict[str, Any]:
|
||||
"""获取调整统计信息"""
|
||||
history: list[dict[str, Any]] = list(key.adjustment_history or [])
|
||||
samples: list[dict[str, Any]] = list(key.utilization_samples or [])
|
||||
|
||||
is_adaptive = key.rpm_limit is None
|
||||
effective_limit = self.get_effective_limit(key) if is_adaptive else int(key.rpm_limit) # type: ignore
|
||||
|
||||
avg_utilization: float | None = None
|
||||
high_util_ratio: float | None = None
|
||||
if samples:
|
||||
avg_utilization = sum(s["util"] for s in samples) / len(samples)
|
||||
high_util_count = sum(1 for s in samples if s["util"] >= self.UTILIZATION_THRESHOLD)
|
||||
high_util_ratio = high_util_count / len(samples)
|
||||
|
||||
last_429_at_str: str | None = None
|
||||
if key.last_429_at:
|
||||
last_429_at_str = cast(datetime, key.last_429_at).isoformat()
|
||||
|
||||
last_probe_at_str: str | None = None
|
||||
if key.last_probe_increase_at:
|
||||
last_probe_at_str = cast(datetime, key.last_probe_increase_at).isoformat()
|
||||
|
||||
known_boundary = key.last_rpm_peak
|
||||
|
||||
# 观察统计
|
||||
observations = [h for h in history if h.get("type") == "429_observation"]
|
||||
header_observations = [
|
||||
o
|
||||
for o in observations
|
||||
if o.get("upstream_limit") is not None and o["upstream_limit"] > 0
|
||||
]
|
||||
latest_upstream = header_observations[-1]["upstream_limit"] if header_observations else None
|
||||
|
||||
confidence = self.get_confidence(key) if is_adaptive else None
|
||||
enforcement_active = (
|
||||
confidence >= self.ENFORCEMENT_CONFIDENCE_THRESHOLD if confidence is not None else None
|
||||
)
|
||||
|
||||
return {
|
||||
"adaptive_mode": is_adaptive,
|
||||
"rpm_limit": key.rpm_limit,
|
||||
"effective_limit": effective_limit,
|
||||
"learned_limit": key.learned_rpm_limit,
|
||||
# 边界记忆相关
|
||||
"known_boundary": known_boundary,
|
||||
"concurrent_429_count": int(key.concurrent_429_count or 0),
|
||||
"rpm_429_count": int(key.rpm_429_count or 0),
|
||||
"last_429_at": last_429_at_str,
|
||||
"last_429_type": key.last_429_type,
|
||||
"adjustment_count": len(history),
|
||||
"recent_adjustments": history[-5:] if history else [],
|
||||
# 滑动窗口相关
|
||||
"window_sample_count": len(samples),
|
||||
"window_avg_utilization": round(avg_utilization, 3) if avg_utilization else None,
|
||||
"window_high_util_ratio": round(high_util_ratio, 3) if high_util_ratio else None,
|
||||
"utilization_threshold": self.UTILIZATION_THRESHOLD,
|
||||
"high_util_ratio_threshold": self.HIGH_UTILIZATION_RATIO,
|
||||
"min_samples_for_decision": self.MIN_SAMPLES_FOR_DECISION,
|
||||
# 探测性扩容相关
|
||||
"last_probe_increase_at": last_probe_at_str,
|
||||
"probe_increase_interval_minutes": self.PROBE_INCREASE_INTERVAL_MINUTES,
|
||||
# 置信度相关
|
||||
"learning_confidence": round(confidence, 3) if confidence is not None else None,
|
||||
"enforcement_active": enforcement_active,
|
||||
"observation_count": len(observations),
|
||||
"header_observation_count": len(header_observations),
|
||||
"latest_upstream_limit": latest_upstream,
|
||||
}
|
||||
|
||||
def reset_learning(self, db: Session, key: ProviderAPIKey) -> None:
|
||||
"""重置学习状态(管理员功能)"""
|
||||
logger.info(f"[RESET] 重置学习状态: Key {key.id[:8]}...")
|
||||
|
||||
key.learned_rpm_limit = None # type: ignore[assignment]
|
||||
key.concurrent_429_count = 0 # type: ignore[assignment]
|
||||
key.rpm_429_count = 0 # type: ignore[assignment]
|
||||
key.last_429_at = None # type: ignore[assignment]
|
||||
key.last_429_type = None # type: ignore[assignment]
|
||||
key.last_rpm_peak = None # type: ignore[assignment]
|
||||
key.adjustment_history = [] # type: ignore[assignment]
|
||||
key.utilization_samples = [] # type: ignore[assignment]
|
||||
key.last_probe_increase_at = None # type: ignore[assignment]
|
||||
|
||||
self._persist_metadata_update(db)
|
||||
|
||||
|
||||
# 全局单例
|
||||
_adaptive_rpm_manager: AdaptiveRPMManager | None = None
|
||||
|
||||
|
||||
def get_adaptive_rpm_manager() -> AdaptiveRPMManager:
|
||||
"""获取全局自适应 RPM 管理器单例"""
|
||||
global _adaptive_rpm_manager
|
||||
if _adaptive_rpm_manager is None:
|
||||
_adaptive_rpm_manager = AdaptiveRPMManager()
|
||||
return _adaptive_rpm_manager
|
||||
|
||||
|
||||
# 向后兼容别名
|
||||
AdaptiveConcurrencyManager = AdaptiveRPMManager
|
||||
get_adaptive_manager = get_adaptive_rpm_manager
|
||||
615
_deprecated_py_src/services/rate_limit/concurrency_manager.py
Normal file
615
_deprecated_py_src/services/rate_limit/concurrency_manager.py
Normal file
@@ -0,0 +1,615 @@
|
||||
"""
|
||||
RPM 限制管理器 - 支持 Redis 或内存的 Key 级别 RPM 限制
|
||||
|
||||
功能:
|
||||
1. ProviderAPIKey 级别的 RPM 限制(按分钟窗口计数)
|
||||
2. 分布式环境下优先使用 Redis,多实例共享
|
||||
3. 在开发/单实例场景下自动降级为内存计数
|
||||
4. 支持缓存用户优先级(预留槽位机制)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import math
|
||||
import os
|
||||
import time
|
||||
from contextlib import asynccontextmanager
|
||||
from typing import Any
|
||||
|
||||
import redis.asyncio as aioredis
|
||||
|
||||
from src.config.constants import RPMDefaults
|
||||
from src.core.logger import logger
|
||||
|
||||
|
||||
class ConcurrencyManager:
|
||||
"""Key RPM 限制管理器"""
|
||||
|
||||
_instance: ConcurrencyManager | None = None
|
||||
_redis: aioredis.Redis | None = None
|
||||
|
||||
def __new__(cls) -> "ConcurrencyManager":
|
||||
"""单例模式"""
|
||||
if cls._instance is None:
|
||||
cls._instance = super().__new__(cls)
|
||||
return cls._instance
|
||||
|
||||
def __init__(self) -> None:
|
||||
"""初始化内存后端结构(只执行一次)"""
|
||||
if hasattr(self, "_memory_initialized"):
|
||||
return
|
||||
|
||||
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
|
||||
|
||||
self._memory_lock: asyncio.Lock = asyncio.Lock()
|
||||
# Key RPM 计数器:{key_id: (bucket, count)},bucket = floor(now / 60)
|
||||
self._memory_key_rpm_counts: dict[str, tuple[int, int]] = {}
|
||||
self._owns_redis: bool = False
|
||||
self._last_cleanup_bucket: int = 0 # 上次清理时的 bucket,用于定期清理过期数据
|
||||
self._last_cleanup_time: float = 0 # 上次清理的时间戳,用于强制定期清理
|
||||
self._cleanup_interval_seconds: int = config.rpm_cleanup_interval_seconds
|
||||
self._cleanup_task: asyncio.Task | None = None # 后台清理任务
|
||||
|
||||
# 内存模式下的最大条目限制,防止内存泄漏(支持环境变量覆盖)
|
||||
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))
|
||||
)
|
||||
self._memory_initialized = True
|
||||
|
||||
async def initialize(self) -> None:
|
||||
"""初始化 Redis 连接"""
|
||||
if self._redis is not None:
|
||||
return
|
||||
|
||||
try:
|
||||
# 复用全局 Redis 客户端(带熔断/降级),避免重复创建连接池
|
||||
from src.clients.redis_client import get_redis_client
|
||||
|
||||
self._redis = await get_redis_client(require_redis=False)
|
||||
self._owns_redis = False
|
||||
if self._redis:
|
||||
logger.info("[OK] ConcurrencyManager 已复用全局 Redis 客户端")
|
||||
else:
|
||||
logger.warning(
|
||||
"[WARN] Redis 不可用,RPM 限制降级为内存模式(仅在单实例环境下安全)"
|
||||
)
|
||||
# 内存模式下启动后台清理任务
|
||||
self._start_background_cleanup()
|
||||
except Exception as e:
|
||||
logger.error("[ERROR] 获取全局 Redis 客户端失败: {}", e)
|
||||
logger.warning("[WARN] RPM 限制将降级为内存模式(仅在单实例环境下安全)")
|
||||
self._redis = None
|
||||
self._owns_redis = False
|
||||
# 内存模式下启动后台清理任务
|
||||
self._start_background_cleanup()
|
||||
|
||||
def _start_background_cleanup(self) -> None:
|
||||
"""启动后台定期清理任务(仅内存模式需要)"""
|
||||
if self._cleanup_task is not None:
|
||||
return # 已经启动
|
||||
|
||||
async def cleanup_loop() -> None:
|
||||
"""后台清理循环"""
|
||||
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:
|
||||
logger.debug("后台清理任务异常: {}", e)
|
||||
|
||||
try:
|
||||
self._cleanup_task = asyncio.create_task(cleanup_loop())
|
||||
logger.debug("[OK] 内存模式后台清理任务已启动")
|
||||
except RuntimeError:
|
||||
# 没有事件循环时忽略
|
||||
pass
|
||||
|
||||
async def close(self) -> None:
|
||||
"""关闭 Redis 连接"""
|
||||
# 停止后台清理任务
|
||||
if self._cleanup_task is not None:
|
||||
self._cleanup_task.cancel()
|
||||
try:
|
||||
await self._cleanup_task
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
self._cleanup_task = None
|
||||
|
||||
if self._redis and self._owns_redis:
|
||||
await self._redis.close()
|
||||
logger.info("ConcurrencyManager Redis 连接已关闭")
|
||||
self._redis = None
|
||||
self._owns_redis = False
|
||||
|
||||
def _get_rpm_bucket(self, now_ts: float | None = None) -> int:
|
||||
"""获取当前 RPM 计数桶(按分钟)"""
|
||||
ts = now_ts if now_ts is not None else time.time()
|
||||
return int(ts // self._key_rpm_bucket_seconds)
|
||||
|
||||
def _get_key_key(self, key_id: str, bucket: int | None = None) -> str:
|
||||
"""获取 ProviderAPIKey RPM 计数的 Redis Key(按分钟桶)"""
|
||||
b = bucket if bucket is not None else self._get_rpm_bucket()
|
||||
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(
|
||||
"[CRITICAL] 内存 RPM 计数器接近上限 ({}/{}),"
|
||||
"强烈建议启用 Redis!继续增长可能导致 RPM 限制失效",
|
||||
current_size,
|
||||
self._max_memory_rpm_entries,
|
||||
)
|
||||
elif current_size >= high_threshold and key_id not in self._memory_key_rpm_counts:
|
||||
# 每 100 个条目告警一次,避免日志过多
|
||||
if current_size % 100 == 0:
|
||||
logger.error(
|
||||
"[ERROR] 内存 RPM 计数器使用率过高 ({}/{}),建议启用 Redis",
|
||||
current_size,
|
||||
self._max_memory_rpm_entries,
|
||||
)
|
||||
elif current_size >= warning_threshold and key_id not in self._memory_key_rpm_counts:
|
||||
if current_size == warning_threshold:
|
||||
logger.warning(
|
||||
"[WARN] 内存 RPM 计数器达到 {:.0%} 阈值 ({}/{}),建议启用 Redis",
|
||||
self._memory_warning_threshold,
|
||||
current_size,
|
||||
self._max_memory_rpm_entries,
|
||||
)
|
||||
|
||||
# 检查是否超过最大条目限制
|
||||
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(
|
||||
self._memory_key_rpm_counts.items(), key=lambda x: x[1][0] # 按 bucket 排序
|
||||
)
|
||||
for k, _ in sorted_keys[:evict_count]:
|
||||
del self._memory_key_rpm_counts[k]
|
||||
logger.warning("[WARN] 内存 RPM 计数器达到上限,已淘汰 {} 个最旧条目", evict_count)
|
||||
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 时调用)
|
||||
|
||||
清理策略:
|
||||
- 常规清理:每分钟最多执行一次(当 bucket 变化时)
|
||||
- 强制清理:每 5 分钟执行一次(防止长时间无请求导致内存泄漏)
|
||||
"""
|
||||
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)
|
||||
|
||||
for key_id in expired_keys:
|
||||
del self._memory_key_rpm_counts[key_id]
|
||||
|
||||
if expired_keys:
|
||||
logger.debug("[CLEANUP] 清理了 {} 个过期的内存 RPM 计数", len(expired_keys))
|
||||
|
||||
async def get_key_rpm_count(self, key_id: str) -> int:
|
||||
"""
|
||||
获取 Key 当前 RPM 计数
|
||||
|
||||
Args:
|
||||
key_id: ProviderAPIKey ID
|
||||
|
||||
Returns:
|
||||
当前分钟窗口内的请求数
|
||||
"""
|
||||
if self._redis is None:
|
||||
async with self._memory_lock:
|
||||
bucket = self._get_rpm_bucket()
|
||||
# 定期清理过期数据,避免内存泄漏
|
||||
self._cleanup_expired_memory_rpm_counts(bucket)
|
||||
return self._get_memory_key_rpm_count(key_id, bucket)
|
||||
|
||||
try:
|
||||
key_key = self._get_key_key(key_id)
|
||||
result = await self._redis.get(key_key)
|
||||
return int(result) if result else 0
|
||||
except Exception as e:
|
||||
logger.error("获取 RPM 计数失败: {}", e)
|
||||
return 0
|
||||
|
||||
async def check_rpm_available(
|
||||
self,
|
||||
key_id: str,
|
||||
key_rpm_limit: int | None,
|
||||
is_cached_user: bool = False,
|
||||
cache_reservation_ratio: float | None = None,
|
||||
) -> bool:
|
||||
"""
|
||||
检查是否可以通过 RPM 限制(不实际增加计数)
|
||||
|
||||
Args:
|
||||
key_id: ProviderAPIKey ID
|
||||
key_rpm_limit: Key RPM 限制(每分钟最大请求数,None 表示不限制)
|
||||
is_cached_user: 是否是缓存用户
|
||||
cache_reservation_ratio: 缓存预留比例
|
||||
|
||||
Returns:
|
||||
是否可用(True/False)
|
||||
"""
|
||||
if key_rpm_limit is None:
|
||||
return True
|
||||
|
||||
# 从配置读取默认值
|
||||
from src.config.settings import config
|
||||
|
||||
if cache_reservation_ratio is None:
|
||||
cache_reservation_ratio = config.cache_reservation_ratio
|
||||
|
||||
key_count = await self.get_key_rpm_count(key_id)
|
||||
|
||||
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
|
||||
|
||||
async def acquire_rpm_slot(
|
||||
self,
|
||||
key_id: str,
|
||||
key_rpm_limit: int | None,
|
||||
is_cached_user: bool = False,
|
||||
cache_reservation_ratio: float | None = None,
|
||||
) -> bool:
|
||||
"""
|
||||
尝试获取 RPM 槽位(支持缓存用户优先级)
|
||||
|
||||
Args:
|
||||
key_id: ProviderAPIKey ID
|
||||
key_rpm_limit: Key RPM 限制(每分钟最大请求数,None 表示不限制)
|
||||
is_cached_user: 是否是缓存用户(缓存用户可使用全部槽位)
|
||||
cache_reservation_ratio: 缓存预留比例,None 时从配置读取
|
||||
|
||||
Returns:
|
||||
是否成功获取(True/False)
|
||||
|
||||
缓存预留机制说明:
|
||||
- 假设 key_rpm_limit = 100, cache_reservation_ratio = 0.3
|
||||
- 新用户最多使用: 70 RPM (100 * (1 - 0.3))
|
||||
- 缓存用户最多使用: 100 RPM(全部)
|
||||
- 预留的 30 RPM 专门给缓存用户,保证他们的请求优先
|
||||
"""
|
||||
# 从配置读取默认值
|
||||
from src.config.settings import config
|
||||
|
||||
if cache_reservation_ratio is None:
|
||||
cache_reservation_ratio = config.cache_reservation_ratio
|
||||
|
||||
if self._redis is None:
|
||||
async with self._memory_lock:
|
||||
bucket = self._get_rpm_bucket()
|
||||
# 定期清理过期数据,避免内存泄漏
|
||||
self._cleanup_expired_memory_rpm_counts(bucket)
|
||||
|
||||
key_count = self._get_memory_key_rpm_count(key_id, bucket)
|
||||
|
||||
# Key RPM 限制,包含缓存预留
|
||||
if key_rpm_limit is not None:
|
||||
if is_cached_user:
|
||||
if key_count >= key_rpm_limit:
|
||||
return False
|
||||
else:
|
||||
# 新用户只能使用 (1 - cache_reservation_ratio) 的槽位
|
||||
available_for_new = max(
|
||||
1, math.floor(key_rpm_limit * (1 - cache_reservation_ratio))
|
||||
)
|
||||
if key_count >= available_for_new:
|
||||
return False
|
||||
|
||||
# 通过限制,更新计数
|
||||
self._set_memory_key_rpm_count(key_id, bucket, key_count + 1)
|
||||
return True
|
||||
|
||||
bucket = self._get_rpm_bucket()
|
||||
key_key = self._get_key_key(key_id, bucket=bucket)
|
||||
|
||||
try:
|
||||
# 使用 Lua 脚本保证原子性(支持缓存预留逻辑)
|
||||
lua_script = """
|
||||
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]) -- 缓存预留比例
|
||||
|
||||
-- 获取当前值
|
||||
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) 的槽位
|
||||
local available_for_new = math.max(1, math.floor(key_max * (1 - cache_ratio)))
|
||||
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)
|
||||
redis.call('EXPIRE', key_key, key_ttl)
|
||||
|
||||
return 1 -- 成功
|
||||
"""
|
||||
|
||||
# 执行脚本
|
||||
result = await self._redis.eval(
|
||||
lua_script,
|
||||
1, # 1 个 KEY
|
||||
key_key,
|
||||
key_rpm_limit if key_rpm_limit is not None else -1,
|
||||
self._key_rpm_key_ttl_seconds,
|
||||
1 if is_cached_user else 0, # 缓存用户标志
|
||||
cache_reservation_ratio, # 预留比例
|
||||
)
|
||||
|
||||
success = result == 1
|
||||
|
||||
if success:
|
||||
user_type = "缓存用户" if is_cached_user else "新用户"
|
||||
logger.debug("[OK] 获取 RPM 槽位成功: key={}, 类型={}", key_id, user_type)
|
||||
else:
|
||||
key_count = await self.get_key_rpm_count(key_id)
|
||||
|
||||
# 计算新用户可用 RPM
|
||||
if key_rpm_limit and not is_cached_user:
|
||||
available_for_new = int(key_rpm_limit * (1 - cache_reservation_ratio))
|
||||
user_info = f"新用户配额={available_for_new}, 当前={key_count}"
|
||||
else:
|
||||
user_info = f"缓存用户, 当前={key_count}/{key_rpm_limit}"
|
||||
|
||||
logger.warning("[WARN] RPM 限制已达上限: key={}({})", key_id, user_info)
|
||||
|
||||
return success
|
||||
|
||||
except Exception as e:
|
||||
logger.error("获取 RPM 槽位失败,降级到内存模式: {}", e)
|
||||
# Redis 异常时降级到内存模式进行保守限流
|
||||
async with self._memory_lock:
|
||||
bucket = self._get_rpm_bucket()
|
||||
self._cleanup_expired_memory_rpm_counts(bucket)
|
||||
|
||||
key_count = self._get_memory_key_rpm_count(key_id, bucket)
|
||||
|
||||
# 降级模式下使用更保守的限制(50%)
|
||||
fallback_rpm_limit = (
|
||||
max(1, key_rpm_limit // 2) if key_rpm_limit is not None else None
|
||||
)
|
||||
|
||||
if fallback_rpm_limit is not None and key_count >= fallback_rpm_limit:
|
||||
logger.warning(
|
||||
"[FALLBACK] Key RPM 达到降级限制: {}/{}",
|
||||
key_count,
|
||||
fallback_rpm_limit,
|
||||
)
|
||||
return False
|
||||
|
||||
# 更新内存计数
|
||||
self._set_memory_key_rpm_count(key_id, bucket, key_count + 1)
|
||||
logger.debug("[FALLBACK] 使用内存模式获取 RPM 槽位: key={}", key_id)
|
||||
return True
|
||||
|
||||
@asynccontextmanager
|
||||
async def rpm_guard(
|
||||
self,
|
||||
key_id: str,
|
||||
key_rpm_limit: int | None,
|
||||
is_cached_user: bool = False,
|
||||
cache_reservation_ratio: float | None = None,
|
||||
) -> Any:
|
||||
"""
|
||||
RPM 限制上下文管理器(支持缓存用户优先级)
|
||||
|
||||
用法:
|
||||
async with manager.rpm_guard(
|
||||
key_id, key_rpm_limit,
|
||||
is_cached_user=True # 缓存用户
|
||||
):
|
||||
# 执行请求
|
||||
response = await send_request(...)
|
||||
|
||||
如果获取失败,会抛出 ConcurrencyLimitError 异常
|
||||
|
||||
注意:RPM 是按分钟窗口计数,不需要在请求结束后释放
|
||||
"""
|
||||
# 从配置读取默认值
|
||||
from src.config.settings import config
|
||||
|
||||
if cache_reservation_ratio is None:
|
||||
cache_reservation_ratio = config.cache_reservation_ratio
|
||||
|
||||
# 尝试获取槽位(传递缓存用户参数)
|
||||
acquired = await self.acquire_rpm_slot(
|
||||
key_id,
|
||||
key_rpm_limit,
|
||||
is_cached_user,
|
||||
cache_reservation_ratio,
|
||||
)
|
||||
|
||||
if not acquired:
|
||||
from src.core.exceptions import ConcurrencyLimitError
|
||||
|
||||
# Keep the client-facing message generic; do not leak internal IDs.
|
||||
raise ConcurrencyLimitError("服务暂时繁忙,请稍后重试")
|
||||
|
||||
# 记录开始时间和状态
|
||||
import time
|
||||
|
||||
slot_acquired_at = time.time()
|
||||
exception_occurred = False
|
||||
|
||||
try:
|
||||
yield # 执行请求
|
||||
except Exception:
|
||||
# 记录异常
|
||||
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(
|
||||
exception=str(exception_occurred),
|
||||
).inc()
|
||||
|
||||
# 告警:槽位占用时间过长(超过 60 秒)
|
||||
if slot_duration > 60:
|
||||
logger.warning(
|
||||
"[WARN] 请求耗时过长: key_id={}..., duration={:.1f}s, exception={}",
|
||||
key_id[:8] if key_id else "unknown",
|
||||
slot_duration,
|
||||
exception_occurred,
|
||||
)
|
||||
|
||||
except Exception as metric_error:
|
||||
# 指标记录失败不应影响业务逻辑
|
||||
logger.debug("记录指标失败: {}", metric_error)
|
||||
|
||||
# 注意:RPM 计数不需要在请求结束后释放,它会在分钟窗口过期后自动重置
|
||||
|
||||
async def reset_key_rpm(self, key_id: str) -> None:
|
||||
"""
|
||||
重置 Key RPM 计数(管理功能,慎用)
|
||||
|
||||
Args:
|
||||
key_id: ProviderAPIKey ID
|
||||
"""
|
||||
if self._redis is None:
|
||||
async with self._memory_lock:
|
||||
self._memory_key_rpm_counts.pop(key_id, None)
|
||||
logger.info("[RESET] 重置 Key RPM 计数(内存): {}", key_id)
|
||||
return
|
||||
|
||||
try:
|
||||
deleted_count = await self._scan_and_delete(f"rpm:key:{key_id}:*")
|
||||
logger.info("[RESET] 重置 Key RPM 计数: {}, 删除 {} 个键", key_id, deleted_count)
|
||||
except Exception as e:
|
||||
logger.error("重置 Key RPM 计数失败: {}", e)
|
||||
|
||||
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:
|
||||
logger.info("[RESET] 重置所有 Key RPM 计数(内存): {} 个", count)
|
||||
return
|
||||
|
||||
try:
|
||||
deleted_count = await self._scan_and_delete("rpm:key:*")
|
||||
if deleted_count:
|
||||
logger.info("[RESET] 重置所有 Key RPM 计数: {} 个", deleted_count)
|
||||
except Exception as e:
|
||||
logger.error("重置所有 Key RPM 计数失败: {}", e)
|
||||
|
||||
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
|
||||
|
||||
|
||||
# 全局单例
|
||||
_concurrency_manager: ConcurrencyManager | None = None
|
||||
|
||||
|
||||
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
|
||||
424
_deprecated_py_src/services/rate_limit/detector.py
Normal file
424
_deprecated_py_src/services/rate_limit/detector.py
Normal file
@@ -0,0 +1,424 @@
|
||||
"""
|
||||
速率限制检测器 - 解析429响应头,区分并发限制和RPM限制
|
||||
"""
|
||||
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from src.core.logger import logger
|
||||
|
||||
|
||||
class RateLimitType:
|
||||
"""速率限制类型"""
|
||||
|
||||
CONCURRENT = "concurrent" # 并发限制
|
||||
RPM = "rpm" # 每分钟请求数限制
|
||||
DAILY = "daily" # 每日限制
|
||||
MONTHLY = "monthly" # 每月限制
|
||||
UNKNOWN = "unknown" # 未知类型
|
||||
|
||||
|
||||
class RateLimitInfo:
|
||||
"""速率限制信息"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
limit_type: str,
|
||||
retry_after: int | None = None,
|
||||
limit_value: int | None = None,
|
||||
remaining: int | None = None,
|
||||
reset_at: datetime | None = None,
|
||||
current_usage: int | None = None,
|
||||
raw_headers: dict[str, str] | None = None,
|
||||
):
|
||||
self.limit_type = limit_type
|
||||
self.retry_after = retry_after # 需要等待的秒数
|
||||
self.limit_value = limit_value # 限制值
|
||||
self.remaining = remaining # 剩余配额
|
||||
self.reset_at = reset_at # 重置时间
|
||||
self.current_usage = current_usage # 当前使用量
|
||||
self.raw_headers = raw_headers or {}
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return (
|
||||
f"RateLimitInfo(type={self.limit_type}, "
|
||||
f"retry_after={self.retry_after}, "
|
||||
f"limit={self.limit_value}, "
|
||||
f"remaining={self.remaining})"
|
||||
)
|
||||
|
||||
|
||||
class RateLimitDetector:
|
||||
"""
|
||||
速率限制检测器
|
||||
|
||||
支持的提供商:
|
||||
- Anthropic Claude API
|
||||
- OpenAI API
|
||||
- 通用 HTTP 标准头
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def detect_from_headers(
|
||||
headers: dict[str, str],
|
||||
provider_name: str = "unknown",
|
||||
current_usage: int | None = None,
|
||||
) -> RateLimitInfo:
|
||||
"""
|
||||
从响应头中检测速率限制类型
|
||||
|
||||
Args:
|
||||
headers: 429响应的HTTP头
|
||||
provider_name: 提供商名称(用于选择解析策略)
|
||||
current_usage: 当前使用量(RPM 计数,用于启发式判断是否为并发限制)
|
||||
|
||||
Returns:
|
||||
RateLimitInfo对象
|
||||
"""
|
||||
# 标准化header key (转小写)
|
||||
headers_lower = {k.lower(): v for k, v in headers.items()}
|
||||
|
||||
# 根据提供商选择解析策略
|
||||
if "anthropic" in provider_name.lower() or "claude" in provider_name.lower():
|
||||
return RateLimitDetector._parse_anthropic_headers(headers_lower, current_usage)
|
||||
elif "openai" in provider_name.lower():
|
||||
return RateLimitDetector._parse_openai_headers(headers_lower, current_usage)
|
||||
else:
|
||||
return RateLimitDetector._parse_generic_headers(headers_lower, current_usage)
|
||||
|
||||
@staticmethod
|
||||
def _parse_anthropic_headers(
|
||||
headers: dict[str, str],
|
||||
current_usage: int | None = None,
|
||||
) -> RateLimitInfo:
|
||||
"""
|
||||
解析 Anthropic Claude API 的速率限制头
|
||||
|
||||
常见头部:
|
||||
- anthropic-ratelimit-requests-limit: 50
|
||||
- anthropic-ratelimit-requests-remaining: 0
|
||||
- anthropic-ratelimit-requests-reset: 2024-01-01T00:00:00Z
|
||||
- anthropic-ratelimit-tokens-limit: 100000
|
||||
- anthropic-ratelimit-tokens-remaining: 50000
|
||||
- retry-after: 60
|
||||
"""
|
||||
retry_after = RateLimitDetector._parse_retry_after(headers)
|
||||
|
||||
# 获取请求限制信息
|
||||
requests_limit = RateLimitDetector._parse_int(
|
||||
headers.get("anthropic-ratelimit-requests-limit")
|
||||
)
|
||||
requests_remaining = RateLimitDetector._parse_int(
|
||||
headers.get("anthropic-ratelimit-requests-remaining")
|
||||
)
|
||||
requests_reset = RateLimitDetector._parse_datetime(
|
||||
headers.get("anthropic-ratelimit-requests-reset")
|
||||
)
|
||||
|
||||
# 判断限制类型
|
||||
# 1. 明确的 RPM 限制:请求数剩余为 0
|
||||
if requests_remaining is not None and requests_remaining == 0:
|
||||
return RateLimitInfo(
|
||||
limit_type=RateLimitType.RPM,
|
||||
retry_after=retry_after,
|
||||
limit_value=requests_limit,
|
||||
remaining=requests_remaining,
|
||||
reset_at=requests_reset,
|
||||
raw_headers=headers,
|
||||
)
|
||||
|
||||
# 2. 并发限制判断(多条件策略)
|
||||
# 注意:current_usage 是 RPM 计数(当前分钟请求数),不是真正的并发数
|
||||
#
|
||||
# 判断条件(满足任一即可):
|
||||
# A. 强判断:remaining > 0 且 retry_after <= 30(Provider 明确告知还有配额但需要等待)
|
||||
# B. 弱判断:只有 retry_after <= 5 且缺少 remaining 头(短等待时间是并发限制的典型特征)
|
||||
#
|
||||
# 选择保守的 retry_after 阈值:
|
||||
# - 强判断用 30 秒(有 remaining 头时)
|
||||
# - 弱判断用 5 秒(无 remaining 头时,更保守)
|
||||
is_likely_concurrent = False
|
||||
concurrent_reason = ""
|
||||
|
||||
# 条件 A:remaining > 0 且 retry_after <= 30
|
||||
if (
|
||||
requests_remaining is not None
|
||||
and requests_remaining > 0
|
||||
and retry_after is not None
|
||||
and retry_after <= 30
|
||||
):
|
||||
is_likely_concurrent = True
|
||||
concurrent_reason = (
|
||||
f"remaining={requests_remaining} > 0, retry_after={retry_after}s <= 30s"
|
||||
)
|
||||
|
||||
# 条件 B:无 remaining 头但 retry_after 很短(<= 5 秒)
|
||||
elif requests_remaining is None and retry_after is not None and retry_after <= 5:
|
||||
is_likely_concurrent = True
|
||||
concurrent_reason = f"no remaining header, retry_after={retry_after}s <= 5s"
|
||||
|
||||
if is_likely_concurrent:
|
||||
logger.info(f"检测到并发限制: {concurrent_reason}")
|
||||
return RateLimitInfo(
|
||||
limit_type=RateLimitType.CONCURRENT,
|
||||
retry_after=retry_after,
|
||||
current_usage=current_usage,
|
||||
raw_headers=headers,
|
||||
)
|
||||
|
||||
# 3. 默认视为 RPM 限制(更保守的处理)
|
||||
# 无法明确区分时,视为 RPM 限制让系统降低 RPM
|
||||
# 这比误判为并发限制(不降 RPM)更安全
|
||||
if retry_after is not None or requests_limit is not None:
|
||||
logger.info(
|
||||
f"无法明确区分限制类型,保守视为 RPM 限制: "
|
||||
f"remaining={requests_remaining}, retry_after={retry_after}"
|
||||
)
|
||||
return RateLimitInfo(
|
||||
limit_type=RateLimitType.RPM,
|
||||
retry_after=retry_after,
|
||||
limit_value=requests_limit,
|
||||
remaining=requests_remaining,
|
||||
reset_at=requests_reset,
|
||||
current_usage=current_usage,
|
||||
raw_headers=headers,
|
||||
)
|
||||
|
||||
# 4. 完全没有信息,标记为未知
|
||||
return RateLimitInfo(
|
||||
limit_type=RateLimitType.UNKNOWN,
|
||||
retry_after=retry_after,
|
||||
raw_headers=headers,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _parse_openai_headers(
|
||||
headers: dict[str, str],
|
||||
current_usage: int | None = None,
|
||||
) -> RateLimitInfo:
|
||||
"""
|
||||
解析 OpenAI API 的速率限制头
|
||||
|
||||
常见头部:
|
||||
- x-ratelimit-limit-requests: 3500
|
||||
- x-ratelimit-remaining-requests: 0
|
||||
- x-ratelimit-reset-requests: 2024-01-01T00:00:00Z
|
||||
- x-ratelimit-limit-tokens: 90000
|
||||
- x-ratelimit-remaining-tokens: 50000
|
||||
- retry-after: 60
|
||||
"""
|
||||
retry_after = RateLimitDetector._parse_retry_after(headers)
|
||||
|
||||
# 获取请求限制信息
|
||||
requests_limit = RateLimitDetector._parse_int(headers.get("x-ratelimit-limit-requests"))
|
||||
requests_remaining = RateLimitDetector._parse_int(
|
||||
headers.get("x-ratelimit-remaining-requests")
|
||||
)
|
||||
requests_reset = RateLimitDetector._parse_datetime(
|
||||
headers.get("x-ratelimit-reset-requests")
|
||||
)
|
||||
|
||||
# 判断限制类型
|
||||
# 1. 明确的 RPM 限制
|
||||
if requests_remaining is not None and requests_remaining == 0:
|
||||
return RateLimitInfo(
|
||||
limit_type=RateLimitType.RPM,
|
||||
retry_after=retry_after,
|
||||
limit_value=requests_limit,
|
||||
remaining=requests_remaining,
|
||||
reset_at=requests_reset,
|
||||
raw_headers=headers,
|
||||
)
|
||||
|
||||
# 2. 并发限制判断(多条件策略)
|
||||
# 判断条件(满足任一即可):
|
||||
# A. 强判断:remaining > 0 且 retry_after <= 30
|
||||
# B. 弱判断:只有 retry_after <= 5 且缺少 remaining 头
|
||||
is_likely_concurrent = False
|
||||
concurrent_reason = ""
|
||||
|
||||
if (
|
||||
requests_remaining is not None
|
||||
and requests_remaining > 0
|
||||
and retry_after is not None
|
||||
and retry_after <= 30
|
||||
):
|
||||
is_likely_concurrent = True
|
||||
concurrent_reason = (
|
||||
f"remaining={requests_remaining} > 0, retry_after={retry_after}s <= 30s"
|
||||
)
|
||||
elif requests_remaining is None and retry_after is not None and retry_after <= 5:
|
||||
is_likely_concurrent = True
|
||||
concurrent_reason = f"no remaining header, retry_after={retry_after}s <= 5s"
|
||||
|
||||
if is_likely_concurrent:
|
||||
logger.info(f"检测到并发限制: {concurrent_reason}")
|
||||
return RateLimitInfo(
|
||||
limit_type=RateLimitType.CONCURRENT,
|
||||
retry_after=retry_after,
|
||||
current_usage=current_usage,
|
||||
raw_headers=headers,
|
||||
)
|
||||
|
||||
# 3. 默认视为 RPM 限制(更保守的处理)
|
||||
if retry_after is not None or requests_limit is not None:
|
||||
logger.info(
|
||||
f"无法明确区分限制类型,保守视为 RPM 限制: "
|
||||
f"remaining={requests_remaining}, retry_after={retry_after}"
|
||||
)
|
||||
return RateLimitInfo(
|
||||
limit_type=RateLimitType.RPM,
|
||||
retry_after=retry_after,
|
||||
limit_value=requests_limit,
|
||||
remaining=requests_remaining,
|
||||
reset_at=requests_reset,
|
||||
current_usage=current_usage,
|
||||
raw_headers=headers,
|
||||
)
|
||||
|
||||
# 4. 完全没有信息,标记为未知
|
||||
return RateLimitInfo(
|
||||
limit_type=RateLimitType.UNKNOWN,
|
||||
retry_after=retry_after,
|
||||
raw_headers=headers,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _parse_generic_headers(
|
||||
headers: dict[str, str],
|
||||
current_usage: int | None = None,
|
||||
) -> RateLimitInfo:
|
||||
"""
|
||||
解析通用的速率限制头
|
||||
|
||||
标准头部:
|
||||
- retry-after: 60
|
||||
- x-ratelimit-limit: 100
|
||||
- x-ratelimit-remaining: 0
|
||||
- x-ratelimit-reset: 1609459200
|
||||
"""
|
||||
retry_after = RateLimitDetector._parse_retry_after(headers)
|
||||
|
||||
limit_value = RateLimitDetector._parse_int(headers.get("x-ratelimit-limit"))
|
||||
remaining = RateLimitDetector._parse_int(headers.get("x-ratelimit-remaining"))
|
||||
|
||||
# 1. 明确的 RPM 限制
|
||||
if remaining is not None and remaining == 0:
|
||||
return RateLimitInfo(
|
||||
limit_type=RateLimitType.RPM,
|
||||
retry_after=retry_after,
|
||||
limit_value=limit_value,
|
||||
remaining=remaining,
|
||||
raw_headers=headers,
|
||||
)
|
||||
|
||||
# 2. 并发限制判断(多条件策略)
|
||||
# 判断条件(满足任一即可):
|
||||
# A. 强判断:remaining > 0 且 retry_after <= 30
|
||||
# B. 弱判断:只有 retry_after <= 5 且缺少 remaining 头
|
||||
is_likely_concurrent = False
|
||||
concurrent_reason = ""
|
||||
|
||||
if (
|
||||
remaining is not None
|
||||
and remaining > 0
|
||||
and retry_after is not None
|
||||
and retry_after <= 30
|
||||
):
|
||||
is_likely_concurrent = True
|
||||
concurrent_reason = f"remaining={remaining} > 0, retry_after={retry_after}s <= 30s"
|
||||
elif remaining is None and retry_after is not None and retry_after <= 5:
|
||||
is_likely_concurrent = True
|
||||
concurrent_reason = f"no remaining header, retry_after={retry_after}s <= 5s"
|
||||
|
||||
if is_likely_concurrent:
|
||||
logger.info(f"检测到并发限制: {concurrent_reason}")
|
||||
return RateLimitInfo(
|
||||
limit_type=RateLimitType.CONCURRENT,
|
||||
retry_after=retry_after,
|
||||
current_usage=current_usage,
|
||||
raw_headers=headers,
|
||||
)
|
||||
|
||||
# 3. 默认视为 RPM 限制(更保守的处理)
|
||||
if retry_after is not None or limit_value is not None:
|
||||
logger.info(
|
||||
f"无法明确区分限制类型,保守视为 RPM 限制: "
|
||||
f"remaining={remaining}, retry_after={retry_after}"
|
||||
)
|
||||
return RateLimitInfo(
|
||||
limit_type=RateLimitType.RPM,
|
||||
retry_after=retry_after,
|
||||
limit_value=limit_value,
|
||||
remaining=remaining,
|
||||
current_usage=current_usage,
|
||||
raw_headers=headers,
|
||||
)
|
||||
|
||||
# 4. 完全没有信息,标记为未知
|
||||
return RateLimitInfo(
|
||||
limit_type=RateLimitType.UNKNOWN,
|
||||
retry_after=retry_after,
|
||||
raw_headers=headers,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _parse_retry_after(headers: dict[str, str]) -> int | None:
|
||||
"""解析 Retry-After 头"""
|
||||
retry_after_str = headers.get("retry-after")
|
||||
if not retry_after_str:
|
||||
return None
|
||||
|
||||
try:
|
||||
# 尝试解析为整数(秒数)
|
||||
return int(retry_after_str)
|
||||
except ValueError:
|
||||
# 尝试解析为HTTP日期格式
|
||||
try:
|
||||
retry_date = datetime.strptime(retry_after_str, "%a, %d %b %Y %H:%M:%S %Z")
|
||||
delta = retry_date - datetime.now(timezone.utc)
|
||||
return max(int(delta.total_seconds()), 0)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _parse_int(value: str | None) -> int | None:
|
||||
"""安全解析整数"""
|
||||
if not value:
|
||||
return None
|
||||
try:
|
||||
return int(value)
|
||||
except (ValueError, TypeError):
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _parse_datetime(value: str | None) -> datetime | None:
|
||||
"""安全解析ISO 8601日期时间"""
|
||||
if not value:
|
||||
return None
|
||||
try:
|
||||
# 尝试解析 ISO 8601 格式
|
||||
if value.endswith("Z"):
|
||||
value = value[:-1] + "+00:00"
|
||||
return datetime.fromisoformat(value)
|
||||
except (ValueError, TypeError):
|
||||
return None
|
||||
|
||||
|
||||
# 便捷函数
|
||||
def detect_rate_limit_type(
|
||||
headers: dict[str, str],
|
||||
provider_name: str = "unknown",
|
||||
current_usage: int | None = None,
|
||||
) -> RateLimitInfo:
|
||||
"""
|
||||
检测速率限制类型(便捷函数)
|
||||
|
||||
Args:
|
||||
headers: 429响应头
|
||||
provider_name: 提供商名称
|
||||
current_usage: 当前使用量(RPM 计数)
|
||||
|
||||
Returns:
|
||||
RateLimitInfo对象
|
||||
"""
|
||||
return RateLimitDetector.detect_from_headers(headers, provider_name, current_usage)
|
||||
354
_deprecated_py_src/services/rate_limit/ip_limiter.py
Normal file
354
_deprecated_py_src/services/rate_limit/ip_limiter.py
Normal file
@@ -0,0 +1,354 @@
|
||||
"""
|
||||
IP 级别的速率限制服务
|
||||
|
||||
提供基于 IP 地址的速率限制,防止暴力破解和 DDoS 攻击
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import ipaddress
|
||||
|
||||
from src.clients.redis_client import get_redis_client
|
||||
from src.core.logger import logger
|
||||
|
||||
|
||||
class IPRateLimiter:
|
||||
"""IP 速率限制服务"""
|
||||
|
||||
# Redis key 前缀
|
||||
RATE_LIMIT_PREFIX = "ip:rate_limit:"
|
||||
BLACKLIST_PREFIX = "ip:blacklist:"
|
||||
WHITELIST_KEY = "ip:whitelist"
|
||||
|
||||
# 默认限制配置(每分钟)
|
||||
DEFAULT_LIMITS = {
|
||||
"default": 100, # 默认限制
|
||||
"login": 20, # 登录接口
|
||||
"register": 10, # 注册接口
|
||||
"api": 60, # API 接口
|
||||
"public": 60, # 公共接口
|
||||
"verification_send": 5, # 发送验证码接口
|
||||
"verification_verify": 20, # 验证验证码接口
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
async def check_limit(
|
||||
ip_address: str, endpoint_type: str = "default", limit: int | None = None
|
||||
) -> tuple[bool, int, int]:
|
||||
"""
|
||||
检查 IP 是否超过速率限制
|
||||
|
||||
Args:
|
||||
ip_address: IP 地址
|
||||
endpoint_type: 端点类型(default, login, register, api, public)
|
||||
limit: 自定义限制值,None 则使用默认值
|
||||
|
||||
Returns:
|
||||
(是否允许, 剩余次数, 重置时间秒数)
|
||||
"""
|
||||
# 检查白名单
|
||||
if await IPRateLimiter.is_whitelisted(ip_address):
|
||||
return True, 999999, 60
|
||||
|
||||
# 检查黑名单
|
||||
if await IPRateLimiter.is_blacklisted(ip_address):
|
||||
logger.warning(f"黑名单 IP 尝试访问: {ip_address}, 类型: {endpoint_type}")
|
||||
return False, 0, 0
|
||||
|
||||
redis_client = await get_redis_client(require_redis=False)
|
||||
|
||||
if redis_client is None:
|
||||
# Redis 不可用时降级:允许访问但记录警告
|
||||
logger.warning("Redis 不可用,跳过 IP 速率限制(降级模式)")
|
||||
return True, 0, 60
|
||||
|
||||
# 确定限制值
|
||||
rate_limit = (
|
||||
limit if limit is not None else IPRateLimiter.DEFAULT_LIMITS.get(endpoint_type, 100)
|
||||
)
|
||||
|
||||
try:
|
||||
# Redis key: ip:rate_limit:{type}:{ip}
|
||||
redis_key = f"{IPRateLimiter.RATE_LIMIT_PREFIX}{endpoint_type}:{ip_address}"
|
||||
|
||||
# 使用 Redis 的滑动窗口计数器
|
||||
# INCR 并设置过期时间
|
||||
count = await redis_client.incr(redis_key)
|
||||
|
||||
# 第一次访问时设置过期时间
|
||||
if count == 1:
|
||||
await redis_client.expire(redis_key, 60) # 60秒窗口
|
||||
|
||||
# 获取 TTL(剩余过期时间)
|
||||
ttl = await redis_client.ttl(redis_key)
|
||||
if ttl < 0:
|
||||
# 如果没有过期时间,重新设置
|
||||
await redis_client.expire(redis_key, 60)
|
||||
ttl = 60
|
||||
|
||||
remaining = max(0, rate_limit - count)
|
||||
allowed = count <= rate_limit
|
||||
|
||||
if not allowed:
|
||||
logger.warning(
|
||||
f"IP 速率限制触发: {ip_address}, 类型: {endpoint_type}, 计数: {count}/{rate_limit}"
|
||||
)
|
||||
|
||||
return allowed, remaining, ttl
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"检查 IP 速率限制失败: {e}")
|
||||
# 发生错误时允许访问,避免误杀
|
||||
return True, 0, 60
|
||||
|
||||
@staticmethod
|
||||
async def add_to_blacklist(
|
||||
ip_address: str, reason: str = "manual", ttl: int | None = None
|
||||
) -> bool:
|
||||
"""
|
||||
将 IP 加入黑名单
|
||||
|
||||
Args:
|
||||
ip_address: IP 地址
|
||||
reason: 加入黑名单的原因
|
||||
ttl: 过期时间(秒),None 表示永久
|
||||
|
||||
Returns:
|
||||
是否成功
|
||||
"""
|
||||
redis_client = await get_redis_client(require_redis=False)
|
||||
|
||||
if redis_client is None:
|
||||
logger.warning("Redis 不可用,无法将 IP 加入黑名单")
|
||||
return False
|
||||
|
||||
try:
|
||||
redis_key = f"{IPRateLimiter.BLACKLIST_PREFIX}{ip_address}"
|
||||
|
||||
if ttl is not None:
|
||||
await redis_client.setex(redis_key, ttl, reason)
|
||||
else:
|
||||
await redis_client.set(redis_key, reason)
|
||||
|
||||
logger.warning(f"IP 已加入黑名单: {ip_address}, 原因: {reason}, TTL: {ttl or '永久'}")
|
||||
return True
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"添加 IP 到黑名单失败: {e}")
|
||||
return False
|
||||
|
||||
@staticmethod
|
||||
async def remove_from_blacklist(ip_address: str) -> bool:
|
||||
"""
|
||||
从黑名单移除 IP
|
||||
|
||||
Args:
|
||||
ip_address: IP 地址
|
||||
|
||||
Returns:
|
||||
是否成功
|
||||
"""
|
||||
redis_client = await get_redis_client(require_redis=False)
|
||||
|
||||
if redis_client is None:
|
||||
logger.warning("Redis 不可用,无法从黑名单移除 IP")
|
||||
return False
|
||||
|
||||
try:
|
||||
redis_key = f"{IPRateLimiter.BLACKLIST_PREFIX}{ip_address}"
|
||||
deleted = await redis_client.delete(redis_key)
|
||||
|
||||
if deleted:
|
||||
logger.info(f"IP 已从黑名单移除: {ip_address}")
|
||||
|
||||
return bool(deleted)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"从黑名单移除 IP 失败: {e}")
|
||||
return False
|
||||
|
||||
@staticmethod
|
||||
async def is_blacklisted(ip_address: str) -> bool:
|
||||
"""
|
||||
检查 IP 是否在黑名单中
|
||||
|
||||
Args:
|
||||
ip_address: IP 地址
|
||||
|
||||
Returns:
|
||||
是否在黑名单中
|
||||
"""
|
||||
redis_client = await get_redis_client(require_redis=False)
|
||||
|
||||
if redis_client is None:
|
||||
return False
|
||||
|
||||
try:
|
||||
redis_key = f"{IPRateLimiter.BLACKLIST_PREFIX}{ip_address}"
|
||||
exists = await redis_client.exists(redis_key)
|
||||
return bool(exists)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"检查 IP 黑名单状态失败: {e}")
|
||||
return False
|
||||
|
||||
@staticmethod
|
||||
async def add_to_whitelist(ip_address: str) -> bool:
|
||||
"""
|
||||
将 IP 加入白名单
|
||||
|
||||
Args:
|
||||
ip_address: IP 地址或 CIDR 格式(如 192.168.1.0/24)
|
||||
|
||||
Returns:
|
||||
是否成功
|
||||
"""
|
||||
redis_client = await get_redis_client(require_redis=False)
|
||||
|
||||
if redis_client is None:
|
||||
logger.warning("Redis 不可用,无法将 IP 加入白名单")
|
||||
return False
|
||||
|
||||
try:
|
||||
# 验证 IP 格式
|
||||
try:
|
||||
ipaddress.ip_network(ip_address, strict=False)
|
||||
except ValueError as e:
|
||||
logger.error(f"无效的 IP 地址格式: {ip_address}, 错误: {e}")
|
||||
return False
|
||||
|
||||
# 使用 Redis Set 存储白名单
|
||||
await redis_client.sadd(IPRateLimiter.WHITELIST_KEY, ip_address)
|
||||
|
||||
logger.info(f"IP 已加入白名单: {ip_address}")
|
||||
return True
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"添加 IP 到白名单失败: {e}")
|
||||
return False
|
||||
|
||||
@staticmethod
|
||||
async def remove_from_whitelist(ip_address: str) -> bool:
|
||||
"""
|
||||
从白名单移除 IP
|
||||
|
||||
Args:
|
||||
ip_address: IP 地址
|
||||
|
||||
Returns:
|
||||
是否成功
|
||||
"""
|
||||
redis_client = await get_redis_client(require_redis=False)
|
||||
|
||||
if redis_client is None:
|
||||
logger.warning("Redis 不可用,无法从白名单移除 IP")
|
||||
return False
|
||||
|
||||
try:
|
||||
removed = await redis_client.srem(IPRateLimiter.WHITELIST_KEY, ip_address)
|
||||
|
||||
if removed:
|
||||
logger.info(f"IP 已从白名单移除: {ip_address}")
|
||||
|
||||
return bool(removed)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"从白名单移除 IP 失败: {e}")
|
||||
return False
|
||||
|
||||
@staticmethod
|
||||
async def is_whitelisted(ip_address: str) -> bool:
|
||||
"""
|
||||
检查 IP 是否在白名单中(支持 CIDR 匹配)
|
||||
|
||||
Args:
|
||||
ip_address: IP 地址
|
||||
|
||||
Returns:
|
||||
是否在白名单中
|
||||
"""
|
||||
redis_client = await get_redis_client(require_redis=False)
|
||||
|
||||
if redis_client is None:
|
||||
return False
|
||||
|
||||
try:
|
||||
# 获取所有白名单条目
|
||||
whitelist = await redis_client.smembers(IPRateLimiter.WHITELIST_KEY)
|
||||
|
||||
if not whitelist:
|
||||
return False
|
||||
|
||||
# 将 IP 地址转换为 ip_address 对象
|
||||
try:
|
||||
ip_obj = ipaddress.ip_address(ip_address)
|
||||
except ValueError:
|
||||
return False
|
||||
|
||||
# 检查是否匹配白名单中的任何条目
|
||||
for entry in whitelist:
|
||||
try:
|
||||
network = ipaddress.ip_network(entry, strict=False)
|
||||
if ip_obj in network:
|
||||
return True
|
||||
except ValueError:
|
||||
# 如果条目格式无效,跳过
|
||||
continue
|
||||
|
||||
return False
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"检查 IP 白名单状态失败: {e}")
|
||||
return False
|
||||
|
||||
@staticmethod
|
||||
async def get_blacklist_stats() -> dict:
|
||||
"""
|
||||
获取黑名单统计信息
|
||||
|
||||
Returns:
|
||||
统计信息字典
|
||||
"""
|
||||
redis_client = await get_redis_client(require_redis=False)
|
||||
|
||||
if redis_client is None:
|
||||
return {"available": False, "total": 0, "error": "Redis 不可用"}
|
||||
|
||||
try:
|
||||
pattern = f"{IPRateLimiter.BLACKLIST_PREFIX}*"
|
||||
cursor = 0
|
||||
total = 0
|
||||
|
||||
while True:
|
||||
cursor, keys = await redis_client.scan(cursor=cursor, match=pattern, count=100)
|
||||
total += len(keys)
|
||||
|
||||
if cursor == 0:
|
||||
break
|
||||
|
||||
return {"available": True, "total": total}
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"获取黑名单统计失败: {e}")
|
||||
return {"available": False, "total": 0, "error": str(e)}
|
||||
|
||||
@staticmethod
|
||||
async def get_whitelist() -> set[str]:
|
||||
"""
|
||||
获取白名单列表
|
||||
|
||||
Returns:
|
||||
白名单 IP 集合
|
||||
"""
|
||||
redis_client = await get_redis_client(require_redis=False)
|
||||
|
||||
if redis_client is None:
|
||||
return set()
|
||||
|
||||
try:
|
||||
whitelist = await redis_client.smembers(IPRateLimiter.WHITELIST_KEY)
|
||||
return whitelist if whitelist else set()
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"获取白名单失败: {e}")
|
||||
return set()
|
||||
354
_deprecated_py_src/services/rate_limit/user_rpm_limiter.py
Normal file
354
_deprecated_py_src/services/rate_limit/user_rpm_limiter.py
Normal file
@@ -0,0 +1,354 @@
|
||||
"""
|
||||
用户/API Key RPM 限制器
|
||||
|
||||
支持两层叠加限流:
|
||||
1. 用户级(或独立 Key 级)总 RPM
|
||||
2. 普通 Key 子限制 RPM
|
||||
|
||||
实现策略:
|
||||
- Redis 可用时使用分钟桶 + Lua 脚本原子检查/消费
|
||||
- Redis 不可用时降级为内存计数(仅适用于单实例)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timezone
|
||||
|
||||
import redis.asyncio as aioredis
|
||||
|
||||
from src.config.settings import config
|
||||
from src.core.logger import logger
|
||||
|
||||
SYSTEM_RPM_CONFIG_KEY = "rate_limit_per_minute"
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class RpmCheckResult:
|
||||
"""RPM 检查结果。"""
|
||||
|
||||
allowed: bool
|
||||
scope: str | None = None
|
||||
limit: int | None = None
|
||||
remaining: int | None = None
|
||||
retry_after: int | None = None
|
||||
|
||||
|
||||
class UserRpmLimiter:
|
||||
"""用户/API Key 双层 RPM 限制器。
|
||||
|
||||
通过模块级 ``get_user_rpm_limiter()`` 工厂函数获取唯一实例,
|
||||
不要直接调用构造函数。
|
||||
"""
|
||||
|
||||
_CHECK_AND_CONSUME_SCRIPT = """
|
||||
local user_key = KEYS[1]
|
||||
local key_key = KEYS[2]
|
||||
local user_limit = tonumber(ARGV[1])
|
||||
local key_limit = tonumber(ARGV[2])
|
||||
local ttl = tonumber(ARGV[3])
|
||||
local retry_after = tonumber(ARGV[4])
|
||||
|
||||
local user_count = 0
|
||||
if user_limit > 0 then
|
||||
user_count = tonumber(redis.call('GET', user_key) or '0')
|
||||
if user_count >= user_limit then
|
||||
return {0, 1, user_limit, 0, retry_after}
|
||||
end
|
||||
end
|
||||
|
||||
local key_count = 0
|
||||
if key_limit > 0 then
|
||||
key_count = tonumber(redis.call('GET', key_key) or '0')
|
||||
if key_count >= key_limit then
|
||||
return {0, 2, key_limit, 0, retry_after}
|
||||
end
|
||||
end
|
||||
|
||||
local remaining = -1
|
||||
if user_limit > 0 then
|
||||
user_count = redis.call('INCR', user_key)
|
||||
redis.call('EXPIRE', user_key, ttl)
|
||||
remaining = user_limit - user_count
|
||||
end
|
||||
|
||||
if key_limit > 0 then
|
||||
key_count = redis.call('INCR', key_key)
|
||||
redis.call('EXPIRE', key_key, ttl)
|
||||
local key_remaining = key_limit - key_count
|
||||
if remaining == -1 or key_remaining < remaining then
|
||||
remaining = key_remaining
|
||||
end
|
||||
end
|
||||
|
||||
return {1, 0, 0, remaining, 0}
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._redis: aioredis.Redis | None = None
|
||||
self._bucket_seconds = int(config.rpm_bucket_seconds)
|
||||
self._key_ttl_seconds = int(config.rpm_key_ttl_seconds)
|
||||
self._cleanup_interval_seconds = int(config.rpm_cleanup_interval_seconds)
|
||||
self._memory_lock: asyncio.Lock = asyncio.Lock()
|
||||
self._memory_counts: dict[str, tuple[int, int]] = {}
|
||||
self._cleanup_task: asyncio.Task | None = None
|
||||
|
||||
async def initialize(self) -> None:
|
||||
if self._redis is not None:
|
||||
return
|
||||
|
||||
try:
|
||||
from src.clients.redis_client import get_redis_client
|
||||
|
||||
self._redis = await get_redis_client(require_redis=False)
|
||||
if self._redis:
|
||||
logger.info("[OK] UserRpmLimiter 已复用全局 Redis 客户端")
|
||||
return
|
||||
except Exception as exc:
|
||||
logger.warning("初始化 UserRpmLimiter Redis 客户端失败,降级为内存模式: {}", exc)
|
||||
|
||||
self._redis = None
|
||||
self._start_background_cleanup()
|
||||
|
||||
async def close(self) -> None:
|
||||
if self._cleanup_task is not None:
|
||||
self._cleanup_task.cancel()
|
||||
try:
|
||||
await self._cleanup_task
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
self._cleanup_task = None
|
||||
|
||||
@property
|
||||
def bucket_seconds(self) -> int:
|
||||
return self._bucket_seconds
|
||||
|
||||
def get_user_rpm_key(self, user_id: str, bucket: int | None = None) -> str:
|
||||
b = bucket if bucket is not None else self._get_rpm_bucket()
|
||||
return f"rpm:user:{user_id}:{b}"
|
||||
|
||||
def get_standalone_rpm_key(self, api_key_id: str, bucket: int | None = None) -> str:
|
||||
b = bucket if bucket is not None else self._get_rpm_bucket()
|
||||
return f"rpm:ukey:{api_key_id}:{b}"
|
||||
|
||||
def get_key_rpm_key(self, api_key_id: str, bucket: int | None = None) -> str:
|
||||
b = bucket if bucket is not None else self._get_rpm_bucket()
|
||||
return f"rpm:key:{api_key_id}:{b}"
|
||||
|
||||
def get_retry_after(self, now_ts: float | None = None) -> int:
|
||||
ts = now_ts if now_ts is not None else time.time()
|
||||
elapsed = int(ts % self._bucket_seconds)
|
||||
return max(1, self._bucket_seconds - elapsed)
|
||||
|
||||
def get_reset_at(self, now_ts: float | None = None) -> datetime:
|
||||
ts = now_ts if now_ts is not None else time.time()
|
||||
bucket = self._get_rpm_bucket(ts)
|
||||
reset_ts = (bucket + 1) * self._bucket_seconds
|
||||
return datetime.fromtimestamp(reset_ts, tz=timezone.utc)
|
||||
|
||||
async def get_scope_count(self, scope_key: str) -> int:
|
||||
await self.initialize()
|
||||
|
||||
if self._redis is None:
|
||||
async with self._memory_lock:
|
||||
bucket = self._get_rpm_bucket()
|
||||
self._cleanup_expired_memory_counts(bucket)
|
||||
return self._get_memory_count(scope_key)
|
||||
|
||||
try:
|
||||
result = await self._redis.get(scope_key)
|
||||
return int(result) if result else 0
|
||||
except Exception as exc:
|
||||
logger.warning("读取 RPM 计数失败,回退内存模式: {}", exc)
|
||||
if config.rate_limit_fail_open:
|
||||
return 0
|
||||
async with self._memory_lock:
|
||||
bucket = self._get_rpm_bucket()
|
||||
self._cleanup_expired_memory_counts(bucket)
|
||||
return self._get_memory_count(scope_key)
|
||||
|
||||
async def check_and_consume(
|
||||
self,
|
||||
*,
|
||||
user_rpm_key: str,
|
||||
user_rpm_limit: int,
|
||||
key_rpm_key: str,
|
||||
key_rpm_limit: int,
|
||||
) -> RpmCheckResult:
|
||||
"""原子检查并消费两层 RPM 配额。"""
|
||||
|
||||
await self.initialize()
|
||||
|
||||
normalized_user_limit = max(int(user_rpm_limit or 0), 0)
|
||||
normalized_key_limit = max(int(key_rpm_limit or 0), 0)
|
||||
|
||||
if normalized_user_limit <= 0 and normalized_key_limit <= 0:
|
||||
return RpmCheckResult(allowed=True)
|
||||
|
||||
if self._redis is None:
|
||||
return await self._check_and_consume_memory(
|
||||
user_rpm_key=user_rpm_key,
|
||||
user_rpm_limit=normalized_user_limit,
|
||||
key_rpm_key=key_rpm_key,
|
||||
key_rpm_limit=normalized_key_limit,
|
||||
)
|
||||
|
||||
retry_after = self.get_retry_after()
|
||||
|
||||
try:
|
||||
raw_result = await self._redis.eval(
|
||||
self._CHECK_AND_CONSUME_SCRIPT,
|
||||
2,
|
||||
user_rpm_key,
|
||||
key_rpm_key,
|
||||
normalized_user_limit,
|
||||
normalized_key_limit,
|
||||
self._key_ttl_seconds,
|
||||
retry_after,
|
||||
)
|
||||
return self._parse_redis_result(raw_result)
|
||||
except Exception as exc:
|
||||
logger.warning("Redis RPM 检查失败: {}", exc)
|
||||
if config.rate_limit_fail_open:
|
||||
return RpmCheckResult(allowed=True)
|
||||
return await self._check_and_consume_memory(
|
||||
user_rpm_key=user_rpm_key,
|
||||
user_rpm_limit=normalized_user_limit,
|
||||
key_rpm_key=key_rpm_key,
|
||||
key_rpm_limit=normalized_key_limit,
|
||||
)
|
||||
|
||||
def _parse_redis_result(self, raw_result: object) -> RpmCheckResult:
|
||||
values = list(raw_result) if isinstance(raw_result, (list, tuple)) else [raw_result]
|
||||
allowed = int(values[0]) == 1
|
||||
scope_code = int(values[1]) if len(values) > 1 else 0
|
||||
limit = int(values[2]) if len(values) > 2 and values[2] is not None else None
|
||||
remaining = int(values[3]) if len(values) > 3 and values[3] is not None else None
|
||||
retry_after = int(values[4]) if len(values) > 4 and values[4] is not None else None
|
||||
scope = {1: "user", 2: "key"}.get(scope_code)
|
||||
return RpmCheckResult(
|
||||
allowed=allowed,
|
||||
scope=scope,
|
||||
limit=limit,
|
||||
remaining=remaining,
|
||||
retry_after=retry_after,
|
||||
)
|
||||
|
||||
async def _check_and_consume_memory(
|
||||
self,
|
||||
*,
|
||||
user_rpm_key: str,
|
||||
user_rpm_limit: int,
|
||||
key_rpm_key: str,
|
||||
key_rpm_limit: int,
|
||||
) -> RpmCheckResult:
|
||||
async with self._memory_lock:
|
||||
bucket = self._get_rpm_bucket()
|
||||
self._cleanup_expired_memory_counts(bucket)
|
||||
|
||||
user_count = self._get_memory_count(user_rpm_key)
|
||||
if user_rpm_limit > 0 and user_count >= user_rpm_limit:
|
||||
return RpmCheckResult(
|
||||
allowed=False,
|
||||
scope="user",
|
||||
limit=user_rpm_limit,
|
||||
remaining=0,
|
||||
retry_after=self.get_retry_after(),
|
||||
)
|
||||
|
||||
key_count = self._get_memory_count(key_rpm_key)
|
||||
if key_rpm_limit > 0 and key_count >= key_rpm_limit:
|
||||
return RpmCheckResult(
|
||||
allowed=False,
|
||||
scope="key",
|
||||
limit=key_rpm_limit,
|
||||
remaining=0,
|
||||
retry_after=self.get_retry_after(),
|
||||
)
|
||||
|
||||
remaining_candidates: list[int] = []
|
||||
|
||||
if user_rpm_limit > 0:
|
||||
user_count += 1
|
||||
self._set_memory_count(user_rpm_key, user_count)
|
||||
remaining_candidates.append(user_rpm_limit - user_count)
|
||||
|
||||
if key_rpm_limit > 0:
|
||||
key_count += 1
|
||||
self._set_memory_count(key_rpm_key, key_count)
|
||||
remaining_candidates.append(key_rpm_limit - key_count)
|
||||
|
||||
remaining = min(remaining_candidates) if remaining_candidates else None
|
||||
return RpmCheckResult(allowed=True, remaining=remaining)
|
||||
|
||||
def _start_background_cleanup(self) -> None:
|
||||
if self._cleanup_task is not None:
|
||||
return
|
||||
|
||||
async def cleanup_loop() -> None:
|
||||
while True:
|
||||
try:
|
||||
await asyncio.sleep(self._bucket_seconds)
|
||||
async with self._memory_lock:
|
||||
self._cleanup_expired_memory_counts(self._get_rpm_bucket())
|
||||
except asyncio.CancelledError:
|
||||
break
|
||||
except Exception as exc:
|
||||
logger.debug("UserRpmLimiter 后台清理异常: {}", exc)
|
||||
|
||||
try:
|
||||
self._cleanup_task = asyncio.create_task(cleanup_loop())
|
||||
except RuntimeError:
|
||||
self._cleanup_task = None
|
||||
|
||||
def _get_rpm_bucket(self, now_ts: float | None = None) -> int:
|
||||
ts = now_ts if now_ts is not None else time.time()
|
||||
return int(ts // self._bucket_seconds)
|
||||
|
||||
def _split_scope_key(self, scope_key: str) -> tuple[str, int]:
|
||||
base_key, bucket_str = scope_key.rsplit(":", 1)
|
||||
return base_key, int(bucket_str)
|
||||
|
||||
def _get_memory_count(self, scope_key: str) -> int:
|
||||
base_key, bucket = self._split_scope_key(scope_key)
|
||||
stored = self._memory_counts.get(base_key)
|
||||
if not stored:
|
||||
return 0
|
||||
stored_bucket, count = stored
|
||||
if stored_bucket != bucket:
|
||||
self._memory_counts.pop(base_key, None)
|
||||
return 0
|
||||
return count
|
||||
|
||||
def _set_memory_count(self, scope_key: str, count: int) -> None:
|
||||
base_key, bucket = self._split_scope_key(scope_key)
|
||||
self._memory_counts[base_key] = (bucket, count)
|
||||
|
||||
def _cleanup_expired_memory_counts(self, current_bucket: int) -> None:
|
||||
expired_keys = [
|
||||
base_key
|
||||
for base_key, (bucket, _count) in self._memory_counts.items()
|
||||
if bucket < current_bucket
|
||||
]
|
||||
for base_key in expired_keys:
|
||||
self._memory_counts.pop(base_key, None)
|
||||
|
||||
if expired_keys:
|
||||
logger.debug(
|
||||
"[CLEANUP] 清理了 {} 个过期的用户/API Key RPM 计数(interval={}s)",
|
||||
len(expired_keys),
|
||||
self._cleanup_interval_seconds,
|
||||
)
|
||||
|
||||
|
||||
_user_rpm_limiter: UserRpmLimiter | None = None
|
||||
|
||||
|
||||
async def get_user_rpm_limiter() -> UserRpmLimiter:
|
||||
global _user_rpm_limiter
|
||||
if _user_rpm_limiter is None:
|
||||
_user_rpm_limiter = UserRpmLimiter()
|
||||
await _user_rpm_limiter.initialize()
|
||||
return _user_rpm_limiter
|
||||
Reference in New Issue
Block a user