feat(pool): 引入多维评分调度策略与账号状态检测

- 新增 multi_score 调度模式,支持 LRU/延迟/健康度/剩余额度多维加权评分
- 新增调度预设维度系统(free_team_first, quota_balanced, recent_refresh, single_account),支持有序对象列表配置格式并兼容旧字符串列表
- 新增 account_state 模块,统一账号封禁/受限检测逻辑,替代分散在 routes 中的判断代码
- 新增 health_cache 模块和 latency 采样(redis_ops.record_latency / batch_get_latency_avgs)
- RequestDispatcher 返回 ttfb_ms,PoolManager.on_request_success 记录延迟样本
- 前端:PoolConfigDialog 替换为 PoolSchedulingDialog,支持预设维度可视化配置;号池管理页增加调度模式标签与账号异常 Badge 显示
- 提取前端 accountBlock 工具函数,ProviderDetailDrawer 复用统一判断
- scheduling_dimensions 增加 account_state 和 latency 维度评估
- 补充 account_state、health_cache、multi_score 策略、preset 维度、redis latency 等测试
This commit is contained in:
fawney19
2026-03-04 22:06:19 +08:00
parent 57b86034cf
commit b2dcf82ca8
40 changed files with 4167 additions and 660 deletions

View File

@@ -28,7 +28,9 @@ from src.core.logger import logger
from src.database import get_db
from src.models.database import Provider, ProviderAPIKey, Usage
from src.services.provider.pool import redis_ops as pool_redis
from src.services.provider.pool.account_state import resolve_pool_account_state
from src.services.provider.pool.config import parse_pool_config
from src.services.provider.pool.dimensions import get_preset_dimension_metas
from src.services.provider.pool.scheduling_dimensions import (
PoolSchedulingSnapshot,
evaluate_pool_scheduling_dimensions,
@@ -47,6 +49,8 @@ from .schemas import (
PoolOverviewResponse,
PoolSchedulingDimension,
PoolSchedulingReason,
PresetDimensionMetaResponse,
PresetModeMetaResponse,
)
router = APIRouter(prefix="/api/admin/pool", tags=["pool-management"])
@@ -68,6 +72,31 @@ async def pool_overview(
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
# ---------------------------------------------------------------------------
# GET /api/admin/pool/scheduling-presets
# ---------------------------------------------------------------------------
def _preset_mode_label(mode: str) -> str:
mapping = {
"free_only": "Free",
"team_only": "Team",
"both": "全部",
}
return mapping.get(mode, mode)
@router.get("/scheduling-presets", response_model=list[PresetDimensionMetaResponse])
async def list_scheduling_presets(
request: Request,
db: Session = Depends(get_db),
) -> list[PresetDimensionMetaResponse]:
"""Return scheduling preset definitions for frontend rendering."""
adapter = AdminListSchedulingPresetsAdapter()
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
# ---------------------------------------------------------------------------
# GET /api/admin/pool/{provider_id}/keys
# ---------------------------------------------------------------------------
@@ -126,24 +155,6 @@ _COOLDOWN_REASON_LABELS: dict[str, str] = {
"server_error_500": "500 错误",
}
_ACCOUNT_BLOCK_REASON_KEYWORDS: tuple[str, ...] = (
"account_block",
"account blocked",
"account has been disabled",
"account disabled",
"organization has been disabled",
"organization_disabled",
"validation_required",
"verify your account",
"forbidden",
"suspended",
"封禁",
"封号",
"被封",
"访问被禁止",
"账号异常",
)
def _to_float(value: Any) -> float | None:
if isinstance(value, bool):
@@ -161,65 +172,15 @@ def _to_float(value: Any) -> float | None:
return None
def _is_truthy_flag(value: Any) -> bool:
if isinstance(value, bool):
return value
if isinstance(value, (int, float)):
return value != 0
if isinstance(value, str):
normalized = value.strip().lower()
return normalized in {"1", "true", "yes", "y"}
return False
def _is_known_banned_reason(reason: str | None) -> bool:
if not reason:
return False
text = str(reason).strip()
if not text:
return False
lowered = text.lower()
# 结构化账号级别封禁标记(如 [ACCOUNT_BLOCK] ...
try:
from src.services.provider.oauth_token import is_account_level_block
if is_account_level_block(text):
return True
except Exception:
pass
return any(keyword in lowered for keyword in _ACCOUNT_BLOCK_REASON_KEYWORDS)
def _is_known_banned_key(key: ProviderAPIKey, provider_type: str) -> bool:
upstream_metadata = getattr(key, "upstream_metadata", None)
normalized_provider = provider_type.strip().lower()
provider_bucket: dict[str, Any] | None = None
if isinstance(upstream_metadata, dict):
maybe_bucket = upstream_metadata.get(normalized_provider)
if isinstance(maybe_bucket, dict):
provider_bucket = maybe_bucket
from src.services.provider.pool.account_state import resolve_pool_account_state
if normalized_provider == "kiro" and provider_bucket:
if _is_truthy_flag(provider_bucket.get("is_banned")):
return True
if normalized_provider == "antigravity" and provider_bucket:
if _is_truthy_flag(provider_bucket.get("is_forbidden")):
return True
for source in (provider_bucket, upstream_metadata):
if not isinstance(source, dict):
continue
if _is_truthy_flag(source.get("is_banned")):
return True
if _is_truthy_flag(source.get("is_forbidden")):
return True
if _is_truthy_flag(source.get("account_disabled")):
return True
return _is_known_banned_reason(getattr(key, "oauth_invalid_reason", None))
state = resolve_pool_account_state(
provider_type=provider_type,
upstream_metadata=getattr(key, "upstream_metadata", None),
oauth_invalid_reason=getattr(key, "oauth_invalid_reason", None),
)
return state.blocked
def _format_percent(value: float) -> str:
@@ -536,6 +497,10 @@ def _format_cooldown_detail(raw: str | None) -> str | None:
def _build_pool_scheduling_state(
*,
is_active: bool,
account_blocked: bool,
account_block_label: str | None,
account_block_reason: str | None,
latency_avg_ms: float | None,
cooldown_reason: str | None,
cooldown_ttl_seconds: int | None,
circuit_breaker_open: bool,
@@ -557,6 +522,10 @@ def _build_pool_scheduling_state(
"""Build unified scheduling state for frontend display."""
snapshot = PoolSchedulingSnapshot(
is_active=is_active,
account_blocked=account_blocked,
account_block_label=account_block_label,
account_block_reason=account_block_reason,
latency_avg_ms=latency_avg_ms,
cooldown_reason=cooldown_reason,
cooldown_ttl_seconds=cooldown_ttl_seconds,
circuit_breaker_open=circuit_breaker_open,
@@ -652,6 +621,39 @@ async def cleanup_banned_keys(
# ---------------------------------------------------------------------------
class AdminListSchedulingPresetsAdapter(AdminApiAdapter):
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
items: list[PresetDimensionMetaResponse] = [
PresetDimensionMetaResponse(
name="lru",
label="LRU 轮转",
description="最久未使用的 Key 优先",
providers=[],
modes=None,
default_mode=None,
)
]
for meta in get_preset_dimension_metas():
modes = None
if meta.modes:
modes = [
PresetModeMetaResponse(value=mode, label=_preset_mode_label(mode))
for mode in meta.modes
]
items.append(
PresetDimensionMetaResponse(
name=meta.name,
label=meta.label,
description=meta.description,
providers=list(meta.providers),
modes=modes,
default_mode=meta.default_mode,
)
)
return items
class AdminPoolOverviewAdapter(AdminApiAdapter):
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
@@ -813,20 +815,34 @@ class AdminListPoolKeysAdapter(AdminApiAdapter):
if pcfg and pcfg.lru_enabled
else asyncio.sleep(0, result={})
)
_latency_coro = (
pool_redis.batch_get_latency_avgs(pid, key_ids, pcfg.latency_window_seconds)
if pcfg and pcfg.scheduling_mode == "multi_score"
else asyncio.sleep(0, result={})
)
_cost_coro = (
pool_redis.batch_get_cost_totals(pid, key_ids, pcfg.cost_window_seconds)
if pcfg
else asyncio.sleep(0, result={})
)
cooldowns, cooldown_ttls, lru_scores, cost_totals, sticky_counts = await asyncio.gather(
(
cooldowns,
cooldown_ttls,
lru_scores,
latency_avgs,
cost_totals,
sticky_counts,
) = await asyncio.gather(
pool_redis.batch_get_cooldowns(pid, key_ids),
pool_redis.batch_get_cooldown_ttls(pid, key_ids),
_lru_coro,
_latency_coro,
_cost_coro,
pool_redis.batch_get_key_sticky_counts(pid, key_ids),
)
else:
cooldowns, cooldown_ttls, lru_scores, cost_totals, sticky_counts = (
cooldowns, cooldown_ttls, lru_scores, latency_avgs, cost_totals, sticky_counts = (
{},
{},
{},
{},
@@ -874,6 +890,13 @@ class AdminListPoolKeysAdapter(AdminApiAdapter):
)
cost_usage = int(cost_totals.get(kid, 0) or 0)
cost_limit = pcfg.cost_limit_per_key_tokens if pcfg else None
latency_avg_raw = latency_avgs.get(kid)
latency_avg_ms = float(latency_avg_raw) if latency_avg_raw is not None else None
account_state = resolve_pool_account_state(
provider_type=provider_type,
upstream_metadata=getattr(k, "upstream_metadata", None),
oauth_invalid_reason=getattr(k, "oauth_invalid_reason", None),
)
(
scheduling_status,
scheduling_reason,
@@ -886,6 +909,10 @@ class AdminListPoolKeysAdapter(AdminApiAdapter):
scheduling_dimensions,
) = _build_pool_scheduling_state(
is_active=bool(k.is_active),
account_blocked=account_state.blocked,
account_block_label=account_state.label,
account_block_reason=account_state.reason,
latency_avg_ms=latency_avg_ms,
cooldown_reason=cd_reason,
cooldown_ttl_seconds=cd_ttl,
circuit_breaker_open=any_circuit_open,
@@ -955,9 +982,7 @@ class AdminListPoolKeysAdapter(AdminApiAdapter):
key_name=k.name or "",
is_active=bool(k.is_active),
auth_type=str(getattr(k, "auth_type", "api_key") or "api_key"),
oauth_expires_at=_derive_oauth_expires_at(
k, auth_config=oauth_auth_config
),
oauth_expires_at=_derive_oauth_expires_at(k, auth_config=oauth_auth_config),
oauth_invalid_at=(
int(k.oauth_invalid_at.timestamp())
if getattr(k, "oauth_invalid_at", None)

View File

@@ -29,6 +29,25 @@ class PoolOverviewResponse(BaseModel):
items: list[PoolOverviewItem] = Field(default_factory=list)
# ---------------------------------------------------------------------------
# Scheduling presets metadata
# ---------------------------------------------------------------------------
class PresetModeMetaResponse(BaseModel):
value: str
label: str
class PresetDimensionMetaResponse(BaseModel):
name: str
label: str
description: str
providers: list[str] = Field(default_factory=list)
modes: list[PresetModeMetaResponse] | None = None
default_mode: str | None = None
# ---------------------------------------------------------------------------
# Paginated key list
# ---------------------------------------------------------------------------

View File

@@ -127,6 +127,80 @@ class FailoverRulesConfig(BaseModel):
)
class ScoringWeightsConfig(BaseModel):
"""多维评分权重配置。"""
lru: float = Field(0.3, ge=0.0, le=1.0)
latency: float = Field(0.25, ge=0.0, le=1.0)
health: float = Field(0.2, ge=0.0, le=1.0)
cost_remaining: float = Field(0.25, ge=0.0, le=1.0)
def _allowed_pool_preset_names() -> set[str]:
from src.services.provider.pool.dimensions import get_preset_names
return get_preset_names() | {"lru"}
def _preset_mode_meta(name: str) -> tuple[set[str], str | None]:
from src.services.provider.pool.dimensions import get_preset_dimension
dim = get_preset_dimension(name)
if dim is None or not dim.modes:
return set(), None
ordered_modes = [str(mode).strip().lower() for mode in dim.modes if str(mode).strip()]
if not ordered_modes:
return set(), None
modes = set(ordered_modes)
default_mode = str(dim.default_mode or "").strip().lower()
if not default_mode or default_mode not in modes:
default_mode = ordered_modes[0]
return modes, default_mode
class SchedulingPresetItem(BaseModel):
"""调度预设条目(新格式:有序对象列表)。"""
preset: str
enabled: bool = True
mode: str | None = None
@field_validator("preset")
@classmethod
def validate_preset(cls, v: str) -> str:
normalized = v.strip().lower()
allowed = _allowed_pool_preset_names()
if normalized not in allowed:
raise ValueError(f"无效的 preset: {normalized}")
return normalized
@field_validator("mode")
@classmethod
def normalize_mode(cls, v: str | None) -> str | None:
if v is None:
return None
normalized = v.strip().lower()
return normalized or None
@model_validator(mode="after")
def validate_mode(self) -> "SchedulingPresetItem":
allowed_modes, default_mode = _preset_mode_meta(self.preset)
if not allowed_modes:
self.mode = None
return self
if self.mode is None:
self.mode = default_mode
return self
if self.mode not in allowed_modes:
raise ValueError(
f"preset={self.preset} 的 mode 必须是: {', '.join(sorted(allowed_modes))}"
)
return self
class PoolAdvancedConfig(BaseModel):
"""通用号池配置(适用于所有 Provider 类型)。"""
@@ -148,7 +222,33 @@ class PoolAdvancedConfig(BaseModel):
le=100,
description="负载率阈值(%),超过时该 Key 被降权。默认 80",
)
# 保留旧字段供向后兼容(新客户端不再发送)
lru_enabled: bool = Field(True, description="LRU 调度(优先选择最久未用的 Key")
scheduling_mode: str | None = Field(
None,
pattern="^(lru|multi_score)$",
description="号池调度模式lru 或 multi_score",
)
scheduling_presets: list[SchedulingPresetItem] | list[str] | None = Field(
None,
description=(
"调度预设列表(新格式:对象列表 [{preset, enabled, mode}]"
"旧格式:字符串列表 ['quota_balanced', ...]"
),
)
scoring_weights: ScoringWeightsConfig | None = Field(None, description="多维评分权重")
latency_window_seconds: int | None = Field(
None,
ge=300,
le=86400,
description="延迟窗口(秒),仅 multi_score 生效",
)
latency_sample_limit: int | None = Field(
None,
ge=10,
le=200,
description="每个 Key 的延迟样本上限,仅 multi_score 生效",
)
cost_window_seconds: int | None = Field(
None,
ge=3600,

View File

@@ -60,7 +60,7 @@ class RequestDispatcher:
attempt_counter: int,
max_attempts: int,
is_stream: bool = False,
) -> tuple[Any, str, str, str, str, str]:
) -> tuple[Any, str, str, str, str, str, int | None]:
"""
执行请求并返回结果
@@ -81,7 +81,7 @@ class RequestDispatcher:
is_stream: 是否为流式请求
Returns:
(response, provider_name, candidate_record_id, provider_id, endpoint_id, key_id)
(response, provider_name, candidate_record_id, provider_id, endpoint_id, key_id, ttfb_ms)
Raises:
ExecutionError: 执行失败时
@@ -144,6 +144,19 @@ class RequestDispatcher:
logger.debug(f" [{request_id}] 请求成功: Provider={provider_name}, 耗时={elapsed_ms}ms")
# Non-stream requests don't have first-byte telemetry in this path.
# Use elapsed latency as a conservative fallback for pool latency sampling.
ttfb_ms: int | None = None
if not is_stream:
raw_ttfb = getattr(execution_result.response, "first_byte_time_ms", None)
try:
if raw_ttfb is not None:
ttfb_ms = max(int(raw_ttfb), 0)
except (TypeError, ValueError):
ttfb_ms = None
if ttfb_ms is None and elapsed_ms >= 0:
ttfb_ms = int(elapsed_ms)
return (
execution_result.response,
provider_name,
@@ -151,4 +164,5 @@ class RequestDispatcher:
provider_id,
endpoint_id,
key_id,
ttfb_ms,
)

View File

@@ -3,12 +3,18 @@
Re-exports the main public API for convenience.
"""
from src.services.provider.pool.config import PoolConfig, UnschedulableRule, parse_pool_config
from src.services.provider.pool.config import (
PoolConfig,
ScoringWeights,
UnschedulableRule,
parse_pool_config,
)
from src.services.provider.pool.manager import PoolManager
__all__ = [
"PoolConfig",
"PoolManager",
"ScoringWeights",
"UnschedulableRule",
"parse_pool_config",
]

View File

@@ -0,0 +1,189 @@
"""Pool account state helpers.
Provides a shared way to classify account-level hard-block states
from upstream metadata and OAuth invalid reasons.
"""
from __future__ import annotations
from dataclasses import dataclass
from typing import Any
OAUTH_ACCOUNT_BLOCK_PREFIX = "[ACCOUNT_BLOCK] "
ACCOUNT_BLOCK_REASON_KEYWORDS: tuple[str, ...] = (
"account_block",
"account blocked",
"account has been disabled",
"account disabled",
"organization has been disabled",
"organization_disabled",
"validation_required",
"verify your account",
"suspended",
# Kiro quota refresher 写入的确切文本
"账户已封禁",
# Antigravity quota refresher 写入的确切文本
"账户访问被禁止",
"封禁",
"封号",
"被封",
"访问被禁止",
"账号异常",
)
@dataclass(frozen=True, slots=True)
class PoolAccountState:
"""Resolved account-level state for one key."""
blocked: bool
code: str | None = None # account_banned / account_forbidden / account_blocked
label: str | None = None
reason: str | None = None
def _is_truthy_flag(value: Any) -> bool:
if isinstance(value, bool):
return value
if isinstance(value, (int, float)):
return value != 0
if isinstance(value, str):
normalized = value.strip().lower()
return normalized in {"1", "true", "yes", "y"}
return False
def _clean_text(value: Any) -> str | None:
if not isinstance(value, str):
return None
text = value.strip()
return text or None
def _extract_reason(source: dict[str, Any] | None, *fields: str) -> str | None:
if not isinstance(source, dict):
return None
for field in fields:
text = _clean_text(source.get(field))
if text:
return text
return None
def _resolve_from_metadata(
provider_type: str | None,
upstream_metadata: Any,
) -> PoolAccountState | None:
if not isinstance(upstream_metadata, dict):
return None
normalized_provider = str(provider_type or "").strip().lower()
provider_bucket: dict[str, Any] | None = None
if normalized_provider:
maybe_bucket = upstream_metadata.get(normalized_provider)
if isinstance(maybe_bucket, dict):
provider_bucket = maybe_bucket
if (
normalized_provider == "kiro"
and provider_bucket
and _is_truthy_flag(provider_bucket.get("is_banned"))
):
reason = _extract_reason(provider_bucket, "ban_reason", "reason", "message")
return PoolAccountState(
blocked=True,
code="account_banned",
label="账号封禁",
reason=reason or "Kiro 账号已封禁",
)
if (
normalized_provider == "antigravity"
and provider_bucket
and _is_truthy_flag(provider_bucket.get("is_forbidden"))
):
reason = _extract_reason(provider_bucket, "forbidden_reason", "reason", "message")
return PoolAccountState(
blocked=True,
code="account_forbidden",
label="访问受限",
reason=reason or "Antigravity 账户访问受限",
)
for source in (provider_bucket, upstream_metadata):
if not isinstance(source, dict):
continue
if _is_truthy_flag(source.get("is_banned")):
reason = _extract_reason(source, "ban_reason", "forbidden_reason", "reason", "message")
return PoolAccountState(
blocked=True,
code="account_banned",
label="账号封禁",
reason=reason or "账号已封禁",
)
if _is_truthy_flag(source.get("is_forbidden")) or _is_truthy_flag(
source.get("account_disabled")
):
reason = _extract_reason(source, "forbidden_reason", "ban_reason", "reason", "message")
return PoolAccountState(
blocked=True,
code="account_forbidden",
label="访问受限",
reason=reason or "账号访问受限",
)
return None
def _resolve_from_oauth_invalid_reason(reason: str | None) -> PoolAccountState | None:
text = _clean_text(reason)
if not text:
return None
if text.startswith(OAUTH_ACCOUNT_BLOCK_PREFIX):
cleaned = text[len(OAUTH_ACCOUNT_BLOCK_PREFIX) :].strip()
return PoolAccountState(
blocked=True,
code="account_blocked",
label="账号异常",
reason=cleaned or "账号异常",
)
lowered = text.lower()
if any(keyword in lowered for keyword in ACCOUNT_BLOCK_REASON_KEYWORDS):
return PoolAccountState(
blocked=True,
code="account_blocked",
label="账号异常",
reason=text,
)
return None
def resolve_pool_account_state(
*,
provider_type: str | None,
upstream_metadata: Any,
oauth_invalid_reason: str | None,
) -> PoolAccountState:
"""Resolve account-level hard-block state for pool scheduling."""
from_metadata = _resolve_from_metadata(provider_type, upstream_metadata)
if from_metadata is not None:
return from_metadata
from_oauth = _resolve_from_oauth_invalid_reason(oauth_invalid_reason)
if from_oauth is not None:
return from_oauth
return PoolAccountState(blocked=False)
__all__ = [
"ACCOUNT_BLOCK_REASON_KEYWORDS",
"OAUTH_ACCOUNT_BLOCK_PREFIX",
"PoolAccountState",
"resolve_pool_account_state",
]

View File

@@ -6,6 +6,26 @@ from dataclasses import dataclass, field
from typing import Any
from src.core.logger import logger
from src.services.provider.pool.dimensions import get_preset_dimension, get_preset_names
@dataclass(frozen=True, slots=True)
class ScoringWeights:
"""Weights used by multi-score scheduling."""
lru: float = 0.3
latency: float = 0.25
health: float = 0.2
cost_remaining: float = 0.25
@dataclass(frozen=True, slots=True)
class SchedulingPreset:
"""Single scheduling preset item with enable/disable and optional sub-config."""
preset: str
enabled: bool = True
mode: str | None = None
@dataclass(frozen=True, slots=True)
@@ -32,8 +52,17 @@ class PoolConfig:
# -- Load-Aware Selection -------------------------------------------------
load_threshold_percent: int = 80
# -- LRU ------------------------------------------------------------------
# -- Scheduling (unified preset list) -------------------------------------
scheduling_presets: tuple[SchedulingPreset, ...] = (
SchedulingPreset(preset="lru", enabled=True),
)
# Derived from scheduling_presets at parse time (backward compat for consumers)
lru_enabled: bool = True
scheduling_mode: str = "lru" # lru | multi_score
scoring_weights: ScoringWeights = field(default_factory=ScoringWeights)
latency_window_seconds: int = 3600
latency_sample_limit: int = 50
# -- Rolling-Window Cost Tracking -----------------------------------------
cost_window_seconds: int = 18000 # 5 hours
@@ -122,11 +151,35 @@ def parse_pool_config(provider_config: Any) -> PoolConfig | None:
except (TypeError, ValueError):
return None
scoring_weights = _parse_scoring_weights(raw_advanced.get("scoring_weights"))
# Parse scheduling presets (new object-list format or legacy string-list)
presets = _parse_scheduling_presets_v2(
raw_advanced.get("scheduling_presets"),
legacy_mode=raw_advanced.get("scheduling_mode"),
legacy_lru=raw_advanced.get("lru_enabled"),
)
# Derive scheduling_mode and lru_enabled from the presets list
enabled = [p for p in presets if p.enabled]
lru_enabled = any(p.preset == "lru" for p in enabled)
non_lru_enabled = [p for p in enabled if p.preset != "lru"]
scheduling_mode = "multi_score" if non_lru_enabled else "lru"
strategies = list(_parse_strategies(raw_advanced.get("strategies")))
if scheduling_mode == "multi_score" and "multi_score" not in strategies:
strategies.append("multi_score")
return PoolConfig(
sticky_session_ttl_seconds=_int_or("sticky_session_ttl_seconds", 3600),
global_priority=_opt_int("global_priority"),
load_threshold_percent=_int_or("load_threshold_percent", 80),
lru_enabled=_bool_or("lru_enabled", True),
scheduling_presets=presets,
lru_enabled=lru_enabled,
scheduling_mode=scheduling_mode,
scoring_weights=scoring_weights,
latency_window_seconds=_int_or("latency_window_seconds", 3600),
latency_sample_limit=_int_or("latency_sample_limit", 50),
cost_window_seconds=_int_or("cost_window_seconds", 18000),
cost_limit_per_key_tokens=_opt_int("cost_limit_per_key_tokens"),
cost_soft_threshold_percent=_int_or("cost_soft_threshold_percent", 80),
@@ -138,12 +191,134 @@ def parse_pool_config(provider_config: Any) -> PoolConfig | None:
stream_timeout_threshold=_int_or("stream_timeout_threshold", 3),
stream_timeout_window_seconds=_int_or("stream_timeout_window_seconds", 1800),
stream_timeout_cooldown_seconds=_int_or("stream_timeout_cooldown_seconds", 300),
strategies=_parse_strategies(raw_advanced.get("strategies")),
strategies=tuple(strategies),
)
# ---------------------------------------------------------------------------
# Internal parsers
# ---------------------------------------------------------------------------
def _allowed_preset_names() -> set[str]:
return get_preset_names() | {"lru"}
def _get_preset_mode_meta(preset_name: str) -> tuple[tuple[str, ...], str | None]:
dim = get_preset_dimension(preset_name)
if dim is None or not dim.modes:
return (), None
modes = tuple(str(mode).strip().lower() for mode in dim.modes if str(mode).strip())
if not modes:
return (), None
raw_default = str(dim.default_mode or "").strip().lower()
default_mode = raw_default if raw_default in modes else modes[0]
return modes, default_mode
def _parse_strategies(raw: Any) -> tuple[str, ...]:
"""Parse strategy names from config (list[str] -> tuple[str, ...])."""
if not isinstance(raw, list):
return ()
return tuple(str(s) for s in raw if isinstance(s, str) and s)
def _parse_scoring_weights(raw: Any) -> ScoringWeights:
"""Parse scoring weights with graceful fallback."""
if not isinstance(raw, dict):
return ScoringWeights()
def _float_or(value: Any, default: float) -> float:
try:
parsed = float(value)
except (TypeError, ValueError):
return default
return max(0.0, min(parsed, 1.0))
return ScoringWeights(
lru=_float_or(raw.get("lru"), 0.3),
latency=_float_or(raw.get("latency"), 0.25),
health=_float_or(raw.get("health"), 0.2),
cost_remaining=_float_or(raw.get("cost_remaining"), 0.25),
)
def _parse_scheduling_presets_v2(
raw: Any,
*,
legacy_mode: Any = None,
legacy_lru: Any = None,
) -> tuple[SchedulingPreset, ...]:
"""Parse scheduling presets, supporting both new and legacy formats.
New format::
[{"preset": "lru", "enabled": true},
{"preset": "free_team_first", "enabled": true, "mode": "free_only"},
...]
Legacy format::
["free_team_first", "recent_refresh"] (with separate scheduling_mode / lru_enabled)
"""
if isinstance(raw, list) and raw:
first = raw[0]
if isinstance(first, dict):
return _parse_preset_object_list(raw)
if isinstance(first, str):
return _convert_legacy_string_list(raw, legacy_mode, legacy_lru)
# No presets at all: derive from legacy fields
return _build_from_legacy_fields(legacy_mode, legacy_lru)
def _parse_preset_object_list(raw: list[Any]) -> tuple[SchedulingPreset, ...]:
"""Parse new-format object list into SchedulingPreset tuple."""
allowed = _allowed_preset_names()
ordered: list[SchedulingPreset] = []
seen: set[str] = set()
for item in raw:
if not isinstance(item, dict):
continue
name = str(item.get("preset", "")).strip().lower()
if name not in allowed or name in seen:
continue
seen.add(name)
enabled = bool(item.get("enabled", True))
mode: str | None = None
modes, default_mode = _get_preset_mode_meta(name)
if modes:
raw_mode = str(item.get("mode", default_mode) or "").strip().lower()
mode = raw_mode if raw_mode in modes else default_mode
ordered.append(SchedulingPreset(preset=name, enabled=enabled, mode=mode))
return tuple(ordered) if ordered else (SchedulingPreset(preset="lru", enabled=True),)
def _convert_legacy_string_list(
raw: list[Any],
legacy_mode: Any,
legacy_lru: Any,
) -> tuple[SchedulingPreset, ...]:
"""Convert legacy string list + mode/lru fields to new format."""
lru_enabled = legacy_lru if isinstance(legacy_lru, bool) else True
allowed_non_lru = _allowed_preset_names() - {"lru"}
items: list[SchedulingPreset] = [SchedulingPreset(preset="lru", enabled=lru_enabled)]
seen: set[str] = {"lru"}
for p in raw:
if not isinstance(p, str):
continue
name = p.strip().lower()
if name not in allowed_non_lru or name in seen:
continue
seen.add(name)
items.append(SchedulingPreset(preset=name, enabled=True))
return tuple(items)
def _build_from_legacy_fields(legacy_mode: Any, legacy_lru: Any) -> tuple[SchedulingPreset, ...]:
"""Build presets from legacy scheduling_mode / lru_enabled only."""
lru_enabled = legacy_lru if isinstance(legacy_lru, bool) else True
return (SchedulingPreset(preset="lru", enabled=lru_enabled),)

View File

@@ -0,0 +1,30 @@
"""Pool scheduling preset dimensions.
Importing this package registers all built-in preset dimensions.
"""
from __future__ import annotations
from . import free_team_first # noqa: F401
from . import quota_balanced # noqa: F401
from . import recent_refresh # noqa: F401
from . import single_account # noqa: F401
from .registry import (
PresetDimensionBase,
PresetDimensionMeta,
get_all_preset_dimensions,
get_preset_dimension,
get_preset_dimension_metas,
get_preset_names,
register_preset_dimension,
)
__all__ = [
"PresetDimensionBase",
"PresetDimensionMeta",
"get_all_preset_dimensions",
"get_preset_dimension",
"get_preset_dimension_metas",
"get_preset_names",
"register_preset_dimension",
]

View File

@@ -0,0 +1,227 @@
"""Shared helpers for pool preset dimensions."""
from __future__ import annotations
import math
import time
from typing import Any
def safe_float(value: Any) -> float | None:
try:
parsed = float(value)
except (TypeError, ValueError):
return None
if math.isnan(parsed) or math.isinf(parsed):
return None
return parsed
def safe_metadata(key_obj: Any) -> dict[str, Any]:
raw = getattr(key_obj, "upstream_metadata", None)
return raw if isinstance(raw, dict) else {}
def normalize_plan(value: Any) -> str | None:
if not isinstance(value, str):
return None
normalized = value.strip().lower()
return normalized or None
def rank_ascending(key_id: str, scores: dict[str, float], all_ids: list[str]) -> float:
"""Rank score within all IDs; lower value means better rank."""
if not all_ids:
return 0.0
valid_count = sum(1 for kid in all_ids if safe_float(scores.get(kid)) is not None)
if valid_count <= 0:
return 0.5
decorated: list[tuple[int, float, int, str]] = []
for idx, kid in enumerate(all_ids):
score_raw = safe_float(scores.get(kid))
if score_raw is None:
decorated.append((1, float("inf"), idx, kid))
else:
decorated.append((0, score_raw, idx, kid))
decorated.sort(key=lambda item: (item[0], item[1], item[2]))
rank_idx = 0
for idx, (_missing, _value, _order, kid) in enumerate(decorated):
if kid == key_id:
rank_idx = idx
break
n = len(all_ids)
if n <= 1:
return 0.0
return rank_idx / float(n - 1)
def rank_descending(key_id: str, scores: dict[str, float], all_ids: list[str]) -> float:
"""Rank score within all IDs; higher value means better rank."""
if not all_ids:
return 0.0
valid_count = sum(1 for kid in all_ids if safe_float(scores.get(kid)) is not None)
if valid_count <= 0:
return 0.5
decorated: list[tuple[int, float, int, str]] = []
for idx, kid in enumerate(all_ids):
score_raw = safe_float(scores.get(kid))
if score_raw is None:
decorated.append((1, float("inf"), idx, kid))
else:
# 排序时取负值使分值越大排名越靠前rank 越小)
decorated.append((0, -score_raw, idx, kid))
decorated.sort(key=lambda item: (item[0], item[1], item[2]))
rank_idx = 0
for idx, (_missing, _value, _order, kid) in enumerate(decorated):
if kid == key_id:
rank_idx = idx
break
n = len(all_ids)
if n <= 1:
return 0.0
return rank_idx / float(n - 1)
def extract_plan_type(key_obj: Any) -> str | None:
direct = normalize_plan(getattr(key_obj, "oauth_plan_type", None))
if direct:
return direct
metadata = safe_metadata(key_obj)
codex = metadata.get("codex")
if isinstance(codex, dict):
codex_plan = normalize_plan(codex.get("plan_type"))
if codex_plan:
return codex_plan
kiro = metadata.get("kiro")
if isinstance(kiro, dict):
subscription_title = normalize_plan(kiro.get("subscription_title"))
if subscription_title:
# Normalize common Kiro labels into free/team buckets used by free_team_first.
if "team" in subscription_title:
return "team"
if "free" in subscription_title:
return "free"
if "pro" in subscription_title:
return "pro"
if "plus" in subscription_title:
return "plus"
return subscription_title
return None
def extract_reset_seconds(key_obj: Any) -> float | None:
metadata = safe_metadata(key_obj)
candidates: list[float] = []
codex = metadata.get("codex")
if isinstance(codex, dict):
for field in ("secondary_reset_seconds", "primary_reset_seconds"):
parsed = safe_float(codex.get(field))
if parsed is None or parsed < 0:
continue
candidates.append(parsed)
kiro = metadata.get("kiro")
if isinstance(kiro, dict):
next_reset_at = safe_float(kiro.get("next_reset_at"))
if next_reset_at is not None and next_reset_at > 0:
candidates.append(max(0.0, next_reset_at - time.time()))
if not candidates:
return None
return min(candidates)
def extract_usage_ratio(key_obj: Any) -> float | None:
metadata = safe_metadata(key_obj)
codex = metadata.get("codex")
if isinstance(codex, dict):
codex_values: list[float] = []
for field in ("primary_used_percent", "secondary_used_percent"):
parsed = safe_float(codex.get(field))
if parsed is None:
continue
codex_values.append(max(0.0, min(parsed, 100.0)) / 100.0)
if codex_values:
return sum(codex_values) / len(codex_values)
kiro = metadata.get("kiro")
if isinstance(kiro, dict):
parsed = safe_float(kiro.get("usage_percentage"))
if parsed is not None:
return max(0.0, min(parsed, 100.0)) / 100.0
antigravity = metadata.get("antigravity")
if isinstance(antigravity, dict):
quota_by_model = antigravity.get("quota_by_model")
if isinstance(quota_by_model, dict):
usage_values: list[float] = []
for model_info in quota_by_model.values():
if not isinstance(model_info, dict):
continue
used_percent = safe_float(model_info.get("used_percent"))
if used_percent is None:
remaining_fraction = safe_float(model_info.get("remaining_fraction"))
if remaining_fraction is not None:
used_percent = (1.0 - remaining_fraction) * 100.0
if used_percent is None:
continue
usage_values.append(max(0.0, min(used_percent, 100.0)) / 100.0)
if usage_values:
return sum(usage_values) / len(usage_values)
return None
def plan_priority_score(plan_type: str | None, mode: str | None = None) -> float:
"""Score a key based on plan type and free_team_first mode."""
effective_mode = (mode or "both").strip().lower()
if effective_mode == "free_only":
if plan_type == "free":
return 0.0
if plan_type == "team":
return 0.5
elif effective_mode == "team_only":
if plan_type == "team":
return 0.0
if plan_type == "free":
return 0.5
else:
# "both" or unrecognized -> original behavior
if plan_type in {"free", "team"}:
return 0.0
if plan_type in {"enterprise", "business"}:
return 0.2
if plan_type in {"plus", "pro"}:
return 0.6
if plan_type:
return 0.7
return 0.8
__all__ = [
"extract_plan_type",
"extract_reset_seconds",
"extract_usage_ratio",
"normalize_plan",
"plan_priority_score",
"rank_ascending",
"rank_descending",
"safe_float",
"safe_metadata",
]

View File

@@ -0,0 +1,52 @@
"""free_team_first preset dimension."""
from __future__ import annotations
from typing import Any
from ._helpers import extract_plan_type, plan_priority_score, rank_ascending
from .registry import PresetDimensionBase, register_preset_dimension
class FreeTeamFirstDimension(PresetDimensionBase):
@property
def name(self) -> str:
return "free_team_first"
@property
def label(self) -> str:
return "Free/Team 优先"
@property
def description(self) -> str:
return "优先消耗低档账号(依赖 plan_type"
@property
def providers(self) -> tuple[str, ...]:
return ("codex", "kiro")
@property
def modes(self) -> tuple[str, ...] | None:
return ("free_only", "team_only", "both")
@property
def default_mode(self) -> str | None:
return "both"
def compute_metric(
self,
*,
key_id: str,
all_key_ids: list[str],
keys_by_id: dict[str, Any],
lru_scores: dict[str, Any],
mode: str | None,
) -> float:
plan_scores = {
kid: plan_priority_score(extract_plan_type(keys_by_id.get(kid)), mode)
for kid in all_key_ids
}
return rank_ascending(key_id, plan_scores, all_key_ids)
register_preset_dimension(FreeTeamFirstDimension())

View File

@@ -0,0 +1,41 @@
"""quota_balanced preset dimension."""
from __future__ import annotations
from typing import Any
from ._helpers import extract_usage_ratio, rank_ascending
from .registry import PresetDimensionBase, register_preset_dimension
class QuotaBalancedDimension(PresetDimensionBase):
@property
def name(self) -> str:
return "quota_balanced"
@property
def label(self) -> str:
return "额度平均"
@property
def description(self) -> str:
return "优先选额度消耗最少的账号"
def compute_metric(
self,
*,
key_id: str,
all_key_ids: list[str],
keys_by_id: dict[str, Any],
lru_scores: dict[str, Any],
mode: str | None,
) -> float:
usage_scores: dict[str, float] = {}
for kid in all_key_ids:
usage_ratio = extract_usage_ratio(keys_by_id.get(kid))
if usage_ratio is not None:
usage_scores[kid] = usage_ratio
return rank_ascending(key_id, usage_scores, all_key_ids)
register_preset_dimension(QuotaBalancedDimension())

View File

@@ -0,0 +1,45 @@
"""recent_refresh preset dimension."""
from __future__ import annotations
from typing import Any
from ._helpers import extract_reset_seconds, rank_ascending
from .registry import PresetDimensionBase, register_preset_dimension
class RecentRefreshDimension(PresetDimensionBase):
@property
def name(self) -> str:
return "recent_refresh"
@property
def label(self) -> str:
return "额度刷新优先"
@property
def description(self) -> str:
return "优先选即将刷新额度的账号"
@property
def providers(self) -> tuple[str, ...]:
return ("codex", "kiro")
def compute_metric(
self,
*,
key_id: str,
all_key_ids: list[str],
keys_by_id: dict[str, Any],
lru_scores: dict[str, Any],
mode: str | None,
) -> float:
reset_scores: dict[str, float] = {}
for kid in all_key_ids:
reset_seconds = extract_reset_seconds(keys_by_id.get(kid))
if reset_seconds is not None:
reset_scores[kid] = reset_seconds
return rank_ascending(key_id, reset_scores, all_key_ids)
register_preset_dimension(RecentRefreshDimension())

View File

@@ -0,0 +1,217 @@
"""Preset dimension registry for pool multi-score scheduling."""
from __future__ import annotations
from abc import ABC, abstractmethod
from dataclasses import dataclass
from threading import RLock
from typing import Any
@dataclass(frozen=True, slots=True)
class PresetDimensionMeta:
"""Serializable metadata for one preset dimension."""
name: str
label: str
description: str
providers: tuple[str, ...]
modes: tuple[str, ...] | None
default_mode: str | None
class PresetDimensionBase(ABC):
"""Base class of one pool scheduling preset dimension."""
@property
@abstractmethod
def name(self) -> str:
"""Stable preset key, e.g. ``free_team_first``."""
@property
@abstractmethod
def label(self) -> str:
"""User-facing label."""
@property
@abstractmethod
def description(self) -> str:
"""User-facing description."""
@property
def providers(self) -> tuple[str, ...]:
"""Supported provider types.
Empty tuple means the dimension is universal and applies to all providers.
"""
return ()
@property
def modes(self) -> tuple[str, ...] | None:
"""Optional sub-modes for this dimension."""
return None
@property
def default_mode(self) -> str | None:
"""Default mode when mode is omitted."""
return None
@abstractmethod
def compute_metric(
self,
*,
key_id: str,
all_key_ids: list[str],
keys_by_id: dict[str, Any],
lru_scores: dict[str, Any],
mode: str | None,
) -> float:
"""Compute normalized metric in [0, 1], lower is better."""
def is_applicable(self, provider_type: str) -> bool:
"""Return whether this dimension applies to the given provider type."""
if not self.providers:
return True
normalized = _normalize_name(provider_type)
return normalized in self.providers
def _normalize_name(value: Any) -> str:
if not isinstance(value, str):
return ""
return value.strip().lower()
def _normalize_names(values: tuple[str, ...] | list[str]) -> tuple[str, ...]:
normalized = [_normalize_name(item) for item in values]
return tuple(item for item in normalized if item)
_registry_lock = RLock()
_registry: dict[str, PresetDimensionBase] = {}
def register_preset_dimension(dim: PresetDimensionBase) -> None:
"""Register or replace one preset dimension by name."""
name = _normalize_name(dim.name)
if not name:
raise ValueError("preset dimension name must be a non-empty string")
providers = _normalize_names(dim.providers)
modes = _normalize_names(dim.modes or ())
default_mode = _normalize_name(dim.default_mode)
if modes and default_mode and default_mode not in modes:
raise ValueError(f"default_mode must be one of modes for preset '{name}'")
class _NormalizedDimension(PresetDimensionBase):
# Lightweight wrapper to keep normalized metadata while preserving compute logic.
def __init__(self, wrapped: PresetDimensionBase) -> None:
self._wrapped = wrapped
@property
def name(self) -> str:
return name
@property
def label(self) -> str:
return self._wrapped.label
@property
def description(self) -> str:
return self._wrapped.description
@property
def providers(self) -> tuple[str, ...]:
return providers
@property
def modes(self) -> tuple[str, ...] | None:
return modes or None
@property
def default_mode(self) -> str | None:
if not modes:
return None
if default_mode:
return default_mode
return modes[0]
def compute_metric(
self,
*,
key_id: str,
all_key_ids: list[str],
keys_by_id: dict[str, Any],
lru_scores: dict[str, Any],
mode: str | None,
) -> float:
return self._wrapped.compute_metric(
key_id=key_id,
all_key_ids=all_key_ids,
keys_by_id=keys_by_id,
lru_scores=lru_scores,
mode=mode,
)
normalized = _NormalizedDimension(dim)
with _registry_lock:
_registry[name] = normalized
def get_preset_dimension(name: str) -> PresetDimensionBase | None:
"""Get one registered preset dimension by name."""
key = _normalize_name(name)
if not key:
return None
with _registry_lock:
return _registry.get(key)
def get_all_preset_dimensions() -> list[PresetDimensionBase]:
"""Get all registered preset dimensions in registration order."""
with _registry_lock:
return list(_registry.values())
def get_preset_names() -> set[str]:
"""Get all registered preset names."""
with _registry_lock:
return set(_registry.keys())
def get_preset_dimension_metas() -> list[PresetDimensionMeta]:
"""Get serializable metadata for all preset dimensions."""
metas: list[PresetDimensionMeta] = []
for dim in get_all_preset_dimensions():
metas.append(
PresetDimensionMeta(
name=dim.name,
label=dim.label,
description=dim.description,
providers=dim.providers,
modes=dim.modes,
default_mode=dim.default_mode,
)
)
return metas
__all__ = [
"PresetDimensionBase",
"PresetDimensionMeta",
"get_all_preset_dimensions",
"get_preset_dimension",
"get_preset_dimension_metas",
"get_preset_names",
"register_preset_dimension",
]

View File

@@ -0,0 +1,36 @@
"""single_account preset dimension."""
from __future__ import annotations
from typing import Any
from ._helpers import rank_descending
from .registry import PresetDimensionBase, register_preset_dimension
class SingleAccountDimension(PresetDimensionBase):
@property
def name(self) -> str:
return "single_account"
@property
def label(self) -> str:
return "单号优先"
@property
def description(self) -> str:
return "集中使用同一账号(反向 LRU"
def compute_metric(
self,
*,
key_id: str,
all_key_ids: list[str],
keys_by_id: dict[str, Any],
lru_scores: dict[str, Any],
mode: str | None,
) -> float:
return rank_descending(key_id, lru_scores, all_key_ids)
register_preset_dimension(SingleAccountDimension())

View File

@@ -0,0 +1,91 @@
"""In-process pool health score cache.
This cache avoids recomputing per-key health aggregation on every request.
It does not replace persistent health storage; source data still comes from
``ProviderAPIKey.health_by_format`` carried on key objects.
"""
from __future__ import annotations
import threading
import time
from typing import Any
_TTL_SECONDS = 30.0
_LOCK = threading.Lock()
_CACHE: dict[str, tuple[float, dict[str, float]]] = {}
def aggregate_health_score(health_by_format: Any) -> float:
"""Aggregate health score from ``health_by_format`` (lower-bound strategy)."""
if not isinstance(health_by_format, dict) or not health_by_format:
return 1.0
scores: list[float] = []
for item in health_by_format.values():
if not isinstance(item, dict):
continue
try:
score = float(item.get("health_score") or 1.0)
except (TypeError, ValueError):
score = 1.0
scores.append(max(0.0, min(score, 1.0)))
if not scores:
return 1.0
return min(scores)
def get_health_scores(provider_id: str, keys: list[Any]) -> dict[str, float]:
"""Return key health scores with per-provider TTL cache.
Uses incremental merge: if the cache is still valid but missing some keys,
only the missing keys are computed and merged into the existing cache entry.
"""
now = time.monotonic()
keys_by_id: dict[str, Any] = {}
for k in keys:
kid = str(getattr(k, "id", "") or "")
if kid:
keys_by_id[kid] = k
if not keys_by_id:
return {}
with _LOCK:
cached = _CACHE.get(provider_id)
if cached is not None:
expires_at, payload = cached
if now < expires_at:
missing_ids = [kid for kid in keys_by_id if kid not in payload]
if not missing_ids:
return {kid: payload[kid] for kid in keys_by_id}
# Compute only for missing keys, merge into existing cache
for kid in missing_ids:
payload[kid] = aggregate_health_score(
getattr(keys_by_id[kid], "health_by_format", None)
)
return {kid: payload[kid] for kid in keys_by_id}
fresh: dict[str, float] = {}
for kid, key in keys_by_id.items():
fresh[kid] = aggregate_health_score(getattr(key, "health_by_format", None))
with _LOCK:
_CACHE[provider_id] = (now + _TTL_SECONDS, fresh)
return dict(fresh)
def invalidate_provider_health_scores(provider_id: str) -> None:
"""Invalidate health-score cache for one provider."""
with _LOCK:
_CACHE.pop(provider_id, None)
def _clear_cache_for_tests() -> None:
with _LOCK:
_CACHE.clear()
__all__ = [
"aggregate_health_score",
"get_health_scores",
"invalidate_provider_health_scores",
]

View File

@@ -20,7 +20,9 @@ from typing import TYPE_CHECKING, Any, TypeVar
from src.core.logger import logger
from src.services.provider.pool import redis_ops
from src.services.provider.pool.account_state import resolve_pool_account_state
from src.services.provider.pool.config import PoolConfig
from src.services.provider.pool.health_cache import get_health_scores
from src.services.provider.pool.trace import PoolCandidateTrace, PoolSchedulingTrace
if TYPE_CHECKING:
@@ -31,11 +33,17 @@ if TYPE_CHECKING:
class PoolManager:
"""Coordinate pool-level scheduling for a single Provider."""
__slots__ = ("provider_id", "config")
__slots__ = ("provider_id", "config", "provider_type")
def __init__(self, provider_id: str, config: PoolConfig) -> None:
def __init__(
self,
provider_id: str,
config: PoolConfig,
provider_type: str | None = None,
) -> None:
self.provider_id = provider_id
self.config = config
self.provider_type = str(provider_type or "").strip().lower() or None
# ------------------------------------------------------------------
# Core scheduling: reorder candidate list for pool-aware selection
@@ -53,7 +61,7 @@ class PoolManager:
1. **Sticky session hit** -- if the session is already bound to a key
and that key appears in *candidates* and is not in cooldown, move it
to position 0.
2. **Filter** out keys in cooldown or cost-exhausted state (mark
2. **Filter** out keys in account-blocked / cooldown / cost-exhausted state (mark
``is_skipped``).
3. **LRU sort** -- among remaining candidates at the same priority
level, sort by least-recently-used.
@@ -102,6 +110,13 @@ class PoolManager:
pid, session_uuid, self.config.sticky_session_ttl_seconds
)
provider_type = self.provider_type
if provider_type is None and candidates:
first_provider = getattr(candidates[0], "provider", None)
provider_type = str(getattr(first_provider, "provider_type", "") or "").strip().lower()
if not provider_type:
provider_type = None
# --- 2. Batch fetch pool state (parallel) ---------------------
all_key_ids = [str(c.key.id) for c in candidates]
@@ -113,17 +128,26 @@ class PoolManager:
else None
)
_lru_coro = redis_ops.get_lru_scores(pid, all_key_ids) if self.config.lru_enabled else None
_latency_coro = (
redis_ops.batch_get_latency_avgs(pid, all_key_ids, self.config.latency_window_seconds)
if self.config.scheduling_mode == "multi_score"
else None
)
# Gather all non-None coroutines in parallel.
coros: list[Any] = [_cooldown_coro]
_cost_idx = -1
_lru_idx = -1
_latency_idx = -1
if _cost_coro is not None:
_cost_idx = len(coros)
coros.append(_cost_coro)
if _lru_coro is not None:
_lru_idx = len(coros)
coros.append(_lru_coro)
if _latency_coro is not None:
_latency_idx = len(coros)
coros.append(_latency_coro)
gathered = await asyncio.gather(*coros)
@@ -158,7 +182,29 @@ class PoolManager:
if _lru_idx >= 0:
lru_scores = gathered[_lru_idx]
# Latency averages
latency_avgs: dict[str, float] = {}
if _latency_idx >= 0:
latency_avgs = gathered[_latency_idx]
# Health scores (TTL cached, no Redis round-trip) -- only needed for multi_score
health_scores: dict[str, float] = {}
if self.config.scheduling_mode == "multi_score":
health_scores = get_health_scores(pid, [c.key for c in candidates])
strategy_context.update(
{
"all_key_ids": all_key_ids,
"lru_scores": lru_scores,
"cost_totals": cost_totals,
"latency_avgs": latency_avgs,
"health_scores": health_scores,
"keys_by_id": {str(c.key.id): c.key for c in candidates},
}
)
# --- Strategy: compute_score ----------------------------------
custom_scores: dict[str, float] = {}
for strategy in strategies:
if hasattr(strategy, "compute_score"):
for kid in all_key_ids:
@@ -169,7 +215,8 @@ class PoolManager:
context=strategy_context,
)
if custom is not None:
lru_scores[kid] = custom
custom_scores[kid] = float(custom)
lru_scores[kid] = float(custom)
except Exception:
pass
@@ -181,6 +228,11 @@ class PoolManager:
for c in candidates:
kid = str(c.key.id)
ct = PoolCandidateTrace(key_id=kid)
ct.scoring_mode = self.config.scheduling_mode
ct.latency_avg_ms = float(latency_avgs.get(kid, 0.0) or 0.0)
ct.health_score = float(health_scores.get(kid, 1.0) or 1.0)
if kid in custom_scores:
ct.composite_score = float(custom_scores[kid])
# Already skipped upstream?
if c.is_skipped:
@@ -191,6 +243,25 @@ class PoolManager:
continue
# Cooldown?
account_state = resolve_pool_account_state(
provider_type=provider_type,
upstream_metadata=getattr(c.key, "upstream_metadata", None),
oauth_invalid_reason=getattr(c.key, "oauth_invalid_reason", None),
)
if account_state.blocked:
c.is_skipped = True
skip_reason = account_state.reason or account_state.label or "account blocked"
c.skip_reason = f"pool account blocked: {skip_reason}"
skipped.append(c)
ct.skipped = True
ct.skip_type = "account_blocked"
ct.account_block_code = account_state.code
ct.account_block_label = account_state.label
ct.account_block_reason = account_state.reason
_attach_pool_extra(c, ct)
trace.candidate_traces[kid] = ct
continue
cd_reason = cooldowns.get(kid)
if cd_reason is not None:
c.is_skipped = True
@@ -225,7 +296,10 @@ class PoolManager:
trace.sticky_session_used = True
else:
available.append(c)
ct.reason = "lru" if lru_scores.get(kid, 0) > 0 else "random"
if kid in custom_scores and self.config.scheduling_mode == "multi_score":
ct.reason = "multi_score"
else:
ct.reason = "lru" if lru_scores.get(kid, 0) > 0 else "random"
ct.lru_score = lru_scores.get(kid, 0.0)
ct.cost_window_usage = cost_totals.get(kid, 0)
@@ -363,7 +437,7 @@ class PoolManager:
:class:`ProviderAPIKey` objects instead of candidates:
1. Sticky session hit (if bound and still healthy).
2. Filter out keys in cooldown or cost-exhausted.
2. Filter out keys in account-blocked / cooldown / cost-exhausted.
3. LRU sort among remaining keys.
4. Random tiebreak for identical LRU scores.
5. Return the first available key, or ``None``.
@@ -390,22 +464,32 @@ class PoolManager:
else None
)
_lru_coro = redis_ops.get_lru_scores(pid, key_ids) if self.config.lru_enabled else None
_latency_coro = (
redis_ops.batch_get_latency_avgs(pid, key_ids, self.config.latency_window_seconds)
if self.config.scheduling_mode == "multi_score"
else None
)
coros_sk: list[Any] = [_cooldown_coro]
_cost_idx_sk = -1
_lru_idx_sk = -1
_latency_idx_sk = -1
if _cost_coro is not None:
_cost_idx_sk = len(coros_sk)
coros_sk.append(_cost_coro)
if _lru_coro is not None:
_lru_idx_sk = len(coros_sk)
coros_sk.append(_lru_coro)
if _latency_coro is not None:
_latency_idx_sk = len(coros_sk)
coros_sk.append(_latency_coro)
gathered_sk = await asyncio.gather(*coros_sk)
cooldowns = gathered_sk[0]
cost_exhausted: set[str] = set()
cost_totals: dict[str, int] = {}
if _cost_idx_sk >= 0:
cost_totals = gathered_sk[_cost_idx_sk]
for kid, total in cost_totals.items():
@@ -416,9 +500,25 @@ class PoolManager:
if _lru_idx_sk >= 0:
lru_scores = gathered_sk[_lru_idx_sk]
latency_avgs: dict[str, float] = {}
if _latency_idx_sk >= 0:
latency_avgs = gathered_sk[_latency_idx_sk]
health_scores: dict[str, float] = {}
if self.config.scheduling_mode == "multi_score":
health_scores = get_health_scores(pid, keys)
# --- Strategy: compute_score ------------------------------------------
strategies = _get_active_strategies(self.config)
strategy_context: dict[str, Any] = {"session_uuid": session_uuid}
strategy_context: dict[str, Any] = {
"session_uuid": session_uuid,
"all_key_ids": key_ids,
"lru_scores": lru_scores,
"cost_totals": cost_totals if _cost_idx_sk >= 0 else {},
"latency_avgs": latency_avgs,
"health_scores": health_scores,
"keys_by_id": {str(k.id): k for k in keys},
}
for strategy in strategies:
if hasattr(strategy, "compute_score"):
for kid in key_ids:
@@ -429,7 +529,7 @@ class PoolManager:
context=strategy_context,
)
if custom is not None:
lru_scores[kid] = custom
lru_scores[kid] = float(custom)
except Exception:
pass
@@ -440,6 +540,14 @@ class PoolManager:
for k in keys:
kid = str(k.id)
account_state = resolve_pool_account_state(
provider_type=self.provider_type,
upstream_metadata=getattr(k, "upstream_metadata", None),
oauth_invalid_reason=getattr(k, "oauth_invalid_reason", None),
)
if account_state.blocked:
continue
if cooldowns.get(kid) is not None:
continue
if kid in cost_exhausted:
@@ -480,6 +588,7 @@ class PoolManager:
session_uuid: str | None,
key_id: str,
tokens_used: int = 0,
ttfb_ms: int | None = None,
) -> None:
"""Called after a successful upstream request."""
pid = self.provider_id
@@ -500,6 +609,16 @@ class PoolManager:
pid, key_id, tokens_used, self.config.cost_window_seconds
)
# Record latency sample for multi-score scheduling.
if self.config.scheduling_mode == "multi_score" and ttfb_ms is not None and ttfb_ms >= 0:
await redis_ops.record_latency(
pid,
key_id,
ttfb_ms,
self.config.latency_window_seconds,
self.config.latency_sample_limit,
)
async def on_request_error(
self,
*,
@@ -570,6 +689,8 @@ def _get_active_strategies(config: PoolConfig) -> list[Any]:
if not config.strategies:
return []
try:
# Import triggers built-in strategy registration via module-level side effects.
from src.services.provider.pool import strategies as _builtin_strategies # noqa: F401
from src.services.provider.pool.strategy import get_active_strategies
return get_active_strategies(config.strategies)

View File

@@ -10,6 +10,7 @@ ap:{pid}:sticky:{session_uuid} STRING -> key_id (TTL: config)
ap:{pid}:lru ZSET member=key_id, score=unix_ts
ap:{pid}:cooldown:{key_id} STRING -> reason (TTL: error-specific)
ap:{pid}:cost:{key_id} ZSET member=req_id, score=unix_ts
ap:{pid}:latency:{key_id} ZSET member=req_id:ttfb_ms, score=unix_ts
provider_oauth_token_cache:{key_id} STRING -> access_token (TTL: expires - 60)
"""
@@ -44,6 +45,10 @@ def _cost_key(provider_id: str, key_id: str) -> str:
return f"{PREFIX}:{provider_id}:cost:{key_id}"
def _latency_key(provider_id: str, key_id: str) -> str:
return f"{PREFIX}:{provider_id}:latency:{key_id}"
def _oauth_cache_key(key_id: str) -> str:
return f"provider_oauth_token_cache:{key_id}"
@@ -91,6 +96,32 @@ end
return total
"""
# Latency window cleanup + average in a single round-trip.
# KEYS[1] = latency zset key, ARGV[1] = window_start timestamp
# Returns nil when there are no samples, or avg(ms) as number.
_LATENCY_WINDOW_AVG_LUA = """
local key = KEYS[1]
local window_start = tonumber(ARGV[1])
redis.call("ZREMRANGEBYSCORE", key, "-inf", window_start)
local members = redis.call("ZRANGEBYSCORE", key, window_start, "+inf")
local total = 0
local count = 0
for _, m in ipairs(members) do
local colon = string.find(m, ":", 1, true)
if colon then
local n = tonumber(string.sub(m, colon + 1))
if n then
total = total + n
count = count + 1
end
end
end
if count == 0 then
return nil
end
return total / count
"""
async def _get_redis() -> "aioredis.Redis | None":
return await get_redis_client(require_redis=False)
@@ -337,6 +368,64 @@ async def batch_get_cost_totals(
return {k: 0 for k in key_ids}
async def record_latency(
provider_id: str,
key_id: str,
ttfb_ms: int,
window_seconds: int,
sample_limit: int,
) -> None:
"""Record one TTFB sample with rolling-window cleanup."""
redis = await _get_redis()
if redis is None:
return
try:
now = time.time()
latency_k = _latency_key(provider_id, key_id)
sample = max(int(ttfb_ms), 0)
member = f"{uuid.uuid4().hex}:{sample}"
pipe = redis.pipeline()
pipe.zadd(latency_k, {member: now})
window_start = now - max(int(window_seconds), 1)
pipe.zremrangebyscore(latency_k, "-inf", window_start)
capped_limit = max(int(sample_limit), 1)
pipe.zremrangebyrank(latency_k, 0, -(capped_limit + 1))
pipe.expire(latency_k, max(int(window_seconds), 1) + 600)
await pipe.execute()
except Exception:
logger.debug("Pool: latency ADD failed for key {}", key_id[:8])
async def batch_get_latency_avgs(
provider_id: str,
key_ids: list[str],
window_seconds: int,
) -> dict[str, float]:
"""Batch-fetch latency averages (ms) for keys in a rolling window."""
redis = await _get_redis()
if redis is None:
return {}
try:
now = time.time()
window_start = now - max(int(window_seconds), 1)
pipe = redis.pipeline()
for kid in key_ids:
pipe.eval(_LATENCY_WINDOW_AVG_LUA, 1, _latency_key(provider_id, kid), str(window_start))
results = await pipe.execute()
out: dict[str, float] = {}
for kid, val in zip(key_ids, results):
if val is None:
continue
try:
out[kid] = float(val)
except Exception:
continue
return out
except Exception:
logger.debug("Pool: batch latency AVG failed for provider {}", provider_id[:8])
return {}
async def clear_cost(provider_id: str, key_id: str) -> None:
redis = await _get_redis()
if redis is None:

View File

@@ -25,6 +25,10 @@ class PoolSchedulingSnapshot:
cost_limit: int | None
cost_soft_threshold_percent: int = 80
health_score: float = 1.0
latency_avg_ms: float | None = None
account_blocked: bool = False
account_block_label: str | None = None
account_block_reason: str | None = None
@dataclass(frozen=True, slots=True)
@@ -67,6 +71,44 @@ class PoolSchedulingDimension(Protocol):
"""Evaluate one dimension from snapshot."""
@dataclass(frozen=True, slots=True)
class _AccountStateDimension:
code: str = "account_state"
label: str = "账号状态"
source: str = "policy"
weight: int = 10
def evaluate(self, snapshot: PoolSchedulingSnapshot) -> PoolSchedulingDimensionResult:
if not snapshot.account_blocked:
return PoolSchedulingDimensionResult(
code=self.code,
label=self.label,
source=self.source,
weight=self.weight,
status="ok",
score=1.0,
)
blocked_label = snapshot.account_block_label or "账号异常"
if blocked_label == "账号封禁":
blocked_code = "account_banned"
elif blocked_label == "访问受限":
blocked_code = "account_forbidden"
else:
blocked_code = "account_blocked"
return PoolSchedulingDimensionResult(
code=blocked_code,
label=blocked_label,
source=self.source,
weight=self.weight,
status="blocked",
blocking=True,
score=0.0,
detail=snapshot.account_block_reason or blocked_label,
)
@dataclass(frozen=True, slots=True)
class _ManualEnableDimension:
code: str = "manual_disabled"
@@ -264,6 +306,59 @@ class _HealthDimension:
)
@dataclass(frozen=True, slots=True)
class _LatencyDimension:
code: str = "latency"
label: str = "延迟"
source: str = "runtime"
weight: int = 3
def evaluate(self, snapshot: PoolSchedulingSnapshot) -> PoolSchedulingDimensionResult:
latency = snapshot.latency_avg_ms
if latency is None:
return PoolSchedulingDimensionResult(
code=self.code,
label=self.label,
source=self.source,
weight=self.weight,
status="ok",
score=1.0,
detail="-",
)
value = max(float(latency), 0.0)
detail = f"{value:.0f}ms"
if value >= 3000:
return PoolSchedulingDimensionResult(
code="latency_high",
label="延迟偏高",
source=self.source,
weight=self.weight,
status="degraded",
score=0.5,
detail=detail,
)
if value >= 1200:
return PoolSchedulingDimensionResult(
code="latency_slow",
label="延迟较慢",
source=self.source,
weight=self.weight,
status="degraded",
score=0.72,
detail=detail,
)
return PoolSchedulingDimensionResult(
code=self.code,
label=self.label,
source=self.source,
weight=self.weight,
status="ok",
score=1.0,
detail=detail,
)
_POOL_DIMENSION_REGISTRY: dict[str, PoolSchedulingDimension] = {}
_POOL_DIMENSION_ORDER: list[str] = []
@@ -363,10 +458,12 @@ def summarize_pool_scheduling_dimensions(
def _register_default_dimensions() -> None:
register_pool_scheduling_dimension("account_state", _AccountStateDimension())
register_pool_scheduling_dimension("manual", _ManualEnableDimension())
register_pool_scheduling_dimension("cooldown", _CooldownDimension())
register_pool_scheduling_dimension("circuit", _CircuitBreakerDimension())
register_pool_scheduling_dimension("cost", _CostDimension())
register_pool_scheduling_dimension("latency", _LatencyDimension())
register_pool_scheduling_dimension("health", _HealthDimension())

View File

@@ -0,0 +1,8 @@
"""Built-in pool strategies."""
# Import side effects: register built-in strategies.
import src.services.provider.pool.dimensions # noqa: F401
from . import multi_score # noqa: F401
__all__ = ["multi_score"]

View File

@@ -0,0 +1,178 @@
"""Multi-dimension pool scoring strategy."""
from __future__ import annotations
from typing import Any
import src.services.provider.pool.dimensions # noqa: F401
from src.services.provider.pool.dimensions import get_preset_dimension, get_preset_names
from src.services.provider.pool.dimensions._helpers import rank_ascending, safe_float
from src.services.provider.pool.strategy import register_pool_strategy
# When LRU is enabled alongside presets, this fraction of the final score
# comes from the LRU rank (tiebreaker to avoid same-score collisions).
_LRU_BLEND_FACTOR = 0.04
# Positional weight decay factor: weight = 1 / (1 + DECAY * index).
_POSITIONAL_DECAY = 0.6
def _normalize_presets_from_config(config: Any) -> tuple[tuple[str, str | None], ...]:
"""Extract enabled (preset_name, mode) tuples from config.scheduling_presets.
Supports both new SchedulingPreset objects and legacy string lists.
Excludes ``lru`` since LRU is handled separately as a blend factor.
"""
raw = getattr(config, "scheduling_presets", ())
if not isinstance(raw, (list, tuple)):
return ()
allowed = get_preset_names() | {"lru"}
ordered: list[tuple[str, str | None]] = []
seen: set[str] = set()
for item in raw:
preset_name: str | None = None
enabled = True
mode: str | None = None
if hasattr(item, "preset"):
preset_name = str(getattr(item, "preset", "")).strip().lower()
enabled = bool(getattr(item, "enabled", True))
raw_mode = getattr(item, "mode", None)
if isinstance(raw_mode, str):
mode = raw_mode.strip().lower() or None
elif isinstance(item, str):
preset_name = item.strip().lower()
else:
continue
if not preset_name or preset_name not in allowed or preset_name in seen:
continue
if not enabled:
continue
if preset_name == "lru":
continue
seen.add(preset_name)
ordered.append((preset_name, mode))
return tuple(ordered)
class MultiScoreStrategy:
name = "multi_score"
def compute_score(
self,
*,
key_id: str,
config: Any,
context: dict[str, Any],
) -> float | None:
mode = str(getattr(config, "scheduling_mode", "lru") or "lru").strip().lower()
if mode != "multi_score":
return None
all_key_ids = [str(k) for k in (context.get("all_key_ids") or []) if str(k)]
if not all_key_ids:
return None
lru_scores = context.get("lru_scores", {})
if not isinstance(lru_scores, dict):
lru_scores = {}
latency_avgs = context.get("latency_avgs", {})
if not isinstance(latency_avgs, dict):
latency_avgs = {}
health_scores = context.get("health_scores", {})
if not isinstance(health_scores, dict):
health_scores = {}
cost_totals = context.get("cost_totals", {})
if not isinstance(cost_totals, dict):
cost_totals = {}
keys_by_id = context.get("keys_by_id", {})
if not isinstance(keys_by_id, dict):
keys_by_id = {}
presets = _normalize_presets_from_config(config)
lru_enabled = bool(getattr(config, "lru_enabled", True))
if presets:
return self._compute_preset_score(
key_id=key_id,
all_key_ids=all_key_ids,
presets=presets,
lru_enabled=lru_enabled,
lru_scores=lru_scores,
keys_by_id=keys_by_id,
)
weights = getattr(config, "scoring_weights", None)
w_lru = safe_float(getattr(weights, "lru", 0.3)) or 0.0
w_latency = safe_float(getattr(weights, "latency", 0.25)) or 0.0
w_health = safe_float(getattr(weights, "health", 0.2)) or 0.0
w_cost = safe_float(getattr(weights, "cost_remaining", 0.25)) or 0.0
lru_rank = rank_ascending(key_id, lru_scores, all_key_ids)
latency_rank = rank_ascending(key_id, latency_avgs, all_key_ids)
health_raw = safe_float(health_scores.get(key_id))
if health_raw is None:
health_raw = 1.0
health_norm = 1.0 - max(0.0, min(health_raw, 1.0))
cost_limit = getattr(config, "cost_limit_per_key_tokens", None)
used = safe_float(cost_totals.get(key_id)) or 0.0
if cost_limit is None or int(cost_limit) <= 0:
cost_norm = 0.0
else:
cost_norm = max(0.0, min(used / float(cost_limit), 1.0))
return (
w_lru * lru_rank
+ w_latency * latency_rank
+ w_health * health_norm
+ w_cost * cost_norm
)
def _compute_preset_score(
self,
*,
key_id: str,
all_key_ids: list[str],
presets: tuple[tuple[str, str | None], ...],
lru_enabled: bool,
lru_scores: dict[str, Any],
keys_by_id: dict[str, Any],
) -> float:
lru_rank_asc = rank_ascending(key_id, lru_scores, all_key_ids)
weighted_sum = 0.0
weight_sum = 0.0
for idx, (preset_name, mode) in enumerate(presets):
metric = 0.5
dim = get_preset_dimension(preset_name)
if dim is not None:
metric = dim.compute_metric(
key_id=key_id,
all_key_ids=all_key_ids,
keys_by_id=keys_by_id,
lru_scores=lru_scores,
mode=mode,
)
weight = 1.0 / (1.0 + _POSITIONAL_DECAY * idx)
weighted_sum += metric * weight
weight_sum += weight
if weight_sum <= 0:
return lru_rank_asc
lru_blend = _LRU_BLEND_FACTOR if lru_enabled else 0.0
preset_blend = 1.0 - lru_blend
blended = (weighted_sum / weight_sum) * preset_blend + lru_rank_asc * lru_blend
return max(0.0, min(blended, 1.0))
register_pool_strategy("multi_score", MultiScoreStrategy())
__all__ = [
"MultiScoreStrategy",
]

View File

@@ -23,9 +23,16 @@ class PoolCandidateTrace:
cost_limit: int | None = None
cost_soft_threshold: bool = False
skipped: bool = False
skip_type: str | None = None # cooldown / cost_exhausted
skip_type: str | None = None # cooldown / cost_exhausted / account_blocked / upstream
cooldown_reason: str | None = None
cooldown_ttl: int | None = None
account_block_code: str | None = None
account_block_label: str | None = None
account_block_reason: str | None = None
latency_avg_ms: float = 0.0
health_score: float = 1.0
composite_score: float = 0.0
scoring_mode: str = "lru"
def to_extra_data(self) -> dict[str, Any]:
"""Build dict to merge into ``RequestCandidate.extra_data``."""
@@ -35,8 +42,16 @@ class PoolCandidateTrace:
skip_info["cooldown_reason"] = self.cooldown_reason
if self.cooldown_ttl is not None:
skip_info["cooldown_ttl"] = self.cooldown_ttl
if self.account_block_code is not None:
skip_info["account_block_code"] = self.account_block_code
if self.account_block_label is not None:
skip_info["account_block_label"] = self.account_block_label
if self.account_block_reason is not None:
skip_info["account_block_reason"] = self.account_block_reason
if self.cost_window_usage:
skip_info["cost_window_usage"] = self.cost_window_usage
if self.scoring_mode:
skip_info["scoring_mode"] = self.scoring_mode
return {"pool_skip": skip_info}
sel: dict[str, Any] = {"reason": self.reason}
@@ -50,6 +65,14 @@ class PoolCandidateTrace:
sel["cost_limit"] = self.cost_limit
if self.cost_soft_threshold:
sel["cost_soft_threshold"] = True
if self.latency_avg_ms > 0:
sel["latency_avg_ms"] = round(self.latency_avg_ms, 2)
if self.health_score < 1.0:
sel["health_score"] = round(self.health_score, 4)
if self.reason == "multi_score":
sel["composite_score"] = round(self.composite_score, 6)
if self.scoring_mode:
sel["scoring_mode"] = self.scoring_mode
return {"pool_selection": sel}
@@ -72,6 +95,7 @@ class PoolSchedulingTrace:
"""Build compact dict for ``Usage.request_metadata["pool_summary"]``."""
skipped_cooldown = 0
skipped_cost = 0
skipped_account_blocked = 0
attempted = 0
for t in self.candidate_traces.values():
if t.skipped:
@@ -79,6 +103,8 @@ class PoolSchedulingTrace:
skipped_cooldown += 1
elif t.skip_type == "cost_exhausted":
skipped_cost += 1
elif t.skip_type == "account_blocked":
skipped_account_blocked += 1
if attempted_key_ids is None:
# Backward-compatible behavior: count all schedulable keys.
@@ -101,6 +127,7 @@ class PoolSchedulingTrace:
"attempted": attempted,
"skipped_cooldown": skipped_cooldown,
"skipped_cost": skipped_cost,
"skipped_account_blocked": skipped_account_blocked,
"sticky_session": self.sticky_session_used,
}
if success_key_id:

View File

@@ -235,7 +235,7 @@ class TaskService:
provider_type = str(getattr(provider, "provider_type", "") or "")
session_uuid = TaskService._extract_session_uuid(provider_type, request_body)
manager = PoolManager(provider_id, pool_cfg)
manager = PoolManager(provider_id, pool_cfg, provider_type=provider_type)
candidate_keys = list(candidate.pool_keys or [])
if not candidate_keys and getattr(candidate, "key", None) is not None:
@@ -334,6 +334,7 @@ class TaskService:
async def _pool_on_success(
candidate: Any,
request_body: dict[str, Any] | None,
ttfb_ms: int | None = None,
) -> None:
"""Notify the pool manager about a successful request (sticky + LRU)."""
try:
@@ -354,10 +355,11 @@ class TaskService:
provider_type = str(getattr(provider, "provider_type", "") or "")
session_uuid = TaskService._extract_session_uuid(provider_type, request_body)
mgr = PoolManager(provider_id, pool_cfg)
mgr = PoolManager(provider_id, pool_cfg, provider_type=provider_type)
await mgr.on_request_success(
session_uuid=session_uuid,
key_id=key_id,
ttfb_ms=ttfb_ms,
)
except Exception:
logger.opt(exception=True).debug("Pool on_request_success failed (non-blocking)")
@@ -571,28 +573,38 @@ class TaskService:
candidate_record_id = str(created.id)
candidate_record_map[(candidate_index, retry_index)] = candidate_record_id
response, _provider_name, attempt_id, _provider_id, _endpoint_id, _key_id = (
await request_dispatcher.dispatch(
candidate=candidate,
candidate_index=candidate_index,
retry_index=retry_index,
candidate_record_id=candidate_record_id,
user_api_key=user_api_key,
request_func=request_func,
request_id=request_id,
api_format=api_format_norm,
model_name=model_name,
affinity_key=affinity_key,
global_model_id=global_model_id,
attempt_counter=attempt_counter,
max_attempts=max_attempts_local,
is_stream=is_stream,
)
(
response,
_provider_name,
attempt_id,
_provider_id,
_endpoint_id,
_key_id,
_first_byte_time_ms,
) = await request_dispatcher.dispatch(
candidate=candidate,
candidate_index=candidate_index,
retry_index=retry_index,
candidate_record_id=candidate_record_id,
user_api_key=user_api_key,
request_func=request_func,
request_id=request_id,
api_format=api_format_norm,
model_name=model_name,
affinity_key=affinity_key,
global_model_id=global_model_id,
attempt_counter=attempt_counter,
max_attempts=max_attempts_local,
is_stream=is_stream,
)
_ = (attempt_id, _provider_name, _provider_id, _endpoint_id, _key_id)
# Account Pool: on success, update sticky binding + LRU.
await self._pool_on_success(candidate, request_body)
await self._pool_on_success(
candidate,
request_body,
ttfb_ms=_first_byte_time_ms,
)
if is_stream:
return AttemptResult(