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

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

View File

@@ -0,0 +1,92 @@
"""Tests for pool account-state resolution helpers."""
from __future__ import annotations
from src.services.provider.pool.account_state import resolve_pool_account_state
def test_resolve_from_kiro_banned_metadata() -> None:
state = resolve_pool_account_state(
provider_type="kiro",
upstream_metadata={"kiro": {"is_banned": True, "ban_reason": "account suspended"}},
oauth_invalid_reason=None,
)
assert state.blocked is True
assert state.code == "account_banned"
assert state.label == "账号封禁"
assert state.reason == "account suspended"
def test_resolve_from_antigravity_forbidden_metadata() -> None:
state = resolve_pool_account_state(
provider_type="antigravity",
upstream_metadata={"antigravity": {"is_forbidden": True, "forbidden_reason": "403"}},
oauth_invalid_reason=None,
)
assert state.blocked is True
assert state.code == "account_forbidden"
assert state.label == "访问受限"
assert state.reason == "403"
def test_resolve_from_structured_oauth_reason() -> None:
state = resolve_pool_account_state(
provider_type="codex",
upstream_metadata=None,
oauth_invalid_reason="[ACCOUNT_BLOCK] Google requires verification",
)
assert state.blocked is True
assert state.code == "account_blocked"
assert state.label == "账号异常"
assert state.reason == "Google requires verification"
def test_resolve_from_keyword_oauth_reason() -> None:
state = resolve_pool_account_state(
provider_type=None,
upstream_metadata={},
oauth_invalid_reason="organization has been disabled by admin",
)
assert state.blocked is True
assert state.code == "account_blocked"
assert state.label == "账号异常"
def test_resolve_healthy_state() -> None:
state = resolve_pool_account_state(
provider_type="codex",
upstream_metadata={"codex": {"primary_used_percent": 30}},
oauth_invalid_reason="Token expired",
)
assert state.blocked is False
assert state.code is None
def test_bare_forbidden_not_treated_as_account_block() -> None:
"""HTTP 403 'Forbidden' from token refresh should not be misclassified."""
state = resolve_pool_account_state(
provider_type="codex",
upstream_metadata={},
oauth_invalid_reason="Forbidden",
)
assert state.blocked is False
def test_kiro_oauth_reason_text_detected_as_block() -> None:
state = resolve_pool_account_state(
provider_type="kiro",
upstream_metadata={},
oauth_invalid_reason="账户已封禁: Terms of Service violation",
)
assert state.blocked is True
assert state.code == "account_blocked"
def test_antigravity_oauth_reason_text_detected_as_block() -> None:
state = resolve_pool_account_state(
provider_type="antigravity",
upstream_metadata={},
oauth_invalid_reason="账户访问被禁止: 403 Forbidden",
)
assert state.blocked is True
assert state.code == "account_blocked"

View File

@@ -4,6 +4,8 @@ from __future__ import annotations
from src.services.provider.pool.config import (
PoolConfig,
SchedulingPreset,
ScoringWeights,
UnschedulableRule,
parse_pool_config,
)
@@ -21,6 +23,14 @@ def test_parse_pool_config_returns_defaults_for_empty_advanced() -> None:
assert cfg.sticky_session_ttl_seconds == 3600
assert cfg.load_threshold_percent == 80
assert cfg.lru_enabled is True
assert cfg.scheduling_mode == "lru"
assert cfg.scoring_weights == ScoringWeights()
# Default: only LRU preset enabled
assert len(cfg.scheduling_presets) == 1
assert cfg.scheduling_presets[0].preset == "lru"
assert cfg.scheduling_presets[0].enabled is True
assert cfg.latency_window_seconds == 3600
assert cfg.latency_sample_limit == 50
assert cfg.cost_window_seconds == 18000
assert cfg.cost_limit_per_key_tokens is None
assert cfg.cost_soft_threshold_percent == 80
@@ -31,13 +41,24 @@ def test_parse_pool_config_returns_defaults_for_empty_advanced() -> None:
assert cfg.unschedulable_rules == []
def test_parse_pool_config_overrides_values() -> None:
def test_parse_pool_config_overrides_values_legacy_string_list() -> None:
"""Legacy string-list format with scheduling_mode/lru_enabled."""
cfg = parse_pool_config(
{
"pool_advanced": {
"sticky_session_ttl_seconds": 7200,
"load_threshold_percent": 90,
"lru_enabled": False,
"scheduling_mode": "multi_score",
"scheduling_presets": ["free_team_first", "recent_refresh", "free_team_first"],
"scoring_weights": {
"lru": 0.1,
"latency": 0.5,
"health": 0.2,
"cost_remaining": 0.2,
},
"latency_window_seconds": 7200,
"latency_sample_limit": 80,
"cost_window_seconds": 36000,
"cost_limit_per_key_tokens": 100000,
"cost_soft_threshold_percent": 70,
@@ -52,6 +73,20 @@ def test_parse_pool_config_overrides_values() -> None:
assert cfg.sticky_session_ttl_seconds == 7200
assert cfg.load_threshold_percent == 90
assert cfg.lru_enabled is False
assert cfg.scheduling_mode == "multi_score"
# Legacy string list → SchedulingPreset objects, deduped
preset_names = tuple(p.preset for p in cfg.scheduling_presets)
assert "free_team_first" in preset_names
assert "recent_refresh" in preset_names
assert cfg.scoring_weights == ScoringWeights(
lru=0.1,
latency=0.5,
health=0.2,
cost_remaining=0.2,
)
assert cfg.latency_window_seconds == 7200
assert cfg.latency_sample_limit == 80
assert "multi_score" in cfg.strategies
assert cfg.cost_window_seconds == 36000
assert cfg.cost_limit_per_key_tokens == 100000
assert cfg.cost_soft_threshold_percent == 70
@@ -61,6 +96,114 @@ def test_parse_pool_config_overrides_values() -> None:
assert cfg.health_policy_enabled is False
def test_parse_pool_config_new_object_list_format() -> None:
"""New object-list format: [{preset, enabled, mode}]."""
cfg = parse_pool_config(
{
"pool_advanced": {
"scheduling_presets": [
{"preset": "lru", "enabled": True},
{"preset": "free_team_first", "enabled": True, "mode": "free_only"},
{"preset": "quota_balanced", "enabled": False},
{"preset": "recent_refresh", "enabled": True},
],
}
}
)
assert cfg is not None
assert len(cfg.scheduling_presets) == 4
lru = cfg.scheduling_presets[0]
assert lru.preset == "lru"
assert lru.enabled is True
ftf = cfg.scheduling_presets[1]
assert ftf.preset == "free_team_first"
assert ftf.enabled is True
assert ftf.mode == "free_only"
qb = cfg.scheduling_presets[2]
assert qb.preset == "quota_balanced"
assert qb.enabled is False
rr = cfg.scheduling_presets[3]
assert rr.preset == "recent_refresh"
assert rr.enabled is True
# Derived fields: lru enabled, non-lru enabled → multi_score
assert cfg.lru_enabled is True
assert cfg.scheduling_mode == "multi_score"
def test_parse_pool_config_new_format_lru_only() -> None:
"""When only LRU is enabled, scheduling_mode should be 'lru'."""
cfg = parse_pool_config(
{
"pool_advanced": {
"scheduling_presets": [
{"preset": "lru", "enabled": True},
{"preset": "quota_balanced", "enabled": False},
],
}
}
)
assert cfg is not None
assert cfg.lru_enabled is True
assert cfg.scheduling_mode == "lru"
def test_parse_pool_config_new_format_lru_disabled() -> None:
"""LRU disabled, other presets enabled → multi_score + lru_enabled=False."""
cfg = parse_pool_config(
{
"pool_advanced": {
"scheduling_presets": [
{"preset": "lru", "enabled": False},
{"preset": "quota_balanced", "enabled": True},
],
}
}
)
assert cfg is not None
assert cfg.lru_enabled is False
assert cfg.scheduling_mode == "multi_score"
def test_parse_pool_config_new_format_free_team_mode_validation() -> None:
"""Invalid mode falls back to 'both'."""
cfg = parse_pool_config(
{
"pool_advanced": {
"scheduling_presets": [
{"preset": "free_team_first", "enabled": True, "mode": "invalid_mode"},
],
}
}
)
assert cfg is not None
ftf = [p for p in cfg.scheduling_presets if p.preset == "free_team_first"][0]
assert ftf.mode == "both"
def test_parse_pool_config_new_format_dedup_presets() -> None:
"""Duplicate presets in object list should be deduplicated."""
cfg = parse_pool_config(
{
"pool_advanced": {
"scheduling_presets": [
{"preset": "lru", "enabled": True},
{"preset": "lru", "enabled": False},
{"preset": "quota_balanced", "enabled": True},
],
}
}
)
assert cfg is not None
lru_presets = [p for p in cfg.scheduling_presets if p.preset == "lru"]
assert len(lru_presets) == 1
assert lru_presets[0].enabled is True # first occurrence wins
def test_parse_pool_config_parses_unschedulable_rules() -> None:
cfg = parse_pool_config(
{
@@ -97,6 +240,47 @@ def test_parse_pool_config_handles_invalid_types_gracefully() -> None:
assert cfg.cost_limit_per_key_tokens is None # default for opt_int
def test_parse_pool_config_invalid_scheduling_mode_falls_back_to_lru() -> None:
cfg = parse_pool_config({"pool_advanced": {"scheduling_mode": "unknown"}})
assert cfg is not None
assert cfg.scheduling_mode == "lru"
def test_parse_pool_config_scoring_weights_invalid_values_are_clamped() -> None:
cfg = parse_pool_config(
{
"pool_advanced": {
"scheduling_mode": "multi_score",
"scoring_weights": {
"lru": 2.0,
"latency": -1.0,
"health": "bad",
},
}
}
)
assert cfg is not None
assert cfg.scoring_weights.lru == 1.0
assert cfg.scoring_weights.latency == 0.0
assert cfg.scoring_weights.health == 0.2
def test_parse_pool_config_invalid_scheduling_presets_are_ignored() -> None:
cfg = parse_pool_config(
{
"pool_advanced": {
"scheduling_mode": "multi_score",
"scheduling_presets": ["quota_balanced", "unknown", 123, "single_account"],
}
}
)
assert cfg is not None
preset_names = tuple(p.preset for p in cfg.scheduling_presets if p.preset != "lru")
assert "quota_balanced" in preset_names
assert "single_account" in preset_names
assert "unknown" not in preset_names
def test_pool_config_is_frozen() -> None:
cfg = PoolConfig()
try:
@@ -104,3 +288,12 @@ def test_pool_config_is_frozen() -> None:
assert False, "Should have raised FrozenInstanceError"
except AttributeError:
pass
def test_scheduling_preset_is_frozen() -> None:
preset = SchedulingPreset(preset="lru", enabled=True)
try:
preset.enabled = False # type: ignore[misc]
assert False, "Should have raised FrozenInstanceError"
except AttributeError:
pass

View File

@@ -0,0 +1,58 @@
"""Tests for pool health cache helpers."""
from __future__ import annotations
from types import SimpleNamespace
from src.services.provider.pool import health_cache
def setup_function() -> None:
health_cache._clear_cache_for_tests()
def teardown_function() -> None:
health_cache._clear_cache_for_tests()
def test_aggregate_health_score_uses_lowest_format_score() -> None:
score = health_cache.aggregate_health_score(
{
"openai:chat": {"health_score": 0.92},
"openai:responses": {"health_score": 0.61},
}
)
assert score == 0.61
def test_get_health_scores_uses_cache_for_same_provider() -> None:
key = SimpleNamespace(id="k1", health_by_format={"f1": {"health_score": 0.7}})
first = health_cache.get_health_scores("p1", [key])
assert first["k1"] == 0.7
key.health_by_format = {"f1": {"health_score": 0.2}}
second = health_cache.get_health_scores("p1", [key])
assert second["k1"] == 0.7
def test_get_health_scores_merges_missing_keys_into_cache() -> None:
k1 = SimpleNamespace(id="k1", health_by_format={"f1": {"health_score": 0.7}})
first = health_cache.get_health_scores("p1", [k1])
assert first == {"k1": 0.7}
# Request with a new key k2 -- k1 should come from cache, k2 freshly computed
k1_stale = SimpleNamespace(id="k1", health_by_format={"f1": {"health_score": 0.1}})
k2 = SimpleNamespace(id="k2", health_by_format={"f1": {"health_score": 0.5}})
second = health_cache.get_health_scores("p1", [k1_stale, k2])
assert second["k1"] == 0.7 # cached, not recomputed
assert second["k2"] == 0.5 # freshly computed
def test_invalidate_provider_health_scores_clears_cache_entry() -> None:
key = SimpleNamespace(id="k1", health_by_format={"f1": {"health_score": 0.8}})
_ = health_cache.get_health_scores("p1", [key])
health_cache.invalidate_provider_health_scores("p1")
key.health_by_format = {"f1": {"health_score": 0.3}}
refreshed = health_cache.get_health_scores("p1", [key])
assert refreshed["k1"] == 0.3

View File

@@ -7,13 +7,23 @@ from unittest.mock import AsyncMock, patch
import pytest
from src.services.provider.pool.config import PoolConfig
from src.services.provider.pool.config import PoolConfig, SchedulingPreset, ScoringWeights
from src.services.provider.pool.manager import PoolManager
def _make_candidate(key_id: str, *, is_skipped: bool = False) -> SimpleNamespace:
def _make_candidate(
key_id: str,
*,
is_skipped: bool = False,
upstream_metadata: dict | None = None,
oauth_invalid_reason: str | None = None,
) -> SimpleNamespace:
return SimpleNamespace(
key=SimpleNamespace(id=key_id),
key=SimpleNamespace(
id=key_id,
upstream_metadata=upstream_metadata,
oauth_invalid_reason=oauth_invalid_reason,
),
is_skipped=is_skipped,
skip_reason=None,
)
@@ -127,6 +137,125 @@ async def test_reorder_cost_exhausted_keys_are_skipped() -> None:
assert result[0].key.id == "key-2"
@pytest.mark.asyncio
async def test_reorder_account_blocked_keys_are_skipped() -> None:
pool = PoolManager("provider-1", PoolConfig(), provider_type="kiro")
c1 = _make_candidate(
"key-1",
upstream_metadata={"kiro": {"is_banned": True, "ban_reason": "account suspended"}},
)
c2 = _make_candidate("key-2")
with (
patch(
"src.services.provider.pool.redis_ops.get_sticky_binding",
new_callable=AsyncMock,
return_value=None,
),
patch(
"src.services.provider.pool.redis_ops.batch_get_cooldowns",
new_callable=AsyncMock,
return_value={"key-1": None, "key-2": None},
),
patch(
"src.services.provider.pool.redis_ops.get_lru_scores",
new_callable=AsyncMock,
return_value={},
),
):
result = await pool.reorder_candidates(None, [c1, c2])
assert c1.is_skipped is True
assert "account blocked" in (c1.skip_reason or "")
assert result[0].key.id == "key-2"
assert c1._pool_extra_data["pool_skip"]["type"] == "account_blocked"
assert c1._pool_extra_data["pool_skip"]["account_block_label"] == "账号封禁"
@pytest.mark.asyncio
async def test_reorder_multi_score_uses_composite_score() -> None:
pool = PoolManager(
"provider-1",
PoolConfig(
scheduling_mode="multi_score",
strategies=("multi_score",),
scoring_weights=ScoringWeights(lru=0.0, latency=1.0, health=0.0, cost_remaining=0.0),
),
)
c1 = _make_candidate("key-1")
c2 = _make_candidate("key-2")
with (
patch(
"src.services.provider.pool.redis_ops.get_sticky_binding",
new_callable=AsyncMock,
return_value=None,
),
patch(
"src.services.provider.pool.redis_ops.batch_get_cooldowns",
new_callable=AsyncMock,
return_value={"key-1": None, "key-2": None},
),
patch(
"src.services.provider.pool.redis_ops.get_lru_scores",
new_callable=AsyncMock,
return_value={"key-1": 10.0, "key-2": 20.0},
),
patch(
"src.services.provider.pool.redis_ops.batch_get_latency_avgs",
new_callable=AsyncMock,
return_value={"key-1": 600.0, "key-2": 120.0},
),
):
result = await pool.reorder_candidates(None, [c1, c2])
assert result[0].key.id == "key-2"
assert result[0]._pool_extra_data["pool_selection"]["reason"] == "multi_score"
assert result[0]._pool_extra_data["pool_selection"]["scoring_mode"] == "multi_score"
@pytest.mark.asyncio
async def test_reorder_multi_score_presets_free_team_first() -> None:
pool = PoolManager(
"provider-1",
PoolConfig(
scheduling_mode="multi_score",
strategies=("multi_score",),
scheduling_presets=(
SchedulingPreset(preset="free_team_first", enabled=True, mode="both"),
),
),
)
c1 = _make_candidate("key-1", upstream_metadata={"codex": {"plan_type": "plus"}})
c2 = _make_candidate("key-2", upstream_metadata={"codex": {"plan_type": "team"}})
with (
patch(
"src.services.provider.pool.redis_ops.get_sticky_binding",
new_callable=AsyncMock,
return_value=None,
),
patch(
"src.services.provider.pool.redis_ops.batch_get_cooldowns",
new_callable=AsyncMock,
return_value={"key-1": None, "key-2": None},
),
patch(
"src.services.provider.pool.redis_ops.get_lru_scores",
new_callable=AsyncMock,
return_value={"key-1": 100.0, "key-2": 100.0},
),
patch(
"src.services.provider.pool.redis_ops.batch_get_latency_avgs",
new_callable=AsyncMock,
return_value={},
),
):
result = await pool.reorder_candidates(None, [c1, c2])
assert result[0].key.id == "key-2"
@pytest.mark.asyncio
async def test_reorder_lru_sorts_least_recently_used_first(
pool: PoolManager,
@@ -216,6 +345,41 @@ async def test_on_request_success_records_cost_when_configured() -> None:
mock_cost.assert_called_once_with("provider-1", "key-1", 500, 18000)
@pytest.mark.asyncio
async def test_on_request_success_records_latency_when_multi_score_enabled() -> None:
pool = PoolManager(
"provider-1",
PoolConfig(
scheduling_mode="multi_score",
latency_window_seconds=7200,
latency_sample_limit=80,
),
)
with (
patch(
"src.services.provider.pool.redis_ops.set_sticky_binding",
new_callable=AsyncMock,
),
patch(
"src.services.provider.pool.redis_ops.touch_lru",
new_callable=AsyncMock,
),
patch(
"src.services.provider.pool.redis_ops.record_latency",
new_callable=AsyncMock,
) as mock_latency,
):
await pool.on_request_success(
session_uuid=None,
key_id="key-1",
tokens_used=0,
ttfb_ms=321,
)
mock_latency.assert_called_once_with("provider-1", "key-1", 321, 7200, 80)
# ---------------------------------------------------------------------------
# on_request_error
# ---------------------------------------------------------------------------

View File

@@ -0,0 +1,233 @@
"""Tests for built-in multi-score pool strategy."""
from __future__ import annotations
from types import SimpleNamespace
from src.services.provider.pool.config import PoolConfig, SchedulingPreset, ScoringWeights
from src.services.provider.pool.strategies.multi_score import MultiScoreStrategy
from src.services.provider.pool.strategy import get_pool_strategy, register_pool_strategy
def _context() -> dict:
return {
"all_key_ids": ["k1", "k2", "k3"],
"lru_scores": {"k1": 100.0, "k2": 200.0, "k3": 300.0},
"latency_avgs": {"k1": 120.0, "k2": 300.0, "k3": 600.0},
"health_scores": {"k1": 0.95, "k2": 0.7, "k3": 0.4},
"cost_totals": {"k1": 100, "k2": 300, "k3": 900},
}
def _key_with_metadata(metadata: dict) -> SimpleNamespace:
return SimpleNamespace(upstream_metadata=metadata)
def test_multi_score_returns_none_when_mode_not_enabled() -> None:
strategy = MultiScoreStrategy()
cfg = PoolConfig(scheduling_mode="lru")
score = strategy.compute_score(key_id="k1", config=cfg, context=_context())
assert score is None
def test_multi_score_prefers_low_latency_when_latency_weight_is_high() -> None:
strategy = MultiScoreStrategy()
cfg = PoolConfig(
scheduling_mode="multi_score",
scoring_weights=ScoringWeights(lru=0.0, latency=1.0, health=0.0, cost_remaining=0.0),
)
ctx = _context()
s1 = strategy.compute_score(key_id="k1", config=cfg, context=ctx)
s2 = strategy.compute_score(key_id="k2", config=cfg, context=ctx)
s3 = strategy.compute_score(key_id="k3", config=cfg, context=ctx)
assert s1 is not None and s2 is not None and s3 is not None
assert s1 < s2 < s3
def test_multi_score_combines_health_and_cost() -> None:
strategy = MultiScoreStrategy()
cfg = PoolConfig(
scheduling_mode="multi_score",
scoring_weights=ScoringWeights(lru=0.0, latency=0.0, health=0.5, cost_remaining=0.5),
cost_limit_per_key_tokens=1000,
)
ctx = _context()
s1 = strategy.compute_score(key_id="k1", config=cfg, context=ctx)
s3 = strategy.compute_score(key_id="k3", config=cfg, context=ctx)
assert s1 is not None and s3 is not None
assert s1 < s3
def test_multi_score_strategy_is_registered() -> None:
register_pool_strategy("multi_score", MultiScoreStrategy())
registered = get_pool_strategy("multi_score")
assert registered is not None
def test_multi_score_preset_free_team_first_prefers_free_or_team() -> None:
strategy = MultiScoreStrategy()
cfg = PoolConfig(
scheduling_mode="multi_score",
scheduling_presets=(SchedulingPreset(preset="free_team_first", enabled=True, mode="both"),),
)
ctx = {
"all_key_ids": ["k1", "k2", "k3"],
"lru_scores": {"k1": 100.0, "k2": 100.0, "k3": 100.0},
"keys_by_id": {
"k1": _key_with_metadata({"codex": {"plan_type": "plus"}}),
"k2": _key_with_metadata({"codex": {"plan_type": "free"}}),
"k3": _key_with_metadata({"codex": {"plan_type": "team"}}),
},
}
s1 = strategy.compute_score(key_id="k1", config=cfg, context=ctx)
s2 = strategy.compute_score(key_id="k2", config=cfg, context=ctx)
s3 = strategy.compute_score(key_id="k3", config=cfg, context=ctx)
assert s1 is not None and s2 is not None and s3 is not None
assert s2 < s1
assert s3 < s1
def test_multi_score_preset_free_team_first_free_only_mode() -> None:
"""free_only mode: free is preferred, team is mid-priority."""
strategy = MultiScoreStrategy()
cfg = PoolConfig(
scheduling_mode="multi_score",
scheduling_presets=(
SchedulingPreset(preset="free_team_first", enabled=True, mode="free_only"),
),
)
ctx = {
"all_key_ids": ["k1", "k2", "k3"],
"lru_scores": {"k1": 100.0, "k2": 100.0, "k3": 100.0},
"keys_by_id": {
"k1": _key_with_metadata({"codex": {"plan_type": "plus"}}),
"k2": _key_with_metadata({"codex": {"plan_type": "free"}}),
"k3": _key_with_metadata({"codex": {"plan_type": "team"}}),
},
}
s1 = strategy.compute_score(key_id="k1", config=cfg, context=ctx)
s2 = strategy.compute_score(key_id="k2", config=cfg, context=ctx)
s3 = strategy.compute_score(key_id="k3", config=cfg, context=ctx)
assert s1 is not None and s2 is not None and s3 is not None
# free < team < plus
assert s2 < s3 < s1
def test_multi_score_preset_free_team_first_team_only_mode() -> None:
"""team_only mode: team is preferred, free is mid-priority."""
strategy = MultiScoreStrategy()
cfg = PoolConfig(
scheduling_mode="multi_score",
scheduling_presets=(
SchedulingPreset(preset="free_team_first", enabled=True, mode="team_only"),
),
)
ctx = {
"all_key_ids": ["k1", "k2", "k3"],
"lru_scores": {"k1": 100.0, "k2": 100.0, "k3": 100.0},
"keys_by_id": {
"k1": _key_with_metadata({"codex": {"plan_type": "plus"}}),
"k2": _key_with_metadata({"codex": {"plan_type": "free"}}),
"k3": _key_with_metadata({"codex": {"plan_type": "team"}}),
},
}
s1 = strategy.compute_score(key_id="k1", config=cfg, context=ctx)
s2 = strategy.compute_score(key_id="k2", config=cfg, context=ctx)
s3 = strategy.compute_score(key_id="k3", config=cfg, context=ctx)
assert s1 is not None and s2 is not None and s3 is not None
# team < free < plus
assert s3 < s2 < s1
def test_multi_score_preset_recent_refresh_prefers_nearer_reset() -> None:
strategy = MultiScoreStrategy()
cfg = PoolConfig(
scheduling_mode="multi_score",
scheduling_presets=(SchedulingPreset(preset="recent_refresh", enabled=True),),
)
ctx = {
"all_key_ids": ["k1", "k2"],
"lru_scores": {"k1": 100.0, "k2": 100.0},
"keys_by_id": {
"k1": _key_with_metadata({"codex": {"primary_reset_seconds": 600}}),
"k2": _key_with_metadata({"codex": {"primary_reset_seconds": 120}}),
},
}
s1 = strategy.compute_score(key_id="k1", config=cfg, context=ctx)
s2 = strategy.compute_score(key_id="k2", config=cfg, context=ctx)
assert s1 is not None and s2 is not None
assert s2 < s1
def test_multi_score_preset_single_account_prefers_latest_used() -> None:
strategy = MultiScoreStrategy()
cfg = PoolConfig(
scheduling_mode="multi_score",
scheduling_presets=(SchedulingPreset(preset="single_account", enabled=True),),
)
ctx = {
"all_key_ids": ["k1", "k2", "k3"],
"lru_scores": {"k1": 100.0, "k2": 900.0, "k3": 400.0},
"keys_by_id": {},
}
s1 = strategy.compute_score(key_id="k1", config=cfg, context=ctx)
s2 = strategy.compute_score(key_id="k2", config=cfg, context=ctx)
s3 = strategy.compute_score(key_id="k3", config=cfg, context=ctx)
assert s1 is not None and s2 is not None and s3 is not None
assert s2 < s3 < s1
def test_multi_score_lru_disabled_no_blend() -> None:
"""When lru_enabled=False, LRU blend factor should be 0."""
strategy = MultiScoreStrategy()
cfg = PoolConfig(
scheduling_mode="multi_score",
lru_enabled=False,
scheduling_presets=(SchedulingPreset(preset="quota_balanced", enabled=True),),
)
ctx = {
"all_key_ids": ["k1", "k2"],
"lru_scores": {"k1": 100.0, "k2": 200.0},
"keys_by_id": {
"k1": _key_with_metadata({"codex": {"primary_used_percent": 80}}),
"k2": _key_with_metadata({"codex": {"primary_used_percent": 20}}),
},
}
s1 = strategy.compute_score(key_id="k1", config=cfg, context=ctx)
s2 = strategy.compute_score(key_id="k2", config=cfg, context=ctx)
assert s1 is not None and s2 is not None
assert s2 < s1
def test_multi_score_disabled_presets_are_skipped() -> None:
"""Disabled presets should not affect scoring."""
strategy = MultiScoreStrategy()
cfg = PoolConfig(
scheduling_mode="multi_score",
scheduling_presets=(
SchedulingPreset(preset="free_team_first", enabled=False),
SchedulingPreset(preset="quota_balanced", enabled=True),
),
)
ctx = {
"all_key_ids": ["k1", "k2"],
"lru_scores": {"k1": 100.0, "k2": 100.0},
"keys_by_id": {
"k1": _key_with_metadata(
{
"codex": {"plan_type": "free", "primary_used_percent": 80},
}
),
"k2": _key_with_metadata(
{
"codex": {"plan_type": "plus", "primary_used_percent": 20},
}
),
},
}
s1 = strategy.compute_score(key_id="k1", config=cfg, context=ctx)
s2 = strategy.compute_score(key_id="k2", config=cfg, context=ctx)
assert s1 is not None and s2 is not None
# quota_balanced only: k2 (20%) should score lower (better) than k1 (80%)
# free_team_first is disabled so plan_type should not matter
assert s2 < s1

View File

@@ -0,0 +1,83 @@
"""Tests for pool preset dimension registry and built-in dimensions."""
from __future__ import annotations
from types import SimpleNamespace
import src.services.provider.pool.dimensions # noqa: F401
from src.services.provider.pool.dimensions import get_preset_dimension, get_preset_names
def _key(metadata: dict, *, plan_type: str | None = None) -> SimpleNamespace:
return SimpleNamespace(upstream_metadata=metadata, oauth_plan_type=plan_type)
def test_registry_discovers_builtin_dimensions() -> None:
names = get_preset_names()
assert {"free_team_first", "recent_refresh", "quota_balanced", "single_account"}.issubset(names)
def test_universal_dimensions_are_applicable_to_any_provider() -> None:
for name in ("quota_balanced", "single_account"):
dim = get_preset_dimension(name)
assert dim is not None
assert dim.is_applicable("openai") is True
assert dim.is_applicable("codex") is True
assert dim.is_applicable("unknown_provider") is True
def test_provider_specific_dimensions_are_filtered() -> None:
for name in ("free_team_first", "recent_refresh"):
dim = get_preset_dimension(name)
assert dim is not None
assert dim.is_applicable("codex") is True
assert dim.is_applicable("kiro") is True
assert dim.is_applicable("openai") is False
def test_builtin_dimensions_compute_metric_in_range() -> None:
all_key_ids = ["k1", "k2", "k3"]
lru_scores = {"k1": 100.0, "k2": 800.0, "k3": 300.0}
keys_by_id = {
"k1": _key(
{
"codex": {
"plan_type": "plus",
"primary_reset_seconds": 900,
"primary_used_percent": 60,
}
}
),
"k2": _key(
{
"codex": {
"plan_type": "free",
"primary_reset_seconds": 120,
"primary_used_percent": 20,
}
}
),
"k3": _key(
{
"kiro": {
"next_reset_at": 4102444800,
"usage_percentage": 45,
"subscription_title": "Kiro Team",
}
},
plan_type="team",
),
}
for name in get_preset_names():
dim = get_preset_dimension(name)
assert dim is not None
mode = dim.default_mode
metric = dim.compute_metric(
key_id="k1",
all_key_ids=all_key_ids,
keys_by_id=keys_by_id,
lru_scores=lru_scores,
mode=mode,
)
assert 0.0 <= metric <= 1.0

View File

@@ -0,0 +1,88 @@
"""Tests for pool redis latency operations."""
from __future__ import annotations
from unittest.mock import AsyncMock, patch
import pytest
from src.services.provider.pool import redis_ops
class _FakePipe:
def __init__(self, execute_result: list[object] | None = None) -> None:
self.ops: list[tuple] = []
self._execute_result = execute_result or []
def zadd(self, key: str, mapping: dict[str, float]):
self.ops.append(("zadd", key, mapping))
return self
def zremrangebyscore(self, key: str, start: str, stop: float):
self.ops.append(("zremrangebyscore", key, start, stop))
return self
def zremrangebyrank(self, key: str, start: int, stop: int):
self.ops.append(("zremrangebyrank", key, start, stop))
return self
def expire(self, key: str, ttl: int):
self.ops.append(("expire", key, ttl))
return self
def eval(self, script: str, numkeys: int, key: str, window_start: str):
self.ops.append(("eval", script, numkeys, key, window_start))
return self
async def execute(self):
return self._execute_result
class _FakeRedis:
def __init__(self, pipe: _FakePipe) -> None:
self._pipe = pipe
def pipeline(self) -> _FakePipe:
return self._pipe
@pytest.mark.asyncio
async def test_record_latency_writes_sample_and_trims() -> None:
pipe = _FakePipe()
fake_redis = _FakeRedis(pipe)
with patch(
"src.services.provider.pool.redis_ops._get_redis",
new_callable=AsyncMock,
return_value=fake_redis,
):
await redis_ops.record_latency(
provider_id="prov-1",
key_id="key-1",
ttfb_ms=250,
window_seconds=3600,
sample_limit=50,
)
op_names = [item[0] for item in pipe.ops]
assert op_names == ["zadd", "zremrangebyscore", "zremrangebyrank", "expire"]
assert pipe.ops[0][1] == "ap:prov-1:latency:key-1"
assert pipe.ops[2][2:] == (0, -51)
@pytest.mark.asyncio
async def test_batch_get_latency_avgs_returns_numeric_results_only() -> None:
pipe = _FakePipe(execute_result=[120.5, None, "330"])
fake_redis = _FakeRedis(pipe)
with patch(
"src.services.provider.pool.redis_ops._get_redis",
new_callable=AsyncMock,
return_value=fake_redis,
):
result = await redis_ops.batch_get_latency_avgs(
provider_id="prov-1",
key_ids=["k1", "k2", "k3"],
window_seconds=3600,
)
assert result == {"k1": 120.5, "k3": 330.0}
assert len([item for item in pipe.ops if item[0] == "eval"]) == 3

View File

@@ -28,7 +28,7 @@ def _snapshot(**overrides: object) -> PoolSchedulingSnapshot:
def test_default_dimension_registry_contains_core_dimensions() -> None:
names = list_pool_scheduling_dimensions()
assert names == ("manual", "cooldown", "circuit", "cost", "health")
assert names == ("account_state", "manual", "cooldown", "circuit", "cost", "latency", "health")
def test_summary_available_when_all_dimensions_ok() -> None:
@@ -54,6 +54,37 @@ def test_summary_blocked_when_manual_disabled() -> None:
assert summary.score < 100.0
def test_summary_blocked_when_account_state_blocked() -> None:
dimensions = evaluate_pool_scheduling_dimensions(
_snapshot(
account_blocked=True,
account_block_label="账号封禁",
account_block_reason="account suspended",
)
)
summary = summarize_pool_scheduling_dimensions(dimensions)
assert summary.status == "blocked"
assert summary.reason == "account_banned"
assert summary.candidate_eligible is False
assert summary.blocked_count >= 1
def test_account_state_takes_priority_over_manual_disabled() -> None:
dimensions = evaluate_pool_scheduling_dimensions(
_snapshot(
is_active=False,
account_blocked=True,
account_block_label="访问受限",
account_block_reason="forbidden",
)
)
summary = summarize_pool_scheduling_dimensions(dimensions)
assert summary.status == "blocked"
assert summary.reason == "account_forbidden"
def test_summary_degraded_when_cost_reaches_soft_threshold() -> None:
dimensions = evaluate_pool_scheduling_dimensions(
_snapshot(cost_window_usage=8200, cost_limit=10000, cost_soft_threshold_percent=80)
@@ -92,3 +123,11 @@ def test_dimension_result_keeps_degraded_health_details() -> None:
assert isinstance(health, PoolSchedulingDimensionResult)
assert health.status == "degraded"
assert health.detail == "0.65"
def test_latency_dimension_degraded_when_latency_high() -> None:
dimensions = evaluate_pool_scheduling_dimensions(_snapshot(latency_avg_ms=3200))
latency = next((item for item in dimensions if item.code == "latency_high"), None)
assert isinstance(latency, PoolSchedulingDimensionResult)
assert latency.status == "degraded"
assert latency.detail == "3200ms"

View File

@@ -48,7 +48,25 @@ class TestPoolCandidateTraceExtraData:
ct = PoolCandidateTrace(key_id="k4", reason="random")
data = ct.to_extra_data()
sel = data["pool_selection"]
assert sel == {"reason": "random"}
assert sel["reason"] == "random"
assert sel["scoring_mode"] == "lru"
def test_selected_multi_score_fields(self) -> None:
ct = PoolCandidateTrace(
key_id="k4b",
reason="multi_score",
scoring_mode="multi_score",
latency_avg_ms=245.7,
health_score=0.82,
composite_score=0.372156,
)
data = ct.to_extra_data()
sel = data["pool_selection"]
assert sel["reason"] == "multi_score"
assert sel["scoring_mode"] == "multi_score"
assert sel["latency_avg_ms"] == 245.7
assert sel["health_score"] == 0.82
assert sel["composite_score"] == 0.372156
def test_skipped_cooldown(self) -> None:
ct = PoolCandidateTrace(
@@ -78,11 +96,28 @@ class TestPoolCandidateTraceExtraData:
assert skip["cost_window_usage"] == 2000
assert "cooldown_reason" not in skip
def test_skipped_account_blocked(self) -> None:
ct = PoolCandidateTrace(
key_id="k6b",
skipped=True,
skip_type="account_blocked",
account_block_code="account_banned",
account_block_label="账号封禁",
account_block_reason="account suspended",
)
data = ct.to_extra_data()
skip = data["pool_skip"]
assert skip["type"] == "account_blocked"
assert skip["account_block_code"] == "account_banned"
assert skip["account_block_label"] == "账号封禁"
assert skip["account_block_reason"] == "account suspended"
def test_skipped_minimal(self) -> None:
ct = PoolCandidateTrace(key_id="k7", skipped=True, skip_type="upstream")
data = ct.to_extra_data()
skip = data["pool_skip"]
assert skip == {"type": "upstream"}
assert skip["type"] == "upstream"
assert skip["scoring_mode"] == "lru"
# ---------------------------------------------------------------------------
@@ -113,6 +148,7 @@ class TestPoolSchedulingTraceSummary:
assert summary["attempted"] == 2
assert summary["skipped_cooldown"] == 2
assert summary["skipped_cost"] == 1
assert summary["skipped_account_blocked"] == 0
assert summary["sticky_session"] is True
assert summary["success_key_id"] == "k1"[:8]
assert summary["success_reason"] == "sticky"
@@ -129,6 +165,7 @@ class TestPoolSchedulingTraceSummary:
assert summary["attempted"] == 2
assert summary["skipped_cooldown"] == 0
assert summary["skipped_cost"] == 0
assert summary["skipped_account_blocked"] == 0
assert "success_key_id" not in summary
assert "success_reason" not in summary
@@ -143,3 +180,15 @@ class TestPoolSchedulingTraceSummary:
assert summary["attempted"] == 0
assert summary["skipped_cooldown"] == 1
assert summary["skipped_cost"] == 1
assert summary["skipped_account_blocked"] == 0
def test_summary_with_account_blocked(self) -> None:
trace = PoolSchedulingTrace(provider_id="prov-4", total_keys=2)
trace.candidate_traces = {
"k1": PoolCandidateTrace(key_id="k1", skipped=True, skip_type="account_blocked"),
"k2": PoolCandidateTrace(key_id="k2", reason="lru"),
}
summary = trace.build_summary(success_key_id="k2")
assert summary["attempted"] == 1
assert summary["skipped_account_blocked"] == 1

View File

@@ -2,13 +2,8 @@
from __future__ import annotations
from types import SimpleNamespace
from src.api.admin.pool.routes import (
_build_pool_scheduling_state,
_is_known_banned_key,
_is_known_banned_reason,
)
from src.api.admin.pool.routes import _build_pool_scheduling_state
from src.services.provider.pool.account_state import resolve_pool_account_state
def test_pool_scheduling_state_manual_disabled_is_blocked() -> None:
@@ -24,6 +19,10 @@ def test_pool_scheduling_state_manual_disabled_is_blocked() -> None:
dimensions,
) = _build_pool_scheduling_state(
is_active=False,
account_blocked=False,
account_block_label=None,
account_block_reason=None,
latency_avg_ms=None,
cooldown_reason=None,
cooldown_ttl_seconds=None,
circuit_breaker_open=False,
@@ -55,6 +54,10 @@ def test_pool_scheduling_state_cooldown_detail_is_mapped() -> None:
dimensions,
) = _build_pool_scheduling_state(
is_active=True,
account_blocked=False,
account_block_label=None,
account_block_reason=None,
latency_avg_ms=None,
cooldown_reason="rate_limited_429",
cooldown_ttl_seconds=180,
circuit_breaker_open=False,
@@ -84,6 +87,10 @@ def test_pool_scheduling_state_cost_soft_is_degraded() -> None:
_dimensions,
) = _build_pool_scheduling_state(
is_active=True,
account_blocked=False,
account_block_label=None,
account_block_reason=None,
latency_avg_ms=None,
cooldown_reason=None,
cooldown_ttl_seconds=None,
circuit_breaker_open=False,
@@ -101,28 +108,36 @@ def test_pool_scheduling_state_cost_soft_is_degraded() -> None:
def test_known_banned_reason_account_block_prefix() -> None:
assert _is_known_banned_reason("[ACCOUNT_BLOCK] Google 要求验证账号") is True
state = resolve_pool_account_state(
provider_type="codex",
upstream_metadata={},
oauth_invalid_reason="[ACCOUNT_BLOCK] Google 要求验证账号",
)
assert state.blocked is True
def test_known_banned_key_detects_kiro_banned_metadata() -> None:
key = SimpleNamespace(
state = resolve_pool_account_state(
provider_type="kiro",
upstream_metadata={"kiro": {"is_banned": True}},
oauth_invalid_reason=None,
)
assert _is_known_banned_key(key, "kiro") is True
assert state.blocked is True
def test_known_banned_key_detects_reason_keywords() -> None:
key = SimpleNamespace(
state = resolve_pool_account_state(
provider_type="antigravity",
upstream_metadata={},
oauth_invalid_reason="AWS account temporarily suspended",
)
assert _is_known_banned_key(key, "antigravity") is True
assert state.blocked is True
def test_known_banned_key_does_not_treat_token_expired_as_banned() -> None:
key = SimpleNamespace(
state = resolve_pool_account_state(
provider_type="kiro",
upstream_metadata={"kiro": {"is_banned": False}},
oauth_invalid_reason="access token expired",
)
assert _is_known_banned_key(key, "kiro") is False
assert state.blocked is False