Files
6krrt/tests/test_classifier_modes_dispatch.py

388 lines
14 KiB
Python

"""classify()'s dispatch on cfg.classifier.mode.
Every test here asserts on WHICH client/function was constructed or called,
not just the returned Classification -- the same style
test_classifier_backoff.py established, because the whole point of a mode
branch is which implementation actually ran.
cloud_llm success is asserted to record source="classifier", specifically
NOT "classifier_cloud" -- that string means "the cascade's backup step
fired" and feeds the /metrics degradation-share warning as a degraded
signal. An intentionally configured cloud primary is not degraded.
"""
from __future__ import annotations
from types import SimpleNamespace
from unittest.mock import MagicMock
import pytest
import dispatcher
import local_encoder
import session_cache
@pytest.fixture(autouse=True)
def _clean_state(monkeypatch):
session_cache.clear()
monkeypatch.setattr(dispatcher, "_last_classifier_failure", 0.0)
dispatcher._provider_refusal_since.clear()
monkeypatch.setattr(dispatcher, "_cached_auto_classifier", None)
monkeypatch.setattr(dispatcher, "_auto_classifier_resolved_at", 0.0)
token = dispatcher._current_session_key.set(None)
yield
dispatcher._current_session_key.reset(token)
dispatcher._provider_refusal_since.clear()
session_cache.clear()
def _answer():
return dispatcher.Classification(
task_category="coding_general",
task_tier=2,
required_context_tokens=100,
confidence=0.9,
source="classifier",
)
# --- mode: local_llm (default) --------------------------------------------
def test_default_mode_is_local_llm_and_behaves_as_before(monkeypatch):
assert dispatcher.cfg.classifier.mode == "local_llm"
built = []
monkeypatch.setattr(
dispatcher, "_classifier_client", lambda: built.append(1) or object()
)
monkeypatch.setattr(dispatcher, "_classify_once", lambda *a, **k: _answer())
got = dispatcher.classify("do a thing", None)
assert built == [1]
assert got.source == "classifier"
# --- mode: cloud_llm, pinned -----------------------------------------------
def test_cloud_llm_pinned_success_records_source_classifier_not_cloud(monkeypatch):
"""The critical distinction this feature introduces: intentional cloud
primary is NOT the same signal as the cascade's degraded cloud step."""
monkeypatch.setattr(dispatcher.cfg.classifier, "mode", "cloud_llm")
monkeypatch.setattr(dispatcher.cfg.classifier, "cloud_primary_auto", False)
monkeypatch.setattr(
dispatcher.cfg.classifier,
"cloud_primary",
SimpleNamespace(
base_url="https://cloud.example/v1",
model="cloud-classifier",
timeout_seconds=2,
api_key_env=None,
max_output_tokens=1024,
),
)
monkeypatch.setattr(
dispatcher, "_classifier_client", lambda: pytest.fail("local classifier dialled")
)
monkeypatch.setattr(dispatcher, "_cloud_classifier_client", lambda cf: object())
monkeypatch.setattr(
session_cache,
"classify_one",
lambda client, model, *a, **k: {
"task_category": "reasoning_math",
"task_tier": 3,
"required_context_tokens": 1234,
"confidence": 0.8,
}
if model == "cloud-classifier"
else None,
)
got = dispatcher.classify("do a thing", None)
assert got.source == "classifier", "primary cloud success must not read as degraded"
assert got.task_category == "reasoning_math"
def test_cloud_llm_pinned_missing_api_key_falls_through_to_cascade(monkeypatch):
monkeypatch.setattr(dispatcher.cfg.classifier, "mode", "cloud_llm")
monkeypatch.setattr(dispatcher.cfg.classifier, "cloud_primary_auto", False)
monkeypatch.setattr(
dispatcher.cfg.classifier,
"cloud_primary",
SimpleNamespace(
base_url="https://cloud.example/v1",
model="cloud-classifier",
timeout_seconds=2,
api_key_env="DEFINITELY_NOT_SET",
max_output_tokens=1024,
),
)
monkeypatch.setattr(dispatcher.cfg.classifier, "cloud_fallback", None)
got = dispatcher.classify("do a thing", None)
assert got.source == "fallback"
# --- mode: cloud_llm, auto_classifier --------------------------------------
def test_cloud_llm_auto_resolves_and_dials_the_cheapest_candidate(monkeypatch):
monkeypatch.setattr(dispatcher.cfg.classifier, "mode", "cloud_llm")
monkeypatch.setattr(dispatcher.cfg.classifier, "cloud_primary_auto", True)
monkeypatch.setattr(dispatcher.cfg.classifier, "cloud_primary", None)
resolved_row = {"model_id": "deepseek-v4-flash", "provider": "neuralwatt"}
monkeypatch.setattr(dispatcher, "_resolve_auto_classifier", lambda: resolved_row)
monkeypatch.setattr(
dispatcher.cfg,
"dispatch_providers",
{
"neuralwatt": SimpleNamespace(
base_url="https://api.neuralwatt.com/v1", api_key_env="NEURALWATT_API_KEY"
)
},
)
dialled = []
def fake_cloud_client(cf):
dialled.append(cf.model)
return object()
monkeypatch.setattr(dispatcher, "_cloud_classifier_client", fake_cloud_client)
monkeypatch.setattr(
session_cache,
"classify_one",
lambda *a, **k: {"task_category": "coding_general", "task_tier": 1},
)
got = dispatcher.classify("do a thing", None)
assert dialled == ["deepseek-v4-flash"]
assert got.source == "classifier"
def test_cloud_llm_auto_with_no_candidate_falls_through_to_cascade(monkeypatch):
monkeypatch.setattr(dispatcher.cfg.classifier, "mode", "cloud_llm")
monkeypatch.setattr(dispatcher.cfg.classifier, "cloud_primary_auto", True)
monkeypatch.setattr(dispatcher.cfg.classifier, "cloud_primary", None)
monkeypatch.setattr(dispatcher, "_resolve_auto_classifier", lambda: None)
monkeypatch.setattr(dispatcher.cfg.classifier, "cloud_fallback", None)
got = dispatcher.classify("do a thing", None)
assert got.source == "fallback"
def test_auto_classifier_resolution_is_cached_across_calls(monkeypatch):
"""A DB scan per request would sit on the latency floor for a value
that only moves as fast as the catalog's own poll cadence."""
monkeypatch.setattr(dispatcher.cfg.classifier, "mode", "cloud_llm")
monkeypatch.setattr(dispatcher.cfg.classifier, "cloud_primary_auto", True)
monkeypatch.setattr(dispatcher.cfg.classifier, "cloud_primary", None)
monkeypatch.setattr(dispatcher.cfg.classifier, "cooldown_seconds", 3600)
resolve_calls = []
def fake_resolve():
resolve_calls.append(1)
return {"model_id": "m", "provider": "neuralwatt"}
monkeypatch.setattr(dispatcher, "_resolve_auto_classifier", fake_resolve)
# Bypass the real caching internals by calling the cached wrapper twice
# through the public entry point instead of the raw resolver:
monkeypatch.setattr(
dispatcher.cfg,
"dispatch_providers",
{"neuralwatt": SimpleNamespace(base_url="https://x/v1", api_key_env=None)},
)
monkeypatch.setattr(dispatcher, "_cloud_classifier_client", lambda cf: object())
monkeypatch.setattr(
session_cache,
"classify_one",
lambda *a, **k: {"task_category": "coding_general", "task_tier": 1},
)
dispatcher.classify("a", None)
dispatcher.classify("b", None)
# _resolve_auto_classifier itself was monkeypatched above, so this only
# proves classify() calls the resolver each time -- the resolver's OWN
# caching is covered by test_resolve_auto_classifier_caches_within_the_cooldown_window.
assert resolve_calls == [1, 1]
def test_resolve_auto_classifier_caches_within_the_cooldown_window(monkeypatch):
monkeypatch.setattr(dispatcher, "_cached_auto_classifier", None)
monkeypatch.setattr(dispatcher, "_auto_classifier_resolved_at", 0.0)
monkeypatch.setattr(dispatcher.cfg.classifier, "cooldown_seconds", 3600)
monkeypatch.setattr(dispatcher.cfg.classifier, "fallback_category", "general_chat")
monkeypatch.setattr(dispatcher.cfg.routing, "allowed_access_levels", ["public"])
load_calls = []
def fake_load_candidates(conn, category):
load_calls.append(category)
return [{"model_id": "m", "provider": "neuralwatt"}]
monkeypatch.setattr(dispatcher, "load_candidates", fake_load_candidates)
monkeypatch.setattr(dispatcher, "cheapest_classifier_candidate", lambda rows, **k: rows[0])
monkeypatch.setattr(dispatcher, "_db", lambda: MagicMock())
first = dispatcher._resolve_auto_classifier()
second = dispatcher._resolve_auto_classifier()
assert first == second == {"model_id": "m", "provider": "neuralwatt"}
assert len(load_calls) == 1, "second call within the cooldown window re-queried the DB"
def test_resolve_auto_classifier_refreshes_after_the_cooldown_window(monkeypatch):
monkeypatch.setattr(dispatcher, "_cached_auto_classifier", {"model_id": "stale"})
monkeypatch.setattr(
dispatcher, "_auto_classifier_resolved_at", dispatcher.time.time() - 9999
)
monkeypatch.setattr(dispatcher.cfg.classifier, "cooldown_seconds", 30)
monkeypatch.setattr(dispatcher, "load_candidates", lambda conn, category: [])
monkeypatch.setattr(
dispatcher, "cheapest_classifier_candidate", lambda rows, **k: {"model_id": "fresh"}
)
monkeypatch.setattr(dispatcher, "_db", lambda: MagicMock())
got = dispatcher._resolve_auto_classifier()
assert got == {"model_id": "fresh"}
# --- mode: local_encoder ----------------------------------------------------
def test_local_encoder_success_records_source_classifier(monkeypatch):
monkeypatch.setattr(dispatcher.cfg.classifier, "mode", "local_encoder")
monkeypatch.setattr(
dispatcher.cfg.classifier,
"encoder",
SimpleNamespace(model="stub-model", device="cpu", confidence_threshold=0.5),
)
monkeypatch.setattr(
dispatcher, "_classifier_client", lambda: pytest.fail("local LLM dialled")
)
monkeypatch.setattr(
local_encoder,
"classify_zero_shot",
lambda task, categories, **k: ("coding_refactor", 0.83),
)
got = dispatcher.classify("refactor this function", None)
assert got.source == "classifier"
assert got.task_category == "coding_refactor"
assert got.confidence == pytest.approx(0.83)
def test_local_encoder_uses_fallback_tier_not_a_second_heuristic(monkeypatch):
"""A documented limitation, not a bug: the encoder produces a category,
not a tier."""
monkeypatch.setattr(dispatcher.cfg.classifier, "mode", "local_encoder")
monkeypatch.setattr(dispatcher.cfg.classifier, "fallback_tier", 2)
monkeypatch.setattr(
dispatcher.cfg.classifier,
"encoder",
SimpleNamespace(model="stub-model", device="cpu", confidence_threshold=0.5),
)
monkeypatch.setattr(
local_encoder, "classify_zero_shot", lambda task, categories, **k: ("general_chat", 0.9)
)
got = dispatcher.classify("hello", None)
assert got.task_tier == 2
def test_local_encoder_below_threshold_confidence_cascades(monkeypatch):
"""Treated identically to a local-LLM parse failure -- the cascade
reuses 100% of its existing machinery, unmodified."""
monkeypatch.setattr(dispatcher.cfg.classifier, "mode", "local_encoder")
monkeypatch.setattr(
dispatcher.cfg.classifier,
"encoder",
SimpleNamespace(model="stub-model", device="cpu", confidence_threshold=0.6),
)
monkeypatch.setattr(
local_encoder,
"classify_zero_shot",
lambda task, categories, **k: ("general_chat", 0.2),
)
monkeypatch.setattr(dispatcher.cfg.classifier, "cloud_fallback", None)
got = dispatcher.classify("do a thing", None)
assert got.source == "fallback", "a low-confidence guess must not be treated as real"
def test_local_encoder_below_threshold_uses_the_real_cascade(monkeypatch):
"""Not just 'degrades somehow' -- specifically walks _classify_cascade,
same as every other mode's failure."""
monkeypatch.setattr(dispatcher.cfg.classifier, "mode", "local_encoder")
monkeypatch.setattr(
dispatcher.cfg.classifier,
"encoder",
SimpleNamespace(model="stub-model", device="cpu", confidence_threshold=0.6),
)
monkeypatch.setattr(
local_encoder, "classify_zero_shot", lambda task, categories, **k: ("x", 0.1)
)
session_cache.put("sess-enc", task_category="coding_refactor", task_tier=3)
token = dispatcher._current_session_key.set("sess-enc")
try:
got = dispatcher.classify("do a thing", None)
finally:
dispatcher._current_session_key.reset(token)
assert got.source == "session_stale"
assert got.task_category == "coding_refactor"
# --- startup: local_encoder mode must fail loudly, not on the first request -
def test_startup_check_is_a_noop_for_other_modes(monkeypatch):
monkeypatch.setattr(dispatcher.cfg.classifier, "mode", "local_llm")
monkeypatch.setattr(
local_encoder,
"ensure_available",
lambda *a, **k: pytest.fail("ensure_available called for a non-encoder mode"),
)
dispatcher._ensure_classifier_mode_ready() # must not raise, must not call ensure_available
def test_startup_check_calls_ensure_available_for_local_encoder_mode(monkeypatch):
monkeypatch.setattr(dispatcher.cfg.classifier, "mode", "local_encoder")
monkeypatch.setattr(
dispatcher.cfg.classifier,
"encoder",
SimpleNamespace(model="stub-model", device="cpu", confidence_threshold=0.5),
)
calls = []
monkeypatch.setattr(
local_encoder, "ensure_available", lambda model, device: calls.append((model, device))
)
dispatcher._ensure_classifier_mode_ready()
assert calls == [("stub-model", "cpu")]
def test_startup_check_surfaces_a_missing_dependency_loudly(monkeypatch):
"""The whole point: fail at boot, not as an opaque error on the first
live classification request."""
monkeypatch.setattr(dispatcher.cfg.classifier, "mode", "local_encoder")
monkeypatch.setattr(
dispatcher.cfg.classifier,
"encoder",
SimpleNamespace(model="stub-model", device="cpu", confidence_threshold=0.5),
)
def fake_ensure_available(model, device):
raise ImportError(
"classifier.mode is 'local_encoder' but the 'transformers'/'torch' "
"packages are not installed. Install them with: "
"pip install -r requirements-encoder.txt"
)
monkeypatch.setattr(local_encoder, "ensure_available", fake_ensure_available)
with pytest.raises(ImportError, match="pip install -r requirements-encoder.txt"):
dispatcher._ensure_classifier_mode_ready()