Files
6krrt/tests/test_classifier_cascade.py

273 lines
9.5 KiB
Python

"""The classifier fallback cascade: local -> stale -> history -> cloud -> guess.
Each step is asserted by making the ones before it fail, so a test that
passes proves that step actually ran rather than that some earlier step
happened to produce the same answer.
"""
from __future__ import annotations
import sqlite3
from datetime import datetime, timedelta, timezone
from pathlib import Path
from types import SimpleNamespace
import pytest
import dispatcher
import session_cache
ROOT = Path(__file__).resolve().parent.parent
SCHEMA_SQL = (ROOT / "config" / "schema.sql").read_text()
def _now() -> datetime:
return datetime.now(timezone.utc)
@pytest.fixture(autouse=True)
def _clean_state(monkeypatch):
"""Each test starts with an empty cache and a closed circuit."""
session_cache.clear()
monkeypatch.setattr(dispatcher, "_last_classifier_failure", 0.0)
dispatcher._provider_refusal_since.clear()
yield
dispatcher._provider_refusal_since.clear()
session_cache.clear()
@pytest.fixture
def db(tmp_path, monkeypatch):
path = tmp_path / "cascade.db"
conn = sqlite3.connect(path)
conn.row_factory = sqlite3.Row
conn.executescript(SCHEMA_SQL)
conn.commit()
def _open():
c = sqlite3.connect(path)
c.row_factory = sqlite3.Row
return c
monkeypatch.setattr(dispatcher, "_db", _open)
yield conn
conn.close()
def test_step1_stale_session_beats_a_generic_guess():
"""A stale REAL classification is used before falling back to a guess."""
session_cache.put("sess-1", task_category="coding_refactor", task_tier=3)
got = dispatcher._classify_cascade("sess-1", "sys", "user")
assert got is not None
assert got.source == "session_stale"
assert got.task_category == "coding_refactor"
assert got.task_tier == 3
def test_step1_ignores_staleness_but_never_renews_it():
"""stale_read must not refresh cached_at — a degraded read earns nothing."""
session_cache.put("sess-2", task_category="debugging", task_tier=2)
before = session_cache.stale_read("sess-2").cached_at
dispatcher._classify_cascade("sess-2", "sys", "user")
assert session_cache.stale_read("sess-2").cached_at == before
def test_step2_session_history_survives_a_restart(db):
"""With an empty cache, the session's own history is read from the DB.
This is the step that matters after a router restart, when the in-memory
cache is gone but route_decisions still remembers the session.
"""
db.execute(
"""
INSERT INTO route_decisions
(kind, task_category, task_tier, required_context_tokens,
selected_model, selected_provider, classification_source,
session_key, observed_at)
VALUES ('chat','translation',1,50,'m','neuralwatt','classifier',?,?)
""",
("sess-3", _now().isoformat()),
)
db.commit()
got = dispatcher._classify_cascade("sess-3", "sys", "user")
assert got is not None
assert got.source == "session_history"
assert got.task_category == "translation"
def test_step2_refuses_to_reuse_a_previous_degraded_answer(db):
"""A prior fallback must not be read back as if it were real.
Otherwise one outage's guess propagates through every later turn of the
session and looks like a genuine classification forever.
"""
for src in ("fallback", "session_stale", "session_history"):
db.execute(
"""
INSERT INTO route_decisions
(kind, task_category, task_tier, required_context_tokens,
selected_model, selected_provider, classification_source,
session_key, observed_at)
VALUES ('chat','general_chat',2,0,'m','neuralwatt',?,?,?)
""",
(src, "sess-4", _now().isoformat()),
)
db.commit()
# No cloud configured -> None means "use the static fallback".
assert dispatcher._classify_cascade("sess-4", "sys", "user") is None
def test_cloud_step_is_skipped_when_unconfigured(db):
"""Zero cost unless someone explicitly configures cloud_fallback."""
assert getattr(dispatcher.cfg.classifier, "cloud_fallback", None) is None
assert dispatcher._classify_cascade("sess-none", "sys", "user") is None
def test_cloud_step_used_when_configured(db, monkeypatch):
"""A configured cloud classifier answers when everything free has failed."""
monkeypatch.setattr(
dispatcher.cfg.classifier,
"cloud_fallback",
SimpleNamespace(
base_url="https://cloud.example/v1",
model="cloud-classifier",
timeout_seconds=2,
api_key_env=None,
max_output_tokens=1024,
),
)
monkeypatch.setattr(dispatcher, "_cloud_classifier_client", lambda cf: object())
monkeypatch.setattr(
session_cache,
"classify_one",
lambda *a, **k: {
"task_category": "reasoning_math",
"task_tier": 3,
"required_context_tokens": 1234,
"confidence": 0.8,
},
)
got = dispatcher._classify_cascade("sess-5", "sys", "user")
assert got is not None
assert got.source == "classifier_cloud"
assert got.task_category == "reasoning_math"
assert got.required_context_tokens == 1234
def test_provider_refusal_skips_the_cloud_step(db, monkeypatch):
"""Out of credit -> a cloud classification is a guaranteed wasted request."""
monkeypatch.setattr(
dispatcher.cfg.classifier,
"cloud_fallback",
SimpleNamespace(
base_url="https://cloud.example/v1",
model="cloud-classifier",
timeout_seconds=2,
api_key_env=None,
max_output_tokens=1024,
),
)
called = []
monkeypatch.setattr(
dispatcher,
"_cloud_classifier_client",
lambda cf: called.append(1) or object(),
)
dispatcher._record_provider_refusal("neuralwatt")
assert dispatcher._classify_cascade("sess-6", "sys", "user") is None
assert called == [], "cloud client built despite a provider-level refusal"
def test_cloud_failure_degrades_to_the_static_guess(db, monkeypatch):
"""An unreachable cloud classifier must not raise — it returns None."""
monkeypatch.setattr(
dispatcher.cfg.classifier,
"cloud_fallback",
SimpleNamespace(
base_url="https://cloud.example/v1",
model="cloud-classifier",
timeout_seconds=2,
api_key_env=None,
max_output_tokens=1024,
),
)
monkeypatch.setattr(dispatcher, "_cloud_classifier_client", lambda cf: object())
monkeypatch.setattr(session_cache, "classify_one", lambda *a, **k: None)
assert dispatcher._classify_cascade("sess-7", "sys", "user") is None
def test_missing_api_key_degrades_rather_than_raising():
"""Unlike the primary classifier, a keyless cloud fallback must not 503.
The primary raising is right — a misconfigured primary should fail loudly.
This one is reached only when things are already broken, so failing the
request would make the outage worse than the guess it replaces.
"""
cf = SimpleNamespace(
base_url="https://cloud.example/v1",
model="cloud-classifier",
timeout_seconds=2,
api_key_env="DEFINITELY_NOT_SET_ANYWHERE",
max_output_tokens=1024,
)
assert dispatcher._cloud_classifier_client(cf) is None
def test_cascade_with_no_session_key_falls_straight_through(db):
"""No session -> steps 1 and 2 are impossible, and that is not an error."""
assert dispatcher._classify_cascade(None, "sys", "user") is None
# ---------------------------------------------------------------------------
# The degradation warning: a survivable failure is the kind that goes unnoticed.
# ---------------------------------------------------------------------------
def _seed_sources(conn, pairs):
for source, count in pairs:
for _ in range(count):
conn.execute(
"""
INSERT INTO route_decisions
(kind, task_category, task_tier, required_context_tokens,
selected_model, selected_provider, classification_source,
observed_at)
VALUES ('chat','general_chat',2,0,'m','neuralwatt',?,?)
""",
(source, _now().isoformat()),
)
conn.commit()
def test_degradation_warning_fires_above_threshold(db):
from metrics import classifier_degradation_warning
cfg = SimpleNamespace(
classifier=SimpleNamespace(degraded_warn_min=20, degraded_warn_threshold=0.5)
)
_seed_sources(db, [("fallback", 15), ("classifier", 10)])
out = classifier_degradation_warning(db, cfg)
assert len(out) == 1
assert "60%" in out[0]
assert "degraded source" in out[0]
def test_degradation_warning_silent_on_thin_traffic(db):
"""A share over three requests is noise; a flapping warning goes unread."""
from metrics import classifier_degradation_warning
cfg = SimpleNamespace(
classifier=SimpleNamespace(degraded_warn_min=20, degraded_warn_threshold=0.5)
)
_seed_sources(db, [("fallback", 3)])
assert classifier_degradation_warning(db, cfg) == []
def test_degradation_warning_silent_when_healthy(db):
from metrics import classifier_degradation_warning
cfg = SimpleNamespace(
classifier=SimpleNamespace(degraded_warn_min=20, degraded_warn_threshold=0.5)
)
_seed_sources(db, [("classifier", 24), ("fallback", 1)])
assert classifier_degradation_warning(db, cfg) == []