388 lines
14 KiB
Python
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()
|