mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
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:
@@ -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)
|
||||
|
||||
@@ -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
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
189
src/services/provider/pool/account_state.py
Normal file
189
src/services/provider/pool/account_state.py
Normal 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",
|
||||
]
|
||||
@@ -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),)
|
||||
|
||||
30
src/services/provider/pool/dimensions/__init__.py
Normal file
30
src/services/provider/pool/dimensions/__init__.py
Normal 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",
|
||||
]
|
||||
227
src/services/provider/pool/dimensions/_helpers.py
Normal file
227
src/services/provider/pool/dimensions/_helpers.py
Normal 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",
|
||||
]
|
||||
52
src/services/provider/pool/dimensions/free_team_first.py
Normal file
52
src/services/provider/pool/dimensions/free_team_first.py
Normal 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())
|
||||
41
src/services/provider/pool/dimensions/quota_balanced.py
Normal file
41
src/services/provider/pool/dimensions/quota_balanced.py
Normal 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())
|
||||
45
src/services/provider/pool/dimensions/recent_refresh.py
Normal file
45
src/services/provider/pool/dimensions/recent_refresh.py
Normal 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())
|
||||
217
src/services/provider/pool/dimensions/registry.py
Normal file
217
src/services/provider/pool/dimensions/registry.py
Normal 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",
|
||||
]
|
||||
36
src/services/provider/pool/dimensions/single_account.py
Normal file
36
src/services/provider/pool/dimensions/single_account.py
Normal 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())
|
||||
91
src/services/provider/pool/health_cache.py
Normal file
91
src/services/provider/pool/health_cache.py
Normal 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",
|
||||
]
|
||||
@@ -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)
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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())
|
||||
|
||||
|
||||
|
||||
8
src/services/provider/pool/strategies/__init__.py
Normal file
8
src/services/provider/pool/strategies/__init__.py
Normal 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"]
|
||||
178
src/services/provider/pool/strategies/multi_score.py
Normal file
178
src/services/provider/pool/strategies/multi_score.py
Normal 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",
|
||||
]
|
||||
@@ -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:
|
||||
|
||||
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user