refactor: 拆分调度器为独立子模块,增强 Sub2API 多认证方式支持

调度器重构:
- 将 CacheAwareScheduler 拆分为 candidate_builder、candidate_sorter、
  concurrency_checker、restriction_checker、scheduling_config、schemas、utils 等独立模块
- 删除旧的 _candidate_builder.py 和 _candidate_sorter.py
- 新增调度并发拒绝 Prometheus 指标

Sub2API 架构增强:
- 支持账号密码登录和 Refresh Token 两种认证方式
- 实现 JWT 自动刷新和 Token Rotation 持久化
- 前端 ProviderAuthDialog 支持多认证方式切换和 credentials_schema 动态渲染
- 验证接口返回 updated_credentials 以同步轮换后的 token

其他改进:
- 并发管理器增加 RPM guard 和动态预留逻辑
- RequestCandidate 支持 mark_skipped 附加 extra_data
- TaskService 增强健壮性
- 补充相关单元测试和契约测试
This commit is contained in:
fawney19
2026-02-15 16:32:23 +08:00
parent 8a670f5524
commit 1c16b77a92
27 changed files with 2159 additions and 645 deletions

View File

@@ -30,10 +30,7 @@
from __future__ import annotations
import hashlib
import math
import time
from dataclasses import dataclass
from typing import Any
from sqlalchemy.orm import Session
@@ -43,7 +40,6 @@ from src.core.logger import logger
from src.core.model_permissions import (
check_model_allowed,
get_allowed_models_preview,
merge_allowed_models,
)
from src.models.database import (
ApiKey,
@@ -51,127 +47,55 @@ from src.models.database import (
ProviderAPIKey,
ProviderEndpoint,
)
from src.services.cache._candidate_builder import (
CandidateBuilder,
)
from src.services.cache._candidate_builder import (
_sort_endpoints_by_family_priority as _sort_endpoints_by_family_priority,
)
from src.services.cache._candidate_sorter import CandidateSorter
from src.services.cache.affinity_manager import (
CacheAffinityManager,
get_affinity_manager,
)
from src.services.cache.candidate_builder import (
CandidateBuilder,
)
from src.services.cache.candidate_builder import (
_sort_endpoints_by_family_priority as _sort_endpoints_by_family_priority,
)
from src.services.cache.candidate_sorter import CandidateSorter
from src.services.cache.concurrency_checker import ConcurrencyChecker
from src.services.cache.model_cache import ModelCacheService
from src.services.cache.restriction_checker import get_effective_restrictions
from src.services.cache.scheduling_config import SchedulingConfig
from src.services.cache.schemas import ConcurrencySnapshot as ConcurrencySnapshot # re-export
from src.services.cache.schemas import ProviderCandidate as ProviderCandidate # re-export
from src.services.cache.utils import affinity_hash as _affinity_hash # re-export compat
from src.services.cache.utils import (
release_db_connection_before_await,
)
from src.services.provider.format import normalize_endpoint_signature
from src.services.rate_limit.adaptive_reservation import (
AdaptiveReservationManager,
get_adaptive_reservation_manager,
)
from src.services.rate_limit.adaptive_rpm import get_adaptive_rpm_manager
from src.services.rate_limit.concurrency_manager import get_concurrency_manager
from src.services.system.config import SystemConfigService
@dataclass
class ProviderCandidate:
"""候选 provider 组合及是否命中缓存"""
provider: Provider
endpoint: ProviderEndpoint
key: ProviderAPIKey
is_cached: bool = False
is_skipped: bool = False # 是否被跳过
skip_reason: str | None = None # 跳过原因
mapping_matched_model: str | None = None # 通过映射匹配到的模型名(用于实际请求)
needs_conversion: bool = False # 是否需要格式转换
provider_api_format: str = "" # Provider 端点实际格式(用于健康度/熔断 bucket
def _stable_order_key(self) -> tuple[int, int, str, str, str]:
"""
为排序/优先队列提供稳定的比较键。
说明:
- 运行时偶发会出现对 ProviderCandidate 做 tuple 排序/heap 排序的场景;
当主键相同需要比较候选本身时,若候选不可比较会触发:
TypeError: '<' not supported between instances of 'ProviderCandidate' and 'ProviderCandidate'
- 这里提供一个与调度逻辑无关、但足够稳定且可比的兜底顺序。
"""
provider_priority_raw = getattr(self.provider, "provider_priority", None)
internal_priority_raw = getattr(self.key, "internal_priority", None)
try:
provider_priority = (
int(provider_priority_raw) if provider_priority_raw is not None else 999999
)
except Exception:
provider_priority = 999999
try:
internal_priority = (
int(internal_priority_raw) if internal_priority_raw is not None else 999999
)
except Exception:
internal_priority = 999999
provider_id = str(getattr(self.provider, "id", "") or "")
endpoint_id = str(getattr(self.endpoint, "id", "") or "")
key_id = str(getattr(self.key, "id", "") or "")
return (provider_priority, internal_priority, provider_id, endpoint_id, key_id)
def __lt__(self, other: object) -> bool:
if not isinstance(other, ProviderCandidate):
return NotImplemented
return self._stable_order_key() < other._stable_order_key()
@dataclass
class ConcurrencySnapshot:
key_current: int
key_limit: int | None
is_cached_user: bool = False
# 动态预留信息
reservation_ratio: float = 0.0
reservation_phase: str = "unknown"
reservation_confidence: float = 0.0
def describe(self) -> str:
key_limit_text = str(self.key_limit) if self.key_limit is not None else "inf"
reservation_text = f"{self.reservation_ratio:.0%}" if self.reservation_ratio > 0 else "N/A"
return (
f"key={self.key_current}/{key_limit_text}, "
f"cached={self.is_cached_user}, "
f"reserve={reservation_text}({self.reservation_phase})"
)
class CacheAwareScheduler:
"""
缓存感知调度器
缓存感知调度器 - 薄协调层
这是Provider选择的核心组件整合了:
- Provider/Endpoint/Key三层架构
- 缓存亲和性管理
- 并发控制(动态预留机制
- 健康度监控
编排以下子组件:
- SchedulingConfig: 调度模式和优先级模式管理
- CandidateBuilder: 候选构建(查询 Provider/Endpoint/Key
- CandidateSorter: 候选排序(优先级/负载均衡
- ConcurrencyChecker: 并发控制RPM + 动态预留)
- CacheAffinityManager: 缓存亲和性管理
"""
# 优先级模式常量
PRIORITY_MODE_PROVIDER = "provider" # 提供商优先模式
PRIORITY_MODE_GLOBAL_KEY = "global_key" # 全局 Key 优先模式
ALLOWED_PRIORITY_MODES = {
PRIORITY_MODE_PROVIDER,
PRIORITY_MODE_GLOBAL_KEY,
}
# 调度模式常量
SCHEDULING_MODE_FIXED_ORDER = "fixed_order" # 固定顺序模式:严格按优先级,忽略缓存
SCHEDULING_MODE_CACHE_AFFINITY = "cache_affinity" # 缓存亲和模式:优先缓存,同优先级哈希分散
SCHEDULING_MODE_LOAD_BALANCE = "load_balance" # 负载均衡模式:忽略缓存,同优先级随机轮换
ALLOWED_SCHEDULING_MODES = {
SCHEDULING_MODE_FIXED_ORDER,
SCHEDULING_MODE_CACHE_AFFINITY,
SCHEDULING_MODE_LOAD_BALANCE,
}
# 类常量 re-export保持外部访问兼容性
PRIORITY_MODE_PROVIDER = SchedulingConfig.PRIORITY_MODE_PROVIDER
PRIORITY_MODE_GLOBAL_KEY = SchedulingConfig.PRIORITY_MODE_GLOBAL_KEY
ALLOWED_PRIORITY_MODES = SchedulingConfig.ALLOWED_PRIORITY_MODES
SCHEDULING_MODE_FIXED_ORDER = SchedulingConfig.SCHEDULING_MODE_FIXED_ORDER
SCHEDULING_MODE_CACHE_AFFINITY = SchedulingConfig.SCHEDULING_MODE_CACHE_AFFINITY
SCHEDULING_MODE_LOAD_BALANCE = SchedulingConfig.SCHEDULING_MODE_LOAD_BALANCE
ALLOWED_SCHEDULING_MODES = SchedulingConfig.ALLOWED_SCHEDULING_MODES
def __init__(
self,
@@ -191,22 +115,13 @@ class CacheAwareScheduler:
scheduling_mode: 调度模式fixed_order | cache_affinity
"""
self.redis = redis_client
self.priority_mode = self._normalize_priority_mode(
priority_mode or self.PRIORITY_MODE_PROVIDER
)
self.scheduling_mode = self._normalize_scheduling_mode(
scheduling_mode or self.SCHEDULING_MODE_CACHE_AFFINITY
)
logger.debug(
f"[CacheAwareScheduler] 初始化优先级模式: {self.priority_mode}, 调度模式: {self.scheduling_mode}"
)
self._config = SchedulingConfig(priority_mode, scheduling_mode)
# 初始化子组件(将在第一次使用时异步初始化)
# 异步子组件(将在第一次使用时初始化)
self._affinity_manager: CacheAffinityManager | None = None
self._concurrency_manager = None
# 动态预留管理器(同步初始化)
self._reservation_manager: AdaptiveReservationManager = get_adaptive_reservation_manager()
self._metrics = {
self._concurrency_checker: ConcurrencyChecker | None = None
self._metrics: dict[str, Any] = {
"total_batches": 0,
"last_batch_size": 0,
"total_candidates": 0,
@@ -224,54 +139,62 @@ class CacheAwareScheduler:
"last_reservation_result": None,
}
# 初始化拆分出的子模块
self._candidate_builder = CandidateBuilder(self)
self._candidate_sorter = CandidateSorter(self)
# 初始化子模块(不传 self解除反向引用
self._candidate_sorter = CandidateSorter(self._config)
self._candidate_builder = CandidateBuilder(self._candidate_sorter)
# ── 属性代理(保持外部访问兼容性)──────────────────────────
@property
def priority_mode(self) -> str:
return self._config.priority_mode
@priority_mode.setter
def priority_mode(self, value: str) -> None:
self._config.priority_mode = value
@property
def scheduling_mode(self) -> str:
return self._config.scheduling_mode
@scheduling_mode.setter
def scheduling_mode(self, value: str) -> None:
self._config.scheduling_mode = value
def set_priority_mode(self, mode: str | None) -> None:
"""运行时更新候选排序策略"""
self._config.set_priority_mode(mode)
def set_scheduling_mode(self, mode: str | None) -> None:
"""运行时更新调度模式"""
self._config.set_scheduling_mode(mode)
# ── 静态方法兼容壳 ───────────────────────────────────────
@staticmethod
def _release_db_connection_before_await(db: Session) -> None:
"""
Best-effort: end a read-only transaction before awaiting async I/O.
release_db_connection_before_await(db)
This scheduler does a lot of async work (cache/Redis) mixed with sync SQLAlchemy reads.
If a SELECT has already started a transaction, the pooled connection can remain checked
out while we await, causing pool pressure under concurrency.
@staticmethod
def _affinity_hash(affinity_key: str, identifier: str) -> int:
return _affinity_hash(affinity_key, identifier)
Safety:
- Only commits when the Session has no ORM pending changes.
- Temporarily disables expire_on_commit to keep already-loaded ORM objects usable.
"""
try:
if db is None:
return
has_pending_changes = bool(db.new) or bool(db.dirty) or bool(db.deleted)
if has_pending_changes:
return
if not db.in_transaction():
return
original_expire_on_commit = getattr(db, "expire_on_commit", True)
db.expire_on_commit = False
try:
db.commit()
except Exception:
try:
db.rollback()
except Exception:
pass
finally:
db.expire_on_commit = original_expire_on_commit
except Exception:
# Never let this optimization break scheduling
return
# ── 异步初始化 ───────────────────────────────────────────
async def _ensure_initialized(self) -> None:
"""确保所有异步组件已初始化"""
if self._affinity_manager is None:
self._affinity_manager = await get_affinity_manager(self.redis)
if self._concurrency_manager is None:
self._concurrency_manager = await get_concurrency_manager()
if self._concurrency_checker is None:
concurrency_manager = await get_concurrency_manager()
reservation_manager = get_adaptive_reservation_manager()
self._concurrency_checker = ConcurrencyChecker(
concurrency_manager=concurrency_manager,
reservation_manager=reservation_manager,
)
# ── 核心编排方法 ─────────────────────────────────────────
async def select_with_cache_affinity(
self,
@@ -308,8 +231,11 @@ class CacheAwareScheduler:
normalized_format = normalize_endpoint_signature(api_format)
logger.debug(
f"[CacheAwareScheduler] select_with_cache_affinity: "
f"affinity_key={affinity_key[:8]}..., api_format={normalized_format}, model={model_name}"
"[CacheAwareScheduler] select_with_cache_affinity: "
"affinity_key={}..., api_format={}, model={}",
affinity_key[:8],
normalized_format,
model_name,
)
self._metrics["last_api_format"] = normalized_format
@@ -337,7 +263,6 @@ class CacheAwareScheduler:
if provider_batch_count == 0:
if provider_offset == 0:
# 没有找到任何候选,提供友好的错误提示(不暴露内部信息)
raise ProviderNotAvailableException("请求的模型当前不可用")
break
@@ -351,28 +276,35 @@ class CacheAwareScheduler:
key = candidate.key
if endpoint.id in excluded_endpoints_set:
logger.debug(f" └─ Endpoint {endpoint.id[:8]}... 在排除列表,跳过")
logger.debug(" └─ Endpoint {}... 在排除列表,跳过", endpoint.id[:8])
continue
if key.id in excluded_keys_set:
logger.debug(f" └─ Key {key.id[:8]}... 在排除列表,跳过")
logger.debug(" └─ Key {}... 在排除列表,跳过", key.id[:8])
continue
is_cached_user = bool(candidate.is_cached)
can_use, snapshot = await self._check_concurrent_available(
can_use, snapshot = await self._concurrency_checker.check_available(
key,
is_cached_user=is_cached_user,
)
# 更新预留指标
self._update_reservation_metrics(snapshot)
if not can_use:
logger.debug(f" └─ Key {key.id[:8]}... 并发已满 ({snapshot.describe()})")
logger.debug(" └─ Key {}... 并发已满 ({})", key.id[:8], snapshot.describe())
self._metrics["concurrency_denied"] += 1
continue
logger.debug(
f" └─ 选择 Provider={provider.name}, Endpoint={endpoint.id[:8]}..., "
f"Key=***{key.api_key[-4:]}, 缓存命中={is_cached_user}, "
f"并发状态[{snapshot.describe()}]"
" └─ 选择 Provider={}, Endpoint={}..., "
"Key=***{}, 缓存命中={}, 并发状态[{}]",
provider.name,
endpoint.id[:8],
key.api_key[-4:],
is_cached_user,
snapshot.describe(),
)
if key.cache_ttl_minutes > 0 and global_model_id:
@@ -400,194 +332,6 @@ class CacheAwareScheduler:
raise ProviderNotAvailableException("服务暂时繁忙,请稍后重试")
def _get_effective_rpm_limit(self, key: ProviderAPIKey) -> int | None:
"""获取有效的 RPM 限制(委托给 AdaptiveRPMManager 统一逻辑)"""
return get_adaptive_rpm_manager().get_effective_limit(key)
async def _check_concurrent_available(
self,
key: ProviderAPIKey,
is_cached_user: bool = False,
) -> tuple[bool, ConcurrencySnapshot]:
"""
检查 RPM 限制是否可用(使用动态预留机制)
核心逻辑 - 动态缓存预留机制:
- 总槽位: 有效 RPM 限制(固定值或学习到的值)
- 预留比例: 由 AdaptiveReservationManager 根据置信度和负载动态计算
- 缓存用户可用: 全部槽位
- 新用户可用: 总槽位 x (1 - 动态预留比例)
Args:
key: ProviderAPIKey对象
is_cached_user: 是否是缓存用户
Returns:
(是否可用, 并发快照)
"""
# 获取有效的并发限制
effective_key_limit = self._get_effective_rpm_limit(key)
logger.debug(
f" -> 并发检查: _concurrency_manager={self._concurrency_manager is not None}, "
f"is_cached_user={is_cached_user}, effective_limit={effective_key_limit}"
)
if not self._concurrency_manager:
# 并发管理器不可用直接返回True
logger.debug(f" -> 无并发管理器,直接通过")
snapshot = ConcurrencySnapshot(
key_current=0,
key_limit=effective_key_limit,
is_cached_user=is_cached_user,
)
return True, snapshot
# 获取当前 RPM 计数
key_count = await self._concurrency_manager.get_key_rpm_count(
key_id=str(key.id),
)
can_use = True
# 计算动态预留比例
reservation_result = self._reservation_manager.calculate_reservation(
key=key,
current_usage=key_count,
effective_limit=effective_key_limit,
)
# 更新指标
if reservation_result.phase == "probe":
self._metrics["reservation_probe_count"] += 1
else:
self._metrics["reservation_stable_count"] += 1
# 计算移动平均预留比例
total_reservations = (
self._metrics["reservation_probe_count"] + self._metrics["reservation_stable_count"]
)
if total_reservations > 0:
# 指数移动平均
alpha = 0.1
self._metrics["avg_reservation_ratio"] = (
alpha * reservation_result.ratio
+ (1 - alpha) * self._metrics["avg_reservation_ratio"]
)
self._metrics["last_reservation_result"] = {
"ratio": reservation_result.ratio,
"phase": reservation_result.phase,
"confidence": reservation_result.confidence,
"load_factor": reservation_result.load_factor,
}
available_for_new = None
reservation_ratio = reservation_result.ratio
# 检查Key级别限制使用动态预留比例
if effective_key_limit is not None:
if is_cached_user:
# 缓存用户: 可以使用全部槽位
if key_count >= effective_key_limit:
can_use = False
else:
# 新用户: 只能使用 (1 - 动态预留比例) 的槽位
# 使用 max 确保至少有 1 个槽位可用
# 与 ConcurrencyManager 的 Lua 脚本保持一致:使用 floor 计算新用户可用槽位
available_for_new = max(
1, math.floor(effective_key_limit * (1 - reservation_ratio))
)
if key_count >= available_for_new:
logger.debug(
f"Key {key.id[:8]}... 新用户配额已满 "
f"({key_count}/{available_for_new}, 总{effective_key_limit}, "
f"预留{reservation_ratio:.0%}[{reservation_result.phase}])"
)
can_use = False
key_limit_for_snapshot: int | None
if is_cached_user:
key_limit_for_snapshot = effective_key_limit
elif effective_key_limit is not None:
key_limit_for_snapshot = (
available_for_new if available_for_new is not None else effective_key_limit
)
else:
key_limit_for_snapshot = None
snapshot = ConcurrencySnapshot(
key_current=key_count,
key_limit=key_limit_for_snapshot,
is_cached_user=is_cached_user,
reservation_ratio=reservation_ratio,
reservation_phase=reservation_result.phase,
reservation_confidence=reservation_result.confidence,
)
return can_use, snapshot
def _get_effective_restrictions(
self,
user_api_key: ApiKey | None,
) -> dict[str, Any]:
"""
获取有效的访问限制(合并 ApiKey 和 User 的限制)
逻辑:
- 如果 ApiKey 和 User 都有限制,取交集
- 如果只有一方有限制,使用该方的限制
- 如果都没有限制,返回 None表示不限制
Args:
user_api_key: 用户 API Key 对象(可能包含 user relationship
Returns:
包含 allowed_providers, allowed_models, allowed_api_formats 的字典
"""
result = {
"allowed_providers": None,
"allowed_models": None,
"allowed_api_formats": None,
}
if not user_api_key:
return result
# 获取 User 的限制
# 注意:这里可能触发 lazy loading需要确保 session 仍然有效
try:
user = user_api_key.user if hasattr(user_api_key, "user") else None
except Exception as e:
logger.warning(f"无法加载 ApiKey 关联的 User: {e},仅使用 ApiKey 级别的限制")
user = None
# 调试日志
logger.debug(
f"[_get_effective_restrictions] ApiKey={user_api_key.id[:8]}..., "
f"User={user.id[:8] if user else 'None'}..., "
f"ApiKey.allowed_models={user_api_key.allowed_models}, "
f"User.allowed_models={user.allowed_models if user else 'N/A'}"
)
# 合并 allowed_providers
result["allowed_providers"] = self._merge_restriction_sets(
user_api_key.allowed_providers, user.allowed_providers if user else None
)
# 合并 allowed_models取交集
result["allowed_models"] = merge_allowed_models(
user_api_key.allowed_models, user.allowed_models if user else None
)
# 合并 allowed_api_formats
result["allowed_api_formats"] = self._merge_restriction_sets(
user_api_key.allowed_api_formats, user.allowed_api_formats if user else None
)
return result
async def list_all_candidates(
self,
db: Session,
@@ -604,10 +348,12 @@ class CacheAwareScheduler:
"""
预先获取所有可用的 Provider/Endpoint/Key 组合
重构后的方法将逻辑拆分为
1. _query_providers: 数据库查询逻辑(委托给 CandidateBuilder
2. _build_candidates: 候选构建逻辑(委托给 CandidateBuilder
3. _apply_cache_affinity: 缓存亲和性处理
编排流程
1. 解析 GlobalModel
2. 检查访问限制
3. 查询 Providers委托给 CandidateBuilder
4. 构建候选列表(委托给 CandidateBuilder
5. 应用排序和缓存亲和性
Args:
db: 数据库会话
@@ -627,7 +373,7 @@ class CacheAwareScheduler:
- provider_batch_count 表示本次查询到的 Provider 数量(未应用 allowed_providers 过滤前)
"""
# If the caller already touched the DB, release the connection before we do async work.
self._release_db_connection_before_await(db)
release_db_connection_before_await(db)
await self._ensure_initialized()
target_format = normalize_endpoint_signature(api_format)
@@ -646,7 +392,7 @@ class CacheAwareScheduler:
global_model = await ModelCacheService.get_global_model_by_name(db, normalized_name)
if not global_model or not global_model.is_active:
logger.warning(f"GlobalModel not found or inactive: {normalized_name}")
logger.warning("GlobalModel not found or inactive: {}", normalized_name)
raise ModelNotSupportedException(model=model_name)
logger.debug(
@@ -664,11 +410,13 @@ class CacheAwareScheduler:
model_mappings: list[str] = (global_model.config or {}).get("model_mappings", [])
if model_mappings:
logger.debug(
f"[Scheduler] GlobalModel={global_model.name} 配置了映射规则: {model_mappings}"
"[Scheduler] GlobalModel={} 配置了映射规则: {}",
global_model.name,
model_mappings,
)
# 获取合并后的访问限制ApiKey + User
restrictions = self._get_effective_restrictions(user_api_key)
restrictions = get_effective_restrictions(user_api_key)
allowed_api_formats = restrictions["allowed_api_formats"]
allowed_providers = restrictions["allowed_providers"]
allowed_models = restrictions["allowed_models"]
@@ -678,8 +426,10 @@ class CacheAwareScheduler:
allowed_norm = {normalize_endpoint_signature(f) for f in allowed_api_formats if f}
if target_format not in allowed_norm:
logger.debug(
f"API Key {user_api_key.id[:8] if user_api_key else 'N/A'}... 不允许使用 API 格式 {target_format}, "
f"允许的格式: {allowed_api_formats}"
"API Key {}... 不允许使用 API 格式 {}, 允许的格式: {}",
user_api_key.id[:8] if user_api_key else "N/A",
target_format,
allowed_api_formats,
)
return [], global_model_id, queried_provider_count
@@ -689,8 +439,9 @@ class CacheAwareScheduler:
allowed_models=allowed_models,
):
logger.debug(
f"用户/API Key 不允许使用模型 {model_name}, "
f"允许的模型: {get_allowed_models_preview(allowed_models)}"
"用户/API Key 不允许使用模型 {}, 允许的模型: {}",
model_name,
get_allowed_models_preview(allowed_models),
)
return [], global_model_id, queried_provider_count
@@ -703,7 +454,7 @@ class CacheAwareScheduler:
queried_provider_count = len(providers)
# Provider query starts a transaction; release connection before entering async candidate build.
self._release_db_connection_before_await(db)
release_db_connection_before_await(db)
logger.debug(
"[Scheduler] Found {} active providers",
@@ -730,7 +481,7 @@ class CacheAwareScheduler:
p for p in providers if p.id in allowed_providers or p.name in allowed_providers
]
if original_count != len(providers):
logger.debug(f"用户/API Key 过滤 Provider: {original_count} -> {len(providers)}")
logger.debug("用户/API Key 过滤 Provider: {} -> {}", original_count, len(providers))
if not providers:
return [], global_model_id, queried_provider_count
@@ -767,8 +518,10 @@ class CacheAwareScheduler:
self._metrics["last_candidate_count"] = len(candidates)
logger.debug(
f"预先获取到 {len(candidates)} 个可用组合 "
f"(api_format={target_format}, model={model_name})"
"预先获取到 {} 个可用组合 (api_format={}, model={})",
len(candidates),
target_format,
model_name,
)
return candidates, global_model_id, queried_provider_count
@@ -889,17 +642,23 @@ class CacheAwareScheduler:
matched_candidate = candidate
matched = True
logger.debug(
f"检测到缓存亲和性: affinity_key={affinity_key[:8]}..., "
f"api_format={api_format_str}, global_model_id={global_model_id[:8]}..., "
f"provider={provider.name}, endpoint={endpoint.id[:8]}..., "
f"provider_key=***{key.api_key[-4:]}, "
f"使用次数={affinity.request_count}"
"检测到缓存亲和性: affinity_key={}..., "
"api_format={}, global_model_id={}..., "
"provider={}, endpoint={}..., "
"provider_key=***{}, 使用次数={}",
affinity_key[:8],
api_format_str,
global_model_id[:8],
provider.name,
endpoint.id[:8],
key.api_key[-4:],
affinity.request_count,
)
else:
candidate.is_cached = False
if not matched:
logger.debug(f"API格式 {api_format_str} 的缓存亲和性存在但组合不可用")
logger.debug("API格式 {} 的缓存亲和性存在但组合不可用", api_format_str)
return candidates
# 缓存亲和性命中且该候选可用(未被跳过)时,无条件优先使用
@@ -912,16 +671,16 @@ class CacheAwareScheduler:
other_candidates = [c for c in candidates if c is not matched_candidate]
result = [matched_candidate] + other_candidates
logger.debug(
f"缓存亲和性命中且健康,无条件优先使用 "
f"(needs_conversion={matched_candidate.needs_conversion})"
"缓存亲和性命中且健康,无条件优先使用 (needs_conversion={})",
matched_candidate.needs_conversion,
)
return result
# 缓存命中但被跳过(不健康),按 exact 优先排序
# 缓存候选在其所属类别内提升到最前面
logger.debug(
f"缓存亲和性命中但不健康 (skip_reason={matched_candidate.skip_reason})"
f"按 exact 优先排序"
"缓存亲和性命中但不健康 (skip_reason={}),按 exact 优先排序",
matched_candidate.skip_reason,
)
matched_should_demote = should_demote(matched_candidate)
@@ -946,60 +705,14 @@ class CacheAwareScheduler:
keep_priority_candidates.insert(0, matched_candidate)
result = keep_priority_candidates + demote_candidates
logger.debug(f"缓存组合已提升至其类别内优先级 (demote={matched_should_demote})")
logger.debug("缓存组合已提升至其类别内优先级 (demote={})", matched_should_demote)
return result
except Exception as e:
logger.warning(f"检查缓存亲和性失败: {e},继续使用默认排序")
logger.warning("检查缓存亲和性失败: {},继续使用默认排序", e)
return candidates
@staticmethod
def _affinity_hash(affinity_key: str, identifier: str) -> int:
"""基于 affinity_key 和标识符的确定性哈希(用于同优先级内分散负载均衡)"""
return int(hashlib.sha256(f"{affinity_key}:{identifier}".encode()).hexdigest()[:16], 16)
@staticmethod
def _merge_restriction_sets(key_restriction: Any, user_restriction: Any) -> set[Any] | None:
"""合并两个限制列表,取交集;任一方为空则使用另一方;均空返回 None"""
key_set = set(key_restriction) if key_restriction else None
user_set = set(user_restriction) if user_restriction else None
if key_set and user_set:
return key_set & user_set
return key_set or user_set
def _normalize_priority_mode(self, mode: str | None) -> str:
normalized = (mode or "").strip().lower()
if normalized not in self.ALLOWED_PRIORITY_MODES:
if normalized:
logger.warning(f"[CacheAwareScheduler] 无效的优先级模式 '{mode}',回退为 provider")
return self.PRIORITY_MODE_PROVIDER
return normalized
def set_priority_mode(self, mode: str | None) -> None:
"""运行时更新候选排序策略"""
normalized = self._normalize_priority_mode(mode)
if normalized == self.priority_mode:
return
self.priority_mode = normalized
logger.debug(f"[CacheAwareScheduler] 切换优先级模式为: {self.priority_mode}")
def _normalize_scheduling_mode(self, mode: str | None) -> str:
normalized = (mode or "").strip().lower()
if normalized not in self.ALLOWED_SCHEDULING_MODES:
if normalized:
logger.warning(
f"[CacheAwareScheduler] 无效的调度模式 '{mode}',回退为 cache_affinity"
)
return self.SCHEDULING_MODE_CACHE_AFFINITY
return normalized
def set_scheduling_mode(self, mode: str | None) -> None:
"""运行时更新调度模式"""
normalized = self._normalize_scheduling_mode(mode)
if normalized == self.scheduling_mode:
return
self.scheduling_mode = normalized
logger.debug(f"[CacheAwareScheduler] 切换调度模式为: {self.scheduling_mode}")
# ── 委托方法(外部 API 兼容)──────────────────────────────
async def invalidate_cache(
self,
@@ -1068,6 +781,33 @@ class CacheAwareScheduler:
ttl=ttl,
)
# ── 指标 ────────────────────────────────────────────────
def _update_reservation_metrics(self, snapshot: ConcurrencySnapshot) -> None:
"""根据并发检查结果更新预留相关指标"""
if snapshot.reservation_phase == "probe":
self._metrics["reservation_probe_count"] += 1
elif snapshot.reservation_phase != "unknown":
self._metrics["reservation_stable_count"] += 1
# 计算移动平均预留比例
total_reservations = (
self._metrics["reservation_probe_count"] + self._metrics["reservation_stable_count"]
)
if total_reservations > 0:
alpha = 0.1
self._metrics["avg_reservation_ratio"] = (
alpha * snapshot.reservation_ratio
+ (1 - alpha) * self._metrics["avg_reservation_ratio"]
)
self._metrics["last_reservation_result"] = {
"ratio": snapshot.reservation_ratio,
"phase": snapshot.reservation_phase,
"confidence": snapshot.reservation_confidence,
"load_factor": snapshot.load_factor,
}
async def get_stats(self) -> dict:
"""获取调度器统计信息"""
await self._ensure_initialized()
@@ -1084,7 +824,7 @@ class CacheAwareScheduler:
)
# 动态预留统计
reservation_stats = self._reservation_manager.get_stats()
reservation_stats = self._concurrency_checker.get_reservation_stats()
total_reservation_checks = (
metrics["reservation_probe_count"] + metrics["reservation_stable_count"]
)

View File

@@ -29,12 +29,14 @@ from src.models.database import (
ProviderEndpoint,
)
from src.services.cache.quota_skipper import is_key_quota_exhausted
from src.services.cache.utils import release_db_connection_before_await
from src.services.health.monitor import health_monitor
from src.services.provider.format import normalize_endpoint_signature
if TYPE_CHECKING:
from src.models.database import GlobalModel
from src.services.cache.aware_scheduler import CacheAwareScheduler, ProviderCandidate
from src.services.cache.candidate_sorter import CandidateSorter
from src.services.cache.schemas import ProviderCandidate
from src.services.cache.model_cache import ModelCacheService
@@ -58,8 +60,8 @@ def _sort_endpoints_by_family_priority(
class CandidateBuilder:
"""候选构建器,负责查询 Provider、检查模型支持和 Key 可用性、构建候选列表。"""
def __init__(self, scheduler: CacheAwareScheduler) -> None:
self._scheduler = scheduler
def __init__(self, candidate_sorter: CandidateSorter) -> None:
self._sorter = candidate_sorter
def _query_providers(
self,
@@ -132,7 +134,7 @@ class CandidateBuilder:
- provider_model_names: Provider 侧可用的模型名称集合主名称 + 映射名称 api_format 过滤
"""
# Avoid holding a DB connection while awaiting cache/Redis inside ModelCacheService.
self._scheduler._release_db_connection_before_await(db)
release_db_connection_before_await(db)
# 仅接受 GlobalModel.name不允许映射名
normalized_name = model_name.strip() if isinstance(model_name, str) else ""
@@ -373,12 +375,12 @@ class CandidateBuilder:
max_candidates: 最大候选数
is_stream: 是否是流式请求如果为 True 则过滤不支持流式的 Provider
capability_requirements: 能力需求可选
global_conversion_enabled: 格式转换开关数据库配置关闭时禁止任何跨格式转换
global_conversion_enabled: 格式转换全局开关数据库配置关闭时回退到 Provider/Endpoint 精细化配置
Returns:
候选列表
"""
from src.services.cache.aware_scheduler import ProviderCandidate
from src.services.cache.schemas import ProviderCandidate
candidates: list[ProviderCandidate] = []
client_format_str = normalize_endpoint_signature(client_format)
@@ -468,14 +470,14 @@ class CandidateBuilder:
str(getattr(endpoint, "endpoint_kind", "")).strip().lower(),
)
# 计算格式转换开关状态(三层优先级)
#
# 1) 全局开关(数据库配置)关闭 -> 禁止任何跨格式转换
# 2) 全局开关开启 -> 允许跨格式转换
# 3) 提供商覆盖Provider.enable_format_conversion开启 -> 强制允许(跳过端点检查)
# 4) 否则 -> 由端点配置 format_acceptance_config 决定是否允许
provider_allows_conversion = getattr(provider, "enable_format_conversion", True)
skip_endpoint_check = global_conversion_enabled or provider_allows_conversion
# 格式转换开关(从高到低):
# 1) 全局开关 enable_format_conversion=ON -> 允许跨格式(跳过端点检查)
# 2) 全局开关 OFF -> Provider.enable_format_conversion=ON -> 允许跨格式(跳过端点检查)
# 3) 否则 -> 需 Endpoint.format_acceptance_config 显式允许
provider_conversion_enabled = bool(
getattr(provider, "enable_format_conversion", False)
)
skip_endpoint_check = global_conversion_enabled or provider_conversion_enabled
is_compatible, needs_conversion, _compat_reason = is_format_compatible(
client_format_str,
@@ -492,7 +494,7 @@ class CandidateBuilder:
endpoint_format_str,
is_compatible,
global_conversion_enabled,
provider_allows_conversion,
provider_conversion_enabled,
skip_endpoint_check,
_compat_reason,
)
@@ -521,8 +523,11 @@ class CandidateBuilder:
)
if not supports_model:
logger.debug(
f"Provider {provider.name} 端点 {endpoint_format_str} "
f"不支持模型 {model_name}: {skip_reason}"
"Provider {} 端点 {} 不支持模型 {}: {}",
provider.name,
endpoint_format_str,
model_name,
skip_reason,
)
continue
@@ -541,11 +546,13 @@ class CandidateBuilder:
use_random = all((key.cache_ttl_minutes or 0) == 0 for key in active_keys)
if use_random and len(active_keys) > 1:
logger.debug(
f" Provider {provider.name} 启用 Key 轮换模式 "
f"(endpoint_format={endpoint_format_str}, {len(active_keys)} keys)"
" Provider {} 启用 Key 轮换模式 (endpoint_format={}, {} keys)",
provider.name,
endpoint_format_str,
len(active_keys),
)
keys = self._scheduler._candidate_sorter._shuffle_keys_by_internal_priority(
keys = self._sorter.shuffle_keys_by_internal_priority(
active_keys, affinity_key, use_random
)

View File

@@ -13,21 +13,22 @@ import random
from collections import defaultdict
from typing import TYPE_CHECKING
from src.core.logger import logger
from src.services.cache.scheduling_config import SchedulingConfig
from src.services.cache.utils import affinity_hash
from src.services.system.config import SystemConfigService
if TYPE_CHECKING:
from sqlalchemy.orm import Session
from src.models.database import ProviderAPIKey
from src.services.cache.aware_scheduler import CacheAwareScheduler, ProviderCandidate
from src.services.cache.schemas import ProviderCandidate
class CandidateSorter:
"""候选排序器,负责优先级模式排序、负载均衡排序和 Key 内部打乱。"""
def __init__(self, scheduler: CacheAwareScheduler) -> None:
self._scheduler = scheduler
def __init__(self, config: SchedulingConfig) -> None:
self._config = config
def _apply_priority_mode_sort(
self,
@@ -51,14 +52,12 @@ class CandidateSorter:
if not candidates:
return candidates
s = self._scheduler
# 全局配置:如果开启,所有候选保持原优先级
global_keep_priority = SystemConfigService.is_keep_priority_on_conversion(db)
if global_keep_priority:
# 全局开启:不分组,直接按优先级模式排序
if s.priority_mode == s.PRIORITY_MODE_GLOBAL_KEY:
if self._config.priority_mode == SchedulingConfig.PRIORITY_MODE_GLOBAL_KEY:
return self._sort_by_global_priority_with_hash(candidates, affinity_key, api_format)
# 提供商优先模式:保持构建时的顺序(已按 provider_priority 排序)
return candidates
@@ -80,7 +79,7 @@ class CandidateSorter:
# convertible 且未配置保持优先级:降级
demote_candidates.append(c)
if s.priority_mode == s.PRIORITY_MODE_GLOBAL_KEY:
if self._config.priority_mode == SchedulingConfig.PRIORITY_MODE_GLOBAL_KEY:
# 全局 Key 优先模式:分别对两组排序后合并
sorted_keep = self._sort_by_global_priority_with_hash(
keep_priority_candidates, affinity_key, api_format
@@ -132,7 +131,7 @@ class CandidateSorter:
scored_candidates = []
for candidate in group:
key_id = candidate.key.id if candidate.key else ""
hash_value = self._scheduler._affinity_hash(affinity_key, key_id)
hash_value = affinity_hash(affinity_key, key_id)
scored_candidates.append((hash_value, candidate))
# 按哈希值排序
@@ -167,11 +166,10 @@ class CandidateSorter:
if not candidates:
return candidates
s = self._scheduler
priority_groups: dict[tuple, list[ProviderCandidate]] = defaultdict(list)
# 根据优先级模式选择分组方式
if s.priority_mode == s.PRIORITY_MODE_GLOBAL_KEY:
if self._config.priority_mode == SchedulingConfig.PRIORITY_MODE_GLOBAL_KEY:
# 全局 Key 优先模式:按格式特定优先级分组
for candidate in candidates:
priority = 999999
@@ -204,6 +202,14 @@ class CandidateSorter:
return result
def shuffle_keys_by_internal_priority(
self,
keys: list[ProviderAPIKey],
affinity_key: str | None = None,
use_random: bool = False,
) -> list[ProviderAPIKey]:
return self._shuffle_keys_by_internal_priority(keys, affinity_key, use_random)
def _shuffle_keys_by_internal_priority(
self,
keys: list[ProviderAPIKey],
@@ -252,7 +258,7 @@ class CandidateSorter:
# 正常模式:使用哈希确定性打乱(保持缓存亲和性)
key_scores = []
for key in group_keys:
hash_value = self._scheduler._affinity_hash(affinity_key, key.id)
hash_value = affinity_hash(affinity_key, key.id)
key_scores.append((hash_value, key))
# 按哈希值排序

View File

@@ -0,0 +1,144 @@
"""
并发控制检查器 (ConcurrencyChecker)
从 CacheAwareScheduler 提取的 RPM 限流和动态预留逻辑。
"""
from __future__ import annotations
import math
from typing import Any
from src.core.logger import logger
from src.models.database import ProviderAPIKey
from src.services.cache.schemas import ConcurrencySnapshot
from src.services.rate_limit.adaptive_reservation import AdaptiveReservationManager
from src.services.rate_limit.adaptive_rpm import get_adaptive_rpm_manager
class ConcurrencyChecker:
"""并发控制检查器,封装 RPM 限流和动态预留逻辑。"""
def __init__(
self,
concurrency_manager: Any,
reservation_manager: AdaptiveReservationManager,
) -> None:
self._concurrency_manager = concurrency_manager
self._reservation_manager = reservation_manager
@staticmethod
def get_effective_rpm_limit(key: ProviderAPIKey) -> int | None:
"""获取有效的 RPM 限制(委托给 AdaptiveRPMManager 统一逻辑)"""
return get_adaptive_rpm_manager().get_effective_limit(key)
async def check_available(
self,
key: ProviderAPIKey,
is_cached_user: bool = False,
) -> tuple[bool, ConcurrencySnapshot]:
"""
检查 RPM 限制是否可用(使用动态预留机制)
核心逻辑 - 动态缓存预留机制:
- 总槽位: 有效 RPM 限制(固定值或学习到的值)
- 预留比例: 由 AdaptiveReservationManager 根据置信度和负载动态计算
- 缓存用户可用: 全部槽位
- 新用户可用: 总槽位 x (1 - 动态预留比例)
Args:
key: ProviderAPIKey对象
is_cached_user: 是否是缓存用户
Returns:
(是否可用, 并发快照)
"""
# 获取有效的并发限制
effective_key_limit = self.get_effective_rpm_limit(key)
logger.debug(
" -> 并发检查: _concurrency_manager={}, "
"is_cached_user={}, effective_limit={}",
self._concurrency_manager is not None,
is_cached_user,
effective_key_limit,
)
if not self._concurrency_manager:
# 并发管理器不可用直接返回True
logger.debug(" -> 无并发管理器,直接通过")
snapshot = ConcurrencySnapshot(
key_current=0,
key_limit=effective_key_limit,
is_cached_user=is_cached_user,
)
return True, snapshot
# 获取当前 RPM 计数
key_count = await self._concurrency_manager.get_key_rpm_count(
key_id=str(key.id),
)
can_use = True
# 计算动态预留比例
reservation_result = self._reservation_manager.calculate_reservation(
key=key,
current_usage=key_count,
effective_limit=effective_key_limit,
)
available_for_new = None
reservation_ratio = reservation_result.ratio
# 检查Key级别限制使用动态预留比例
if effective_key_limit is not None:
if is_cached_user:
# 缓存用户: 可以使用全部槽位
if key_count >= effective_key_limit:
can_use = False
else:
# 新用户: 只能使用 (1 - 动态预留比例) 的槽位
# 使用 max 确保至少有 1 个槽位可用
# 与 ConcurrencyManager 的 Lua 脚本保持一致:使用 floor 计算新用户可用槽位
available_for_new = max(
1, math.floor(effective_key_limit * (1 - reservation_ratio))
)
if key_count >= available_for_new:
logger.debug(
"Key {}... 新用户配额已满 " "({}/{}, 总{}, 预留{:.0%}[{}])",
key.id[:8],
key_count,
available_for_new,
effective_key_limit,
reservation_ratio,
reservation_result.phase,
)
can_use = False
key_limit_for_snapshot: int | None
if is_cached_user:
key_limit_for_snapshot = effective_key_limit
elif effective_key_limit is not None:
key_limit_for_snapshot = (
available_for_new if available_for_new is not None else effective_key_limit
)
else:
key_limit_for_snapshot = None
snapshot = ConcurrencySnapshot(
key_current=key_count,
key_limit=key_limit_for_snapshot,
is_cached_user=is_cached_user,
reservation_ratio=reservation_ratio,
reservation_phase=reservation_result.phase,
reservation_confidence=reservation_result.confidence,
load_factor=reservation_result.load_factor,
)
return can_use, snapshot
def get_reservation_stats(self) -> dict[str, Any]:
"""获取动态预留管理器的统计信息"""
return self._reservation_manager.get_stats()

View File

@@ -0,0 +1,82 @@
"""
访问限制检查
从 CacheAwareScheduler 提取的 ApiKey + User 访问限制合并逻辑。
"""
from __future__ import annotations
from typing import Any
from src.core.logger import logger
from src.core.model_permissions import merge_allowed_models
from src.models.database import ApiKey
def get_effective_restrictions(user_api_key: ApiKey | None) -> dict[str, Any]:
"""
获取有效的访问限制(合并 ApiKey 和 User 的限制)
逻辑:
- 如果 ApiKey 和 User 都有限制,取交集
- 如果只有一方有限制,使用该方的限制
- 如果都没有限制,返回 None表示不限制
Args:
user_api_key: 用户 API Key 对象(可能包含 user relationship
Returns:
包含 allowed_providers, allowed_models, allowed_api_formats 的字典
"""
result: dict[str, Any] = {
"allowed_providers": None,
"allowed_models": None,
"allowed_api_formats": None,
}
if not user_api_key:
return result
# 获取 User 的限制
# 注意:这里可能触发 lazy loading需要确保 session 仍然有效
try:
user = user_api_key.user if hasattr(user_api_key, "user") else None
except Exception as e:
logger.warning("无法加载 ApiKey 关联的 User: {},仅使用 ApiKey 级别的限制", e)
user = None
# 调试日志
logger.debug(
"[_get_effective_restrictions] ApiKey={}..., User={}..., "
"ApiKey.allowed_models={}, User.allowed_models={}",
user_api_key.id[:8],
user.id[:8] if user else "None",
user_api_key.allowed_models,
user.allowed_models if user else "N/A",
)
# 合并 allowed_providers
result["allowed_providers"] = merge_restriction_sets(
user_api_key.allowed_providers, user.allowed_providers if user else None
)
# 合并 allowed_models取交集
result["allowed_models"] = merge_allowed_models(
user_api_key.allowed_models, user.allowed_models if user else None
)
# 合并 allowed_api_formats
result["allowed_api_formats"] = merge_restriction_sets(
user_api_key.allowed_api_formats, user.allowed_api_formats if user else None
)
return result
def merge_restriction_sets(key_restriction: Any, user_restriction: Any) -> set[Any] | None:
"""合并两个限制列表,取交集;任一方为空则使用另一方;均空返回 None"""
key_set = set(key_restriction) if key_restriction else None
user_set = set(user_restriction) if user_restriction else None
if key_set and user_set:
return key_set & user_set
return key_set or user_set

82
src/services/cache/scheduling_config.py vendored Normal file
View File

@@ -0,0 +1,82 @@
"""
调度配置 (SchedulingConfig)
从 CacheAwareScheduler 提取的常量定义和模式管理逻辑。
"""
from __future__ import annotations
from src.core.logger import logger
class SchedulingConfig:
"""调度配置:管理优先级模式和调度模式的常量、归一化和运行时更新。"""
# 优先级模式常量
PRIORITY_MODE_PROVIDER = "provider" # 提供商优先模式
PRIORITY_MODE_GLOBAL_KEY = "global_key" # 全局 Key 优先模式
ALLOWED_PRIORITY_MODES = {
PRIORITY_MODE_PROVIDER,
PRIORITY_MODE_GLOBAL_KEY,
}
# 调度模式常量
SCHEDULING_MODE_FIXED_ORDER = "fixed_order" # 固定顺序模式:严格按优先级,忽略缓存
SCHEDULING_MODE_CACHE_AFFINITY = "cache_affinity" # 缓存亲和模式:优先缓存,同优先级哈希分散
SCHEDULING_MODE_LOAD_BALANCE = "load_balance" # 负载均衡模式:忽略缓存,同优先级随机轮换
ALLOWED_SCHEDULING_MODES = {
SCHEDULING_MODE_FIXED_ORDER,
SCHEDULING_MODE_CACHE_AFFINITY,
SCHEDULING_MODE_LOAD_BALANCE,
}
def __init__(
self,
priority_mode: str | None = None,
scheduling_mode: str | None = None,
) -> None:
self.priority_mode = self._normalize_priority_mode(
priority_mode or self.PRIORITY_MODE_PROVIDER
)
self.scheduling_mode = self._normalize_scheduling_mode(
scheduling_mode or self.SCHEDULING_MODE_CACHE_AFFINITY
)
logger.debug(
"[SchedulingConfig] 初始化优先级模式: {}, 调度模式: {}",
self.priority_mode,
self.scheduling_mode,
)
def _normalize_priority_mode(self, mode: str | None) -> str:
normalized = (mode or "").strip().lower()
if normalized not in self.ALLOWED_PRIORITY_MODES:
if normalized:
logger.warning("[SchedulingConfig] 无效的优先级模式 '{}',回退为 provider", mode)
return self.PRIORITY_MODE_PROVIDER
return normalized
def _normalize_scheduling_mode(self, mode: str | None) -> str:
normalized = (mode or "").strip().lower()
if normalized not in self.ALLOWED_SCHEDULING_MODES:
if normalized:
logger.warning(
"[SchedulingConfig] 无效的调度模式 '{}',回退为 cache_affinity", mode
)
return self.SCHEDULING_MODE_CACHE_AFFINITY
return normalized
def set_priority_mode(self, mode: str | None) -> None:
"""运行时更新候选排序策略"""
normalized = self._normalize_priority_mode(mode)
if normalized == self.priority_mode:
return
self.priority_mode = normalized
logger.debug("[SchedulingConfig] 切换优先级模式为: {}", self.priority_mode)
def set_scheduling_mode(self, mode: str | None) -> None:
"""运行时更新调度模式"""
normalized = self._normalize_scheduling_mode(mode)
if normalized == self.scheduling_mode:
return
self.scheduling_mode = normalized
logger.debug("[SchedulingConfig] 切换调度模式为: {}", self.scheduling_mode)

88
src/services/cache/schemas.py vendored Normal file
View File

@@ -0,0 +1,88 @@
"""
调度器核心数据类型
从 CacheAwareScheduler 提取的共享数据结构,被 24+ 个模块使用。
"""
from __future__ import annotations
from dataclasses import dataclass
from src.models.database import (
Provider,
ProviderAPIKey,
ProviderEndpoint,
)
@dataclass
class ProviderCandidate:
"""候选 provider 组合及是否命中缓存"""
provider: Provider
endpoint: ProviderEndpoint
key: ProviderAPIKey
is_cached: bool = False
is_skipped: bool = False # 是否被跳过
skip_reason: str | None = None # 跳过原因
mapping_matched_model: str | None = None # 通过映射匹配到的模型名(用于实际请求)
needs_conversion: bool = False # 是否需要格式转换
provider_api_format: str = "" # Provider 端点实际格式(用于健康度/熔断 bucket
def _stable_order_key(self) -> tuple[int, int, str, str, str]:
"""
为排序/优先队列提供稳定的比较键。
说明:
- 运行时偶发会出现对 ProviderCandidate 做 tuple 排序/heap 排序的场景;
当主键相同需要比较候选本身时,若候选不可比较会触发:
TypeError: '<' not supported between instances of 'ProviderCandidate' and 'ProviderCandidate'
- 这里提供一个与调度逻辑无关、但足够稳定且可比的兜底顺序。
"""
provider_priority_raw = getattr(self.provider, "provider_priority", None)
internal_priority_raw = getattr(self.key, "internal_priority", None)
try:
provider_priority = (
int(provider_priority_raw) if provider_priority_raw is not None else 999999
)
except Exception:
provider_priority = 999999
try:
internal_priority = (
int(internal_priority_raw) if internal_priority_raw is not None else 999999
)
except Exception:
internal_priority = 999999
provider_id = str(getattr(self.provider, "id", "") or "")
endpoint_id = str(getattr(self.endpoint, "id", "") or "")
key_id = str(getattr(self.key, "id", "") or "")
return (provider_priority, internal_priority, provider_id, endpoint_id, key_id)
def __lt__(self, other: object) -> bool:
if not isinstance(other, ProviderCandidate):
return NotImplemented
return self._stable_order_key() < other._stable_order_key()
@dataclass
class ConcurrencySnapshot:
key_current: int
key_limit: int | None
is_cached_user: bool = False
# 动态预留信息
reservation_ratio: float = 0.0
reservation_phase: str = "unknown"
reservation_confidence: float = 0.0
load_factor: float = 0.0
def describe(self) -> str:
key_limit_text = str(self.key_limit) if self.key_limit is not None else "inf"
reservation_text = f"{self.reservation_ratio:.0%}" if self.reservation_ratio > 0 else "N/A"
return (
f"key={self.key_current}/{key_limit_text}, "
f"cached={self.is_cached_user}, "
f"reserve={reservation_text}({self.reservation_phase})"
)

53
src/services/cache/utils.py vendored Normal file
View File

@@ -0,0 +1,53 @@
"""
调度器工具函数
从 CacheAwareScheduler 提取的静态工具方法。
"""
from __future__ import annotations
import hashlib
from sqlalchemy.orm import Session
def affinity_hash(affinity_key: str, identifier: str) -> int:
"""基于 affinity_key 和标识符的确定性哈希(用于同优先级内分散负载均衡)"""
return int(hashlib.sha256(f"{affinity_key}:{identifier}".encode()).hexdigest()[:16], 16)
def release_db_connection_before_await(db: Session) -> None:
"""
Best-effort: end a read-only transaction before awaiting async I/O.
This scheduler does a lot of async work (cache/Redis) mixed with sync SQLAlchemy reads.
If a SELECT has already started a transaction, the pooled connection can remain checked
out while we await, causing pool pressure under concurrency.
Safety:
- Only commits when the Session has no ORM pending changes.
- Temporarily disables expire_on_commit to keep already-loaded ORM objects usable.
"""
try:
if db is None:
return
has_pending_changes = bool(db.new) or bool(db.dirty) or bool(db.deleted)
if has_pending_changes:
return
if not db.in_transaction():
return
original_expire_on_commit = getattr(db, "expire_on_commit", True)
db.expire_on_commit = False
try:
db.commit()
except Exception:
try:
db.rollback()
except Exception:
pass
finally:
db.expire_on_commit = original_expire_on_commit
except Exception:
# Never let this optimization break scheduling
return

View File

@@ -2,10 +2,14 @@
Sub2API 余额查询操作
"""
import asyncio
import time
from typing import Any
import httpx
from src.services.provider_ops.actions.balance import BalanceAction
from src.services.provider_ops.types import BalanceInfo
from src.services.provider_ops.types import ActionResult, ActionStatus, BalanceInfo
class Sub2ApiBalanceAction(BalanceAction):
@@ -13,33 +17,126 @@ class Sub2ApiBalanceAction(BalanceAction):
Sub2API 余额查询
特点:
- 使用 /api/v1/auth/me 端点
- balance 为充值余额points 为赠送余额,均以美元为单位
- 并发调用 /api/v1/auth/me 和 /api/v1/subscriptions/summary
- auth/me 提供基础余额balance + points
- subscriptions/summary 提供订阅详情(各订阅的日/周/月用量和额度)
- 响应格式: {"code": 0, "message": "success", "data": {...}}
"""
display_name = "查询余额"
description = "查询 Sub2API 账户余额信息"
description = "查询 Sub2API 账户余额和订阅信息"
def _parse_balance(self, data: Any) -> BalanceInfo:
"""解析 Sub2API 余额信息"""
# Sub2API 使用 {"code": 0, ...} 表示成功,非 0 表示业务错误
if isinstance(data, dict) and data.get("code") is not None and data.get("code") != 0:
message = data.get("message", "查询失败")
raise ValueError(f"Sub2API 业务错误: {message}")
"""本类完全重写了 _do_query_balance绕过基类默认流程故此方法不会被调用"""
raise NotImplementedError(
"Sub2API 重写了 _do_query_balance不走基类的 _parse_balance 路径"
)
user_data = data.get("data", {}) if isinstance(data, dict) else {}
async def _do_query_balance(self, client: httpx.AsyncClient) -> ActionResult:
"""并发查询 auth/me 和 subscriptions/summary"""
start_time = time.time()
balance = self._to_float(user_data.get("balance")) or 0.0
points = self._to_float(user_data.get("points")) or 0.0
me_endpoint = self.config.get("endpoint", "/api/v1/auth/me?timezone=Asia/Shanghai")
sub_endpoint = self.config.get("subscription_endpoint", "/api/v1/subscriptions/summary")
total_available = balance + points
try:
me_resp, sub_resp = await asyncio.gather(
client.get(me_endpoint),
client.get(sub_endpoint),
return_exceptions=True,
)
response_time_ms = int((time.time() - start_time) * 1000)
return self._create_balance_info(
total_available=total_available,
currency="USD",
extra={
# 解析 auth/me
me_data: dict[str, Any] = {}
me_ok = False
if isinstance(me_resp, httpx.Response):
if me_resp.status_code in (401, 403):
return self._make_error_result(
ActionStatus.AUTH_FAILED, "认证失败,请检查凭据配置"
)
if me_resp.status_code == 200:
try:
me_json = me_resp.json()
if me_json.get("code") == 0:
me_data = me_json.get("data", {})
me_ok = True
except Exception:
pass
if not me_ok:
return self._make_error_result(ActionStatus.UNKNOWN_ERROR, "查询用户信息失败")
# 基础余额
balance = self._to_float(me_data.get("balance")) or 0.0
points = self._to_float(me_data.get("points")) or 0.0
total_available = balance + points
extra: dict[str, Any] = {
"balance": balance,
"points": points,
},
)
}
# 解析 subscriptions/summary可选失败不影响主流程
if isinstance(sub_resp, httpx.Response) and sub_resp.status_code == 200:
try:
sub_json = sub_resp.json()
if sub_json.get("code") == 0:
summary = sub_json.get("data", {})
extra["active_subscriptions"] = summary.get("active_count", 0)
extra["total_used_usd"] = summary.get("total_used_usd", 0)
extra["subscriptions"] = self._parse_subscriptions(
summary.get("subscriptions", [])
)
except Exception:
pass
balance_info = self._create_balance_info(
total_available=total_available,
currency="USD",
extra=extra,
)
return self._make_success_result(
data=balance_info,
response_time_ms=response_time_ms,
raw_response={"me": me_data},
)
except httpx.TimeoutException:
return self._make_error_result(
ActionStatus.NETWORK_ERROR, "请求超时", retry_after_seconds=30
)
except httpx.RequestError as e:
return self._make_error_result(
ActionStatus.NETWORK_ERROR, f"网络错误: {e}", retry_after_seconds=30
)
except Exception as e:
return self._make_error_result(ActionStatus.UNKNOWN_ERROR, f"未知错误: {e}")
@staticmethod
def _parse_subscriptions(items: list[dict[str, Any]]) -> list[dict[str, Any]]:
"""将订阅列表精简为前端需要的字段"""
result = []
for item in items:
sub: dict[str, Any] = {
"group_name": item.get("group_name", ""),
"status": item.get("status", ""),
}
# 只保留非 None 的限额字段0 是有效值,表示未使用/无限额)
for field in (
"daily_used_usd",
"daily_limit_usd",
"weekly_used_usd",
"weekly_limit_usd",
"monthly_used_usd",
"monthly_limit_usd",
):
val = item.get(field)
if val is not None:
sub[field] = val
expires_at = item.get("expires_at")
if expires_at:
sub["expires_at"] = expires_at
result.append(sub)
return result

View File

@@ -3,7 +3,7 @@ Provider 架构抽象基类
"""
from abc import ABC, abstractmethod
from collections.abc import AsyncIterator
from collections.abc import AsyncIterator, Callable
from contextlib import asynccontextmanager
from dataclasses import dataclass
from datetime import datetime, timezone
@@ -59,6 +59,9 @@ class ProviderConnector(ABC):
self._timeout = self.config.get("timeout", 30)
self._headers: dict[str, str] = {}
# 凭据更新回调Token Rotation 等场景需要持久化新凭据)
self._on_credentials_updated: Callable[[dict[str, Any]], None] | None = None
@abstractmethod
async def connect(self, credentials: dict[str, Any]) -> bool:
"""
@@ -389,12 +392,11 @@ class ProviderArchitecture(ABC):
base_url: str,
config: dict[str, Any],
credentials: dict[str, Any],
) -> dict[str, Any]:
) -> dict[str, Any] | tuple[dict[str, Any], dict[str, Any]]:
"""
验证前的异步预处理(可选)
子类可重写以执行异步操作(如获取动态 Cookie
返回的配置会传递给 build_verify_headers。
子类可重写以执行异步操作(如获取动态 Cookie、登录获取 Token)。
Args:
base_url: API 基础地址
@@ -402,7 +404,10 @@ class ProviderArchitecture(ABC):
credentials: 凭据信息
Returns:
处理后的配置(会与原 config 合并)
- dict: 额外配置(会与原 config 合并传递给 build_verify_headers
- tuple[dict, dict]: (额外配置, 需持久化的凭据更新)
当预处理过程中凭据发生变更时(如 Token Rotation
通过第二个 dict 显式返回需要持久化的字段。
"""
return {}
@@ -516,7 +521,11 @@ class ProviderArchitecture(ABC):
"credentials_schema": self.get_credentials_schema(),
"verify_endpoint": self.get_verify_endpoint(),
"supported_auth_types": [
{"type": c.auth_type.value, "display_name": c.display_name}
{
"type": c.auth_type.value,
"display_name": c.display_name,
"credentials_schema": c.get_credentials_schema(),
}
for c in self.supported_connectors
],
"supported_actions": [

View File

@@ -2,12 +2,19 @@
Sub2API 架构
针对 Sub2API 风格的中转站优化的预设配置。
支持两种认证方式:
1. 账号密码登录(自动获取 JWT过期自动刷新refresh 失败自动重新登录)
2. Refresh Token从浏览器 localStorage 获取,自动续期,适合 OAuth 用户)
"""
import time
from collections.abc import AsyncIterator, Callable
from contextlib import asynccontextmanager
from typing import Any
import httpx
from src.core.logger import logger
from src.services.provider_ops.actions import ProviderAction
from src.services.provider_ops.actions.sub2api_balance import Sub2ApiBalanceAction
from src.services.provider_ops.architectures.base import (
@@ -16,51 +23,238 @@ from src.services.provider_ops.architectures.base import (
VerifyResult,
)
from src.services.provider_ops.types import ConnectorAuthType, ProviderActionType
from src.utils.ssl_utils import get_ssl_context
class Sub2ApiConnector(ProviderConnector):
"""
Sub2API 连接器
def _calc_expires_at(token_data: dict[str, Any]) -> float:
"""从 token 响应数据计算过期时间(秒级时间戳,提前 60s"""
token_expires_at = token_data.get("token_expires_at")
if token_expires_at is not None:
# Sub2API 返回毫秒级绝对时间戳
return token_expires_at / 1000 - 60
expires_in = token_data.get("expires_in", 900)
return time.time() + expires_in - 60
使用 Bearer Token 认证。
async def _do_login(
client: httpx.AsyncClient,
email: str,
password: str,
) -> dict[str, Any]:
"""调用 Sub2API 登录接口,返回 token_data。失败抛 ValueError。"""
resp = await client.post(
"/api/v1/auth/login",
json={"email": email, "password": password},
)
data = resp.json()
if resp.status_code != 200 or data.get("code", -1) != 0:
raise ValueError(data.get("message", f"登录失败 (HTTP {resp.status_code})"))
return data.get("data", {})
async def _do_refresh(
client: httpx.AsyncClient,
refresh_token: str,
) -> dict[str, Any]:
"""调用 Sub2API refresh 接口,返回 token_data。失败抛 ValueError。"""
resp = await client.post(
"/api/v1/auth/refresh",
json={"refresh_token": refresh_token},
)
data = resp.json()
if resp.status_code != 200 or data.get("code", -1) != 0:
raise ValueError(data.get("message", "Refresh Token 无效或已过期"))
return data.get("data", {})
def _collect_updated_credentials(
token_data: dict[str, Any],
old_refresh_token: str | None = None,
) -> dict[str, Any]:
"""从 token 响应中提取需要持久化的凭据变更"""
updated: dict[str, Any] = {}
new_refresh_token = token_data.get("refresh_token")
if new_refresh_token and new_refresh_token != old_refresh_token:
updated["refresh_token"] = new_refresh_token
access_token = token_data.get("access_token")
if access_token:
updated["_cached_access_token"] = access_token
updated["_cached_token_expires_at"] = _calc_expires_at(token_data)
return updated
class _Sub2ApiTokenMixin:
"""Sub2API JWT token 管理公共逻辑
与 ProviderConnector 配合使用MRO 中由 ProviderConnector 提供实际属性初始化)。
以下类型注解声明 mixin 依赖的协议属性,不会创建新的实例属性。
"""
auth_type = ConnectorAuthType.API_KEY
display_name = "Sub2API Key"
# Mixin 自身管理的 token 状态(提供默认值防止子类遗漏初始化)
_access_token: str | None = None
_refresh_token: str | None = None
_token_expires_at: float = 0
# 以下属性由 ProviderConnector.__init__ 初始化,仅作协议声明
_on_credentials_updated: Callable[[dict[str, Any]], None] | None
base_url: str
_timeout: int | float
_proxy: str | httpx.Proxy | None
@asynccontextmanager
async def _get_raw_client(self) -> AsyncIterator[httpx.AsyncClient]:
"""获取不带 auth hook 的裸 HTTP 客户端(用于登录/刷新 token"""
transport = None
if self._proxy:
transport = httpx.AsyncHTTPTransport(proxy=self._proxy)
async with httpx.AsyncClient(
base_url=self.base_url,
timeout=self._timeout,
transport=transport,
verify=get_ssl_context(),
) as client:
yield client
def _update_tokens(self, token_data: dict[str, Any]) -> None:
"""更新实例 token 状态并通过回调持久化变更"""
old_refresh_token = self._refresh_token
self._access_token = token_data.get("access_token")
self._refresh_token = token_data.get("refresh_token", self._refresh_token)
self._token_expires_at = _calc_expires_at(token_data)
if self._on_credentials_updated:
updated = _collect_updated_credentials(token_data, old_refresh_token)
if updated:
self._on_credentials_updated(updated)
async def _refresh(self) -> bool:
"""使用 refresh_token 续期"""
if not self._refresh_token:
return False
try:
async with self._get_raw_client() as client:
token_data = await _do_refresh(client, self._refresh_token)
self._update_tokens(token_data)
logger.debug("Sub2API token 续期成功")
return True
except ValueError as e:
logger.warning("Sub2API refresh_token 续期失败: {}", e)
return False
except Exception as e:
logger.warning("Sub2API refresh_token 续期异常: {}", e)
return False
class Sub2ApiConnector(_Sub2ApiTokenMixin, ProviderConnector):
"""
Sub2API 连接器(账号密码模式)
使用 email + password 登录获取 JWT Token支持自动刷新
- 登录后获取 access_token + refresh_token
- access_token 过期前自动使用 refresh_token 续期
- refresh_token 也过期时自动重新登录
"""
auth_type = ConnectorAuthType.SESSION_LOGIN
display_name = "账号密码"
def __init__(self, base_url: str, config: dict[str, Any] | None = None):
super().__init__(base_url, config)
self._api_key: str | None = None
@staticmethod
def _strip_bearer(value: str) -> str:
"""去掉用户可能粘贴的 Bearer 前缀"""
stripped = value.strip()
if stripped.lower().startswith("bearer "):
stripped = stripped[7:].strip()
return stripped
self._access_token: str | None = None
self._refresh_token: str | None = None
self._token_expires_at: float = 0
self._email: str | None = None
self._password: str | None = None
async def connect(self, credentials: dict[str, Any]) -> bool:
api_key = credentials.get("api_key")
if not api_key:
self._set_error("JWT Token 不能为空")
email = credentials.get("email", "").strip()
password = credentials.get("password", "").strip()
if not email or not password:
self._set_error("邮箱和密码不能为空")
return False
self._api_key = self._strip_bearer(api_key)
self._set_connected()
return True
self._email = email
self._password = password
# 如果有缓存的 access_token 且未过期,直接复用,避免不必要的登录
cached_access_token = credentials.get("_cached_access_token", "")
cached_expires_at = credentials.get("_cached_token_expires_at", 0)
if cached_access_token and time.time() < cached_expires_at:
self._access_token = cached_access_token
self._token_expires_at = cached_expires_at
# 恢复 refresh_token 以便 access_token 过期后可刷新而非重新登录
self._refresh_token = credentials.get("refresh_token", "").strip() or None
self._set_connected()
return True
return await self._login()
async def _login(self) -> bool:
"""使用 email + password 登录获取 token pair"""
try:
async with self._get_raw_client() as client:
token_data = await _do_login(client, self._email or "", self._password or "")
self._update_tokens(token_data)
self._set_connected()
logger.debug(
"Sub2API 登录成功: {}", self._email[:3] + "***" if self._email else "N/A"
)
return True
except ValueError as e:
self._set_error(str(e))
return False
except httpx.TimeoutException:
self._set_error("登录请求超时")
return False
except httpx.RequestError as e:
self._set_error(f"登录网络错误: {e}")
return False
except Exception as e:
self._set_error(f"登录失败: {e}")
return False
async def _ensure_token(self) -> None:
"""确保 access_token 有效,过期则自动刷新或重新登录"""
if self._access_token and time.time() < self._token_expires_at:
return
if self._refresh_token and await self._refresh():
return
logger.info("Sub2API token 已过期,尝试重新登录")
if not await self._login():
logger.error("Sub2API 重新登录失败: {}", self._last_error)
async def disconnect(self) -> None:
self._api_key = None
self._access_token = None
self._refresh_token = None
self._token_expires_at = 0
self._email = None
self._password = None
self._set_disconnected()
async def is_authenticated(self) -> bool:
return self._api_key is not None
if not self._access_token:
return False
return self._refresh_token is not None or self._email is not None
def _apply_auth(self, request: httpx.Request) -> httpx.Request:
if self._api_key:
request.headers["Authorization"] = f"Bearer {self._api_key}"
if self._access_token:
request.headers["Authorization"] = f"Bearer {self._access_token}"
return request
async def _auth_hook(self, request: httpx.Request) -> None:
await self._ensure_token()
self._apply_auth(request)
async def refresh_auth(self, credentials: dict[str, Any]) -> bool:
if await self._refresh():
return True
self._email = credentials.get("email", self._email)
self._password = credentials.get("password", self._password)
return await self._login()
@classmethod
def get_credentials_schema(cls) -> dict[str, Any]:
return {
@@ -71,26 +265,144 @@ class Sub2ApiConnector(ProviderConnector):
"title": "站点地址",
"description": "API 基础地址",
},
"api_key": {
"email": {
"type": "string",
"title": "JWT Token",
"description": "Sub2API 的访问令牌",
"title": "邮箱",
"description": "Sub2API 登录邮箱",
},
"password": {
"type": "string",
"title": "密码",
"description": "Sub2API 登录密码",
"x-sensitive": True,
"x-input-type": "password",
},
},
"required": ["api_key"],
"required": ["email", "password"],
"x-field-groups": [
{"fields": ["base_url"]},
{"fields": ["api_key"]},
{"fields": ["email"]},
{"fields": ["password"]},
],
"x-auth-type": "session_login",
"x-auth-method": "jwt",
"x-validation": [
{
"type": "required",
"fields": ["email", "password"],
"message": "请填写邮箱和密码",
},
],
}
class Sub2ApiRefreshTokenConnector(_Sub2ApiTokenMixin, ProviderConnector):
"""
Sub2API 连接器Refresh Token 模式)
适合 OAuth 登录用户(如 LinuxDo从浏览器 localStorage 获取 refresh_token。
- 首次连接时用 refresh_token 换取 access_token
- access_token 过期前自动续期
- refresh_token 过期后需手动更新(无法自动重新登录)
"""
auth_type = ConnectorAuthType.API_KEY
display_name = "Refresh Token"
def __init__(self, base_url: str, config: dict[str, Any] | None = None):
super().__init__(base_url, config)
self._access_token: str | None = None
self._refresh_token: str | None = None
self._token_expires_at: float = 0
async def connect(self, credentials: dict[str, Any]) -> bool:
refresh_token = credentials.get("refresh_token", "").strip()
if not refresh_token:
self._set_error("请填写 Refresh Token")
return False
self._refresh_token = refresh_token
# 如果有缓存的 access_token 且未过期,直接复用,不消耗 refresh_token
cached_access_token = credentials.get("_cached_access_token", "")
cached_expires_at = credentials.get("_cached_token_expires_at", 0)
if cached_access_token and time.time() < cached_expires_at:
self._access_token = cached_access_token
self._token_expires_at = cached_expires_at
self._set_connected()
return True
# 首次连接或 access_token 已过期,用 refresh_token 换取
if not await self._refresh():
self._refresh_token = None # 清理,避免残留无效状态
self._set_error("Refresh Token 无效或已过期")
return False
self._set_connected()
return True
async def _ensure_token(self) -> None:
"""确保 access_token 有效,有 refresh_token 时自动续期"""
if self._access_token and time.time() < self._token_expires_at:
return
if self._refresh_token:
await self._refresh()
async def disconnect(self) -> None:
self._access_token = None
self._refresh_token = None
self._token_expires_at = 0
self._set_disconnected()
async def is_authenticated(self) -> bool:
return self._refresh_token is not None
def _apply_auth(self, request: httpx.Request) -> httpx.Request:
if self._access_token:
request.headers["Authorization"] = f"Bearer {self._access_token}"
return request
async def _auth_hook(self, request: httpx.Request) -> None:
await self._ensure_token()
self._apply_auth(request)
async def refresh_auth(self, credentials: dict[str, Any]) -> bool:
if self._refresh_token and await self._refresh():
return True
return await self.connect(credentials)
@classmethod
def get_credentials_schema(cls) -> dict[str, Any]:
return {
"type": "object",
"properties": {
"base_url": {
"type": "string",
"title": "站点地址",
"description": "API 基础地址",
},
"refresh_token": {
"type": "string",
"title": "Refresh Token",
"description": ("从浏览器 F12 > Application > Local Storage 获取"),
"x-sensitive": True,
"x-input-type": "password",
"x-help": "浏览器控制台执行 localStorage.getItem('refresh_token') 获取",
},
},
"required": ["refresh_token"],
"x-field-groups": [
{"fields": ["base_url"]},
{"fields": ["refresh_token"]},
],
"x-auth-type": "api_key",
"x-auth-method": "bearer",
"x-validation": [
{
"type": "required",
"fields": ["api_key"],
"message": "请填写 JWT Token",
"fields": ["refresh_token"],
"message": "请填写 Refresh Token",
},
],
}
@@ -101,22 +413,27 @@ class Sub2ApiArchitecture(ProviderArchitecture):
Sub2API 架构
特点:
- 使用 Bearer Token 认证
- 支持两种认证方式:账号密码 / Refresh Token
- 验证端点: /api/v1/auth/me
- balance 为充值余额points 为赠送余额
- 余额查询同时获取订阅概览信息
"""
architecture_id = "sub2api"
display_name = "Sub2API"
description = "Sub2API 风格中转站的预设配置"
supported_connectors: list[type[ProviderConnector]] = [Sub2ApiConnector]
supported_connectors: list[type[ProviderConnector]] = [
Sub2ApiConnector,
Sub2ApiRefreshTokenConnector,
]
supported_actions: list[type[ProviderAction]] = [Sub2ApiBalanceAction]
default_action_configs: dict[ProviderActionType, dict[str, Any]] = {
ProviderActionType.QUERY_BALANCE: {
"endpoint": "/api/v1/auth/me?timezone=Asia/Shanghai",
"subscription_endpoint": "/api/v1/subscriptions/summary",
"method": "GET",
},
}
@@ -133,21 +450,76 @@ class Sub2ApiArchitecture(ProviderArchitecture):
credentials: dict[str, Any],
) -> dict[str, str]:
headers: dict[str, str] = {}
api_key = credentials.get("api_key", "")
if api_key:
headers["Authorization"] = f"Bearer {Sub2ApiConnector._strip_bearer(api_key)}"
access_token = credentials.get("_access_token", "")
if access_token:
headers["Authorization"] = f"Bearer {access_token}"
return headers
async def prepare_verify_config(
self,
base_url: str,
config: dict[str, Any],
credentials: dict[str, Any],
) -> tuple[dict[str, Any], dict[str, Any]]:
"""
验证前预处理:根据凭据类型选择登录方式
- 有 email + password -> 账号密码登录
- 有 refresh_token -> 用 refresh_token 换 access_token
Returns:
(extra_config, updated_credentials):
extra_config 为空updated_credentials 包含需持久化的凭据变更
Token Rotation 后的新 refresh_token、缓存的 access_token 等)
"""
base_url = base_url.rstrip("/")
from src.services.proxy_node.resolver import resolve_ops_proxy
proxy = resolve_ops_proxy(config)
client_kwargs: dict[str, Any] = {
"base_url": base_url,
"timeout": 30.0,
"verify": get_ssl_context(),
}
if proxy:
client_kwargs["proxy"] = proxy
email = credentials.get("email", "").strip()
password = credentials.get("password", "").strip()
refresh_token = credentials.get("refresh_token", "").strip()
try:
if email and password:
async with httpx.AsyncClient(**client_kwargs) as client:
token_data = await _do_login(client, email, password)
elif refresh_token:
async with httpx.AsyncClient(**client_kwargs) as client:
token_data = await _do_refresh(client, refresh_token)
else:
raise ValueError("请填写账号密码或 Refresh Token")
except ValueError:
raise
except Exception as e:
raise ValueError(f"验证失败: {e}") from e
access_token = token_data.get("access_token", "")
credentials["_access_token"] = access_token
updated_credentials = _collect_updated_credentials(
token_data, old_refresh_token=refresh_token or None
)
return {}, updated_credentials
def parse_verify_response(
self,
status_code: int,
data: dict[str, Any],
) -> VerifyResult:
"""
解析 Sub2API 验证响应
Sub2API 使用 {"code": 0, "message": "success", "data": {...}} 格式。
"""
if status_code == 401:
return VerifyResult(success=False, message=self._auth_fail_message(401))
if status_code == 403:
@@ -166,7 +538,6 @@ class Sub2ApiArchitecture(ProviderArchitecture):
def _build_verify_result(
self, user_data: dict[str, Any], raw_data: dict[str, Any] | None = None
) -> VerifyResult:
# or 0 防御上游返回 None / "" 等 falsy 值
balance = float(user_data.get("balance") or 0)
points = float(user_data.get("points") or 0)

View File

@@ -13,6 +13,7 @@ from datetime import datetime, timezone
from typing import Any
from sqlalchemy.orm import Session
from sqlalchemy.orm.attributes import flag_modified
from src.config import config
from src.core.cache_service import CacheService
@@ -96,6 +97,7 @@ class ProviderOpsService:
SENSITIVE_FIELDS = {
"api_key",
"password",
"refresh_token",
"session_token",
"session_cookie",
"token_cookie",
@@ -193,10 +195,11 @@ class ProviderOpsService:
# 加密敏感凭据
encrypted_credentials = self._encrypt_credentials(config.connector_credentials)
logger.debug(
f"加密凭据: provider_id={provider_id}, "
f"input_keys={list(config.connector_credentials.keys())}, "
f"output_keys={list(encrypted_credentials.keys())}, "
f"has_api_key={bool(config.connector_credentials.get('api_key'))}"
"加密凭据: provider_id={}, input_keys={}, output_keys={}, has_api_key={}",
provider_id,
list(config.connector_credentials.keys()),
list(encrypted_credentials.keys()),
bool(config.connector_credentials.get("api_key")),
)
# 构建配置
@@ -214,7 +217,7 @@ class ProviderOpsService:
if provider_id in self._connectors:
del self._connectors[provider_id]
logger.info(f"保存 Provider 操作配置: provider_id={provider_id}")
logger.info("保存 Provider 操作配置: provider_id={}", provider_id)
return True
def delete_config(self, provider_id: str) -> bool:
@@ -301,8 +304,14 @@ class ProviderOpsService:
# 建立连接
logger.info(
f"尝试连接: provider_id={provider_id}, "
f"credentials_keys={list(actual_credentials.keys())}"
"尝试连接: provider_id={}, credentials_keys={}",
provider_id,
list(actual_credentials.keys()),
)
# 注册凭据更新回调Token Rotation 场景持久化新 refresh_token
# 必须在 connect 之前注册,因为 connect 内部可能已经触发 Token Rotation
connector._on_credentials_updated = (
lambda updated, pid=provider_id: self._persist_updated_credentials(pid, updated)
)
success = await connector.connect(actual_credentials)
if success:
@@ -491,11 +500,11 @@ class ProviderOpsService:
# 没有缓存
if allow_sync_query:
# 同步查询一次(首次访问)
logger.info(f"余额缓存未命中,同步查询: provider_id={provider_id}")
logger.info("余额缓存未命中,同步查询: provider_id={}", provider_id)
return await self.query_balance(provider_id)
else:
# 仅触发异步刷新,立即返回
logger.debug(f"余额缓存未命中,触发异步刷新: provider_id={provider_id}")
logger.debug("余额缓存未命中,触发异步刷新: provider_id={}", provider_id)
asyncio.create_task(self._refresh_balance_async(provider_id))
return ActionResult(
status=ActionStatus.PENDING,
@@ -520,7 +529,7 @@ class ProviderOpsService:
# 使用 wait_for 设置超时,避免无限等待
await asyncio.wait_for(semaphore.acquire(), timeout=5.0)
except asyncio.TimeoutError:
logger.debug(f"异步刷新余额跳过(并发限制): provider_id={provider_id}")
logger.debug("异步刷新余额跳过(并发限制): provider_id={}", provider_id)
return
db = None
@@ -530,7 +539,7 @@ class ProviderOpsService:
service = ProviderOpsService(db)
await service.query_balance(provider_id)
except Exception as e:
logger.warning(f"异步刷新余额失败: provider_id={provider_id}, error={e}")
logger.warning("异步刷新余额失败: provider_id={}, error={}", provider_id, e)
finally:
# 确保 session 被关闭,归还连接到连接池
if db is not None:
@@ -545,7 +554,7 @@ class ProviderOpsService:
"""清除余额缓存"""
cache_key = f"provider_ops:balance:{provider_id}"
await CacheService.delete(cache_key)
logger.info(f"余额缓存已清除: provider_id={provider_id}")
logger.info("余额缓存已清除: provider_id={}", provider_id)
async def _cache_auth_failed(self, provider_id: str, result: ActionResult) -> None:
"""
@@ -564,7 +573,9 @@ class ProviderOpsService:
}
await CacheService.set(cache_key, cache_data, AUTH_FAILED_CACHE_TTL)
logger.info(
f"余额缓存已写入(认证失败): provider_id={provider_id}, message={result.message}"
"余额缓存已写入(认证失败): provider_id={}, message={}",
provider_id,
result.message,
)
async def _cache_balance(self, provider_id: str, result: ActionResult) -> None:
@@ -617,7 +628,7 @@ class ProviderOpsService:
}
await CacheService.set(cache_key, cache_data, BALANCE_CACHE_TTL)
logger.debug(f"验证成功,缓存余额: provider_id={provider_id}, quota_usd={quota_usd}")
logger.debug("验证成功,缓存余额: provider_id={}, quota_usd={}", provider_id, quota_usd)
async def _get_cached_balance(self, provider_id: str) -> ActionResult | None:
"""获取缓存的余额"""
@@ -660,7 +671,7 @@ class ProviderOpsService:
cache_ttl_seconds=ttl,
)
except Exception as e:
logger.warning(f"解析缓存余额失败: provider_id={provider_id}, error={e}")
logger.warning("解析缓存余额失败: provider_id={}, error={}", provider_id, e)
return None
async def checkin(
@@ -705,6 +716,78 @@ class ProviderOpsService:
return None
def _persist_updated_credentials(self, provider_id: str, updated: dict[str, Any]) -> None:
"""
持久化连接器运行时更新的凭据(如 Token Rotation 后的新 refresh_token
通过 run_in_executor 将同步 DB 操作 offload 到线程池,避免在异步事件循环中阻塞。
"""
try:
loop = asyncio.get_running_loop()
future = loop.run_in_executor(
None, self._persist_updated_credentials_sync, provider_id, updated
)
def _log_persist_error(f: asyncio.Future) -> None: # type: ignore[type-arg]
if not f.cancelled() and f.exception():
logger.warning("异步持久化凭据失败: {}", f.exception())
future.add_done_callback(_log_persist_error)
except RuntimeError:
# 没有运行中的事件循环(测试等场景),直接同步执行
self._persist_updated_credentials_sync(provider_id, updated)
def _persist_updated_credentials_sync(self, provider_id: str, updated: dict[str, Any]) -> None:
"""同步执行凭据持久化(在线程池中运行)"""
import copy
db = None
try:
db = create_session()
provider = db.query(Provider).filter(Provider.id == provider_id).first()
if not provider:
return
# 深拷贝整个 config避免原地修改导致 SQLAlchemy 变更检测失败
provider_config = copy.deepcopy(dict(provider.config or {}))
config_data = provider_config.get("provider_ops")
if not config_data:
return
credentials = config_data.get("connector", {}).get("credentials", {})
# 更新凭据(敏感字段加密)
for key, value in updated.items():
if key in self.SENSITIVE_FIELDS and isinstance(value, str) and value:
credentials[key] = self.crypto.encrypt(value)
else:
credentials[key] = value
config_data["connector"]["credentials"] = credentials
provider.config = provider_config
flag_modified(provider, "config")
db.commit()
logger.info(
"凭据已持久化更新: provider_id={}, updated_keys={}",
provider_id,
list(updated.keys()),
)
except Exception as e:
logger.warning("持久化凭据更新失败: provider_id={}, error={}", provider_id, e)
if db:
try:
db.rollback()
except Exception:
pass
finally:
if db:
try:
db.close()
except Exception:
pass
def _encrypt_credentials(self, credentials: dict[str, Any]) -> dict[str, Any]:
"""加密凭据中的敏感字段"""
encrypted = {}
@@ -713,10 +796,13 @@ class ProviderOpsService:
if value: # 只加密非空值
encrypted[key] = self.crypto.encrypt(value)
logger.debug(
f"加密字段 {key}: 原始长度={len(value)}, 加密后长度={len(encrypted[key])}"
"加密字段 {}: 原始长度={}, 加密后长度={}",
key,
len(value),
len(encrypted[key]),
)
else:
logger.warning(f"跳过空值字段 {key}")
logger.warning("跳过空值字段 {}", key)
encrypted[key] = value
else:
encrypted[key] = value
@@ -730,17 +816,22 @@ class ProviderOpsService:
try:
decrypted[key] = self.crypto.decrypt(value)
except Exception as e:
logger.warning(f"解密字段 {key} 失败: {e}")
logger.warning("解密字段 {} 失败: {}", key, e)
decrypted[key] = value # 解密失败则保持原值
else:
decrypted[key] = value
return decrypted
# 密码类字段:脱敏时全部遮盖,不显示任何明文字符
FULLY_MASKED_FIELDS = {"password"}
def get_masked_credentials(self, credentials: dict[str, Any]) -> dict[str, Any]:
"""
获取脱敏后的凭据
解密凭据并对敏感字段进行脱敏处理(显示部分字符)
解密凭据并对敏感字段进行脱敏处理。
- 密码类字段:全部遮盖为 ********
- 其他敏感字段:显示部分字符(如 sk-x****a12k
Args:
credentials: 加密的凭据
@@ -753,8 +844,10 @@ class ProviderOpsService:
for field in self.SENSITIVE_FIELDS:
if field in decrypted and decrypted[field]:
value = str(decrypted[field])
# 显示前4位和后4位中间固定4个 *(如 sk-x****a12k
if len(value) > 12:
if field in self.FULLY_MASKED_FIELDS:
decrypted[field] = "********"
elif len(value) > 12:
# 显示前4位和后4位中间固定4个 *(如 sk-x****a12k
decrypted[field] = value[:4] + "****" + value[-4:]
elif len(value) > 8:
decrypted[field] = value[:2] + "****" + value[-2:]
@@ -788,6 +881,7 @@ class ProviderOpsService:
sensitive_fields = [
"api_key",
"password",
"refresh_token",
"session_token",
"cookie_string",
"cookie",
@@ -802,7 +896,12 @@ class ProviderOpsService:
if not req_value or (isinstance(req_value, str) and set(req_value) <= {"*"}):
if field in saved_credentials:
merged[field] = saved_credentials[field]
logger.debug(f"合并凭据 - 使用已保存的 {field}")
logger.debug("合并凭据 - 使用已保存的 {}", field)
# 保留内部缓存字段(如 _cached_access_token前端不感知这些字段
for key, value in saved_credentials.items():
if key.startswith("_") and key not in merged:
merged[key] = value
return merged
@@ -847,7 +946,7 @@ class ProviderOpsService:
)
return provider_id, result
except Exception as e:
logger.warning(f"查询余额失败: provider_id={provider_id}, error={e}")
logger.warning("查询余额失败: provider_id={}, error={}", provider_id, e)
return provider_id, ActionResult(
status=ActionStatus.UNKNOWN_ERROR,
action_type=ProviderActionType.QUERY_BALANCE,
@@ -905,15 +1004,33 @@ class ProviderOpsService:
# 使用架构的方法构建请求
verify_endpoint = f"{base_url}{architecture.get_verify_endpoint()}"
# 执行异步预处理(如获取动态 Cookie
extra_config = await architecture.prepare_verify_config(base_url, config, credentials)
# 执行异步预处理(如获取动态 Cookie、登录获取 Token
# 返回值可以是 dict仅额外配置或 tuple[dict, dict](额外配置 + 凭据更新)
prepare_result = await architecture.prepare_verify_config(base_url, config, credentials)
if isinstance(prepare_result, tuple):
extra_config, updated_creds = prepare_result
else:
extra_config = prepare_result
updated_creds = {}
merged_config = {**config, **extra_config}
# Token Rotation: prepare_verify_config 可能已消耗旧 refresh_token 并获取新值,
# 无论后续验证是否成功都需要立即持久化,否则旧 token 已失效但数据库未更新。
if updated_creds and provider_id:
logger.info(
"验证过程检测到凭据变更: provider_id={}, updated_keys={}",
provider_id,
list(updated_creds.keys()),
)
self._persist_updated_credentials(provider_id, updated_creds)
headers = architecture.build_verify_headers(merged_config, credentials)
logger.debug(
f"验证认证: architecture={architecture_id}, "
f"endpoint={verify_endpoint}, headers={list(headers.keys())}"
"验证认证: architecture={}, endpoint={}, headers={}",
architecture_id,
verify_endpoint,
list(headers.keys()),
)
# 获取代理配置(支持 proxy_node_id 和旧的 proxy URL
@@ -929,14 +1046,15 @@ class ProviderOpsService:
}
if proxy:
client_kwargs["proxy"] = proxy
logger.debug(f"使用代理: {proxy}")
logger.debug("使用代理: {}", proxy)
async with httpx.AsyncClient(**client_kwargs) as client:
response = await client.get(verify_endpoint, headers=headers)
logger.debug(
f"验证响应: status={response.status_code}, "
f"content_type={response.headers.get('content-type')}"
"验证响应: status={}, content_type={}",
response.status_code,
response.headers.get("content-type"),
)
# 尝试解析 JSON
@@ -955,6 +1073,12 @@ class ProviderOpsService:
result = architecture.parse_verify_response(response.status_code, data)
result_dict = result.to_dict()
# 将凭据更新信息附加到响应,供前端同步更新表单
# 过滤掉内部缓存字段(以 _ 开头),前端不需要这些
frontend_creds = {k: v for k, v in updated_creds.items() if not k.startswith("_")}
if frontend_creds:
result_dict["updated_credentials"] = frontend_creds
# 验证成功且有 provider_id 时,缓存余额
if result.success and provider_id and result.quota is not None:
# 从架构的默认配置获取 quota_divisor
@@ -974,5 +1098,5 @@ class ProviderOpsService:
except httpx.ConnectError as e:
return {"success": False, "message": f"连接失败: {str(e)}"}
except Exception as e:
logger.error(f"验证认证失败: {e}")
logger.error("验证认证失败: {}", e)
return {"success": False, "message": f"验证失败: {str(e)}"}

View File

@@ -81,7 +81,7 @@ class ConcurrencyManager:
# 内存模式下启动后台清理任务
self._start_background_cleanup()
except Exception as e:
logger.error(f"[ERROR] 获取全局 Redis 客户端失败: {e}")
logger.error("[ERROR] 获取全局 Redis 客户端失败: {}", e)
logger.warning("[WARN] RPM 限制将降级为内存模式(仅在单实例环境下安全)")
self._redis = None
self._owns_redis = False
@@ -104,7 +104,7 @@ class ConcurrencyManager:
except asyncio.CancelledError:
break
except Exception as e:
logger.debug(f"后台清理任务异常: {e}")
logger.debug("后台清理任务异常: {}", e)
try:
self._cleanup_task = asyncio.create_task(cleanup_loop())
@@ -164,21 +164,26 @@ class ConcurrencyManager:
# 分级告警:根据使用率记录不同级别的日志
if current_size >= critical_threshold and key_id not in self._memory_key_rpm_counts:
logger.critical(
f"[CRITICAL] 内存 RPM 计数器接近上限 ({current_size}/{self._max_memory_rpm_entries})"
f"强烈建议启用 Redis继续增长可能导致 RPM 限制失效"
"[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(
f"[ERROR] 内存 RPM 计数器使用率过高 ({current_size}/{self._max_memory_rpm_entries})"
f"建议启用 Redis"
"[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(
f"[WARN] 内存 RPM 计数器达到 {self._memory_warning_threshold:.0%} 阈值 "
f"({current_size}/{self._max_memory_rpm_entries}),建议启用 Redis"
"[WARN] 内存 RPM 计数器达到 {:.0%} 阈值 ({}/{}),建议启用 Redis",
self._memory_warning_threshold,
current_size,
self._max_memory_rpm_entries,
)
# 检查是否超过最大条目限制
@@ -197,7 +202,7 @@ class ConcurrencyManager:
)
for k, _ in sorted_keys[:evict_count]:
del self._memory_key_rpm_counts[k]
logger.warning(f"[WARN] 内存 RPM 计数器达到上限,已淘汰 {evict_count} 个最旧条目")
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:
@@ -232,7 +237,7 @@ class ConcurrencyManager:
del self._memory_key_rpm_counts[key_id]
if expired_keys:
logger.debug(f"[CLEANUP] 清理了 {len(expired_keys)} 个过期的内存 RPM 计数")
logger.debug("[CLEANUP] 清理了 {} 个过期的内存 RPM 计数", len(expired_keys))
async def get_key_rpm_count(self, key_id: str) -> int:
"""
@@ -256,7 +261,7 @@ class ConcurrencyManager:
result = await self._redis.get(key_key)
return int(result) if result else 0
except Exception as e:
logger.error(f"获取 RPM 计数失败: {e}")
logger.error("获取 RPM 计数失败: {}", e)
return 0
async def check_rpm_available(
@@ -405,7 +410,7 @@ class ConcurrencyManager:
if success:
user_type = "缓存用户" if is_cached_user else "新用户"
logger.debug(f"[OK] 获取 RPM 槽位成功: key={key_id}, 类型={user_type}")
logger.debug("[OK] 获取 RPM 槽位成功: key={}, 类型={}", key_id, user_type)
else:
key_count = await self.get_key_rpm_count(key_id)
@@ -416,12 +421,12 @@ class ConcurrencyManager:
else:
user_info = f"缓存用户, 当前={key_count}/{key_rpm_limit}"
logger.warning(f"[WARN] RPM 限制已达上限: key={key_id}({user_info})")
logger.warning("[WARN] RPM 限制已达上限: key={}({})", key_id, user_info)
return success
except Exception as e:
logger.error(f"获取 RPM 槽位失败,降级到内存模式: {e}")
logger.error("获取 RPM 槽位失败,降级到内存模式: {}", e)
# Redis 异常时降级到内存模式进行保守限流
async with self._memory_lock:
bucket = self._get_rpm_bucket()
@@ -436,13 +441,15 @@ class ConcurrencyManager:
if fallback_rpm_limit is not None and key_count >= fallback_rpm_limit:
logger.warning(
f"[FALLBACK] Key RPM 达到降级限制: {key_count}/{fallback_rpm_limit}"
"[FALLBACK] Key RPM 达到降级限制: {}/{}",
key_count,
fallback_rpm_limit,
)
return False
# 更新内存计数
self._set_memory_key_rpm_count(key_id, bucket, key_count + 1)
logger.debug(f"[FALLBACK] 使用内存模式获取 RPM 槽位: key={key_id}")
logger.debug("[FALLBACK] 使用内存模式获取 RPM 槽位: key={}", key_id)
return True
@asynccontextmanager
@@ -485,8 +492,8 @@ class ConcurrencyManager:
if not acquired:
from src.core.exceptions import ConcurrencyLimitError
user_type = "缓存用户" if is_cached_user else "新用户"
raise ConcurrencyLimitError(f"RPM 限制已达上限: key={key_id}, 类型={user_type}")
# Keep the client-facing message generic; do not leak internal IDs.
raise ConcurrencyLimitError("服务暂时繁忙,请稍后重试")
# 记录开始时间和状态
import time
@@ -525,15 +532,15 @@ class ConcurrencyManager:
# 告警:槽位占用时间过长(超过 60 秒)
if slot_duration > 60:
logger.warning(
f"[WARN] 请求耗时过长: "
f"key_id={key_id[:8] if key_id else 'unknown'}..., "
f"duration={slot_duration:.1f}s, "
f"exception={exception_occurred}"
"[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(f"记录指标失败: {metric_error}")
logger.debug("记录指标失败: {}", metric_error)
# 注意RPM 计数不需要在请求结束后释放,它会在分钟窗口过期后自动重置
@@ -547,14 +554,14 @@ class ConcurrencyManager:
if self._redis is None:
async with self._memory_lock:
self._memory_key_rpm_counts.pop(key_id, None)
logger.info(f"[RESET] 重置 Key RPM 计数(内存): {key_id}")
logger.info("[RESET] 重置 Key RPM 计数(内存): {}", key_id)
return
try:
deleted_count = await self._scan_and_delete(f"rpm:key:{key_id}:*")
logger.info(f"[RESET] 重置 Key RPM 计数: {key_id}, 删除 {deleted_count} 个键")
logger.info("[RESET] 重置 Key RPM 计数: {}, 删除 {} 个键", key_id, deleted_count)
except Exception as e:
logger.error(f"重置 Key RPM 计数失败: {e}")
logger.error("重置 Key RPM 计数失败: {}", e)
async def reset_all_rpm(self) -> None:
"""重置所有 Key RPM 计数(管理功能,慎用)"""
@@ -563,15 +570,15 @@ class ConcurrencyManager:
count = len(self._memory_key_rpm_counts)
self._memory_key_rpm_counts.clear()
if count:
logger.info(f"[RESET] 重置所有 Key RPM 计数(内存): {count}")
logger.info("[RESET] 重置所有 Key RPM 计数(内存): {}", count)
return
try:
deleted_count = await self._scan_and_delete("rpm:key:*")
if deleted_count:
logger.info(f"[RESET] 重置所有 Key RPM 计数: {deleted_count}")
logger.info("[RESET] 重置所有 Key RPM 计数: {}", deleted_count)
except Exception as e:
logger.error(f"重置所有 Key RPM 计数失败: {e}")
logger.error("重置所有 Key RPM 计数失败: {}", e)
async def _scan_and_delete(self, pattern: str, batch_size: int = 100) -> int:
"""使用 SCAN 遍历并分批删除匹配的键,避免阻塞 Redis"""

View File

@@ -247,7 +247,13 @@ class RequestCandidateService:
@staticmethod
def mark_candidate_skipped(
db: Session, candidate_id: str, skip_reason: str | None = None
db: Session,
candidate_id: str,
skip_reason: str | None = None,
*,
status_code: int | None = None,
concurrent_requests: int | None = None,
extra_data: dict | None = None,
) -> None:
"""
标记候选为已跳过
@@ -256,12 +262,25 @@ class RequestCandidateService:
db: 数据库会话
candidate_id: 候选ID
skip_reason: 跳过原因
status_code: HTTP 状态码(可选)
concurrent_requests: 并发请求数(这里实际记录 RPM 计数)
extra_data: 额外数据(合并写入)
"""
candidate = db.query(RequestCandidate).filter(RequestCandidate.id == candidate_id).first()
if candidate:
candidate.status = "skipped"
candidate.skip_reason = skip_reason
candidate.finished_at = datetime.now(timezone.utc)
if status_code is not None:
candidate.status_code = int(status_code)
if concurrent_requests is not None:
candidate.concurrent_requests = int(concurrent_requests)
if extra_data:
base = candidate.extra_data if isinstance(candidate.extra_data, dict) else {}
candidate.extra_data = {**base, **extra_data}
db.flush() # 只 flush不立即 commit
get_batch_committer().mark_dirty(db)
@@ -364,5 +383,5 @@ class RequestCandidateService:
first_byte_epoch_ms = request_start_epoch_ms + global_first_byte_time_ms
return max(0, int(first_byte_epoch_ms - started_at_epoch_ms))
except Exception as e:
logger.debug(f"计算候选 TTFB 失败: {e}")
logger.debug("计算候选 TTFB 失败: {}", e)
return global_first_byte_time_ms

View File

@@ -4,6 +4,7 @@
from __future__ import annotations
import math
import time
from collections.abc import Callable
from dataclasses import dataclass
@@ -34,6 +35,13 @@ class ExecutionContext:
start_time: float | None = None
elapsed_ms: int | None = None
concurrent_requests: int | None = None
rpm_current: int | None = None
rpm_limit: int | None = None
rpm_available_for_new: int | None = None
reservation_ratio: float | None = None
reservation_phase: str | None = None
reservation_confidence: float | None = None
reservation_load_factor: float | None = None
@dataclass
@@ -100,9 +108,13 @@ class RequestExecutor:
key_id=key.id,
)
except Exception as e:
logger.debug(f"获取 RPM 计数失败(用于预留计算): {e}")
logger.debug("获取 RPM 计数失败(用于预留计算): {}", e)
current_key_rpm = 0
# 在获取 guard 之前记录当前 RPM 计数,便于并发拒绝场景落库
context.concurrent_requests = current_key_rpm
context.rpm_current = current_key_rpm
# 获取有效的 RPM 限制(自适应或固定)
effective_key_limit = get_adaptive_rpm_manager().get_effective_limit(key)
@@ -113,10 +125,23 @@ class RequestExecutor:
)
dynamic_reservation_ratio = reservation_result.ratio
context.rpm_limit = effective_key_limit
context.reservation_ratio = dynamic_reservation_ratio
context.reservation_phase = reservation_result.phase
context.reservation_confidence = reservation_result.confidence
context.reservation_load_factor = reservation_result.load_factor
if effective_key_limit is not None and not is_cached_user:
context.rpm_available_for_new = max(
1, math.floor(effective_key_limit * (1 - dynamic_reservation_ratio))
)
logger.debug(
f"[Executor] 动态预留: key={key.id[:8]}..., "
f"ratio={dynamic_reservation_ratio:.0%}, phase={reservation_result.phase}, "
f"confidence={reservation_result.confidence:.0%}"
"[Executor] 动态预留: key={}..., ratio={:.0%}, phase={}, confidence={:.0%}",
key.id[:8],
dynamic_reservation_ratio,
reservation_result.phase,
reservation_result.confidence,
)
async with self.concurrency_manager.rpm_guard(
@@ -131,10 +156,11 @@ class RequestExecutor:
key_id=key.id,
)
except Exception as e:
logger.debug(f"获取 RPM 计数失败guard 内): {e}")
logger.debug("获取 RPM 计数失败guard 内): {}", e)
key_rpm_count = None
context.concurrent_requests = key_rpm_count # 用于记录,实际是 RPM 计数
if key_rpm_count is not None:
context.concurrent_requests = key_rpm_count # 用于记录,实际是 RPM 计数
context.start_time = time.time()
response = await request_func(provider, endpoint, key, candidate)

View File

@@ -740,17 +740,99 @@ class TaskService:
has_retry_left = retry_index < (max_retries_for_candidate - 1)
if isinstance(cause, ConcurrencyLimitError):
rpm_current = context.rpm_current
if rpm_current is None:
rpm_current = captured_key_concurrent
rpm_limit = context.rpm_limit
rpm_available_for_new = context.rpm_available_for_new
reservation_ratio = context.reservation_ratio
reservation_phase = context.reservation_phase or "unknown"
reservation_confidence = context.reservation_confidence
reservation_load_factor = context.reservation_load_factor
reason_code = "unknown"
if rpm_limit is not None and rpm_current is not None:
if context.is_cached_user:
if rpm_current >= rpm_limit:
reason_code = "total_limit"
else:
if rpm_available_for_new is not None and rpm_current >= rpm_available_for_new:
reason_code = (
"reserved_for_cached" if rpm_current < rpm_limit else "total_limit"
)
elif rpm_current >= rpm_limit:
reason_code = "total_limit"
reason_text = "并发限制"
if reason_code == "reserved_for_cached":
reason_text = "并发限制: 新用户配额已满(预留给缓存用户)"
elif reason_code == "total_limit":
reason_text = "并发限制: 总配额已满"
parts: list[str] = []
if rpm_current is not None:
parts.append(f"current={rpm_current}")
if rpm_limit is not None:
parts.append(f"limit={rpm_limit}")
if rpm_available_for_new is not None and not context.is_cached_user:
parts.append(f"new={rpm_available_for_new}")
if reservation_ratio is not None:
parts.append(f"reserve={reservation_ratio:.0%}")
if reservation_phase:
parts.append(f"phase={reservation_phase}")
skip_reason = reason_text
if parts:
skip_reason = f"{reason_text} ({', '.join(parts)})"
logger.warning(
" [{}] 并发限制 (attempt={}/{}): {}",
" [{}] 并发限制 (attempt={}/{}): provider={}, key={}, cached={}, reason={}, {}",
request_id,
attempt,
max_attempts,
str(cause),
provider.name,
str(key.id)[:8],
bool(context.is_cached_user),
reason_code,
", ".join(parts) if parts else "N/A",
)
extra_data: dict[str, Any] = {
"concurrency_denied": True,
"concurrency_reason": reason_code,
"rpm_current": rpm_current,
"rpm_limit": rpm_limit,
"rpm_available_for_new": rpm_available_for_new,
"reservation_ratio": reservation_ratio,
"reservation_phase": reservation_phase,
"reservation_confidence": reservation_confidence,
"reservation_load_factor": reservation_load_factor,
"attempt": attempt,
"max_attempts": max_attempts,
}
extra_data = {k: v for k, v in extra_data.items() if v is not None}
if _proxy_extra:
extra_data = {**_proxy_extra, **extra_data}
try:
from src.core.metrics import scheduler_concurrency_denied_total
scheduler_concurrency_denied_total.labels(
is_cached_user=str(bool(context.is_cached_user)).lower(),
reason=reason_code,
reservation_phase=str(reservation_phase or "unknown"),
).inc()
except Exception:
pass
RequestCandidateService.mark_candidate_skipped(
db=self.db,
candidate_id=candidate_record_id,
skip_reason=f"并发限制: {str(cause)}",
skip_reason=skip_reason,
status_code=429,
concurrent_requests=rpm_current,
extra_data=extra_data,
)
return "break"