mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 17:30:23 +08:00
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:
604
src/services/cache/aware_scheduler.py
vendored
604
src/services/cache/aware_scheduler.py
vendored
@@ -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"]
|
||||
)
|
||||
|
||||
@@ -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
|
||||
)
|
||||
|
||||
@@ -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))
|
||||
|
||||
# 按哈希值排序
|
||||
144
src/services/cache/concurrency_checker.py
vendored
Normal file
144
src/services/cache/concurrency_checker.py
vendored
Normal 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()
|
||||
82
src/services/cache/restriction_checker.py
vendored
Normal file
82
src/services/cache/restriction_checker.py
vendored
Normal 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
82
src/services/cache/scheduling_config.py
vendored
Normal 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
88
src/services/cache/schemas.py
vendored
Normal 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
53
src/services/cache/utils.py
vendored
Normal 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
|
||||
@@ -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
|
||||
|
||||
@@ -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": [
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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)}"}
|
||||
|
||||
@@ -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"""
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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"
|
||||
|
||||
|
||||
Reference in New Issue
Block a user