Files
6krrt/tests/test_route_decisions.py

1463 lines
53 KiB
Python

"""Tests for the route_decisions table, its inline-create helper, and the gate.
The monitoring TUI (see .omo/plans/router-monitoring-tui.md) needs a record of
every routing decision — which model was picked and why — that survives in the
datbase rather than only in the journal. This file pins the three pieces todo #1
adds:
- the `route_decisions` table in schema.sql (columns and the guarded index),
- `dispatcher.ensure_route_decisions(conn)` — the idempotent inline-create
helper that is the *only* way the table appears on a live router.db (the
live DB is never recreated; schema.sql alone is CREATE TABLE IF NOT EXISTS
and silently does nothing to an existing DB),
- the `logging.log_route_decisions` config gate.
The database tests follow the same offline temp-DB pattern as
tests/test_chat_completions.py: a throwaway SQLite file seeded from schema.sql,
never the live router.db.
"""
import json
import random
import sqlite3
from pathlib import Path
import pytest
from starlette.testclient import TestClient
import config
import dispatcher
import metrics
import prefix_probe
from dispatcher import Classification, app
ROOT = Path(__file__).resolve().parent.parent
SCHEMA_SQL = (ROOT / "config" / "schema.sql").read_text()
ROUTE_DECISIONS_COLUMNS = [
"id",
"observed_at",
"kind",
"task_category",
"task_tier",
"required_context_tokens",
"confidence",
"classifier_ms",
"classification_source",
"latency_tolerance",
"candidates_considered",
"selected_model",
"selected_provider",
"runner_up_models",
"est_cost_usd",
"est_proficiency",
"rejected_reason",
"session_key",
"tools",
"images",
"json_mode",
"streamed",
"flex_preference",
"flex_swapped",
"flex_forced",
# Enforcement for this list lives in tests/test_tui_schema_drift.py —
# keep new route_decisions columns registered there, not just here.
"request_id",
"exploration",
"pinch_original_tokens",
"pinch_final_tokens",
"profile",
"prefix_divergence_index",
"prefix_tokens_after_divergence",
"prefix_prev_message_count",
"agent",
"parent_key",
]
def _table_exists(conn: sqlite3.Connection, table: str) -> bool:
row = conn.execute(
"SELECT name FROM sqlite_master WHERE type='table' AND name=?",
(table,),
).fetchone()
return row is not None
def _index_exists(conn: sqlite3.Connection, index: str) -> bool:
row = conn.execute(
"SELECT name FROM sqlite_master WHERE type='index' AND name=?",
(index,),
).fetchone()
return row is not None
def _schema_minus_route_decisions() -> str:
"""schema.sql with the route_decisions block removed, for the failure case."""
lines = []
skipping = False
for line in SCHEMA_SQL.splitlines():
stripped = line.strip()
if stripped.startswith("CREATE TABLE IF NOT EXISTS route_decisions"):
skipping = True
continue
if skipping:
# The create ends at the closing paren + semicolon of the table
# statement. Anything still in the table body is skipped.
if stripped == ");":
skipping = False
continue
# Drop EVERY index on the table, not one named index: a hand-kept
# list of names drifts the moment another index is added, and the
# failure is an opaque "no such table" from inside executescript.
if "ON route_decisions" in line or "idx_route_decisions" in line:
continue
lines.append(line)
return "\n".join(lines)
# --- schema round-trips -----------------------------------------------------
def test_schema_defines_route_decisions_table():
"""schema.sql declares the table, so a fresh DB from it has it already."""
assert "CREATE TABLE IF NOT EXISTS route_decisions" in SCHEMA_SQL
def test_schema_index_is_guarded():
"""The observed_at index must be IF NOT EXISTS so re-applying is a no-op."""
assert "idx_route_decisions_observed" in SCHEMA_SQL
assert (
"CREATE INDEX IF NOT EXISTS idx_route_decisions_observed "
"ON route_decisions (observed_at)" in SCHEMA_SQL
)
# --- happy path: fresh DB from full schema ----------------------------------
def test_fresh_schema_already_has_table(tmp_path):
conn = sqlite3.connect(tmp_path / "fresh.db")
conn.executescript(SCHEMA_SQL)
assert _table_exists(conn, "route_decisions")
for col in ROUTE_DECISIONS_COLUMNS:
assert col in {r[1] for r in conn.execute("PRAGMA table_info(route_decisions)")}
assert _index_exists(conn, "idx_route_decisions_observed")
conn.close()
def test_ensure_route_decisions_is_idempotent(tmp_path):
"""Fresh DB already has the table; calling the helper twice no-ops."""
conn = sqlite3.connect(tmp_path / "idem.db")
conn.executescript(SCHEMA_SQL)
# Seed some rows in another table so we can prove nothing is dropped.
conn.execute(
"INSERT INTO models (model_id, provider, last_updated) "
"VALUES ('m1', 'neuralwatt', '2026-01-01T00:00:00+00:00')"
)
conn.commit()
dispatcher.ensure_route_decisions(conn) # first call
dispatcher.ensure_route_decisions(conn) # second call: must no-op cleanly
assert _table_exists(conn, "route_decisions")
count = conn.execute("SELECT COUNT(*) FROM models").fetchone()[0]
assert count == 1 # pre-existing rows survived
conn.close()
# --- failure path: pre-existing DB WITHOUT the table ------------------------
def test_ensure_route_decisions_adds_table_without_dropping_rows(tmp_path):
"""A DB that predates the table gets it added; existing rows survive."""
conn = sqlite3.connect(tmp_path / "old.db")
conn.executescript(_schema_minus_route_decisions())
assert not _table_exists(conn, "route_decisions")
# A row in a genuinely existing table, to prove it survives the upgrade.
conn.execute(
"INSERT INTO models (model_id, provider, last_updated) "
"VALUES ('legacy', 'neuralwatt', '2026-01-01T00:00:00+00:00')"
)
conn.commit()
dispatcher.ensure_route_decisions(conn)
assert _table_exists(conn, "route_decisions")
assert _index_exists(conn, "idx_route_decisions_observed")
legacy = conn.execute(
"SELECT model_id FROM models WHERE model_id='legacy'"
).fetchone()
assert legacy is not None # the existing row was not dropped
conn.close()
def test_ensure_route_decisions_allows_insert(tmp_path):
"""After the helper runs, the table actually accepts the documented shape."""
conn = sqlite3.connect(tmp_path / "insert.db")
conn.executescript(_schema_minus_route_decisions())
dispatcher.ensure_route_decisions(conn)
conn.execute(
"""
INSERT INTO route_decisions (
observed_at, kind, task_category, task_tier,
required_context_tokens, confidence, classifier_ms,
classification_source, latency_tolerance, candidates_considered,
selected_model, selected_provider, runner_up_models,
est_cost_usd, est_proficiency, rejected_reason, session_key,
tools, images, json_mode, streamed
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
""",
(
"2026-01-01T00:00:00+00:00", "route", "coding_general", 2, 500,
0.95, 1868, "classifier", "interactive", 8, "deepseek-v4-flash",
"neuralwatt", '[{"model_id": "gemma-4-31b", "provider": "neuralwatt"}]',
0.00016296, 1.0, None, "sess-hash", 0, 0, 0, 1,
),
)
conn.commit()
kind = conn.execute(
"SELECT kind FROM route_decisions WHERE selected_model='deepseek-v4-flash'"
).fetchone()
assert kind is not None and kind[0] == "route"
conn.close()
def test_persist_writes_pinch_columns(tmp_path, monkeypatch):
"""persist_route_decision writes pinch_original_tokens and pinch_final_tokens
when provided, and defaults them to NULL when omitted."""
db_path = tmp_path / "pinch.db"
conn = sqlite3.connect(db_path)
conn.executescript(_schema_minus_route_decisions())
conn.close()
monkeypatch.setattr(dispatcher.cfg.database, "path", str(db_path))
monkeypatch.setattr(dispatcher.cfg.logging, "log_route_decisions", True)
dispatcher.persist_route_decision(
"route",
classification=Classification(
task_category="coding_general", task_tier=2,
required_context_tokens=100, confidence=0.9,
),
latency_tolerance="interactive",
selected_model="deepseek-v4-flash",
selected_provider="neuralwatt",
pinch_original_tokens=100,
pinch_final_tokens=50,
)
dispatcher.persist_route_decision(
"route",
classification=Classification(
task_category="coding_general", task_tier=2,
required_context_tokens=100, confidence=0.9,
),
latency_tolerance="interactive",
selected_model="deepseek-v4-flash",
selected_provider="neuralwatt",
)
conn = sqlite3.connect(db_path)
conn.row_factory = sqlite3.Row
rows = conn.execute(
"SELECT pinch_original_tokens, pinch_final_tokens FROM route_decisions ORDER BY id"
).fetchall()
assert len(rows) == 2
assert rows[0]["pinch_original_tokens"] == 100
assert rows[0]["pinch_final_tokens"] == 50
assert rows[1]["pinch_original_tokens"] is None
assert rows[1]["pinch_final_tokens"] is None
conn.close()
def test_persist_route_decision_writes_profile_from_response(tmp_path, monkeypatch):
"""A RouteResponse with profile='locality' writes that profile to the row."""
db_path = tmp_path / "profile.db"
conn = sqlite3.connect(db_path)
conn.executescript(SCHEMA_SQL)
conn.close()
monkeypatch.setattr(dispatcher.cfg.database, "path", str(db_path))
monkeypatch.setattr(dispatcher.cfg.logging, "log_route_decisions", True)
from dispatcher import Candidate
dispatcher.persist_route_decision(
"route",
classification=dispatcher.RouteResponse(
classification=Classification(
task_category="coding_general", task_tier=2,
required_context_tokens=100, confidence=0.9,
),
latency_tolerance="interactive",
profile="locality",
selected=Candidate(
model_id="deepseek-v4-flash", provider="neuralwatt",
tier=2, latency_class="standard", reasoning_mode="default",
context_variant="full", effective_context_window=128000,
composite=0.5, cost_score=0.5, proficiency_score=0.9,
),
candidates_considered=1,
),
latency_tolerance="interactive",
selected_model="deepseek-v4-flash",
selected_provider="neuralwatt",
)
conn = sqlite3.connect(db_path)
conn.row_factory = sqlite3.Row
row = conn.execute("SELECT profile FROM route_decisions").fetchone()
assert row is not None
assert row["profile"] == "locality"
# A bare Classification (no RouteResponse) should leave profile at None.
dispatcher.persist_route_decision(
"route",
classification=Classification(
task_category="coding_general", task_tier=2,
required_context_tokens=100, confidence=0.9,
),
latency_tolerance="interactive",
selected_model="deepseek-v4-flash",
selected_provider="neuralwatt",
)
rows = conn.execute("SELECT profile FROM route_decisions ORDER BY id").fetchall()
assert len(rows) == 2
assert rows[1]["profile"] is None
conn.close()
def _schema_minus_profile_column() -> str:
"""schema.sql with the `profile` column and its comments stripped from
route_decisions, yielding the previous-todo schema for migration tests."""
lines = []
inside_route_decisions = False
skip_profile_block = False
for line in SCHEMA_SQL.splitlines():
stripped = line.strip()
if stripped.startswith("CREATE TABLE IF NOT EXISTS route_decisions"):
inside_route_decisions = True
if inside_route_decisions:
if stripped.startswith("pinch_final_tokens"):
lines.append(" pinch_final_tokens INTEGER")
skip_profile_block = True
continue
if skip_profile_block:
if stripped == ");":
skip_profile_block = False
inside_route_decisions = False
lines.append(line)
continue
continue
lines.append(line)
return "\n".join(lines)
def test_ensure_route_decisions_migrates_profile_column(tmp_path):
"""A live table without `profile` gets the column once, idempotently."""
conn = sqlite3.connect(tmp_path / "migrate.db")
conn.executescript(_schema_minus_profile_column())
assert "profile" not in {
r[1] for r in conn.execute("PRAGMA table_info(route_decisions)")
}
dispatcher.ensure_route_decisions(conn)
assert "profile" in {
r[1] for r in conn.execute("PRAGMA table_info(route_decisions)")
}
# Second call must be a no-op and not duplicate the column.
dispatcher.ensure_route_decisions(conn)
cols = [r[1] for r in conn.execute("PRAGMA table_info(route_decisions)")]
assert cols.count("profile") == 1
# The new column accepts NULL and string values.
conn.execute(
"""
INSERT INTO route_decisions (
observed_at, kind, profile
) VALUES (?, ?, ?)
""",
("2026-01-01T00:00:00+00:00", "route", "locality"),
)
conn.execute(
"""
INSERT INTO route_decisions (
observed_at, kind
) VALUES (?, ?)
""",
("2026-01-01T00:00:00+00:00", "route"),
)
conn.commit()
profiles = [r[0] for r in conn.execute("SELECT profile FROM route_decisions ORDER BY id")]
assert profiles == ["locality", None]
conn.close()
def test_persist_ensure_on_write_fixes_live_db_missing_table(tmp_path, monkeypatch):
"""A live router.db without route_decisions gets it on the WRITE path.
F4 scope-fidelity regression: ``persist_route_decision`` INSERTs without
ever calling ``ensure_route_decisions``, so if a live DB lacks the table
and the module-load startup hook did not run (a test harness, a process
that calls persist first, a future lazy-import refactor), every decision
row is silently swallowed (the INSERT raises no-such-table) and /metrics'
recent_decisions 500s. Mirroring proficiency_store._write -> ensure_columns,
the migration must be guaranteed on the write path too, not only at module
load. The temp DB is deliberately the schema-minus-route_decisions shape —
a DB that predates the feature.
"""
db_path = tmp_path / "live-no-table.db"
conn = sqlite3.connect(db_path)
conn.executescript(_schema_minus_route_decisions())
assert not _table_exists(conn, "route_decisions")
conn.close()
monkeypatch.setattr(dispatcher.cfg.database, "path", str(db_path))
monkeypatch.setattr(dispatcher.cfg.logging, "log_route_decisions", True)
# The write itself. Best-effort means it must not raise even on a missing
# table; then the row must actually land and the table must exist.
dispatcher.persist_route_decision(
"route",
classification=Classification(
task_category="coding_general", task_tier=2,
required_context_tokens=100, confidence=0.9,
),
latency_tolerance="interactive",
)
conn = sqlite3.connect(db_path)
conn.row_factory = sqlite3.Row
assert _table_exists(conn, "route_decisions"), \
"the migration must run on the write path so the table exists"
rows = conn.execute(
"SELECT kind FROM route_decisions WHERE kind='route'"
).fetchall()
assert len(rows) == 1, "the decision row must persist once the table exists"
# /metrics / recent_decisions against the same live-DB shape must not 500:
# reading must succeed now that the table exists.
rec = metrics.recent_decisions(conn)
assert len(rec) == 1
assert rec[0]["kind"] == "route"
conn.close()
# --- config gate ------------------------------------------------------------
def test_log_route_decisions_gate_defaults_on():
"""The key is declared and defaults to on, matching the config.yaml value."""
cfg = config.load_config(str(ROOT / "config" / "config.yaml"))
assert cfg.logging.log_route_decisions is True
def test_config_has_log_route_decisions_key():
"""The strict config accepts the key — it must be declared or load fails."""
assert "log_route_decisions" in (ROOT / "config" / "config.yaml").read_text()
# =============================================================================
# Todo #2: persist_route_decision wired into every decision path.
# =============================================================================
CHEAP = "cheap-model"
DEAR = "dear-model"
CHEAP_FLEX = "cheap-model-flex"
CHEAP_STALE_FLEX = "cheap-model-flex-stale"
EXPLORABLE_DEAR = "dear-but-explorable"
class FakeResponse:
"""Just enough of requests.Response for the dispatcher's provider calls."""
def __init__(self, payload=None, *, status_code=200, lines=None):
self.status_code = status_code
self._payload = payload or {}
self._lines = lines or []
self.text = json.dumps(self._payload)
self.closed = False
def json(self):
return self._payload
@property
def encoding(self):
"""The charset requests would pick, by requests' own rule.
Derived rather than asserted, so this fake tracks requests instead of
restating a belief about it. The default is a charset-less
``text/event-stream`` on purpose: that is what OpenRouter actually
sends, and `get_encoding_from_headers` answers ISO-8859-1 for any
``text/*`` without a charset. Decoding a UTF-8 stream with that and
re-encoding it produced mojibake on every streamed non-ASCII
character -- so the hostile case is the DEFAULT here, and a fake can
no longer make a broken passthrough look correct.
"""
from requests.structures import CaseInsensitiveDict
from requests.utils import get_encoding_from_headers
headers = CaseInsensitiveDict(
getattr(self, "headers", None) or {"Content-Type": "text/event-stream"}
)
return get_encoding_from_headers(headers) or "utf-8"
def iter_lines(self, decode_unicode=False):
# Faithful to requests in both modes: BYTES unless decode_unicode is
# set, and when it is set the charset comes from Content-Type.
#
# The `wire` line is the load-bearing one. Whatever a test wrote into
# `lines`, what a provider actually puts on the wire is UTF-8 bytes,
# so that is what gets decoded. An earlier version of this fake passed
# str lines straight through under decode_unicode, which modelled a
# stream that had ALREADY been decoded correctly -- and a fake that
# hands back the right answer cannot reproduce a charset bug. It let
# the whole streaming suite pass against a proxy that was mangling
# every non-ASCII character.
for line in self._lines:
wire = line.encode("utf-8") if isinstance(line, str) else line
yield wire.decode(self.encoding) if decode_unicode else wire
def close(self):
self.closed = True
def _completion(model, content="hello there"):
return {
"id": "chatcmpl-dec-1",
"model": model,
"choices": [
{"message": {"role": "assistant", "content": content},
"finish_reason": "stop"}
],
"usage": {"prompt_tokens": 31, "completion_tokens": 12},
"energy": {"energy_kwh": 5.0e-05, "carbon_g_co2eq": 2.4e-03},
"cost": {"request_cost_usd": 4.0e-04},
}
_STREAM_LINES = [
'data: {"id":"chatcmpl-stream-dec","choices":[{"delta":{"content":"hel"}}]}',
"",
'data: {"id":"chatcmpl-stream-dec","choices":[{"delta":{"content":"lo"},'
'"finish_reason":"stop"}],"usage":{"prompt_tokens":31,'
'"completion_tokens":9}}',
"",
"data: [DONE]",
"",
]
def _messages(text="write me a function"):
return [{"role": "user", "content": text}]
def _image_messages():
return [
{
"role": "user",
"content": [
{"type": "text", "text": "what is in this image?"},
{"type": "image_url", "image_url": {"url": "data:image/png;base64,AAAA"}},
],
}
]
@pytest.fixture
def decision_router(tmp_path, monkeypatch):
"""A routable dispatcher over a throwaway DB; nothing dials out."""
db_path = tmp_path / "decisions.db"
conn = sqlite3.connect(db_path)
conn.executescript(SCHEMA_SQL)
for model_id, completion_price, vision in (
(CHEAP, 0.30, 1),
(DEAR, 9.00, 0),
(EXPLORABLE_DEAR, 0.30, 0),
):
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 (?, 'neuralwatt', ?, 2, 262128, 192500, 16384, ?, ?,
?, 1, 'standard', 'default', 'full', 'public', 'active',
'2026-08-22T00:00:00+00:00')
""",
(model_id, model_id, completion_price / 3, completion_price, vision),
)
conn.execute("INSERT INTO proficiency (model_id, provider, category, blended_score, source, last_updated, outcome_samples) VALUES (?, 'neuralwatt', 'coding_general', 0.9, 'leaderboard', '2026-08-22T00:00:00+00:00', ?)", (CHEAP, 50))
conn.execute("INSERT INTO proficiency (model_id, provider, category, blended_score, source, last_updated, outcome_samples) VALUES (?, 'neuralwatt', 'coding_general', 0.9, 'leaderboard', '2026-08-22T00:00:00+00:00', ?)", (DEAR, 0))
conn.execute("INSERT INTO proficiency (model_id, provider, category, blended_score, source, last_updated, outcome_samples) VALUES (?, 'neuralwatt', 'coding_general', 0.9, 'leaderboard', '2026-08-22T00:00:00+00:00', ?)", (EXPLORABLE_DEAR, 0))
# A flex twin of CHEAP (same base_model_id/reasoning_mode/context_variant,
# latency_class='flex') with its own pricing so a swap's re-estimated cost
# is distinguishable from the standard row's.
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 (?, 'neuralwatt', ?, 2, 262128, 192500, 16384, ?, ?,
1, 1, 'flex', 'default', 'full', 'public', 'active',
'2026-08-22T00:00:00+00:00')
""",
(CHEAP_FLEX, CHEAP, 0.20, 0.60),
)
# A stale flex twin of CHEAP: same identity dims, latency_class='flex' but
# availability='stale'. It must never be swapped to even under force-flex.
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 (?, 'neuralwatt', ?, 2, 262128, 192500, 16384, ?, ?,
1, 1, 'flex', 'default', 'full', 'public', 'stale',
'2026-08-22T00:00:00+00:00')
""",
(CHEAP_STALE_FLEX, CHEAP, 0.20, 0.60),
)
conn.commit()
conn.close()
monkeypatch.setattr(dispatcher.cfg.database, "path", str(db_path))
monkeypatch.setattr(dispatcher.cfg.verification, "local_llm_enabled", False)
monkeypatch.setattr(dispatcher.cfg.local_vision, "enabled", False)
monkeypatch.setenv("NEURALWATT_API_KEY", "test-key")
# ensure the gate is on for the happy-path tests (default, but pin it).
monkeypatch.setattr(dispatcher.cfg.logging, "log_route_decisions", True)
# Exploration is on by default; pin epsilon to 0 so existing flex tests
# keep deterministic winner/flex behavior (they assert on CHEAP/CHEAP_FLEX).
monkeypatch.setattr(dispatcher.cfg.exploration, "enabled", True)
monkeypatch.setattr(dispatcher.cfg.exploration, "epsilon", 0.0)
calls = []
def fake_post(url, headers=None, json=None, stream=False, timeout=None):
calls.append({"url": url, "body": json, "stream": stream})
if stream:
return FakeResponse(lines=_STREAM_LINES)
return FakeResponse(_completion(json["model"]))
monkeypatch.setattr(dispatcher.requests, "post", fake_post)
monkeypatch.setattr(
dispatcher, "classify",
lambda task, context: Classification(
task_category="coding_general", task_tier=2,
required_context_tokens=100, confidence=0.9,
),
)
class _Raw:
text = json.dumps(_completion(CHEAP))
class _Completions:
@property
def with_raw_response(self):
return self
def create(self, **kwargs):
return _Raw()
class _FakeClient:
chat = type("_Chat", (), {"completions": _Completions()})()
monkeypatch.setattr(dispatcher, "_provider_client", lambda provider: _FakeClient())
yield TestClient(app), db_path
def _rows(db_path):
conn = sqlite3.connect(db_path)
conn.row_factory = sqlite3.Row
rows = conn.execute(
"SELECT * FROM route_decisions ORDER BY id ASC"
).fetchall()
conn.close()
return rows
def _drop_cheap(db_path):
"""Make CHEAP ineligible so DEAR (or nothing) remains."""
conn = sqlite3.connect(db_path)
conn.execute("UPDATE models SET tier = 1 WHERE model_id = ?", (CHEAP,))
conn.commit()
conn.close()
# --- happy paths ------------------------------------------------------------
def test_route_endpoint_persists_one_row(decision_router):
client, db_path = decision_router
resp = client.post(
"/route", json={"task": "write me a function"}
)
assert resp.status_code == 200
rows = _rows(db_path)
assert len(rows) == 1
r = rows[0]
assert r["kind"] == "route"
assert r["selected_model"] == CHEAP
assert r["selected_provider"] == "neuralwatt"
assert r["classification_source"] == "classifier"
assert r["task_category"] == "coding_general"
assert r["task_tier"] == 2
assert r["latency_tolerance"] == "interactive"
assert r["session_key"] is None
assert "session_dir" not in r.keys()
def test_route_endpoint_returns_named_profile(decision_router):
client, db_path = decision_router
resp = client.post(
"/route", json={"task": "write me a function", "profile": "locality"}
)
assert resp.status_code == 200
data = resp.json()
assert data["profile"] == "locality"
rows = _rows(db_path)
assert len(rows) == 1
assert rows[0]["profile"] == "locality"
assert rows[0]["kind"] == "route"
def test_route_endpoint_unknown_profile_returns_422(decision_router):
client, db_path = decision_router
resp = client.post(
"/route", json={"task": "write me a function", "profile": "nosuchprofile"}
)
assert resp.status_code == 422
detail = resp.json()["detail"]
assert "nosuchprofile" in detail
assert "default" in detail
assert "batch" in detail
rows = _rows(db_path)
assert len(rows) == 0
def test_dispatch_endpoint_returns_named_profile(decision_router):
client, db_path = decision_router
resp = client.post(
"/dispatch",
json={
"task": "write me a function",
"profile": "batch",
"task_category": "coding_general",
"task_tier": 2,
"required_context_tokens": 100,
},
)
assert resp.status_code == 200
data = resp.json()
assert data["route"]["profile"] == "batch"
rows = _rows(db_path)
assert len(rows) == 1
assert rows[0]["profile"] == "batch"
assert rows[0]["kind"] == "dispatch"
def test_route_endpoint_override_has_classifier_ms_null(decision_router):
client, db_path = decision_router
resp = client.post(
"/route",
json={
"task": "x",
"task_category": "coding_refactor",
"task_tier": 3,
"required_context_tokens": 5000,
},
)
assert resp.status_code == 200
rows = _rows(db_path)
assert len(rows) == 1
r = rows[0]
assert r["classification_source"] == "override"
assert r["classifier_ms"] is None, "an override never consulted the classifier"
assert r["task_category"] == "coding_refactor"
def test_dispatch_endpoint_persists_one_row(decision_router):
client, db_path = decision_router
resp = client.post(
"/dispatch",
json={"task": "write me a function",
"task_category": "coding_general",
"task_tier": 2,
"required_context_tokens": 100},
)
assert resp.status_code == 200
rows = _rows(db_path)
assert len(rows) == 1
r = rows[0]
assert r["kind"] == "dispatch"
assert r["selected_model"] == CHEAP
assert r["classification_source"] == "override"
def test_routed_chat_persists_one_row(decision_router):
client, db_path = decision_router
resp = client.post(
"/v1/chat/completions",
json={"model": "auto", "messages": _messages()},
)
assert resp.status_code == 200
rows = _rows(db_path)
assert len(rows) == 1
r = rows[0]
assert r["kind"] == "chat"
assert r["selected_model"] == CHEAP
assert r["selected_provider"] == "neuralwatt"
assert r["classification_source"] == "classifier"
assert r["session_key"] is not None
assert r["profile"] == "default"
# The session key is a hash — never a directory, never content.
assert len(r["session_key"]) == 16
assert "/" not in (r["session_key"] or "")
def test_routed_chat_reroute_keeps_classifier_source(decision_router):
"""Re-routing for measured context must not persist source='override'."""
client, db_path = decision_router
# >100 measured tokens (measured = chars/3), so chat_completions reroutes.
resp = client.post(
"/v1/chat/completions",
json={"model": "auto", "messages": _messages("refactor " + "x " * 600)},
)
assert resp.status_code == 200
rows = _rows(db_path)
assert len(rows) == 1
assert rows[0]["classification_source"] == "classifier", \
"the re-route's source='override' must not overwrite the classifier's"
assert rows[0]["classifier_ms"] is not None
def test_streamed_routed_chat_persists_one_row(decision_router):
client, db_path = decision_router
resp = client.post(
"/v1/chat/completions",
json={"model": "auto", "messages": _messages(), "stream": True},
)
assert resp.status_code == 200
rows = _rows(db_path)
assert len(rows) == 1
assert rows[0]["kind"] == "chat"
assert rows[0]["selected_model"] == CHEAP
assert rows[0]["streamed"] == 1
assert rows[0]["profile"] == "default"
def test_streamed_routed_chat_writes_request_id_back(decision_router):
"""The primary traffic path must persist the provider request_id on the decision row."""
client, db_path = decision_router
resp = client.post(
"/v1/chat/completions",
json={"model": "auto", "messages": _messages(), "stream": True},
)
assert resp.status_code == 200
rows = _rows(db_path)
assert len(rows) == 1
assert rows[0]["request_id"] == "chatcmpl-stream-dec", \
"streamed response id must be written back to route_decisions"
def test_passthrough_persists_one_row_with_no_nameerror(decision_router):
"""The pre-existing pass-through NameError must stay gone, and a row lands."""
client, db_path = decision_router
resp = client.post(
"/v1/chat/completions",
json={"model": DEAR, "messages": _messages()},
)
assert resp.status_code == 200
rows = _rows(db_path)
assert len(rows) == 1
r = rows[0]
assert r["kind"] == "passthrough"
assert r["selected_model"] == DEAR
assert r["selected_provider"] == "neuralwatt"
assert r["classification_source"] is None
assert r["session_key"] is not None
assert r["profile"] is None, "non-routed passthrough must write profile=None"
def test_passthrough_records_pinch_columns_when_pruned(decision_router, monkeypatch):
"""Regression: passthrough pruning must run before persistence so pinch stats
are recorded rather than NULL."""
client, db_path = decision_router
monkeypatch.setattr(dispatcher.cfg.pinch, "enabled", True)
monkeypatch.setattr(dispatcher.cfg.pinch, "budget_tokens", 1)
monkeypatch.setattr(dispatcher.cfg.pinch, "keep_last_turns", 1)
# Long conversation plus tools to exceed the tiny budget.
long_text = "word " * 500
messages = [
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": long_text},
{"role": "assistant", "content": "acknowledged"},
{"role": "user", "content": "now summarize"},
]
resp = client.post(
"/v1/chat/completions",
json={
"model": DEAR,
"messages": messages,
"stream": False,
"tools": [
{
"type": "function",
"function": {
"name": "noop",
"description": "A tool with a long description to inflate overhead",
"parameters": {"type": "object", "properties": {}},
},
}
],
},
)
# Routing succeeds because the fixture stubs the provider call; persistence
# happens before forwarding, so the row is already written.
assert resp.status_code == 200
rows = _rows(db_path)
passthrough_rows = [r for r in rows if r["kind"] == "passthrough"]
assert len(passthrough_rows) == 1
r = passthrough_rows[0]
assert r["pinch_original_tokens"] is not None
assert r["pinch_final_tokens"] is not None
assert r["pinch_original_tokens"] > r["pinch_final_tokens"]
def _two_turn_session():
"""Two turns of one session: same opening message, one message appended.
The opening message is what ``session_fingerprint`` hashes, so holding it
constant is what makes these two requests the same session as far as the
probe is concerned.
"""
first = [
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": "write me a function"},
]
return first, first + [
{"role": "assistant", "content": "here you go"},
{"role": "user", "content": "now add a test"},
]
def test_routed_chat_records_the_prefix_probe_across_two_turns(
decision_router, monkeypatch
):
"""Turn 1 has nothing to compare against; turn 2 reports its divergence."""
client, db_path = decision_router
prefix_probe.reset()
monkeypatch.setattr(dispatcher.cfg.pinch, "enabled", True)
monkeypatch.setattr(dispatcher.cfg.pinch, "prefix_probe", True)
turn_one, turn_two = _two_turn_session()
for messages in (turn_one, turn_two):
resp = client.post(
"/v1/chat/completions", json={"model": "auto", "messages": messages}
)
assert resp.status_code == 200
rows = _rows(db_path)
assert len(rows) == 2
assert rows[0]["prefix_divergence_index"] is None, (
"the first turn of a session has nothing to compare against, and a "
"stored 0 would read as total cache loss at message 0"
)
assert rows[0]["prefix_prev_message_count"] is None
# This conversation only grew, so it diverges at exactly the position the
# previous turn did not have: a healthy append, not rewritten history.
assert rows[1]["prefix_prev_message_count"] == len(turn_one)
assert rows[1]["prefix_divergence_index"] == len(turn_one)
assert rows[1]["prefix_tokens_after_divergence"] > 0
def test_passthrough_records_the_prefix_probe_too(decision_router, monkeypatch):
"""Both prune call sites feed the probe, not just the routed one."""
client, db_path = decision_router
prefix_probe.reset()
monkeypatch.setattr(dispatcher.cfg.pinch, "enabled", True)
monkeypatch.setattr(dispatcher.cfg.pinch, "prefix_probe", True)
turn_one, turn_two = _two_turn_session()
for messages in (turn_one, turn_two):
resp = client.post(
"/v1/chat/completions", json={"model": DEAR, "messages": messages}
)
assert resp.status_code == 200
rows = _rows(db_path)
assert [r["kind"] for r in rows] == ["passthrough", "passthrough"]
assert rows[0]["prefix_divergence_index"] is None
assert rows[1]["prefix_divergence_index"] == len(turn_one)
def test_prefix_probe_off_writes_nulls(decision_router, monkeypatch):
"""The knob is a real gate: no hashing, and three NULL columns."""
client, db_path = decision_router
prefix_probe.reset()
monkeypatch.setattr(dispatcher.cfg.pinch, "enabled", True)
monkeypatch.setattr(dispatcher.cfg.pinch, "prefix_probe", False)
turn_one, turn_two = _two_turn_session()
for messages in (turn_one, turn_two):
resp = client.post(
"/v1/chat/completions", json={"model": "auto", "messages": messages}
)
assert resp.status_code == 200
rows = _rows(db_path)
assert len(rows) == 2
for row in rows:
assert row["prefix_divergence_index"] is None
assert row["prefix_tokens_after_divergence"] is None
assert row["prefix_prev_message_count"] is None
assert prefix_probe._store == {}, "nothing should have been fingerprinted"
def test_local_vision_success_persists_one_local_row(decision_router, monkeypatch):
client, db_path = decision_router
_drop_cheap(db_path)
monkeypatch.setattr(dispatcher.cfg.local_vision, "enabled", True)
def fake_post(url, headers=None, json=None, stream=False, timeout=None):
return FakeResponse(
{"choices": [{"message": {"role": "assistant",
"content": "local caption"},
"finish_reason": "stop"}]}
)
monkeypatch.setattr(dispatcher.requests, "post", fake_post)
resp = client.post(
"/v1/chat/completions",
json={"model": "auto", "messages": _image_messages()},
)
assert resp.status_code == 200
rows = _rows(db_path)
assert len(rows) == 1, "one decision, and only one: the local_vision row"
r = rows[0]
assert r["kind"] == "local_vision"
assert r["selected_model"] == dispatcher.cfg.local_vision.model
assert r["selected_provider"] == "local"
assert r["rejected_reason"] is None
assert r["images"] == 1
assert r["profile"] is None, "non-routed local_vision must write profile=None"
def test_no_candidate_422_still_persists_a_rejection_row(decision_router):
client, db_path = decision_router
# No vision cloud candidate, local fallback disabled -> 422.
_drop_cheap(db_path)
resp = client.post(
"/v1/chat/completions",
json={"model": "auto", "messages": _image_messages()},
)
assert resp.status_code == 422
rows = _rows(db_path)
assert len(rows) == 1
r = rows[0]
assert r["kind"] == "chat"
assert r["selected_model"] is None
assert r["rejected_reason"] is not None
assert "vision" in r["rejected_reason"]
assert r["profile"] == "default"
# --- failure modes: best-effort, config-gated -------------------------------
def test_gate_off_writes_nothing_but_routing_still_200(decision_router):
client, db_path = decision_router
dispatcher.cfg.logging.log_route_decisions = False
try:
resp = client.post("/route", json={"task": "write me a function"})
assert resp.status_code == 200
finally:
dispatcher.cfg.logging.log_route_decisions = True
assert _rows(db_path) == [], "the gate off must leave the table untouched"
def _raise_on_route_decisions_insert(db_path):
real_db = dispatcher._db
class GuardingConn:
def __init__(self, conn):
self._conn = conn
def __getattr__(self, name):
return getattr(self._conn, name)
def execute(self, sql, parameters=()):
if isinstance(sql, str) and "INSERT INTO route_decisions" in sql:
raise sqlite3.OperationalError("database is locked")
return self._conn.execute(sql, parameters)
def wrapped_db():
return GuardingConn(real_db())
return wrapped_db
def test_db_write_failure_never_fails_routing(decision_router, monkeypatch):
"""A locked/read-only DB must not error the request; persistence is best-effort."""
client, db_path = decision_router
monkeypatch.setattr(dispatcher, "_db", _raise_on_route_decisions_insert(db_path))
resp = client.post("/route", json={"task": "write me a function"})
assert resp.status_code == 200, "a failed decision write must never fail routing"
# And the same holds for a routed completion.
resp2 = client.post(
"/v1/chat/completions", json={"model": "auto", "messages": _messages()}
)
assert resp2.status_code == 200
# =============================================================================
# flex_preference: request override, swap behavior, cost re-estimation.
# =============================================================================
def test_route_with_force_flex_swaps_under_interactive(decision_router):
"""force-flex bypasses the interactive latency filter and rides the flex twin."""
client, db_path = decision_router
resp = client.post(
"/route", json={"task": "write me a function", "flex_preference": "force-flex"}
)
assert resp.status_code == 200
data = resp.json()
assert data["selected"]["model_id"] == CHEAP_FLEX
assert data["flex_preference"] == "force-flex"
assert data["flex_swapped"] is True
assert data["flex_forced"] is True
r = _rows(db_path)[0]
assert r["flex_preference"] == "force-flex"
assert r["flex_swapped"] == 1
assert r["flex_forced"] == 1
assert r["selected_model"] == CHEAP_FLEX
assert r["exploration"] == 0
def test_route_with_prefer_flex_does_not_swap_under_interactive(decision_router):
"""prefer-flex defers to the latency filter, so interactive keeps the standard row."""
client, db_path = decision_router
resp = client.post(
"/route", json={"task": "write me a function", "flex_preference": "prefer-flex"}
)
assert resp.status_code == 200
data = resp.json()
assert data["selected"]["model_id"] == CHEAP
assert data["flex_preference"] == "prefer-flex"
assert data["flex_swapped"] is False
assert data["flex_forced"] is False
r = _rows(db_path)[0]
assert r["flex_preference"] == "prefer-flex"
assert r["flex_swapped"] == 0
assert r["selected_model"] == CHEAP
def test_route_with_prefer_flex_swaps_under_batch(decision_router):
"""Under batch the latency filter admits flex, so prefer-flex rides the twin."""
client, db_path = decision_router
resp = client.post(
"/route",
json={"task": "write me a function", "flex_preference": "prefer-flex",
"latency_tolerance": "batch"},
)
assert resp.status_code == 200
data = resp.json()
assert data["selected"]["model_id"] == CHEAP_FLEX
assert data["flex_preference"] == "prefer-flex"
assert data["flex_swapped"] is True
assert data["flex_forced"] is False
r = _rows(db_path)[0]
assert r["flex_preference"] == "prefer-flex"
assert r["flex_swapped"] == 1
assert r["flex_forced"] == 0
assert r["latency_tolerance"] == "batch"
assert r["selected_model"] == CHEAP_FLEX
def test_route_with_no_flex_never_swaps(decision_router):
"""no-flex never routes to a flex row, even under batch."""
client, db_path = decision_router
resp = client.post(
"/route",
json={"task": "write me a function", "flex_preference": "no-flex",
"latency_tolerance": "batch"},
)
assert resp.status_code == 200
data = resp.json()
assert data["selected"]["model_id"] == CHEAP
assert data["flex_preference"] == "no-flex"
assert data["flex_swapped"] is False
assert data["flex_forced"] is False
r = _rows(db_path)[0]
assert r["flex_preference"] == "no-flex"
assert r["flex_swapped"] == 0
assert r["selected_model"] == CHEAP
def test_route_with_auto_uses_default_behavior(decision_router):
"""No flex_preference -> auto: no swap, interactive latency, flex_preference='auto'."""
client, db_path = decision_router
resp = client.post("/route", json={"task": "write me a function"})
assert resp.status_code == 200
data = resp.json()
assert data["selected"]["model_id"] == CHEAP
assert data["flex_preference"] == "auto"
assert data["flex_swapped"] is False
assert data["flex_forced"] is False
r = _rows(db_path)[0]
assert r["flex_preference"] == "auto"
assert r["flex_swapped"] == 0
assert r["selected_model"] == CHEAP
def test_route_flex_preference_override_beats_config_default(decision_router, monkeypatch):
"""A request override wins over routing.default_flex_preference."""
client, db_path = decision_router
monkeypatch.setattr(
dispatcher.cfg.routing, "default_flex_preference", config.FlexPreference.force_flex
)
# No override: the config default (force-flex) applies and swaps.
resp = client.post("/route", json={"task": "write me a function"})
assert resp.status_code == 200
assert resp.json()["selected"]["model_id"] == CHEAP_FLEX
assert resp.json()["flex_swapped"] is True
assert resp.json()["flex_preference"] == "force-flex"
# Explicit override: no-flex beats the config default and does not swap.
resp = client.post(
"/route", json={"task": "write me a function", "flex_preference": "no-flex"}
)
assert resp.status_code == 200
data = resp.json()
assert data["selected"]["model_id"] == CHEAP
assert data["flex_preference"] == "no-flex"
assert data["flex_swapped"] is False
rows = _rows(db_path)
assert rows[0]["selected_model"] == CHEAP_FLEX
assert rows[1]["selected_model"] == CHEAP
def test_route_force_flex_without_sibling_falls_back(decision_router):
"""force-flex on a model with no flex sibling keeps the standard row, no swap."""
client, db_path = decision_router
_drop_cheap(db_path) # DEAR (no flex sibling) becomes the winner.
conn = sqlite3.connect(db_path)
conn.execute("DELETE FROM proficiency WHERE model_id = ?", (EXPLORABLE_DEAR,))
conn.execute("UPDATE proficiency SET outcome_samples = 50 WHERE model_id = ?", (DEAR,))
conn.commit()
conn.close()
resp = client.post(
"/route", json={"task": "write me a function", "flex_preference": "force-flex"}
)
assert resp.status_code == 200
data = resp.json()
assert data["selected"]["model_id"] == DEAR
assert data["flex_preference"] == "force-flex"
assert data["flex_swapped"] is False
assert data["flex_forced"] is False
r = _rows(db_path)[0]
assert r["flex_swapped"] == 0
assert r["selected_model"] == DEAR
def test_route_swapped_cost_reflects_flex_variant(decision_router):
"""A flex swap re-estimates cost from the flex row's pricing, not the standard's."""
client, db_path = decision_router
resp = client.post(
"/route", json={"task": "write me a function", "flex_preference": "force-flex"}
)
assert resp.status_code == 200
data = resp.json()
assert data["selected"]["model_id"] == CHEAP_FLEX
flex_cost = data["selected"]["cost"]
r = _rows(db_path)[0]
assert r["est_cost_usd"] == pytest.approx(flex_cost)
cache_rate = dispatcher.cfg.objective.assumed_cache_rate
completion_tokens = dispatcher.cfg.objective.assumed_completion_tokens
prompt_tokens = 100 # the seeded classifier's required_context_tokens
expected = (
prompt_tokens * cache_rate * 0.20
+ prompt_tokens * (1.0 - cache_rate) * 0.20
+ completion_tokens * 0.60
) / 1_000_000
assert flex_cost == pytest.approx(expected)
def test_route_force_flex_does_not_swap_to_stale_sibling(decision_router):
"""force-flex must not dispatch to a stale flex sibling; the active standard wins."""
client, db_path = decision_router
resp = client.post(
"/route", json={"task": "write me a function", "flex_preference": "force-flex"}
)
assert resp.status_code == 200
data = resp.json()
# The active CHEAP_FLEX sibling is still the healthy swap target.
assert data["selected"]["model_id"] == CHEAP_FLEX
assert data["flex_swapped"] is True
# Make CHEAP_FLEX unavailable so the only remaining flex sibling is stale;
# force-flex must refuse it and keep the standard CHEAP winner.
conn = sqlite3.connect(db_path)
conn.execute(
"UPDATE models SET availability = 'stale' WHERE model_id = ?", (CHEAP_FLEX,)
)
conn.commit()
conn.close()
resp = client.post(
"/route", json={"task": "write me a function", "flex_preference": "force-flex"}
)
assert resp.status_code == 200
data = resp.json()
assert data["selected"]["model_id"] == CHEAP
assert data["flex_swapped"] is False
assert data["flex_forced"] is False
def _seed_proficiency_outcomes(db_path, counts):
conn = sqlite3.connect(db_path)
for model_id, samples in counts.items():
conn.execute(
"UPDATE proficiency SET outcome_samples = ? WHERE model_id = ?",
(samples, model_id),
)
conn.commit()
conn.close()
def test_exploration_epsilon_one_picks_fewest_samples(decision_router, monkeypatch):
client, db_path = decision_router
_seed_proficiency_outcomes(db_path, {CHEAP: 50, EXPLORABLE_DEAR: 0})
monkeypatch.setattr(dispatcher.cfg.exploration, "enabled", True)
monkeypatch.setattr(dispatcher.cfg.exploration, "epsilon", 1.0)
monkeypatch.setattr(dispatcher.cfg.exploration, "max_tier", 2)
monkeypatch.setattr(dispatcher, "_router_rng", random.Random(0))
resp = client.post("/route", json={"task": "write me a function"})
data = resp.json()
assert resp.status_code == 200
assert data["explored"] is True
assert data["selected"]["model_id"] == EXPLORABLE_DEAR
assert data["selected"]["model_id"] not in {
c["model_id"] for c in data["runners_up"]
}
assert any(c["model_id"] == CHEAP for c in data["runners_up"])
r = _rows(db_path)[0]
assert r["exploration"] == 1
assert r["selected_model"] == EXPLORABLE_DEAR
def test_exploration_disabled_keeps_winner_and_zero_flag(decision_router, monkeypatch):
client, db_path = decision_router
_seed_proficiency_outcomes(db_path, {CHEAP: 50, EXPLORABLE_DEAR: 0})
monkeypatch.setattr(dispatcher.cfg.exploration, "enabled", False)
resp = client.post("/route", json={"task": "write me a function"})
data = resp.json()
assert resp.status_code == 200
assert data["explored"] is False
assert data["selected"]["model_id"] == CHEAP
r = _rows(db_path)[0]
assert r["exploration"] == 0
assert r["selected_model"] == CHEAP
def test_exploration_max_tier_excludes_tier_three(decision_router, monkeypatch):
_, db_path = decision_router
_seed_proficiency_outcomes(db_path, {CHEAP: 50, EXPLORABLE_DEAR: 0})
monkeypatch.setattr(dispatcher.cfg.exploration, "enabled", True)
monkeypatch.setattr(dispatcher.cfg.exploration, "epsilon", 1.0)
monkeypatch.setattr(dispatcher.cfg.exploration, "max_tier", 2)
# Direct route() call lets us inspect explored without the 422 behavior.
from dispatcher import TaskRequest
from dispatcher import route as route_fn
decision = route_fn(
TaskRequest(
task="write me a function",
task_category="coding_general",
task_tier=3,
required_context_tokens=100,
)
)
assert decision.explored is False
def test_exploration_runners_up_coherent_after_flex_swap(decision_router, monkeypatch):
client, db_path = decision_router
_seed_proficiency_outcomes(db_path, {CHEAP: 50, EXPLORABLE_DEAR: 0, CHEAP_FLEX: 0})
monkeypatch.setattr(dispatcher.cfg.exploration, "enabled", True)
monkeypatch.setattr(dispatcher.cfg.exploration, "epsilon", 1.0)
monkeypatch.setattr(dispatcher.cfg.exploration, "max_tier", 2)
monkeypatch.setattr(dispatcher, "_router_rng", random.Random(0))
resp = client.post(
"/route",
json={"task": "write me a function", "flex_preference": "force-flex"},
)
data = resp.json()
assert resp.status_code == 200
assert data["explored"] is True
assert data["selected"]["model_id"] == EXPLORABLE_DEAR
selected_id = data["selected"]["model_id"]
runner_ids = {c["model_id"] for c in data["runners_up"]}
assert selected_id not in runner_ids
assert CHEAP in runner_ids