Files
6krrt/tests/test_classifier_attempt.py
adlee-was-taken 4c00f2fee5 refactor(config): name the degraded-warn bounds so a runtime write can import them
A runtime admin write goes straight past Pydantic, so the registry's bounds
are the only check between a request body and a field the warning reads. The
validators inlined `0 < v <= 1` and `v >= 1`, so a registry entry would have
had to restate them and could drift: it would accept 0.0 where load refuses.

Name the three edges (exclusive low, high, min floor), use them in the
validators, and pin them to the validators' actual boundary behaviour.
No behaviour change.

Co-Authored-By: Claude Sonnet 5.5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01KkCGRantZsSwmcFpet6FTa
2026-10-04 22:39:34 -04:00

533 lines
20 KiB
Python

"""The classifier's own attempt reaches route_decisions.
Before this, an abstention left one number behind and it was in a log line:
``local_decision confidence 0.431 is below classifier.decision.confidence_min
(0.5)``. The database said only ``classification_source = session_history``, and
``route_decisions.confidence`` read 1.0 on 99.9% of chat rows because the chat
path re-routes through the override branch, which hard-codes it. So
``confidence_min`` could not be tuned from data: nobody could say whether the
abstentions sat just under the floor or nowhere near it.
These tests pin the three columns end to end: ``classifier_confidence``,
``classifier_coverage`` and ``classifier_reject``. The chat-path test runs the
REAL ``classify()`` (only the Ollama call is stubbed), because the stub that
``test_no_header_snapshot`` installs replaces ``classify`` itself and would hide
exactly the stamping this file is about.
"""
import json
import sqlite3
from pathlib import Path
from types import SimpleNamespace
import openai
import pytest
import requests
import test_no_header_snapshot as snap
from test_route_decisions import _schema_minus_profile_column
import dispatcher
import local_decision
import prefix_probe
import session_cache
from classifier_rejection import (
REASON_BELOW_CONFIDENCE,
REASON_BELOW_COVERAGE,
REASON_NO_LOGPROBS,
REASON_PARSE,
REASON_PRIMARY_FAILED,
REASON_TIMEOUT,
REASON_TRANSPORT,
ClassifierRejected,
)
from dispatcher import Classification, ClassifierAttempt
ROOT = Path(__file__).resolve().parent.parent
SCHEMA_SQL = (ROOT / "config" / "schema.sql").read_text()
ATTEMPT_COLUMNS = ("classifier_confidence", "classifier_coverage", "classifier_reject")
@pytest.fixture(autouse=True)
def _clean_state(monkeypatch):
session_cache.clear()
prefix_probe._store.clear()
monkeypatch.setattr(dispatcher, "_last_classifier_failure", 0.0)
dispatcher._provider_refusal_since.clear()
token = dispatcher._current_session_key.set(None)
yield
dispatcher._current_session_key.reset(token)
dispatcher._provider_refusal_since.clear()
session_cache.clear()
# The chat-path test posts the snapshot's own payload, and the probe keeps
# per-session fingerprints in process memory: left behind, the headerless
# snapshot test (which runs after this file) sees a "previous turn" and
# records a non-NULL prefix divergence.
prefix_probe._store.clear()
def _use_local_decision(monkeypatch, *, confidence_min=0.5):
"""local_decision as the primary, unmetered, with nothing skipped."""
monkeypatch.setattr(dispatcher.cfg.classifier, "mode", "local_decision")
monkeypatch.setattr(
dispatcher.cfg.classifier,
"decision",
SimpleNamespace(
base_url="http://localhost:11434",
model="stub-model",
num_ctx=8192,
timeout_s=10,
confidence_min=confidence_min,
coverage_min=0.3,
tier_enabled=False,
),
)
monkeypatch.setattr(dispatcher, "_local_classifier_skip_reason", lambda: None)
monkeypatch.setattr(dispatcher.cfg.local_energy, "enabled", False)
monkeypatch.setattr(dispatcher.cfg.local_compute, "enabled", True)
def _answers(monkeypatch, result):
"""Stub the Ollama call: return ``result`` or raise it if it is an exception."""
def fake(*args, **kwargs):
if isinstance(result, BaseException):
raise result
return result
monkeypatch.setattr(local_decision, "classify_category", fake)
# --- the raise sites carry the numbers -------------------------------------
def test_a_floor_miss_records_the_confidence_that_missed_it(monkeypatch):
_use_local_decision(monkeypatch)
_answers(monkeypatch, ("debugging", 0.431, 0.874))
got = dispatcher.classify("fix the failing test", None)
assert got.source == "fallback"
assert got.attempt == ClassifierAttempt(
confidence=0.431, coverage=0.874, reject_reason=REASON_BELOW_CONFIDENCE
)
def test_a_coverage_miss_records_the_coverage_and_the_would_be_confidence(monkeypatch):
"""Both gates are recorded, because tuning one without the other misleads."""
_use_local_decision(monkeypatch)
# Total option mass 0.0498 sits under coverage_min (0.3); letter A holds all of it.
low_mass = {
"logprobs": [
{"token": "A", "logprob": -3.0, "top_logprobs": [{"token": "A", "logprob": -3.0}]}
]
}
monkeypatch.setattr(
local_decision,
"classify_category",
lambda *a, **k: local_decision.parse_logprobs(low_mass, ["A", "B"], coverage_min=0.3),
)
got = dispatcher.classify("fix the failing test", None)
assert got.attempt.reject_reason == REASON_BELOW_COVERAGE
assert got.attempt.coverage == pytest.approx(0.0498, abs=1e-4)
assert got.attempt.confidence == pytest.approx(1.0)
def test_parse_logprobs_raises_a_runtime_error_with_its_original_message():
"""Subclassing RuntimeError keeps every existing handler and log line intact."""
with pytest.raises(ClassifierRejected, match="no logprobs found") as no_lp:
local_decision.parse_logprobs({}, ["A", "B"])
assert isinstance(no_lp.value, RuntimeError)
assert no_lp.value.reason == REASON_NO_LOGPROBS
thin = {"logprobs": [{"token": "A", "top_logprobs": [{"token": "A", "logprob": -4.0}]}]}
with pytest.raises(RuntimeError, match=r"coverage 0\.0183 below minimum 0\.3") as thin_err:
local_decision.parse_logprobs(thin, ["A", "B"], coverage_min=0.3)
assert thin_err.value.reason == REASON_BELOW_COVERAGE
def test_the_rejection_message_is_unchanged(monkeypatch):
"""The `fallback` journal line prints this text; operators grep for it."""
_use_local_decision(monkeypatch)
_answers(monkeypatch, ("debugging", 0.431, 0.874))
with pytest.raises(ClassifierRejected) as err:
dispatcher._classify_via_local_decision("sys", "do a thing")
assert str(err.value) == (
"local_decision confidence 0.431 is below classifier.decision.confidence_min (0.5)"
)
# --- every failure class gets a code ---------------------------------------
@pytest.mark.parametrize(
"failure, reason",
[
(requests.ConnectionError("refused"), REASON_TRANSPORT),
(requests.Timeout("slow"), REASON_TIMEOUT),
# The SDK's own timeout type, which is an OpenAIError subclass and so
# would be filed as a transport error if the order of the checks slipped.
(
openai.APITimeoutError(request=SimpleNamespace(method="POST", url="http://x")),
REASON_TIMEOUT,
),
(ValueError("not json"), REASON_PARSE),
(KeyError("choices"), REASON_PARSE),
(json.JSONDecodeError("bad", "{", 0), REASON_PARSE),
(RuntimeError("something else"), REASON_PRIMARY_FAILED),
],
ids=lambda p: type(p).__name__ if isinstance(p, BaseException) else p,
)
def test_each_failure_class_has_a_stable_code_and_no_invented_numbers(
monkeypatch, failure, reason
):
_use_local_decision(monkeypatch)
_answers(monkeypatch, failure)
got = dispatcher.classify("fix the failing test", None)
assert got.source == "fallback"
assert got.attempt == ClassifierAttempt(reject_reason=reason)
@pytest.mark.parametrize("why", ["gaming_mode", "backoff"])
def test_a_skip_is_recorded_as_a_skip_not_a_failure(monkeypatch, why):
_use_local_decision(monkeypatch)
monkeypatch.setattr(dispatcher, "_local_classifier_skip_reason", lambda: why)
got = dispatcher.classify("fix the failing test", None)
assert got.attempt == ClassifierAttempt(reject_reason=f"skipped_{why}")
def test_an_accepted_answer_carries_its_numbers_and_no_reason(monkeypatch):
_use_local_decision(monkeypatch)
_answers(monkeypatch, ("debugging", 0.83, 0.91))
got = dispatcher.classify("fix the failing test", None)
assert got.source == "classifier"
assert got.attempt == ClassifierAttempt(confidence=0.83, coverage=0.91, reject_reason=None)
def test_modes_without_a_coverage_still_record_their_confidence(monkeypatch):
"""The generic stamp: local_llm, cloud_llm and the encoder have no coverage."""
monkeypatch.setattr(
dispatcher,
"_classify_via_configured_mode",
lambda *a, **k: Classification(
task_category="coding_general",
task_tier=2,
required_context_tokens=100,
confidence=0.9,
),
)
got = dispatcher.classify("fix the failing test", None)
assert got.attempt == ClassifierAttempt(confidence=0.9)
# --- the degraded builder keeps what the cascade forgot ---------------------
def test_the_cascade_result_is_stamped_without_changing_what_it_routes_on(monkeypatch):
borrowed = Classification(
task_category="debugging",
task_tier=3,
required_context_tokens=0,
confidence=0.0,
source="session_history",
)
monkeypatch.setattr(dispatcher, "_classify_cascade", lambda *a: borrowed)
attempt = ClassifierAttempt(confidence=0.431, reject_reason=REASON_BELOW_CONFIDENCE)
got = dispatcher._degraded_classification("k", "sys", "user", attempt=attempt)
assert got.attempt == attempt
assert (got.source, got.task_category, got.task_tier, got.confidence) == (
"session_history",
"debugging",
3,
0.0,
)
assert borrowed.attempt is None, "the cascade's own object must not be mutated"
def test_the_static_guess_is_stamped_too(monkeypatch):
monkeypatch.setattr(dispatcher, "_classify_cascade", lambda *a: None)
attempt = ClassifierAttempt(reject_reason=REASON_TIMEOUT)
got = dispatcher._degraded_classification("k", "sys", "user", attempt=attempt)
assert got.source == "fallback"
assert got.attempt == attempt
assert dispatcher._degraded_classification("k", "sys", "user").attempt is None
# --- the row ----------------------------------------------------------------
def _fresh_db(tmp_path, monkeypatch):
db_path = tmp_path / "attempt.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)
return db_path
def _rows(db_path):
conn = sqlite3.connect(db_path)
conn.row_factory = sqlite3.Row
try:
return [dict(r) for r in conn.execute("SELECT * FROM route_decisions ORDER BY id")]
finally:
conn.close()
def _degraded(attempt):
return Classification(
task_category="debugging",
task_tier=2,
required_context_tokens=0,
confidence=0.0,
source="session_history",
attempt=attempt,
)
def test_persist_writes_the_attempt_from_the_classification(tmp_path, monkeypatch):
db_path = _fresh_db(tmp_path, monkeypatch)
attempt = ClassifierAttempt(
confidence=0.431, coverage=0.874, reject_reason=REASON_BELOW_CONFIDENCE
)
dispatcher.persist_route_decision(
"route", classification=_degraded(attempt), selected_model="m", selected_provider="p"
)
(row,) = _rows(db_path)
assert row["classifier_confidence"] == pytest.approx(0.431)
assert row["classifier_coverage"] == pytest.approx(0.874)
assert row["classifier_reject"] == REASON_BELOW_CONFIDENCE
assert row["confidence"] == 0.0, "the routing classification's own confidence is untouched"
def test_persist_prefers_attempt_of_when_the_classification_was_rerouted(tmp_path, monkeypatch):
"""The chat path rewrites `classification` to an override that remembers nothing."""
db_path = _fresh_db(tmp_path, monkeypatch)
override = Classification(
task_category="debugging",
task_tier=2,
required_context_tokens=90000,
confidence=1.0,
source="override",
)
first = _degraded(
ClassifierAttempt(confidence=0.431, reject_reason=REASON_BELOW_CONFIDENCE)
)
dispatcher.persist_route_decision(
"chat",
classification=override,
classification_source="session_history",
attempt_of=first,
selected_model="m",
selected_provider="p",
)
(row,) = _rows(db_path)
assert row["confidence"] == 1.0, "legacy column: the override's hard-coded value"
assert row["classifier_confidence"] == pytest.approx(0.431)
assert row["classifier_reject"] == REASON_BELOW_CONFIDENCE
def test_persist_without_an_attempt_leaves_all_three_null(tmp_path, monkeypatch):
db_path = _fresh_db(tmp_path, monkeypatch)
override = Classification(
task_category="debugging",
task_tier=2,
required_context_tokens=100,
confidence=1.0,
source="override",
)
dispatcher.persist_route_decision(
"route", classification=override, selected_model="m", selected_provider="p"
)
(row,) = _rows(db_path)
assert [row[c] for c in ATTEMPT_COLUMNS] == [None, None, None]
def test_the_live_event_carries_the_same_three_keys(tmp_path, monkeypatch):
_fresh_db(tmp_path, monkeypatch)
published = []
monkeypatch.setattr(dispatcher.events, "publish_decision", published.append)
attempt = ClassifierAttempt(confidence=0.431, reject_reason=REASON_BELOW_CONFIDENCE)
dispatcher.persist_route_decision(
"route", classification=_degraded(attempt), selected_model="m", selected_provider="p"
)
(event,) = published
assert event["classifier_confidence"] == pytest.approx(0.431)
assert event["classifier_coverage"] is None
assert event["classifier_reject"] == REASON_BELOW_CONFIDENCE
# --- migration ---------------------------------------------------------------
def test_a_live_table_gains_the_columns_once_and_keeps_its_rows(tmp_path):
conn = sqlite3.connect(tmp_path / "live.db")
conn.executescript(_schema_minus_profile_column())
conn.execute(
"INSERT INTO route_decisions (observed_at, kind) VALUES ('2026-10-01T00:00:00+00:00', 'chat')"
)
conn.commit()
before = {r[1] for r in conn.execute("PRAGMA table_info(route_decisions)")}
assert not set(ATTEMPT_COLUMNS) & before
dispatcher.ensure_route_decisions(conn)
dispatcher.ensure_route_decisions(conn) # idempotent
cols = [r[1] for r in conn.execute("PRAGMA table_info(route_decisions)")]
for column in ATTEMPT_COLUMNS:
assert cols.count(column) == 1
old = conn.execute(f"SELECT {', '.join(ATTEMPT_COLUMNS)} FROM route_decisions").fetchone()
assert tuple(old) == (None, None, None), "an old row cannot be backfilled, and says so"
conn.close()
@pytest.mark.skipif(
sqlite3.sqlite_version_info < (3, 35),
reason="ALTER TABLE ... DROP COLUMN needs SQLite 3.35",
)
def test_recent_decisions_reads_a_database_that_has_not_been_migrated(tmp_path):
"""/metrics and the admin page share metrics.py; the live DB lacks the columns
until the router restarts, and that must read as NULL rather than a 500.
The shape under test is a CURRENT schema minus only these three columns,
which is what the live router.db is between the deploy and the restart.
"""
import metrics
conn = sqlite3.connect(tmp_path / "old.db")
conn.row_factory = sqlite3.Row
conn.executescript(SCHEMA_SQL)
for column in ATTEMPT_COLUMNS:
conn.execute(f"ALTER TABLE route_decisions DROP COLUMN {column}")
conn.execute(
"INSERT INTO route_decisions (observed_at, kind) VALUES ('2026-10-01T00:00:00+00:00', 'chat')"
)
conn.commit()
assert not set(ATTEMPT_COLUMNS) & {r[1] for r in conn.execute("PRAGMA table_info(route_decisions)")}
(row,) = metrics.recent_decisions(conn, limit=5)
conn.close()
for column in ATTEMPT_COLUMNS:
assert column in row and row[column] is None
# --- the chat path, with the real classify() ---------------------------------
def test_chat_rows_record_the_attempt_even_though_the_reroute_hides_it(tmp_path, monkeypatch):
"""Turn 1 is accepted; turn 2 of the same session is declined and replays turn 1's label.
Both rows must say what the classifier did. The re-route to measured
context replaces the Classification with an override in between, which is
why ``confidence`` could never have shown this.
"""
real_classify = dispatcher.classify
client, db_path = snap.make_router(tmp_path, monkeypatch)
monkeypatch.setattr(dispatcher, "classify", real_classify)
_use_local_decision(monkeypatch)
answers = iter([("coding_general", 0.91, 0.95), ("coding_general", 0.431, 0.874)])
monkeypatch.setattr(local_decision, "classify_category", lambda *a, **k: next(answers))
for _ in range(2):
resp = snap.post_no_header(client)
assert resp.status_code == 200, resp.text
first, second = _rows(db_path)
assert first["classification_source"] == "classifier"
assert first["classifier_confidence"] == pytest.approx(0.91)
assert first["classifier_coverage"] == pytest.approx(0.95)
assert first["classifier_reject"] is None
assert second["classification_source"] == "session_history"
assert second["classifier_confidence"] == pytest.approx(0.431)
assert second["classifier_coverage"] == pytest.approx(0.874)
assert second["classifier_reject"] == REASON_BELOW_CONFIDENCE
# --- the warning knobs refuse values that could never behave -------------------
def _classifier_with(**overrides):
"""The real loaded classifier block with only the knob under test changed.
ClassifierConfig has required fields, so building one bare would fail on
those and bury what the test is about.
"""
from config import ClassifierConfig
return ClassifierConfig(**{**dispatcher.cfg.classifier.model_dump(), **overrides})
@pytest.mark.parametrize("threshold", [0.0, -0.1, 1.5, 5, 50])
def test_degraded_warn_threshold_outside_a_share_is_refused_at_load(threshold):
"""`80` meant as 80% once shipped as a raw number into another knob and made
every score read as below threshold; a share above 1 here is the same
mistake and would make the warning inert instead."""
from pydantic import ValidationError
with pytest.raises(ValidationError, match="degraded_warn_threshold"):
_classifier_with(degraded_warn_threshold=threshold)
@pytest.mark.parametrize("threshold", [0.01, 0.2, 0.5, 1.0])
def test_degraded_warn_threshold_accepts_any_real_share(threshold):
got = _classifier_with(degraded_warn_threshold=threshold)
assert got.degraded_warn_threshold == threshold
@pytest.mark.parametrize("minimum", [0, -3])
def test_degraded_warn_min_must_be_a_real_sample_size(minimum):
from pydantic import ValidationError
with pytest.raises(ValidationError, match="degraded_warn_min"):
_classifier_with(degraded_warn_min=minimum)
def test_degraded_warn_bound_constants_are_the_validators_boundaries():
"""A runtime admin write skips Pydantic, so its registry imports these named
bounds. They are only trustworthy if they ARE the load validator's edges:
the exclusive low and the floor are refused, the high and floor+0 accepted."""
from pydantic import ValidationError
from config import (
DEGRADED_WARN_MIN_FLOOR,
DEGRADED_WARN_THRESHOLD_MAX,
DEGRADED_WARN_THRESHOLD_MIN_EXCLUSIVE,
)
with pytest.raises(ValidationError):
_classifier_with(degraded_warn_threshold=DEGRADED_WARN_THRESHOLD_MIN_EXCLUSIVE)
assert (
_classifier_with(degraded_warn_threshold=DEGRADED_WARN_THRESHOLD_MAX).degraded_warn_threshold
== DEGRADED_WARN_THRESHOLD_MAX
)
with pytest.raises(ValidationError):
_classifier_with(degraded_warn_min=DEGRADED_WARN_MIN_FLOOR - 1)
assert (
_classifier_with(degraded_warn_min=DEGRADED_WARN_MIN_FLOOR).degraded_warn_min
== DEGRADED_WARN_MIN_FLOOR
)