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:
fawney19
2026-04-03 16:26:16 +08:00
parent 8f26e1a31f
commit 1d9c77522a
868 changed files with 1735 additions and 2433 deletions

View 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",
]

View 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

View 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:
当前 confidence0.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

View 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

View 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 <= 30Provider 明确告知还有配额但需要等待)
# B. 弱判断:只有 retry_after <= 5 且缺少 remaining 头(短等待时间是并发限制的典型特征)
#
# 选择保守的 retry_after 阈值:
# - 强判断用 30 秒(有 remaining 头时)
# - 弱判断用 5 秒(无 remaining 头时,更保守)
is_likely_concurrent = False
concurrent_reason = ""
# 条件 Aremaining > 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)

View 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()

View 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