mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
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:
106
tests/services/test_pool_config.py
Normal file
106
tests/services/test_pool_config.py
Normal 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
|
||||
54
tests/services/test_pool_config_backward_compat.py
Normal file
54
tests/services/test_pool_config_backward_compat.py
Normal 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
|
||||
141
tests/services/test_pool_cost_tracker.py
Normal file
141
tests/services/test_pool_cost_tracker.py
Normal 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
|
||||
279
tests/services/test_pool_health_policy.py
Normal file
279
tests/services/test_pool_health_policy.py
Normal 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()
|
||||
274
tests/services/test_pool_manager.py
Normal file
274
tests/services/test_pool_manager.py
Normal 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
|
||||
202
tests/services/test_pool_manager_trace.py
Normal file
202
tests/services/test_pool_manager_trace.py
Normal 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
|
||||
65
tests/services/test_pool_strategy.py
Normal file
65
tests/services/test_pool_strategy.py
Normal 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(()) == []
|
||||
145
tests/services/test_pool_trace.py
Normal file
145
tests/services/test_pool_trace.py
Normal 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
|
||||
Reference in New Issue
Block a user