"""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