feat(pool): 增加 OAuth 账号池管理功能

- 新增 pool manager / strategy / health_policy / cost_tracker / redis_ops 等核心模块
- 新增 pool admin API 路由与 schemas
- 新增 OAuth 账号类型解析 (oauth_plan)
- 前端增加 PoolManagement 页面、PoolConfigDialog、PoolImportDialog、PoolStatusCard 组件
- 补充 pool config / cost tracker / health policy / manager / strategy / trace 等测试
This commit is contained in:
fawney19
2026-02-27 13:58:58 +08:00
parent 579b5e4623
commit 76a6d0ce8e
35 changed files with 5761 additions and 0 deletions

View File

@@ -0,0 +1,106 @@
"""Tests for pool_config.py — configuration parsing and defaults."""
from __future__ import annotations
from src.services.provider.pool.config import (
PoolConfig,
UnschedulableRule,
parse_pool_config,
)
def test_parse_pool_config_returns_none_when_no_advanced_section() -> None:
assert parse_pool_config({}) is None
assert parse_pool_config(None) is None
assert parse_pool_config({"other_key": 1}) is None
def test_parse_pool_config_returns_defaults_for_empty_advanced() -> None:
cfg = parse_pool_config({"pool_advanced": {}})
assert cfg is not None
assert cfg.sticky_session_ttl_seconds == 3600
assert cfg.load_threshold_percent == 80
assert cfg.lru_enabled is True
assert cfg.cost_window_seconds == 18000
assert cfg.cost_limit_per_key_tokens is None
assert cfg.cost_soft_threshold_percent == 80
assert cfg.rate_limit_cooldown_seconds == 300
assert cfg.overload_cooldown_seconds == 30
assert cfg.proactive_refresh_seconds == 180
assert cfg.health_policy_enabled is True
assert cfg.unschedulable_rules == []
def test_parse_pool_config_overrides_values() -> None:
cfg = parse_pool_config(
{
"pool_advanced": {
"sticky_session_ttl_seconds": 7200,
"load_threshold_percent": 90,
"lru_enabled": False,
"cost_window_seconds": 36000,
"cost_limit_per_key_tokens": 100000,
"cost_soft_threshold_percent": 70,
"rate_limit_cooldown_seconds": 600,
"overload_cooldown_seconds": 60,
"proactive_refresh_seconds": 300,
"health_policy_enabled": False,
}
}
)
assert cfg is not None
assert cfg.sticky_session_ttl_seconds == 7200
assert cfg.load_threshold_percent == 90
assert cfg.lru_enabled is False
assert cfg.cost_window_seconds == 36000
assert cfg.cost_limit_per_key_tokens == 100000
assert cfg.cost_soft_threshold_percent == 70
assert cfg.rate_limit_cooldown_seconds == 600
assert cfg.overload_cooldown_seconds == 60
assert cfg.proactive_refresh_seconds == 300
assert cfg.health_policy_enabled is False
def test_parse_pool_config_parses_unschedulable_rules() -> None:
cfg = parse_pool_config(
{
"pool_advanced": {
"unschedulable_rules": [
{"keyword": "rate_limit", "duration_minutes": 10},
{"keyword": "overloaded"},
{"invalid": "entry"}, # should be skipped
"not_a_dict", # should be skipped
]
}
}
)
assert cfg is not None
assert len(cfg.unschedulable_rules) == 2
assert cfg.unschedulable_rules[0] == UnschedulableRule(
keyword="rate_limit", duration_minutes=10
)
assert cfg.unschedulable_rules[1] == UnschedulableRule(keyword="overloaded", duration_minutes=5)
def test_parse_pool_config_handles_invalid_types_gracefully() -> None:
# Invalid int values should fall back to defaults
cfg = parse_pool_config(
{
"pool_advanced": {
"sticky_session_ttl_seconds": "not_a_number",
"cost_limit_per_key_tokens": "bad",
}
}
)
assert cfg is not None
assert cfg.sticky_session_ttl_seconds == 3600 # default
assert cfg.cost_limit_per_key_tokens is None # default for opt_int
def test_pool_config_is_frozen() -> None:
cfg = PoolConfig()
try:
cfg.sticky_session_ttl_seconds = 999 # type: ignore[misc]
assert False, "Should have raised FrozenInstanceError"
except AttributeError:
pass

View File

@@ -0,0 +1,54 @@
"""Tests for pool config backward compatibility.
Verifies that:
1. Only ``pool_advanced`` key activates pool mode
2. ``claude_code_advanced`` alone does NOT activate pool mode
3. Old import paths via shim modules still resolve correctly
"""
from __future__ import annotations
from src.services.provider.pool.config import PoolConfig, parse_pool_config
def test_pool_advanced_key_activates_pool() -> None:
"""pool_advanced key activates pool mode."""
cfg = parse_pool_config({"pool_advanced": {"sticky_session_ttl_seconds": 999}})
assert cfg is not None
assert cfg.sticky_session_ttl_seconds == 999
def test_claude_code_advanced_alone_does_not_activate_pool() -> None:
"""claude_code_advanced alone does NOT activate pool mode."""
cfg = parse_pool_config({"claude_code_advanced": {"lru_enabled": False}})
assert cfg is None
def test_no_config_returns_none() -> None:
cfg = parse_pool_config({"other": 1})
assert cfg is None
def test_empty_pool_advanced_returns_defaults() -> None:
cfg = parse_pool_config({"pool_advanced": {}})
assert cfg is not None
defaults = PoolConfig()
assert cfg.sticky_session_ttl_seconds == defaults.sticky_session_ttl_seconds
assert cfg.lru_enabled is True
def test_shim_imports_resolve() -> None:
"""Old import paths through shim modules still work."""
from src.services.provider.adapters.claude_code.pool_config import PoolConfig as ShimPoolConfig
from src.services.provider.adapters.claude_code.pool_config import (
parse_pool_config as shim_parse,
)
from src.services.provider.adapters.claude_code.pool_manager import (
ClaudeCodePoolManager,
)
from src.services.provider.pool.manager import PoolManager
# ShimPoolConfig should be the same class
assert ShimPoolConfig is PoolConfig
assert shim_parse is parse_pool_config
assert ClaudeCodePoolManager is PoolManager

View File

@@ -0,0 +1,141 @@
"""Tests for pool_cost_tracker.py — rolling window cost tracking."""
from __future__ import annotations
from unittest.mock import AsyncMock, patch
import pytest
from src.services.provider.pool.config import PoolConfig
from src.services.provider.pool.cost_tracker import (
get_window_usage,
is_approaching_limit,
is_at_limit,
record_usage,
)
PID = "provider-test"
KID = "key-test"
@pytest.fixture()
def config() -> PoolConfig:
return PoolConfig(
cost_limit_per_key_tokens=10000,
cost_soft_threshold_percent=80,
cost_window_seconds=18000,
)
@pytest.fixture()
def config_no_limit() -> PoolConfig:
return PoolConfig(cost_limit_per_key_tokens=None)
# ---------------------------------------------------------------------------
# record_usage
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_record_usage_calls_redis(config: PoolConfig) -> None:
with patch(
"src.services.provider.pool.redis_ops.add_cost_entry",
new_callable=AsyncMock,
) as mock_add:
await record_usage(PID, KID, 500, config)
mock_add.assert_called_once_with(PID, KID, 500, 18000)
@pytest.mark.asyncio
async def test_record_usage_skips_zero_tokens(config: PoolConfig) -> None:
with patch(
"src.services.provider.pool.redis_ops.add_cost_entry",
new_callable=AsyncMock,
) as mock_add:
await record_usage(PID, KID, 0, config)
mock_add.assert_not_called()
@pytest.mark.asyncio
async def test_record_usage_skips_when_no_limit(config_no_limit: PoolConfig) -> None:
with patch(
"src.services.provider.pool.redis_ops.add_cost_entry",
new_callable=AsyncMock,
) as mock_add:
await record_usage(PID, KID, 500, config_no_limit)
mock_add.assert_not_called()
# ---------------------------------------------------------------------------
# is_at_limit
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_is_at_limit_true(config: PoolConfig) -> None:
with patch(
"src.services.provider.pool.redis_ops.get_cost_window_total",
new_callable=AsyncMock,
return_value=10000,
):
assert await is_at_limit(PID, KID, config) is True
@pytest.mark.asyncio
async def test_is_at_limit_false(config: PoolConfig) -> None:
with patch(
"src.services.provider.pool.redis_ops.get_cost_window_total",
new_callable=AsyncMock,
return_value=5000,
):
assert await is_at_limit(PID, KID, config) is False
@pytest.mark.asyncio
async def test_is_at_limit_always_false_when_no_limit(config_no_limit: PoolConfig) -> None:
assert await is_at_limit(PID, KID, config_no_limit) is False
# ---------------------------------------------------------------------------
# is_approaching_limit
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_is_approaching_limit_true(config: PoolConfig) -> None:
# 80% of 10000 = 8000
with patch(
"src.services.provider.pool.redis_ops.get_cost_window_total",
new_callable=AsyncMock,
return_value=8500,
):
assert await is_approaching_limit(PID, KID, config) is True
@pytest.mark.asyncio
async def test_is_approaching_limit_false(config: PoolConfig) -> None:
with patch(
"src.services.provider.pool.redis_ops.get_cost_window_total",
new_callable=AsyncMock,
return_value=5000,
):
assert await is_approaching_limit(PID, KID, config) is False
# ---------------------------------------------------------------------------
# get_window_usage
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_get_window_usage(config: PoolConfig) -> None:
with patch(
"src.services.provider.pool.redis_ops.get_cost_window_total",
new_callable=AsyncMock,
return_value=4200,
):
assert await get_window_usage(PID, KID, config) == 4200

View File

@@ -0,0 +1,279 @@
"""Tests for pool_health_policy.py — error code classification."""
from __future__ import annotations
import json
from unittest.mock import AsyncMock, patch
import pytest
from src.services.provider.pool.config import (
PoolConfig,
UnschedulableRule,
)
from src.services.provider.pool.health_policy import (
_extract_error_message,
_parse_retry_after,
apply_health_policy,
)
PID = "provider-test"
KID = "key-test"
@pytest.fixture()
def config() -> PoolConfig:
return PoolConfig()
# ---------------------------------------------------------------------------
# Helper tests
# ---------------------------------------------------------------------------
def test_parse_retry_after_none() -> None:
assert _parse_retry_after(None) is None
assert _parse_retry_after({}) is None
def test_parse_retry_after_valid() -> None:
assert _parse_retry_after({"retry-after": "60"}) == 60
assert _parse_retry_after({"Retry-After": "120"}) == 120
def test_parse_retry_after_clamped() -> None:
assert _parse_retry_after({"retry-after": "0"}) == 1
assert _parse_retry_after({"retry-after": "9999"}) == 3600
def test_parse_retry_after_invalid() -> None:
assert _parse_retry_after({"retry-after": "not-a-number"}) is None
def test_extract_error_message_json_error_object() -> None:
body = json.dumps({"error": {"message": "bad request"}})
assert _extract_error_message(body) == "bad request"
def test_extract_error_message_json_error_string() -> None:
body = json.dumps({"error": "something went wrong"})
assert _extract_error_message(body) == "something went wrong"
def test_extract_error_message_plain_text() -> None:
assert _extract_error_message("plain text error") == "plain text error"
def test_extract_error_message_empty() -> None:
assert _extract_error_message(None) == ""
assert _extract_error_message("") == ""
# ---------------------------------------------------------------------------
# Status code handling
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_401_invalidates_oauth_cache_and_sets_cooldown(config: PoolConfig) -> None:
with (
patch(
"src.services.provider.pool.redis_ops.invalidate_oauth_token_cache",
new_callable=AsyncMock,
) as mock_inv,
patch(
"src.services.provider.pool.redis_ops.set_cooldown",
new_callable=AsyncMock,
) as mock_cd,
):
await apply_health_policy(
provider_id=PID,
key_id=KID,
status_code=401,
error_body=None,
response_headers=None,
config=config,
)
mock_inv.assert_called_once_with(KID)
mock_cd.assert_called_once_with(PID, KID, "auth_failed_401", ttl=60)
@pytest.mark.asyncio
async def test_402_sets_long_cooldown(config: PoolConfig) -> None:
with patch(
"src.services.provider.pool.redis_ops.set_cooldown",
new_callable=AsyncMock,
) as mock_cd:
await apply_health_policy(
provider_id=PID,
key_id=KID,
status_code=402,
error_body=None,
response_headers=None,
config=config,
)
mock_cd.assert_called_once_with(PID, KID, "payment_required_402", ttl=3600)
@pytest.mark.asyncio
async def test_403_sets_long_cooldown(config: PoolConfig) -> None:
with patch(
"src.services.provider.pool.redis_ops.set_cooldown",
new_callable=AsyncMock,
) as mock_cd:
await apply_health_policy(
provider_id=PID,
key_id=KID,
status_code=403,
error_body=None,
response_headers=None,
config=config,
)
mock_cd.assert_called_once_with(PID, KID, "forbidden_403", ttl=3600)
@pytest.mark.asyncio
async def test_400_with_org_disabled_pattern(config: PoolConfig) -> None:
body = json.dumps({"error": {"message": "Your organization has been disabled"}})
with patch(
"src.services.provider.pool.redis_ops.set_cooldown",
new_callable=AsyncMock,
) as mock_cd:
await apply_health_policy(
provider_id=PID,
key_id=KID,
status_code=400,
error_body=body,
response_headers=None,
config=config,
)
mock_cd.assert_called_once()
assert "account_disabled_400" in mock_cd.call_args.args[2]
@pytest.mark.asyncio
async def test_400_without_pattern_does_nothing(config: PoolConfig) -> None:
body = json.dumps({"error": {"message": "invalid json field"}})
with patch(
"src.services.provider.pool.redis_ops.set_cooldown",
new_callable=AsyncMock,
) as mock_cd:
await apply_health_policy(
provider_id=PID,
key_id=KID,
status_code=400,
error_body=body,
response_headers=None,
config=config,
)
mock_cd.assert_not_called()
@pytest.mark.asyncio
async def test_429_uses_retry_after_header(config: PoolConfig) -> None:
with patch(
"src.services.provider.pool.redis_ops.set_cooldown",
new_callable=AsyncMock,
) as mock_cd:
await apply_health_policy(
provider_id=PID,
key_id=KID,
status_code=429,
error_body=None,
response_headers={"retry-after": "120"},
config=config,
)
mock_cd.assert_called_once_with(PID, KID, "rate_limited_429", ttl=120)
@pytest.mark.asyncio
async def test_429_falls_back_to_config_default(config: PoolConfig) -> None:
with patch(
"src.services.provider.pool.redis_ops.set_cooldown",
new_callable=AsyncMock,
) as mock_cd:
await apply_health_policy(
provider_id=PID,
key_id=KID,
status_code=429,
error_body=None,
response_headers=None,
config=config,
)
mock_cd.assert_called_once_with(PID, KID, "rate_limited_429", ttl=300)
@pytest.mark.asyncio
async def test_529_uses_overload_cooldown(config: PoolConfig) -> None:
with patch(
"src.services.provider.pool.redis_ops.set_cooldown",
new_callable=AsyncMock,
) as mock_cd:
await apply_health_policy(
provider_id=PID,
key_id=KID,
status_code=529,
error_body=None,
response_headers=None,
config=config,
)
mock_cd.assert_called_once_with(PID, KID, "overloaded_529", ttl=30)
# ---------------------------------------------------------------------------
# Unschedulable keyword rules
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_keyword_rule_matches_and_sets_cooldown() -> None:
cfg = PoolConfig(
unschedulable_rules=[UnschedulableRule(keyword="capacity", duration_minutes=10)]
)
body = json.dumps({"error": {"message": "Server at capacity, try later"}})
with patch(
"src.services.provider.pool.redis_ops.set_cooldown",
new_callable=AsyncMock,
) as mock_cd:
await apply_health_policy(
provider_id=PID,
key_id=KID,
status_code=500,
error_body=body,
response_headers=None,
config=cfg,
)
mock_cd.assert_called_once_with(PID, KID, "rule:capacity", ttl=600)
# ---------------------------------------------------------------------------
# health_policy_enabled = False
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_disabled_health_policy_does_nothing() -> None:
cfg = PoolConfig(health_policy_enabled=False)
with patch(
"src.services.provider.pool.redis_ops.set_cooldown",
new_callable=AsyncMock,
) as mock_cd:
await apply_health_policy(
provider_id=PID,
key_id=KID,
status_code=429,
error_body=None,
response_headers=None,
config=cfg,
)
mock_cd.assert_not_called()

View File

@@ -0,0 +1,274 @@
"""Tests for pool_manager.py — candidate reordering, success/error hooks."""
from __future__ import annotations
from types import SimpleNamespace
from unittest.mock import AsyncMock, patch
import pytest
from src.services.provider.pool.config import PoolConfig
from src.services.provider.pool.manager import PoolManager
def _make_candidate(key_id: str, *, is_skipped: bool = False) -> SimpleNamespace:
return SimpleNamespace(
key=SimpleNamespace(id=key_id),
is_skipped=is_skipped,
skip_reason=None,
)
@pytest.fixture()
def pool() -> PoolManager:
return PoolManager("provider-1", PoolConfig())
# ---------------------------------------------------------------------------
# reorder_candidates
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_reorder_empty_candidates(pool: PoolManager) -> None:
result = await pool.reorder_candidates("sess-1", [])
assert result == []
@pytest.mark.asyncio
async def test_reorder_sticky_hit_moves_to_front(pool: PoolManager) -> None:
c1 = _make_candidate("key-1")
c2 = _make_candidate("key-2")
c3 = _make_candidate("key-3")
with (
patch(
"src.services.provider.pool.redis_ops.get_sticky_binding",
new_callable=AsyncMock,
return_value="key-2",
),
patch(
"src.services.provider.pool.redis_ops.batch_get_cooldowns",
new_callable=AsyncMock,
return_value={"key-1": None, "key-2": None, "key-3": None},
),
patch(
"src.services.provider.pool.redis_ops.get_lru_scores",
new_callable=AsyncMock,
return_value={"key-1": 100.0, "key-2": 200.0, "key-3": 50.0},
),
):
result = await pool.reorder_candidates("sess-1", [c1, c2, c3])
assert result[0].key.id == "key-2" # sticky hit first
@pytest.mark.asyncio
async def test_reorder_cooldown_keys_are_skipped(pool: PoolManager) -> None:
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": "rate_limited_429", "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])
# key-1 should be skipped
assert c1.is_skipped is True
assert "cooldown" in (c1.skip_reason or "")
# key-2 first in available
assert result[0].key.id == "key-2"
@pytest.mark.asyncio
async def test_reorder_cost_exhausted_keys_are_skipped() -> None:
pool = PoolManager(
"provider-1",
PoolConfig(cost_limit_per_key_tokens=1000, lru_enabled=False),
)
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.batch_get_cost_totals",
new_callable=AsyncMock,
return_value={"key-1": 1500, "key-2": 200},
),
):
result = await pool.reorder_candidates(None, [c1, c2])
assert c1.is_skipped is True
assert "cost" in (c1.skip_reason or "")
assert result[0].key.id == "key-2"
@pytest.mark.asyncio
async def test_reorder_lru_sorts_least_recently_used_first(
pool: PoolManager,
) -> None:
c1 = _make_candidate("key-1")
c2 = _make_candidate("key-2")
c3 = _make_candidate("key-3")
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, "key-3": None},
),
patch(
"src.services.provider.pool.redis_ops.get_lru_scores",
new_callable=AsyncMock,
return_value={"key-1": 300.0, "key-2": 100.0, "key-3": 200.0},
),
):
result = await pool.reorder_candidates(None, [c1, c2, c3])
# Least recently used (lowest score) first
available = [r for r in result if not r.is_skipped]
assert available[0].key.id == "key-2"
assert available[1].key.id == "key-3"
assert available[2].key.id == "key-1"
# ---------------------------------------------------------------------------
# on_request_success
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_on_request_success_binds_sticky_and_touches_lru(
pool: PoolManager,
) -> None:
with (
patch(
"src.services.provider.pool.redis_ops.set_sticky_binding",
new_callable=AsyncMock,
) as mock_sticky,
patch(
"src.services.provider.pool.redis_ops.touch_lru",
new_callable=AsyncMock,
) as mock_lru,
patch(
"src.services.provider.pool.redis_ops.add_cost_entry",
new_callable=AsyncMock,
) as mock_cost,
):
await pool.on_request_success(session_uuid="sess-1", key_id="key-1", tokens_used=0)
mock_sticky.assert_called_once_with("provider-1", "sess-1", "key-1", 3600)
mock_lru.assert_called_once_with("provider-1", "key-1")
mock_cost.assert_not_called() # tokens_used=0
@pytest.mark.asyncio
async def test_on_request_success_records_cost_when_configured() -> None:
pool = PoolManager(
"provider-1",
PoolConfig(cost_limit_per_key_tokens=50000),
)
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.add_cost_entry",
new_callable=AsyncMock,
) as mock_cost,
):
await pool.on_request_success(session_uuid="sess-1", key_id="key-1", tokens_used=500)
mock_cost.assert_called_once_with("provider-1", "key-1", 500, 18000)
# ---------------------------------------------------------------------------
# on_request_error
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_on_request_error_delegates_to_health_policy(
pool: PoolManager,
) -> None:
with patch(
"src.services.provider.pool.health_policy.apply_health_policy",
new_callable=AsyncMock,
) as mock_hp:
await pool.on_request_error(
key_id="key-1", status_code=429, error_body=None, response_headers=None
)
mock_hp.assert_called_once()
call_kwargs = mock_hp.call_args.kwargs
assert call_kwargs["key_id"] == "key-1"
assert call_kwargs["status_code"] == 429
# ---------------------------------------------------------------------------
# is_key_schedulable
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_is_key_schedulable_returns_false_when_in_cooldown(
pool: PoolManager,
) -> None:
with patch(
"src.services.provider.pool.redis_ops.get_cooldown",
new_callable=AsyncMock,
return_value="rate_limited_429",
):
ok, reason = await pool.is_key_schedulable("key-1")
assert ok is False
assert "cooldown" in (reason or "")
@pytest.mark.asyncio
async def test_is_key_schedulable_returns_true_when_healthy(
pool: PoolManager,
) -> None:
with patch(
"src.services.provider.pool.redis_ops.get_cooldown",
new_callable=AsyncMock,
return_value=None,
):
ok, reason = await pool.is_key_schedulable("key-1")
assert ok is True
assert reason is None

View File

@@ -0,0 +1,202 @@
"""Tests for pool manager trace data collection and attachment."""
from __future__ import annotations
from types import SimpleNamespace
from unittest.mock import AsyncMock, patch
import pytest
from src.services.provider.pool.config import PoolConfig
from src.services.provider.pool.manager import PoolManager
from src.services.provider.pool.trace import PoolSchedulingTrace
def _make_candidate(key_id: str, *, is_skipped: bool = False) -> SimpleNamespace:
return SimpleNamespace(
key=SimpleNamespace(id=key_id),
is_skipped=is_skipped,
skip_reason=None,
)
# ---------------------------------------------------------------------------
# Trace attachment to candidates
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_trace_attached_to_first_candidate() -> None:
pool = PoolManager("prov-1", PoolConfig())
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, None), "key-2": (None, None)},
),
patch(
"src.services.provider.pool.redis_ops.get_lru_scores",
new_callable=AsyncMock,
return_value={"key-1": 100.0, "key-2": 50.0},
),
):
result = await pool.reorder_candidates("sess-1", [c1, c2])
# First candidate should have _pool_scheduling_trace
trace = getattr(result[0], "_pool_scheduling_trace", None)
assert trace is not None
assert isinstance(trace, PoolSchedulingTrace)
assert trace.total_keys == 2
assert trace.provider_id == "prov-1"
@pytest.mark.asyncio
async def test_pool_extra_data_on_selected_candidate() -> None:
pool = PoolManager("prov-1", PoolConfig())
c1 = _make_candidate("key-1")
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, None)},
),
patch(
"src.services.provider.pool.redis_ops.get_lru_scores",
new_callable=AsyncMock,
return_value={},
),
):
result = await pool.reorder_candidates(None, [c1])
extra = getattr(result[0], "_pool_extra_data", None)
assert extra is not None
assert "pool_selection" in extra
@pytest.mark.asyncio
async def test_pool_extra_data_on_skipped_candidate() -> None:
pool = PoolManager("prov-1", PoolConfig())
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": ("rate_limited_429", 120),
"key-2": (None, None),
},
),
patch(
"src.services.provider.pool.redis_ops.get_lru_scores",
new_callable=AsyncMock,
return_value={},
),
):
result = await pool.reorder_candidates(None, [c1, c2])
# key-1 is skipped with pool_skip extra data
skipped = [r for r in result if r.is_skipped]
assert len(skipped) == 1
extra = getattr(skipped[0], "_pool_extra_data", None)
assert extra is not None
assert extra["pool_skip"]["type"] == "cooldown"
assert extra["pool_skip"]["cooldown_reason"] == "rate_limited_429"
assert extra["pool_skip"]["cooldown_ttl"] == 120
@pytest.mark.asyncio
async def test_trace_build_summary_matches() -> None:
pool = PoolManager("prov-1", PoolConfig(cost_limit_per_key_tokens=1000))
c1 = _make_candidate("key-1")
c2 = _make_candidate("key-2")
c3 = _make_candidate("key-3")
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": ("overloaded_529", 60),
"key-2": (None, None),
"key-3": (None, None),
},
),
patch(
"src.services.provider.pool.redis_ops.batch_get_cost_totals",
new_callable=AsyncMock,
return_value={"key-2": 1500, "key-3": 200},
),
):
result = await pool.reorder_candidates(None, [c1, c2, c3])
trace = getattr(result[0], "_pool_scheduling_trace", None)
assert trace is not None
summary = trace.build_summary(success_key_id="key-3")
assert summary["total_keys"] == 3
assert summary["skipped_cooldown"] == 1 # key-1
assert summary["skipped_cost"] == 1 # key-2
assert summary["attempted"] == 1 # key-3
assert summary["success_key_id"] == "key-3"[:8]
@pytest.mark.asyncio
async def test_sticky_trace_info() -> None:
pool = PoolManager("prov-1", PoolConfig(sticky_session_ttl_seconds=3600))
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="key-2",
),
patch(
"src.services.provider.pool.redis_ops.batch_get_cooldowns",
new_callable=AsyncMock,
return_value={"key-1": (None, None), "key-2": (None, None)},
),
patch(
"src.services.provider.pool.redis_ops.get_lru_scores",
new_callable=AsyncMock,
return_value={},
),
):
result = await pool.reorder_candidates("sess-1", [c1, c2])
trace = getattr(result[0], "_pool_scheduling_trace", None)
assert trace is not None
assert trace.sticky_session_used is True
# key-2 (sticky hit) should be first
assert result[0].key.id == "key-2"
extra = getattr(result[0], "_pool_extra_data", None)
assert extra is not None
assert extra["pool_selection"]["sticky_hit"] is True

View File

@@ -0,0 +1,65 @@
"""Tests for pool scheduling strategy registry."""
from __future__ import annotations
from src.services.provider.pool.strategy import (
_strategy_registry,
get_active_strategies,
get_pool_strategy,
register_pool_strategy,
)
class _DummyStrategy:
name = "dummy"
def compute_score(self, *, key_id, config, context):
return 42.0
class _AnotherStrategy:
name = "another"
def on_before_select(self, *, provider_id, key_ids, config, context):
return key_ids[:1]
def setup_function():
_strategy_registry.clear()
def teardown_function():
_strategy_registry.clear()
def test_register_and_get() -> None:
s = _DummyStrategy()
register_pool_strategy("dummy", s)
assert get_pool_strategy("dummy") is s
def test_get_nonexistent() -> None:
assert get_pool_strategy("nonexistent") is None
def test_get_active_strategies_filters_by_name() -> None:
s1 = _DummyStrategy()
s2 = _AnotherStrategy()
register_pool_strategy("dummy", s1)
register_pool_strategy("another", s2)
active = get_active_strategies(["dummy"])
assert len(active) == 1
assert active[0] is s1
active_both = get_active_strategies(["dummy", "another"])
assert len(active_both) == 2
active_none = get_active_strategies(["missing"])
assert len(active_none) == 0
def test_get_active_strategies_empty_names() -> None:
register_pool_strategy("dummy", _DummyStrategy())
assert get_active_strategies([]) == []
assert get_active_strategies(()) == []

View File

@@ -0,0 +1,145 @@
"""Tests for pool scheduling trace dataclasses."""
from __future__ import annotations
from src.services.provider.pool.trace import PoolCandidateTrace, PoolSchedulingTrace
# ---------------------------------------------------------------------------
# PoolCandidateTrace.to_extra_data
# ---------------------------------------------------------------------------
class TestPoolCandidateTraceExtraData:
def test_selected_sticky(self) -> None:
ct = PoolCandidateTrace(key_id="k1", reason="sticky", sticky_hit=True)
data = ct.to_extra_data()
assert "pool_selection" in data
assert data["pool_selection"]["reason"] == "sticky"
assert data["pool_selection"]["sticky_hit"] is True
def test_selected_lru_with_cost(self) -> None:
ct = PoolCandidateTrace(
key_id="k2",
reason="lru",
lru_score=1234.5,
cost_window_usage=500,
cost_limit=1000,
)
data = ct.to_extra_data()
sel = data["pool_selection"]
assert sel["reason"] == "lru"
assert sel["lru_score"] == 1234.5
assert sel["cost_window_usage"] == 500
assert sel["cost_limit"] == 1000
assert "sticky_hit" not in sel
def test_selected_with_soft_threshold(self) -> None:
ct = PoolCandidateTrace(
key_id="k3",
reason="lru",
cost_window_usage=850,
cost_limit=1000,
cost_soft_threshold=True,
)
data = ct.to_extra_data()
assert data["pool_selection"]["cost_soft_threshold"] is True
def test_selected_random_minimal(self) -> None:
ct = PoolCandidateTrace(key_id="k4", reason="random")
data = ct.to_extra_data()
sel = data["pool_selection"]
assert sel == {"reason": "random"}
def test_skipped_cooldown(self) -> None:
ct = PoolCandidateTrace(
key_id="k5",
skipped=True,
skip_type="cooldown",
cooldown_reason="rate_limited_429",
cooldown_ttl=120,
)
data = ct.to_extra_data()
assert "pool_skip" in data
skip = data["pool_skip"]
assert skip["type"] == "cooldown"
assert skip["cooldown_reason"] == "rate_limited_429"
assert skip["cooldown_ttl"] == 120
def test_skipped_cost_exhausted(self) -> None:
ct = PoolCandidateTrace(
key_id="k6",
skipped=True,
skip_type="cost_exhausted",
cost_window_usage=2000,
)
data = ct.to_extra_data()
skip = data["pool_skip"]
assert skip["type"] == "cost_exhausted"
assert skip["cost_window_usage"] == 2000
assert "cooldown_reason" not in skip
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"}
# ---------------------------------------------------------------------------
# PoolSchedulingTrace.build_summary
# ---------------------------------------------------------------------------
class TestPoolSchedulingTraceSummary:
def test_basic_summary(self) -> None:
trace = PoolSchedulingTrace(provider_id="prov-1", total_keys=5)
trace.candidate_traces = {
"k1": PoolCandidateTrace(key_id="k1", reason="sticky", sticky_hit=True),
"k2": PoolCandidateTrace(key_id="k2", reason="lru"),
"k3": PoolCandidateTrace(
key_id="k3", skipped=True, skip_type="cooldown", cooldown_reason="429"
),
"k4": PoolCandidateTrace(key_id="k4", skipped=True, skip_type="cost_exhausted"),
"k5": PoolCandidateTrace(
key_id="k5", skipped=True, skip_type="cooldown", cooldown_reason="500"
),
}
trace.sticky_session_used = True
summary = trace.build_summary(success_key_id="k1")
assert summary["enabled"] is True
assert summary["total_keys"] == 5
assert summary["attempted"] == 2
assert summary["skipped_cooldown"] == 2
assert summary["skipped_cost"] == 1
assert summary["sticky_session"] is True
assert summary["success_key_id"] == "k1"[:8]
assert summary["success_reason"] == "sticky"
def test_summary_no_success_key(self) -> None:
trace = PoolSchedulingTrace(provider_id="prov-2", total_keys=2)
trace.candidate_traces = {
"k1": PoolCandidateTrace(key_id="k1", reason="random"),
"k2": PoolCandidateTrace(key_id="k2", reason="lru"),
}
summary = trace.build_summary(success_key_id=None)
assert summary["attempted"] == 2
assert summary["skipped_cooldown"] == 0
assert summary["skipped_cost"] == 0
assert "success_key_id" not in summary
assert "success_reason" not in summary
def test_summary_all_skipped(self) -> None:
trace = PoolSchedulingTrace(provider_id="prov-3", total_keys=2)
trace.candidate_traces = {
"k1": PoolCandidateTrace(key_id="k1", skipped=True, skip_type="cooldown"),
"k2": PoolCandidateTrace(key_id="k2", skipped=True, skip_type="cost_exhausted"),
}
summary = trace.build_summary()
assert summary["attempted"] == 0
assert summary["skipped_cooldown"] == 1
assert summary["skipped_cost"] == 1