Files
6krrt/tests/test_measured_cache_rates.py

319 lines
11 KiB
Python

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