2026-02-27 13:58:58 +08:00
|
|
|
|
"""Account Pool Manager (provider-agnostic).
|
|
|
|
|
|
|
|
|
|
|
|
Stateless facade that coordinates pool operations for any Provider with
|
|
|
|
|
|
pool configuration enabled. All state lives in Redis via :mod:`redis_ops`.
|
|
|
|
|
|
|
|
|
|
|
|
Usage::
|
|
|
|
|
|
|
|
|
|
|
|
mgr = PoolManager(provider_id, pool_config)
|
|
|
|
|
|
reordered = await mgr.reorder_candidates(session_uuid, candidates)
|
|
|
|
|
|
# ... execute request ...
|
|
|
|
|
|
await mgr.on_request_success(session_uuid=..., key_id=..., tokens_used=...)
|
|
|
|
|
|
"""
|
|
|
|
|
|
|
|
|
|
|
|
from __future__ import annotations
|
|
|
|
|
|
|
|
|
|
|
|
import asyncio
|
|
|
|
|
|
import random
|
|
|
|
|
|
from collections.abc import Callable
|
|
|
|
|
|
from typing import TYPE_CHECKING, Any, TypeVar
|
|
|
|
|
|
|
|
|
|
|
|
from src.core.logger import logger
|
|
|
|
|
|
from src.services.provider.pool import redis_ops
|
2026-03-04 22:06:19 +08:00
|
|
|
|
from src.services.provider.pool.account_state import resolve_pool_account_state
|
2026-02-27 13:58:58 +08:00
|
|
|
|
from src.services.provider.pool.config import PoolConfig
|
2026-03-04 22:06:19 +08:00
|
|
|
|
from src.services.provider.pool.health_cache import get_health_scores
|
2026-02-27 13:58:58 +08:00
|
|
|
|
from src.services.provider.pool.trace import PoolCandidateTrace, PoolSchedulingTrace
|
|
|
|
|
|
|
|
|
|
|
|
if TYPE_CHECKING:
|
|
|
|
|
|
from src.models.database import ProviderAPIKey
|
|
|
|
|
|
from src.services.scheduling.schemas import ProviderCandidate
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
class PoolManager:
|
|
|
|
|
|
"""Coordinate pool-level scheduling for a single Provider."""
|
|
|
|
|
|
|
2026-03-04 22:06:19 +08:00
|
|
|
|
__slots__ = ("provider_id", "config", "provider_type")
|
2026-02-27 13:58:58 +08:00
|
|
|
|
|
2026-03-04 22:06:19 +08:00
|
|
|
|
def __init__(
|
|
|
|
|
|
self,
|
|
|
|
|
|
provider_id: str,
|
|
|
|
|
|
config: PoolConfig,
|
|
|
|
|
|
provider_type: str | None = None,
|
|
|
|
|
|
) -> None:
|
2026-02-27 13:58:58 +08:00
|
|
|
|
self.provider_id = provider_id
|
|
|
|
|
|
self.config = config
|
2026-03-04 22:06:19 +08:00
|
|
|
|
self.provider_type = str(provider_type or "").strip().lower() or None
|
2026-02-27 13:58:58 +08:00
|
|
|
|
|
|
|
|
|
|
# ------------------------------------------------------------------
|
|
|
|
|
|
# Core scheduling: reorder candidate list for pool-aware selection
|
|
|
|
|
|
# ------------------------------------------------------------------
|
|
|
|
|
|
|
|
|
|
|
|
async def reorder_candidates(
|
|
|
|
|
|
self,
|
|
|
|
|
|
session_uuid: str | None,
|
|
|
|
|
|
candidates: list[ProviderCandidate],
|
|
|
|
|
|
) -> list[ProviderCandidate]:
|
|
|
|
|
|
"""Reorder *candidates* according to pool rules.
|
|
|
|
|
|
|
|
|
|
|
|
The returned list keeps the same elements but in a new order:
|
|
|
|
|
|
|
|
|
|
|
|
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.
|
2026-03-04 22:06:19 +08:00
|
|
|
|
2. **Filter** out keys in account-blocked / cooldown / cost-exhausted state (mark
|
2026-02-27 13:58:58 +08:00
|
|
|
|
``is_skipped``).
|
|
|
|
|
|
3. **LRU sort** -- among remaining candidates at the same priority
|
|
|
|
|
|
level, sort by least-recently-used.
|
|
|
|
|
|
4. **Random tiebreak** -- among candidates with identical LRU score.
|
|
|
|
|
|
|
|
|
|
|
|
Also builds a :class:`PoolSchedulingTrace` and attaches per-candidate
|
|
|
|
|
|
trace data via ``_pool_extra_data`` / ``_pool_scheduling_trace``
|
|
|
|
|
|
attributes on candidate objects.
|
|
|
|
|
|
"""
|
|
|
|
|
|
if not candidates:
|
|
|
|
|
|
return candidates
|
|
|
|
|
|
|
|
|
|
|
|
pid = self.provider_id
|
|
|
|
|
|
|
|
|
|
|
|
# Build trace
|
|
|
|
|
|
trace = PoolSchedulingTrace(
|
|
|
|
|
|
provider_id=pid,
|
|
|
|
|
|
total_keys=len(candidates),
|
|
|
|
|
|
session_uuid=session_uuid[:8] if session_uuid else None,
|
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
# --- Strategy: before_select ----------------------------------
|
|
|
|
|
|
strategies = _get_active_strategies(self.config)
|
|
|
|
|
|
key_ids = [str(c.key.id) for c in candidates]
|
|
|
|
|
|
strategy_context: dict[str, Any] = {"session_uuid": session_uuid}
|
|
|
|
|
|
for strategy in strategies:
|
|
|
|
|
|
if hasattr(strategy, "on_before_select"):
|
|
|
|
|
|
try:
|
|
|
|
|
|
filtered = strategy.on_before_select(
|
|
|
|
|
|
provider_id=pid,
|
|
|
|
|
|
key_ids=key_ids,
|
|
|
|
|
|
config=self.config,
|
|
|
|
|
|
context=strategy_context,
|
|
|
|
|
|
)
|
|
|
|
|
|
if filtered is not None:
|
|
|
|
|
|
key_ids = filtered
|
|
|
|
|
|
except Exception:
|
|
|
|
|
|
logger.opt(exception=True).debug(
|
|
|
|
|
|
"Pool[{}]: strategy before_select failed", pid[:8]
|
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
# --- 1. Sticky session ----------------------------------------
|
|
|
|
|
|
sticky_key_id: str | None = None
|
|
|
|
|
|
if session_uuid and self.config.sticky_session_ttl_seconds > 0:
|
|
|
|
|
|
sticky_key_id = await redis_ops.get_sticky_binding(
|
|
|
|
|
|
pid, session_uuid, self.config.sticky_session_ttl_seconds
|
|
|
|
|
|
)
|
|
|
|
|
|
|
2026-03-04 22:06:19 +08:00
|
|
|
|
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
|
|
|
|
|
|
|
2026-02-27 13:58:58 +08:00
|
|
|
|
# --- 2. Batch fetch pool state (parallel) ---------------------
|
|
|
|
|
|
all_key_ids = [str(c.key.id) for c in candidates]
|
|
|
|
|
|
|
|
|
|
|
|
# Fire independent Redis queries concurrently.
|
2026-03-09 14:54:45 +08:00
|
|
|
|
# Only fetch reason (no TTL) on the scheduling hot path -- TTL is only
|
|
|
|
|
|
# used for trace display and costs an extra pipeline command per key.
|
|
|
|
|
|
_cooldown_coro = redis_ops.batch_get_cooldowns(pid, all_key_ids, include_ttl=False)
|
2026-02-27 13:58:58 +08:00
|
|
|
|
_cost_coro = (
|
|
|
|
|
|
redis_ops.batch_get_cost_totals(pid, all_key_ids, self.config.cost_window_seconds)
|
2026-03-05 09:53:23 +08:00
|
|
|
|
if (
|
|
|
|
|
|
self.config.cost_limit_per_key_tokens is not None
|
|
|
|
|
|
or self.config.scheduling_mode == "multi_score"
|
|
|
|
|
|
)
|
2026-02-27 13:58:58 +08:00
|
|
|
|
else None
|
|
|
|
|
|
)
|
2026-03-09 15:22:17 +08:00
|
|
|
|
# LRU scores are needed both for plain LRU sorting and for multi_score
|
|
|
|
|
|
# dimensions (e.g. cache_affinity / single_account) that rely on
|
|
|
|
|
|
# lru_scores data.
|
|
|
|
|
|
_need_lru = self.config.lru_enabled or self.config.scheduling_mode == "multi_score"
|
|
|
|
|
|
_lru_coro = redis_ops.get_lru_scores(pid, all_key_ids) if _need_lru else None
|
2026-03-04 22:06:19 +08:00
|
|
|
|
_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
|
|
|
|
|
|
)
|
2026-02-27 13:58:58 +08:00
|
|
|
|
|
|
|
|
|
|
# Gather all non-None coroutines in parallel.
|
|
|
|
|
|
coros: list[Any] = [_cooldown_coro]
|
|
|
|
|
|
_cost_idx = -1
|
|
|
|
|
|
_lru_idx = -1
|
2026-03-04 22:06:19 +08:00
|
|
|
|
_latency_idx = -1
|
2026-02-27 13:58:58 +08:00
|
|
|
|
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)
|
2026-03-04 22:06:19 +08:00
|
|
|
|
if _latency_coro is not None:
|
|
|
|
|
|
_latency_idx = len(coros)
|
|
|
|
|
|
coros.append(_latency_coro)
|
2026-02-27 13:58:58 +08:00
|
|
|
|
|
|
|
|
|
|
gathered = await asyncio.gather(*coros)
|
|
|
|
|
|
|
2026-03-09 14:54:45 +08:00
|
|
|
|
cooldowns: dict[str, str | None] = gathered[0]
|
2026-02-27 13:58:58 +08:00
|
|
|
|
|
|
|
|
|
|
# Cost check
|
|
|
|
|
|
cost_exhausted: set[str] = set()
|
|
|
|
|
|
cost_soft: set[str] = set()
|
|
|
|
|
|
cost_totals: dict[str, int] = {}
|
|
|
|
|
|
if _cost_idx >= 0:
|
|
|
|
|
|
cost_totals = gathered[_cost_idx]
|
|
|
|
|
|
limit = self.config.cost_limit_per_key_tokens
|
2026-03-05 09:53:23 +08:00
|
|
|
|
if limit is not None:
|
|
|
|
|
|
for kid, total in cost_totals.items():
|
|
|
|
|
|
if total >= limit:
|
|
|
|
|
|
cost_exhausted.add(kid)
|
|
|
|
|
|
elif total >= limit * self.config.cost_soft_threshold_percent / 100:
|
|
|
|
|
|
cost_soft.add(kid)
|
2026-02-27 13:58:58 +08:00
|
|
|
|
|
|
|
|
|
|
# LRU scores
|
|
|
|
|
|
lru_scores: dict[str, float] = {}
|
|
|
|
|
|
if _lru_idx >= 0:
|
|
|
|
|
|
lru_scores = gathered[_lru_idx]
|
|
|
|
|
|
|
2026-03-04 22:06:19 +08:00
|
|
|
|
# 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(
|
|
|
|
|
|
{
|
2026-03-10 14:55:48 +08:00
|
|
|
|
"provider_type": provider_type,
|
2026-03-04 22:06:19 +08:00
|
|
|
|
"all_key_ids": all_key_ids,
|
|
|
|
|
|
"lru_scores": lru_scores,
|
|
|
|
|
|
"cost_totals": cost_totals,
|
2026-03-05 09:53:23 +08:00
|
|
|
|
"cost_limit_per_key_tokens": self.config.cost_limit_per_key_tokens,
|
2026-03-04 22:06:19 +08:00
|
|
|
|
"latency_avgs": latency_avgs,
|
|
|
|
|
|
"health_scores": health_scores,
|
|
|
|
|
|
"keys_by_id": {str(c.key.id): c.key for c in candidates},
|
|
|
|
|
|
}
|
|
|
|
|
|
)
|
|
|
|
|
|
|
2026-02-27 13:58:58 +08:00
|
|
|
|
# --- Strategy: compute_score ----------------------------------
|
2026-03-04 22:06:19 +08:00
|
|
|
|
custom_scores: dict[str, float] = {}
|
2026-02-27 13:58:58 +08:00
|
|
|
|
for strategy in strategies:
|
|
|
|
|
|
if hasattr(strategy, "compute_score"):
|
|
|
|
|
|
for kid in all_key_ids:
|
|
|
|
|
|
try:
|
|
|
|
|
|
custom = strategy.compute_score(
|
|
|
|
|
|
key_id=kid,
|
|
|
|
|
|
config=self.config,
|
|
|
|
|
|
context=strategy_context,
|
|
|
|
|
|
)
|
|
|
|
|
|
if custom is not None:
|
2026-03-04 22:06:19 +08:00
|
|
|
|
custom_scores[kid] = float(custom)
|
|
|
|
|
|
lru_scores[kid] = float(custom)
|
2026-02-27 13:58:58 +08:00
|
|
|
|
except Exception:
|
|
|
|
|
|
pass
|
|
|
|
|
|
|
|
|
|
|
|
# --- 3. Classify candidates -----------------------------------
|
2026-03-10 12:33:15 +08:00
|
|
|
|
# Use precomputed account states from CandidateBuilder when available
|
|
|
|
|
|
# (upstream_metadata is deferred on pool keys to save memory).
|
|
|
|
|
|
# Fall back to on-the-fly resolution for non-pool candidates.
|
2026-03-09 14:54:45 +08:00
|
|
|
|
account_states: dict[str, Any] = {}
|
|
|
|
|
|
for c in candidates:
|
|
|
|
|
|
kid = str(c.key.id)
|
|
|
|
|
|
if kid not in account_states:
|
2026-03-10 12:33:15 +08:00
|
|
|
|
precomputed = getattr(c.key, "_pool_account_state", None)
|
|
|
|
|
|
if precomputed is not None:
|
|
|
|
|
|
account_states[kid] = precomputed
|
|
|
|
|
|
else:
|
|
|
|
|
|
account_states[kid] = 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),
|
|
|
|
|
|
)
|
2026-03-09 14:54:45 +08:00
|
|
|
|
|
2026-02-27 13:58:58 +08:00
|
|
|
|
sticky_candidate: ProviderCandidate | None = None
|
|
|
|
|
|
available: list[ProviderCandidate] = []
|
|
|
|
|
|
skipped: list[ProviderCandidate] = []
|
|
|
|
|
|
|
|
|
|
|
|
for c in candidates:
|
|
|
|
|
|
kid = str(c.key.id)
|
|
|
|
|
|
ct = PoolCandidateTrace(key_id=kid)
|
2026-03-04 22:06:19 +08:00
|
|
|
|
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])
|
2026-02-27 13:58:58 +08:00
|
|
|
|
|
|
|
|
|
|
# Already skipped upstream?
|
|
|
|
|
|
if c.is_skipped:
|
|
|
|
|
|
skipped.append(c)
|
|
|
|
|
|
ct.skipped = True
|
|
|
|
|
|
ct.skip_type = "upstream"
|
|
|
|
|
|
trace.candidate_traces[kid] = ct
|
|
|
|
|
|
continue
|
|
|
|
|
|
|
2026-03-09 14:54:45 +08:00
|
|
|
|
# Account blocked?
|
|
|
|
|
|
account_state = account_states[kid]
|
2026-03-04 22:06:19 +08:00
|
|
|
|
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
|
|
|
|
|
|
|
2026-02-27 13:58:58 +08:00
|
|
|
|
cd_reason = cooldowns.get(kid)
|
|
|
|
|
|
if cd_reason is not None:
|
|
|
|
|
|
c.is_skipped = True
|
|
|
|
|
|
c.skip_reason = f"pool cooldown: {cd_reason}"
|
|
|
|
|
|
skipped.append(c)
|
|
|
|
|
|
ct.skipped = True
|
|
|
|
|
|
ct.skip_type = "cooldown"
|
|
|
|
|
|
ct.cooldown_reason = cd_reason
|
2026-03-09 14:54:45 +08:00
|
|
|
|
ct.cooldown_ttl = None # TTL skipped on hot path for perf
|
2026-02-27 13:58:58 +08:00
|
|
|
|
_attach_pool_extra(c, ct)
|
|
|
|
|
|
trace.candidate_traces[kid] = ct
|
|
|
|
|
|
continue
|
|
|
|
|
|
|
|
|
|
|
|
# Cost exhausted?
|
|
|
|
|
|
if kid in cost_exhausted:
|
|
|
|
|
|
c.is_skipped = True
|
|
|
|
|
|
c.skip_reason = "pool cost limit reached"
|
|
|
|
|
|
skipped.append(c)
|
|
|
|
|
|
ct.skipped = True
|
|
|
|
|
|
ct.skip_type = "cost_exhausted"
|
|
|
|
|
|
ct.cost_window_usage = cost_totals.get(kid, 0)
|
|
|
|
|
|
ct.cost_limit = self.config.cost_limit_per_key_tokens
|
|
|
|
|
|
_attach_pool_extra(c, ct)
|
|
|
|
|
|
trace.candidate_traces[kid] = ct
|
|
|
|
|
|
continue
|
|
|
|
|
|
|
|
|
|
|
|
# Sticky hit?
|
|
|
|
|
|
if sticky_key_id and kid == sticky_key_id:
|
|
|
|
|
|
sticky_candidate = c
|
|
|
|
|
|
ct.reason = "sticky"
|
|
|
|
|
|
ct.sticky_hit = True
|
|
|
|
|
|
trace.sticky_session_used = True
|
|
|
|
|
|
else:
|
|
|
|
|
|
available.append(c)
|
2026-03-04 22:06:19 +08:00
|
|
|
|
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"
|
2026-02-27 13:58:58 +08:00
|
|
|
|
|
|
|
|
|
|
ct.lru_score = lru_scores.get(kid, 0.0)
|
|
|
|
|
|
ct.cost_window_usage = cost_totals.get(kid, 0)
|
|
|
|
|
|
ct.cost_limit = self.config.cost_limit_per_key_tokens
|
|
|
|
|
|
if kid in cost_soft:
|
|
|
|
|
|
ct.cost_soft_threshold = True
|
|
|
|
|
|
_attach_pool_extra(c, ct)
|
|
|
|
|
|
trace.candidate_traces[kid] = ct
|
|
|
|
|
|
|
|
|
|
|
|
# --- 4. Sort available by LRU ---------------------------------
|
|
|
|
|
|
if lru_scores and available:
|
|
|
|
|
|
available.sort(key=lambda c: lru_scores.get(str(c.key.id), 0.0))
|
|
|
|
|
|
|
|
|
|
|
|
# Random tiebreak among candidates with the same LRU score
|
|
|
|
|
|
if len(available) > 1 and lru_scores:
|
|
|
|
|
|
_shuffle_same_score_groups(available, lru_scores)
|
|
|
|
|
|
|
|
|
|
|
|
# --- 5. Assemble final order ----------------------------------
|
|
|
|
|
|
result: list[ProviderCandidate] = []
|
|
|
|
|
|
if sticky_candidate is not None:
|
|
|
|
|
|
result.append(sticky_candidate)
|
|
|
|
|
|
result.extend(available)
|
|
|
|
|
|
result.extend(skipped)
|
|
|
|
|
|
|
|
|
|
|
|
if sticky_candidate:
|
|
|
|
|
|
logger.debug(
|
|
|
|
|
|
"Pool[{}]: sticky hit key={}",
|
|
|
|
|
|
pid[:8],
|
|
|
|
|
|
sticky_key_id and sticky_key_id[:8],
|
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
# --- Strategy: after_select -----------------------------------
|
|
|
|
|
|
if result:
|
|
|
|
|
|
first_kid = str(result[0].key.id)
|
|
|
|
|
|
first_trace = trace.candidate_traces.get(first_kid)
|
|
|
|
|
|
for strategy in strategies:
|
|
|
|
|
|
if hasattr(strategy, "on_after_select") and first_trace:
|
|
|
|
|
|
try:
|
|
|
|
|
|
strategy.on_after_select(
|
|
|
|
|
|
provider_id=pid,
|
|
|
|
|
|
selected_key_id=first_kid,
|
|
|
|
|
|
trace=first_trace,
|
|
|
|
|
|
config=self.config,
|
|
|
|
|
|
context=strategy_context,
|
|
|
|
|
|
)
|
|
|
|
|
|
except Exception:
|
|
|
|
|
|
pass
|
|
|
|
|
|
|
|
|
|
|
|
# Attach the full trace to the first candidate for downstream use.
|
|
|
|
|
|
if result:
|
|
|
|
|
|
setattr(result[0], "_pool_scheduling_trace", trace)
|
|
|
|
|
|
|
|
|
|
|
|
return result
|
|
|
|
|
|
|
2026-03-03 17:24:22 +08:00
|
|
|
|
async def select_pool_keys(
|
|
|
|
|
|
self,
|
|
|
|
|
|
session_uuid: str | None,
|
|
|
|
|
|
keys: list[ProviderAPIKey],
|
2026-03-09 13:26:54 +08:00
|
|
|
|
*,
|
|
|
|
|
|
availability_checker: (
|
|
|
|
|
|
Callable[[ProviderAPIKey], tuple[bool, str | None, str | None]] | None
|
|
|
|
|
|
) = None,
|
|
|
|
|
|
page_size: int = 50,
|
2026-03-03 17:24:22 +08:00
|
|
|
|
) -> tuple[list[ProviderAPIKey], PoolSchedulingTrace]:
|
|
|
|
|
|
"""Select and order pool keys with trace output.
|
|
|
|
|
|
|
|
|
|
|
|
Reuses :meth:`reorder_candidates` logic by adapting keys to lightweight
|
|
|
|
|
|
candidate-like wrappers, then propagates skip/trace metadata back onto
|
|
|
|
|
|
each key object for downstream execution/recording.
|
2026-03-09 13:26:54 +08:00
|
|
|
|
|
|
|
|
|
|
When *availability_checker* is provided, post-sort availability checks
|
|
|
|
|
|
are performed lazily: only the top *page_size* non-skipped keys are
|
|
|
|
|
|
checked at a time; if all fail, the next page is checked, and so on.
|
|
|
|
|
|
Keys beyond the last checked page are marked as ``deferred`` (skipped
|
|
|
|
|
|
without checking) to avoid unnecessary CPU work on large pools.
|
2026-03-03 17:24:22 +08:00
|
|
|
|
"""
|
|
|
|
|
|
if not keys:
|
|
|
|
|
|
return (
|
|
|
|
|
|
[],
|
|
|
|
|
|
PoolSchedulingTrace(
|
|
|
|
|
|
provider_id=self.provider_id,
|
|
|
|
|
|
total_keys=0,
|
|
|
|
|
|
session_uuid=session_uuid[:8] if session_uuid else None,
|
|
|
|
|
|
),
|
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
class _KeyCandidate:
|
2026-03-04 02:24:46 +08:00
|
|
|
|
__slots__ = (
|
|
|
|
|
|
"key",
|
|
|
|
|
|
"is_skipped",
|
|
|
|
|
|
"skip_reason",
|
|
|
|
|
|
"_pool_extra_data",
|
|
|
|
|
|
"_pool_scheduling_trace",
|
|
|
|
|
|
)
|
2026-03-03 17:24:22 +08:00
|
|
|
|
|
|
|
|
|
|
def __init__(self, key: ProviderAPIKey) -> None:
|
|
|
|
|
|
self.key = key
|
|
|
|
|
|
self.is_skipped = False
|
|
|
|
|
|
self.skip_reason: str | None = None
|
2026-03-04 02:24:46 +08:00
|
|
|
|
self._pool_extra_data: dict | None = None
|
|
|
|
|
|
self._pool_scheduling_trace: PoolSchedulingTrace | None = None
|
2026-03-03 17:24:22 +08:00
|
|
|
|
|
|
|
|
|
|
wrappers = [_KeyCandidate(k) for k in keys]
|
|
|
|
|
|
reordered_wrappers = await self.reorder_candidates(session_uuid, wrappers) # type: ignore[arg-type]
|
|
|
|
|
|
|
|
|
|
|
|
trace: PoolSchedulingTrace | None = None
|
|
|
|
|
|
if reordered_wrappers:
|
|
|
|
|
|
maybe_trace = getattr(reordered_wrappers[0], "_pool_scheduling_trace", None)
|
|
|
|
|
|
if isinstance(maybe_trace, PoolSchedulingTrace):
|
|
|
|
|
|
trace = maybe_trace
|
|
|
|
|
|
if trace is None:
|
|
|
|
|
|
trace = PoolSchedulingTrace(
|
|
|
|
|
|
provider_id=self.provider_id,
|
|
|
|
|
|
total_keys=len(keys),
|
|
|
|
|
|
session_uuid=session_uuid[:8] if session_uuid else None,
|
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
ordered_keys: list[ProviderAPIKey] = []
|
|
|
|
|
|
for order_idx, wrapped in enumerate(reordered_wrappers):
|
|
|
|
|
|
key = wrapped.key
|
|
|
|
|
|
is_skipped = bool(getattr(wrapped, "is_skipped", False))
|
|
|
|
|
|
skip_reason = str(getattr(wrapped, "skip_reason", "") or "")
|
|
|
|
|
|
setattr(key, "_pool_skipped", is_skipped)
|
|
|
|
|
|
setattr(key, "_pool_skip_reason", skip_reason if skip_reason else None)
|
|
|
|
|
|
setattr(key, "_pool_order_index", order_idx)
|
|
|
|
|
|
pool_extra = getattr(wrapped, "_pool_extra_data", None)
|
|
|
|
|
|
setattr(
|
|
|
|
|
|
key, "_pool_extra_data", dict(pool_extra) if isinstance(pool_extra, dict) else {}
|
|
|
|
|
|
)
|
|
|
|
|
|
ordered_keys.append(key)
|
|
|
|
|
|
|
2026-03-09 13:26:54 +08:00
|
|
|
|
# -- 分页可用性检查 --
|
|
|
|
|
|
# 排序后对非 skipped key 分页调用 availability_checker,
|
|
|
|
|
|
# 找到 page_size 个可用 key 后停止检查,剩余标记 deferred。
|
|
|
|
|
|
if availability_checker is not None:
|
|
|
|
|
|
available_count = 0
|
|
|
|
|
|
found_enough = False
|
|
|
|
|
|
for key in ordered_keys:
|
|
|
|
|
|
if getattr(key, "_pool_skipped", False):
|
|
|
|
|
|
continue
|
|
|
|
|
|
if found_enough:
|
|
|
|
|
|
setattr(key, "_pool_skipped", True)
|
|
|
|
|
|
setattr(key, "_pool_skip_reason", "deferred")
|
|
|
|
|
|
continue
|
|
|
|
|
|
is_available, skip_reason_check, mapping_model = availability_checker(key)
|
|
|
|
|
|
if not is_available:
|
|
|
|
|
|
setattr(key, "_pool_skipped", True)
|
|
|
|
|
|
setattr(key, "_pool_skip_reason", skip_reason_check)
|
|
|
|
|
|
else:
|
|
|
|
|
|
if mapping_model:
|
|
|
|
|
|
setattr(key, "_pool_mapping_matched_model", mapping_model)
|
|
|
|
|
|
available_count += 1
|
|
|
|
|
|
if available_count >= page_size:
|
|
|
|
|
|
found_enough = True
|
|
|
|
|
|
|
2026-03-03 17:24:22 +08:00
|
|
|
|
return ordered_keys, trace
|
|
|
|
|
|
|
2026-02-27 13:58:58 +08:00
|
|
|
|
# ------------------------------------------------------------------
|
|
|
|
|
|
# Post-request hooks
|
|
|
|
|
|
# ------------------------------------------------------------------
|
|
|
|
|
|
|
|
|
|
|
|
async def on_request_success(
|
|
|
|
|
|
self,
|
|
|
|
|
|
*,
|
|
|
|
|
|
session_uuid: str | None,
|
|
|
|
|
|
key_id: str,
|
|
|
|
|
|
tokens_used: int = 0,
|
2026-03-04 22:06:19 +08:00
|
|
|
|
ttfb_ms: int | None = None,
|
2026-02-27 13:58:58 +08:00
|
|
|
|
) -> None:
|
|
|
|
|
|
"""Called after a successful upstream request."""
|
|
|
|
|
|
pid = self.provider_id
|
|
|
|
|
|
|
|
|
|
|
|
# Bind sticky session
|
|
|
|
|
|
if session_uuid and self.config.sticky_session_ttl_seconds > 0:
|
|
|
|
|
|
await redis_ops.set_sticky_binding(
|
|
|
|
|
|
pid, session_uuid, key_id, self.config.sticky_session_ttl_seconds
|
|
|
|
|
|
)
|
|
|
|
|
|
|
2026-03-09 15:22:17 +08:00
|
|
|
|
# Touch LRU -- needed for both plain LRU mode and multi_score dimensions
|
|
|
|
|
|
# (e.g. cache_affinity) that rely on LRU timestamps.
|
|
|
|
|
|
if self.config.lru_enabled or self.config.scheduling_mode == "multi_score":
|
2026-02-27 13:58:58 +08:00
|
|
|
|
await redis_ops.touch_lru(pid, key_id)
|
|
|
|
|
|
|
|
|
|
|
|
# Record cost
|
|
|
|
|
|
if tokens_used > 0 and self.config.cost_limit_per_key_tokens is not None:
|
|
|
|
|
|
await redis_ops.add_cost_entry(
|
|
|
|
|
|
pid, key_id, tokens_used, self.config.cost_window_seconds
|
|
|
|
|
|
)
|
|
|
|
|
|
|
2026-03-04 22:06:19 +08:00
|
|
|
|
# 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,
|
|
|
|
|
|
)
|
|
|
|
|
|
|
2026-02-27 13:58:58 +08:00
|
|
|
|
async def on_request_error(
|
|
|
|
|
|
self,
|
|
|
|
|
|
*,
|
|
|
|
|
|
key_id: str,
|
|
|
|
|
|
status_code: int,
|
|
|
|
|
|
error_body: str | None = None,
|
|
|
|
|
|
response_headers: dict[str, str] | None = None,
|
|
|
|
|
|
) -> None:
|
|
|
|
|
|
"""Called after an upstream error. Delegates to health policy."""
|
|
|
|
|
|
# Import lazily to avoid circular deps
|
|
|
|
|
|
from src.services.provider.pool.health_policy import apply_health_policy
|
|
|
|
|
|
|
|
|
|
|
|
await apply_health_policy(
|
|
|
|
|
|
provider_id=self.provider_id,
|
|
|
|
|
|
key_id=key_id,
|
|
|
|
|
|
status_code=status_code,
|
|
|
|
|
|
error_body=error_body,
|
|
|
|
|
|
response_headers=response_headers,
|
|
|
|
|
|
config=self.config,
|
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
# ------------------------------------------------------------------
|
|
|
|
|
|
# Key schedulability check (used by candidate_builder)
|
|
|
|
|
|
# ------------------------------------------------------------------
|
|
|
|
|
|
|
|
|
|
|
|
async def is_key_schedulable(self, key_id: str) -> tuple[bool, str | None]:
|
|
|
|
|
|
"""Check if *key_id* is currently schedulable (not in cooldown, not
|
|
|
|
|
|
cost-exhausted). Returns ``(True, None)`` or ``(False, reason)``.
|
|
|
|
|
|
"""
|
|
|
|
|
|
pid = self.provider_id
|
|
|
|
|
|
|
|
|
|
|
|
# Cooldown check
|
|
|
|
|
|
cd = await redis_ops.get_cooldown(pid, key_id)
|
|
|
|
|
|
if cd is not None:
|
|
|
|
|
|
return False, f"pool cooldown: {cd}"
|
|
|
|
|
|
|
|
|
|
|
|
# Cost check
|
|
|
|
|
|
if self.config.cost_limit_per_key_tokens is not None:
|
|
|
|
|
|
total = await redis_ops.get_cost_window_total(
|
|
|
|
|
|
pid, key_id, self.config.cost_window_seconds
|
|
|
|
|
|
)
|
|
|
|
|
|
if total >= self.config.cost_limit_per_key_tokens:
|
|
|
|
|
|
return False, "pool cost limit reached"
|
|
|
|
|
|
|
|
|
|
|
|
return True, None
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
# Backward-compatible alias
|
|
|
|
|
|
ClaudeCodePoolManager = PoolManager
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
# Helpers
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
|
|
|
|
|
|
_T = TypeVar("_T")
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def _attach_pool_extra(candidate: Any, ct: PoolCandidateTrace) -> None:
|
|
|
|
|
|
"""Attach pool trace extra_data onto a candidate object."""
|
|
|
|
|
|
existing = getattr(candidate, "_pool_extra_data", None) or {}
|
|
|
|
|
|
existing.update(ct.to_extra_data())
|
|
|
|
|
|
setattr(candidate, "_pool_extra_data", existing)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def _get_active_strategies(config: PoolConfig) -> list[Any]:
|
|
|
|
|
|
"""Get active strategies for the given config (lazy import)."""
|
|
|
|
|
|
if not config.strategies:
|
|
|
|
|
|
return []
|
|
|
|
|
|
try:
|
2026-03-04 22:06:19 +08:00
|
|
|
|
# Import triggers built-in strategy registration via module-level side effects.
|
|
|
|
|
|
from src.services.provider.pool import strategies as _builtin_strategies # noqa: F401
|
2026-02-27 13:58:58 +08:00
|
|
|
|
from src.services.provider.pool.strategy import get_active_strategies
|
|
|
|
|
|
|
|
|
|
|
|
return get_active_strategies(config.strategies)
|
|
|
|
|
|
except Exception:
|
|
|
|
|
|
return []
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def _shuffle_same_score(
|
|
|
|
|
|
items: list[_T],
|
|
|
|
|
|
lru_scores: dict[str, float],
|
|
|
|
|
|
key_fn: Callable[[_T], str],
|
|
|
|
|
|
) -> None:
|
|
|
|
|
|
"""In-place random shuffle within groups that share the same LRU score."""
|
|
|
|
|
|
if len(items) <= 1:
|
|
|
|
|
|
return
|
|
|
|
|
|
|
|
|
|
|
|
i = 0
|
|
|
|
|
|
while i < len(items):
|
|
|
|
|
|
score_i = lru_scores.get(key_fn(items[i]), 0.0)
|
|
|
|
|
|
j = i + 1
|
|
|
|
|
|
while j < len(items) and lru_scores.get(key_fn(items[j]), 0.0) == score_i:
|
|
|
|
|
|
j += 1
|
|
|
|
|
|
if j - i > 1:
|
|
|
|
|
|
group = items[i:j]
|
|
|
|
|
|
random.shuffle(group)
|
|
|
|
|
|
items[i:j] = group
|
|
|
|
|
|
i = j
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def _shuffle_same_score_groups(
|
|
|
|
|
|
candidates: list[ProviderCandidate],
|
|
|
|
|
|
lru_scores: dict[str, float],
|
|
|
|
|
|
) -> None:
|
|
|
|
|
|
_shuffle_same_score(candidates, lru_scores, lambda c: str(c.key.id))
|