1463 lines
53 KiB
Python
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
|