273 lines
9.5 KiB
Python
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) == []
|