"""Tests for ``dispatcher._measured_cache_rates()`` — TTL, pricing floor, fail-open.""" from __future__ import annotations import sqlite3 import time from datetime import datetime, timedelta, timezone from pathlib import Path from types import SimpleNamespace import pytest import dispatcher import metrics ROOT = Path(__file__).resolve().parent.parent SCHEMA_SQL = (ROOT / "config" / "schema.sql").read_text() def _now_iso() -> str: return datetime.now(timezone.utc).isoformat() def _make_db_file(tmp_path: Path) -> Path: """Create a DB file seeded from schema.sql and return its path.""" db_file = tmp_path / "cache-rates.db" conn = sqlite3.connect(str(db_file)) conn.executescript(SCHEMA_SQL) conn.close() return db_file def _seed_on_file( db_file: Path, *, model_id: str = "gpt-4o", provider: str = "neuralwatt", prompt_tokens: int, cached_prompt_tokens: int | None, observations: int = 1, age_hours: float = 1.0, task_category: str = "coding_general", ) -> None: """Seed *observations* energy_observations rows into *db_file*.""" at = (datetime.now(timezone.utc) - timedelta(hours=age_hours)).isoformat() with sqlite3.connect(str(db_file)) as conn: for _ in range(max(1, observations)): conn.execute( """INSERT INTO energy_observations ( model_id, provider, task_category, prompt_tokens, cached_prompt_tokens, cached_tokens_source, observed_at ) VALUES (?, ?, ?, ?, ?, 'reported', ?)""", (model_id, provider, task_category, prompt_tokens, cached_prompt_tokens, at), ) def _fresh_conn_factory(db_file: Path): """Return a callable that opens a fresh Row-backed connection to *db_file*.""" def factory(): conn = sqlite3.connect(str(db_file)) conn.row_factory = sqlite3.Row return conn return factory def _reset_cache_globals() -> None: """Forget any cached measured rates so the next call hits the DB.""" dispatcher._measured_rate_cache = None dispatcher._measured_rate_cached_at = 0.0 # ── Fixtures ────────────────────────────────────────────────────────────── @pytest.fixture(autouse=True) def _purge_cache(): """Reset module-level cache state between every test.""" _reset_cache_globals() yield _reset_cache_globals() # ── (a) TTL ────────────────────────────────────────────────────────────── def test_ttl_caches_for_refresh_window(tmp_path, monkeypatch): """Two calls within the window issue exactly ONE cache_rate_series invocation.""" db_file = _make_db_file(tmp_path) _seed_on_file(db_file, model_id="gpt-4o", prompt_tokens=1000, cached_prompt_tokens=900, observations=26) call_count = 0 factory = _fresh_conn_factory(db_file) def counting_wrapper(): nonlocal call_count call_count += 1 return factory() # First call fires (cache empty), second call is served from TTL cache. monkeypatch.setattr(dispatcher, "_db", counting_wrapper) monkeypatch.setattr(dispatcher, "cache_rate_series", lambda conn, cfg: metrics.cache_rate_series(conn, cfg)) monkeypatch.setattr( dispatcher.cfg.objective, "incumbent_rate_refresh_seconds", 300 ) result1 = dispatcher._measured_cache_rates() result2 = dispatcher._measured_cache_rates() assert call_count == 1, ( f"cache_rate_series called {call_count} times within TTL window — " "expected exactly 1" ) assert len(result1) == 1 assert len(result2) == 1 def test_ttl_expiry_forces_refresh(tmp_path, monkeypatch): """A stale cache forces a fresh call to cache_rate_series.""" db_file = _make_db_file(tmp_path) _seed_on_file(db_file, model_id="gpt-4o", prompt_tokens=1000, cached_prompt_tokens=900, observations=26) call_count = 0 factory = _fresh_conn_factory(db_file) def counting_wrapper(): nonlocal call_count call_count += 1 return factory() monkeypatch.setattr(dispatcher, "_db", counting_wrapper) monkeypatch.setattr( dispatcher.cfg.objective, "incumbent_rate_refresh_seconds", 300 ) # First call: cache is stale → fires. dispatcher._measured_cache_rates() assert call_count == 1 # Second call: served from cache (cached_at was set by first call). dispatcher._measured_cache_rates() assert call_count == 1 # Simulate expiry. _reset_cache_globals() dispatcher._measured_cache_rates() assert call_count == 2, "cache expiry should trigger a second call" # ── (b) Pricing floor ──────────────────────────────────────────────────── def test_pricing_floor_excludes_low_observation_rows(tmp_path, monkeypatch): """observations=5 excluded at incumbent_rate_min_observations=25; 26 admitted.""" db_file = _make_db_file(tmp_path) _seed_on_file(db_file, model_id="low_obs", prompt_tokens=5000, cached_prompt_tokens=4500, observations=5) _seed_on_file(db_file, model_id="high_obs", prompt_tokens=26000, cached_prompt_tokens=23400, observations=26) monkeypatch.setattr(dispatcher, "_db", _fresh_conn_factory(db_file)) monkeypatch.setattr( dispatcher.cfg.objective, "incumbent_rate_min_observations", 25 ) monkeypatch.setattr( dispatcher.cfg.objective, "incumbent_rate_refresh_seconds", 300 ) rates = dispatcher._measured_cache_rates() assert ("neuralwatt", "low_obs") not in rates, ( "5-observation row must be excluded at floor=25" ) assert ("neuralwatt", "high_obs") in rates assert rates[("neuralwatt", "high_obs")] == pytest.approx(0.9) def test_pricing_floor_independent_of_warning_floor(tmp_path, monkeypatch): """Lowering the pricing floor admits a 5-obs row, proving independence.""" db_file = _make_db_file(tmp_path) _seed_on_file(db_file, model_id="low_obs", prompt_tokens=5000, cached_prompt_tokens=4500, observations=5) factory = _fresh_conn_factory(db_file) # With a high floor, the row is excluded. monkeypatch.setattr(dispatcher, "_db", factory) monkeypatch.setattr( dispatcher.cfg.objective, "incumbent_rate_min_observations", 25 ) monkeypatch.setattr( dispatcher.cfg.objective, "incumbent_rate_refresh_seconds", 300 ) rates_high = dispatcher._measured_cache_rates() assert ("neuralwatt", "low_obs") not in rates_high # Clear cache so the second call re-reads the DB. _reset_cache_globals() # With a lowered floor, the same row is included — the pricing floor # knob drives the decision, not some hardcoded value. monkeypatch.setattr( dispatcher.cfg.objective, "incumbent_rate_min_observations", 1 ) rates_low = dispatcher._measured_cache_rates() assert ("neuralwatt", "low_obs") in rates_low, ( "lowering incumbent_rate_min_observations must admit the row — " "proving the pricing floor is independent of the warning floor" ) assert rates_low[("neuralwatt", "low_obs")] == pytest.approx(0.9) # ── (c) Empty DB → {} ──────────────────────────────────────────────────── def test_empty_db_returns_empty_dict(tmp_path, monkeypatch): """An empty energy_observations table yields {}.""" db_file = _make_db_file(tmp_path) monkeypatch.setattr(dispatcher, "_db", _fresh_conn_factory(db_file)) monkeypatch.setattr( dispatcher.cfg.objective, "incumbent_rate_refresh_seconds", 300 ) rates = dispatcher._measured_cache_rates() assert rates == {} # ── (d) cache_rate_series raising → {} + no exception propagates ────────── def test_exception_returns_empty_dict(tmp_path, monkeypatch): """When cache_rate_series raises, the function returns {} with no propagation.""" db_file = _make_db_file(tmp_path) monkeypatch.setattr(dispatcher, "_db", _fresh_conn_factory(db_file)) raise_error = RuntimeError("boom") def raising_series(conn, cfg): raise raise_error monkeypatch.setattr(dispatcher, "cache_rate_series", raising_series) monkeypatch.setattr( dispatcher.cfg.objective, "incumbent_rate_refresh_seconds", 300 ) rates = dispatcher._measured_cache_rates() assert rates == {} # Cache stays None (exception path never assigns it). assert dispatcher._measured_rate_cache is None def test_db_error_returns_empty_dict(monkeypatch): """When _db() itself raises, the function returns {}.""" monkeypatch.setattr( dispatcher, "_db", lambda: (_ for _ in ()).throw(ConnectionError("no db")) ) monkeypatch.setattr( dispatcher.cfg.objective, "incumbent_rate_refresh_seconds", 300 ) rates = dispatcher._measured_cache_rates() assert rates == {} # ── (e) cache_rate_warnings divergence guard ───────────────────────────── def test_cache_rate_warnings_fires_on_divergence(tmp_path): """cache_rate_warnings fires when measured diverges from assumed_cache_rate.""" db_file = _make_db_file(tmp_path) # Seed 30 rows with 50% cache rate — far from the default 0.917. _seed_on_file(db_file, model_id="divergent", prompt_tokens=10000, cached_prompt_tokens=5000, observations=30, task_category="coding_general") div_cfg = SimpleNamespace( routing=SimpleNamespace( allowed_access_levels=["public"], default_flex_preference=SimpleNamespace(value="auto"), ), objective=SimpleNamespace( plan_kwh_per_period=6.25, billing_reset_day=None, selection_coverage_window_hours=168, assumed_cache_rate=0.917, cache_rate_window_hours=168, cache_rate_warn_margin=0.10, cache_rate_warn_min_observations=5, ), dispatch_providers={ "neuralwatt": SimpleNamespace(has_energy_telemetry=True, balance_url=None) }, escalation=SimpleNamespace(enabled=True), classifier=SimpleNamespace( degraded_warn_min=20, degraded_warn_threshold=0.5 ), ) conn = sqlite3.connect(str(db_file)) conn.row_factory = sqlite3.Row warnings = metrics.cache_rate_warnings(conn, div_cfg) conn.close() assert warnings, ( "cache_rate_warnings must fire when measured (0.500) diverges " "from assumed (0.917) beyond margin (0.10)" ) assert any("cache rate: measured" in w for w in warnings), ( f"expected a 'cache rate: measured' warning, got: {warnings}" )