"""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) == []