Files
6krrt/tests/test_credit_attenuation_routing.py

422 lines
14 KiB
Python

"""Dispatcher-level credit attenuation wiring tests.
These tests exercise the cached resolver and its call sites on the real
``dispatcher.route_endpoint`` TestClient path. They are intentionally
end-to-end for the wiring: pure ``rank_candidates`` cases live in
``tests/test_routing.py``.
"""
from __future__ import annotations
import sqlite3
from datetime import datetime, timedelta, timezone
from pathlib import Path
import pytest
from starlette.testclient import TestClient
import dispatcher
ROOT = Path(__file__).resolve().parent.parent
SCHEMA_SQL = (ROOT / "config" / "schema.sql").read_text()
def _hours_ago(hours: float) -> str:
return (datetime.now(timezone.utc) - timedelta(hours=hours)).isoformat()
def _make_db(tmp_path: Path) -> sqlite3.Connection:
conn = sqlite3.connect(str(tmp_path / "test.db"))
conn.row_factory = sqlite3.Row
conn.executescript(SCHEMA_SQL)
return conn
def _seed_two_provider_models(conn: sqlite3.Connection) -> None:
"""Two providers, two models: near-identical profile so the tiebreak decides.
OpenRouter is priced slightly BELOW NeuralWatt (0.29 vs 0.30 per 1M), so
the raw-cost order picks or-cheap. The attenuation multipliers must
overcome that gap to flip the selection to neuralwatt — which is what
makes tests (b)/(c)/(c2) observable rather than tautological: with the
wiring broken (multipliers never reaching rank_candidates), openrouter
wins and tests (c)/(c2) FAIL.
"""
rows = [
# (model_id, provider, tier, context_window,
# cost_per_1m_prompt, cost_per_1m_completion)
(
"nw-cheap",
"neuralwatt",
2,
131072,
0.30,
0.30,
),
(
"or-cheap",
"openrouter",
2,
131072,
0.29,
0.29,
),
]
for model_id, provider, tier, ctx, prompt_cost, completion_cost in rows:
conn.execute(
"""
INSERT INTO models (
model_id, provider, base_model_id, tier, context_window,
effective_context_window, max_output_tokens,
cost_per_1m_prompt, cost_per_1m_completion,
supports_vision, supports_json_mode,
latency_class, reasoning_mode, context_variant,
access_level, availability, last_updated
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?,
1, 'standard', 'default', 'full', 'public', 'active',
'2026-08-22T00:00:00+00:00')
""",
(
model_id,
provider,
model_id,
tier,
ctx,
ctx,
16384,
prompt_cost,
completion_cost,
1,
),
)
conn.execute(
"""
INSERT INTO proficiency
(model_id, provider, category, blended_score, source, last_updated)
VALUES (?, ?, 'coding_general', 0.80, 'self_eval_thin',
'2026-01-01T00:00:00+00:00')
""",
(model_id, provider),
)
conn.commit()
def _seed_provider_balance(conn: sqlite3.Connection, provider: str, balance: float) -> None:
conn.execute(
"INSERT INTO provider_balance_observations (provider, balance_usd, observed_at) "
"VALUES (?, ?, ?)",
(provider, balance, _hours_ago(0.1)),
)
conn.commit()
def _seed_energy_allowance(
conn: sqlite3.Connection,
provider: str,
model_id: str,
allowance_remaining_usd: float,
) -> None:
conn.execute(
"INSERT INTO energy_observations "
"(model_id, provider, task_category, completion_tokens, energy_kwh, "
"cost_usd, carbon_g_co2eq, attribution_ratio, observed_at, "
"allowance_remaining_usd) "
"VALUES (?, ?, 'coding_general', 100, 5.0e-05, 0.001, "
"2.4e-03, 0.25, ?, ?)",
(model_id, provider, _hours_ago(0.1), allowance_remaining_usd),
)
conn.commit()
def _route_response(client: TestClient, task: str = "fix a bug") -> dict:
resp = client.post(
"/route",
json={
"task": task,
"task_category": "coding_general",
"task_tier": 2,
"required_context_tokens": 1000,
},
)
assert resp.status_code == 200, resp.text
return resp.json()
@pytest.fixture(autouse=True)
def _reset_provider_multiplier_cache():
"""Every dispatcher-attenuation test starts with a cold cache.
This also prevents the resolver cache from leaking between tests.
"""
setattr(dispatcher, "_provider_cost_multipliers_cache", None)
setattr(dispatcher, "_provider_cost_multipliers_resolved_at", 0.0)
yield
setattr(dispatcher, "_provider_cost_multipliers_cache", None)
setattr(dispatcher, "_provider_cost_multipliers_resolved_at", 0.0)
@pytest.fixture
def route_client(tmp_path, monkeypatch):
"""A TestClient wired to a temp DB seeded with two cross-provider models."""
conn = _make_db(tmp_path)
_seed_two_provider_models(conn)
conn.close()
monkeypatch.setattr(dispatcher.cfg.database, "path", str(tmp_path / "test.db"))
monkeypatch.setattr(dispatcher.cfg.verification, "local_llm_enabled", False)
monkeypatch.setattr(dispatcher.cfg.routing, "require_vision", False)
monkeypatch.setattr(dispatcher.cfg.exploration, "enabled", False)
monkeypatch.setenv("NEURALWATT_API_KEY", "test-key")
monkeypatch.setenv("OPENROUTER_API_KEY", "test-key")
with TestClient(dispatcher.app) as client:
yield client
# (a) OFF (default): the resolver never touches the DB, even if a call would raise.
def test_credit_attenuation_off_by_default_no_db_read(
route_client, monkeypatch
):
"""With attenuation disabled, quota_balance_and_burn is never called."""
calls = []
def explode(*args, **kwargs):
calls.append((args, kwargs))
raise RuntimeError("resolver should not have been called")
monkeypatch.setattr(dispatcher, "quota_balance_and_burn", explode)
decision = _route_response(route_client)
assert decision["selected"] is not None
assert calls == []
# (b) ON + all balances healthy → selection matches the no-attenuation baseline.
def test_credit_attenuation_on_healthy_balances_matches_baseline(
route_client, monkeypatch
):
"""Healthy OpenRouter balance keeps the raw-cost order (OpenRouter first).
OpenRouter is priced slightly cheaper, so the OFF baseline selects it.
Enabling attenuation with a healthy balance must NOT change that — the
multiplier is 1.0 and the winner stays or-cheap.
"""
db_path = dispatcher.cfg.database.path
conn = sqlite3.connect(db_path)
conn.row_factory = sqlite3.Row
_seed_provider_balance(conn, "openrouter", 50.0)
conn.close()
# OFF baseline: the cheaper model wins.
decision = _route_response(route_client)
assert decision["selected"]["provider"] == "openrouter"
assert decision["selected"]["model_id"] == "or-cheap"
# ON + healthy balance: multiplier 1.0, same winner.
monkeypatch.setattr(dispatcher.cfg.objective.credit_attenuation, "enabled", True)
monkeypatch.setattr(
dispatcher.cfg.objective.credit_attenuation, "soft_floor_usd", 5.0
)
decision = _route_response(route_client)
assert decision["selected"]["provider"] == "openrouter"
# (c) ON + synthetic low OpenRouter balance flips the tiebreak toward the
# healthier NeuralWatt provider.
def test_credit_attenuation_on_low_openrouter_balance_flips_to_healthy_provider(
route_client, monkeypatch
):
"""A near-zero OpenRouter balance must flip the pick to NeuralWatt.
OpenRouter is priced cheaper than NeuralWatt, so the only way neuralwatt
can win is a >1.0 multiplier applied to or-cheap's comparison cost inside
route(). If the dispatcher stopped passing provider_cost_multipliers into
rank_candidates, openrouter (raw-cheaper) would win and this test FAILS.
"""
monkeypatch.setattr(dispatcher.cfg.objective.credit_attenuation, "enabled", True)
monkeypatch.setattr(
dispatcher.cfg.objective.credit_attenuation, "soft_floor_usd", 5.0
)
monkeypatch.setattr(
dispatcher.cfg.objective.credit_attenuation, "zero_floor_usd", 0.0
)
monkeypatch.setattr(
dispatcher.cfg.objective.credit_attenuation, "max_multiplier", 5.0
)
db_path = dispatcher.cfg.database.path
conn = sqlite3.connect(db_path)
conn.row_factory = sqlite3.Row
# 0.01 is just above the zero floor: multiplier ≈ 4.99, well past the
# ~3% raw-price gap between or-cheap and nw-cheap.
_seed_provider_balance(conn, "openrouter", 0.01)
conn.close()
decision = _route_response(route_client)
assert decision["selected"]["provider"] == "neuralwatt"
assert decision["selected"]["model_id"] == "nw-cheap"
# (c2) Billing-shape guard: low NeuralWatt allowance (telemetry provider) is
# ignored; OpenRouter stays attenuated.
def test_credit_attenuation_ignores_telemetry_provider_allowance(
route_client, monkeypatch
):
"""NeuralWatt's per-completion allowance is never attenuated.
OpenRouter is raw-cheaper, so neuralwatt winning here requires BOTH:
OpenRouter attenuated (low polled balance → ~5x comparison cost) AND
NeuralWatt left at multiplier 1.0 despite its negative telemetry
allowance. If the balance_url billing-shape guard broke and telemetry
providers were attenuated too, nw-cheap (~1.5 effective) would lose to
or-cheap (~1.45 effective) and this test FAILS. If the multiplier
passthrough broke, or-cheap wins raw and this test FAILS.
"""
monkeypatch.setattr(dispatcher.cfg.objective.credit_attenuation, "enabled", True)
monkeypatch.setattr(
dispatcher.cfg.objective.credit_attenuation, "soft_floor_usd", 5.0
)
monkeypatch.setattr(
dispatcher.cfg.objective.credit_attenuation, "zero_floor_usd", 0.0
)
monkeypatch.setattr(
dispatcher.cfg.objective.credit_attenuation, "max_multiplier", 5.0
)
db_path = dispatcher.cfg.database.path
conn = sqlite3.connect(db_path)
conn.row_factory = sqlite3.Row
# NeuralWatt reading near zero (mirrors live overage-invoice noise).
_seed_energy_allowance(conn, "neuralwatt", "nw-cheap", -0.004)
# OpenRouter prepaid pool nearly empty.
_seed_provider_balance(conn, "openrouter", 0.01)
conn.close()
decision = _route_response(route_client)
assert decision["selected"]["provider"] == "neuralwatt"
assert decision["selected"]["model_id"] == "nw-cheap"
# (d) Resolver caches: monkeypatch the binding; only the first call resolves,
# the immediate second call returns cached value.
def test_credit_attenuation_resolver_caches_within_refresh_window(
route_client, monkeypatch
):
"""Two route calls within refresh_seconds resolve balances once."""
monkeypatch.setattr(dispatcher.cfg.objective.credit_attenuation, "enabled", True)
calls = []
def counting_quota(*args, **kwargs):
calls.append((args, kwargs))
return {
"by_provider": {
"openrouter": {
"balance_usd": 10.0,
"balance_source": "polled",
},
"neuralwatt": {
"balance_usd": 10.0,
"balance_source": "telemetry",
},
},
"total_balance_usd": 20.0,
}
monkeypatch.setattr(dispatcher, "quota_balance_and_burn", counting_quota)
_route_response(route_client)
_route_response(route_client)
assert len(calls) == 1
# (e) Only balance_url providers appear in the multiplier dict.
def test_provider_cost_multipliers_only_for_balance_url_providers(
route_client, monkeypatch
):
"""The resolver returns a dict containing only providers with balance_url."""
monkeypatch.setattr(dispatcher.cfg.objective.credit_attenuation, "enabled", True)
def fake_quota(*args, **kwargs):
return {
"by_provider": {
"openrouter": {
"balance_usd": 1.0,
"balance_source": "polled",
},
"neuralwatt": {
"balance_usd": -0.004,
"balance_source": "telemetry",
},
},
"total_balance_usd": 0.996,
}
monkeypatch.setattr(dispatcher, "quota_balance_and_burn", fake_quota)
multipliers = dispatcher._provider_cost_multipliers()
assert "openrouter" in multipliers
assert "neuralwatt" not in multipliers
# (f) Disabled feature returns None.
def test_provider_cost_multipliers_none_when_disabled(
route_client, monkeypatch
):
"""When credit_attenuation.enabled is False the resolver returns None."""
monkeypatch.setattr(dispatcher.cfg.objective.credit_attenuation, "enabled", False)
def explode(*args, **kwargs):
raise RuntimeError("should not resolve when disabled")
monkeypatch.setattr(dispatcher, "quota_balance_and_burn", explode)
assert dispatcher._provider_cost_multipliers() is None
# (g) Stale cache refreshes outside refresh_seconds.
def test_credit_attenuation_resolver_refreshes_after_refresh_seconds(
route_client, monkeypatch
):
"""A long-past resolution timestamp triggers a fresh DB read."""
monkeypatch.setattr(dispatcher.cfg.objective.credit_attenuation, "enabled", True)
monkeypatch.setattr(
dispatcher.cfg.objective.credit_attenuation, "refresh_seconds", 60
)
calls = []
def counting_quota(*args, **kwargs):
calls.append((args, kwargs))
return {
"by_provider": {
"openrouter": {
"balance_usd": 10.0,
"balance_source": "polled",
},
"neuralwatt": {
"balance_usd": 10.0,
"balance_source": "telemetry",
},
},
"total_balance_usd": 20.0,
}
monkeypatch.setattr(dispatcher, "quota_balance_and_burn", counting_quota)
# Warm the cache, then age it past the refresh window.
dispatcher._provider_cost_multipliers()
assert len(calls) == 1
dispatcher._provider_cost_multipliers_resolved_at = (
dispatcher._provider_cost_multipliers_resolved_at - 120
)
dispatcher._provider_cost_multipliers()
assert len(calls) == 2