319 lines
11 KiB
Python
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}"
|
|
)
|