"""Tests for the local-dispatch fallback trigger helpers and the integration path in chat_completions (Todo 8). This file started as the Todo 7 helper tests (the _account_level_refusal and _local_dispatch_fallback_entry unit tests at the top) and grew the chat_completions integration tests that exercise the try_local_fallback closure end to end with a stubbed dispatch backend. Nothing here touches the network. The classifier, the cloud provider and the local dispatch backend are all stubbed, so the tests run in milliseconds and pin the fail-through contract rather than reachability. """ import json import sqlite3 from pathlib import Path import pytest from openai import APIStatusError from starlette.testclient import TestClient import dispatcher from dispatcher import ( Classification, _account_level_refusal, _local_dispatch_fallback_entry, ) from config import LocalDispatchModel ROOT = Path(__file__).resolve().parent.parent SCHEMA_SQL = (ROOT / "config" / "schema.sql").read_text() CHEAP = "cheap-model" DEAR = "dear-model" class FakeResponse: """Just enough of requests.Response for both dispatcher paths.""" def __init__( self, payload=None, *, status_code=200, lines=None, reason=None, headers=None, request=None, ): self.status_code = status_code self._payload = payload or {} self._lines = lines or [] self.reason = reason if reason is not None else "" self.text = json.dumps(self._payload) self.headers = headers or {} self.request = request self.closed = False def json(self): return self._payload @property def encoding(self): """The charset requests would pick, by requests' own rule. Derived rather than asserted, so this fake tracks requests instead of restating a belief about it. The default is a charset-less ``text/event-stream`` on purpose: that is what OpenRouter actually sends, and `get_encoding_from_headers` answers ISO-8859-1 for any ``text/*`` without a charset. Decoding a UTF-8 stream with that and re-encoding it produced mojibake on every streamed non-ASCII character -- so the hostile case is the DEFAULT here, and a fake can no longer make a broken passthrough look correct. """ from requests.structures import CaseInsensitiveDict from requests.utils import get_encoding_from_headers headers = CaseInsensitiveDict( getattr(self, "headers", None) or {"Content-Type": "text/event-stream"} ) return get_encoding_from_headers(headers) or "utf-8" def iter_lines(self, decode_unicode=False): # Faithful to requests in both modes: BYTES unless decode_unicode is # set, and when it is set the charset comes from Content-Type. # # The `wire` line is the load-bearing one. Whatever a test wrote into # `lines`, what a provider actually puts on the wire is UTF-8 bytes, # so that is what gets decoded. An earlier version of this fake passed # str lines straight through under decode_unicode, which modelled a # stream that had ALREADY been decoded correctly -- and a fake that # hands back the right answer cannot reproduce a charset bug. It let # the whole streaming suite pass against a proxy that was mangling # every non-ASCII character. for line in self._lines: wire = line.encode("utf-8") if isinstance(line, str) else line yield wire.decode(self.encoding) if decode_unicode else wire def close(self): self.closed = True def completion(model, content="hello there", *, completion_tokens=12): """A provider response shaped like NeuralWatt's, telemetry blocks included.""" return { "id": "chatcmpl-test-1", "model": model, "choices": [ {"message": {"role": "assistant", "content": content}, "finish_reason": "stop"} ], "usage": {"prompt_tokens": 31, "completion_tokens": completion_tokens}, "energy": {"energy_kwh": 5.0e-05, "avg_power_watts": 400.0, "duration_seconds": 1.4, "attribution_ratio": 0.25, "carbon_g_co2eq": 2.4e-03, "carbon_source": "agent_cache", "grid_id": "FI"}, "cost": {"request_cost_usd": 4.0e-04}, } # --- Todo 7 helper unit tests --------------------------------------------- @pytest.mark.parametrize( "status_code, expected", [ (400, False), (404, False), (422, False), (401, True), (402, True), (403, True), (405, True), (409, True), (429, True), (302, False), (500, False), (502, False), (503, False), ], ) def test_account_level_refusal(status_code, expected): assert _account_level_refusal(status_code) is expected def test_local_dispatch_fallback_entry_none_category(): assert _local_dispatch_fallback_entry(None) is None def test_local_dispatch_fallback_entry_returns_matching_entry(monkeypatch): entry = LocalDispatchModel( model_id="local-m", base_url="http://localhost:11434/v1", context_window=32768, tier=1, eligible_categories=["file_summarization"], ) monkeypatch.setattr(dispatcher.cfg, "local_dispatch_models", [entry]) assert _local_dispatch_fallback_entry("file_summarization") == entry assert _local_dispatch_fallback_entry("coding_general") is None def test_local_dispatch_fallback_entry_first_match(monkeypatch): first = LocalDispatchModel( model_id="first", base_url="http://localhost:11434/v1", context_window=32768, tier=1, eligible_categories=["file_summarization", "diff_checking"], ) second = LocalDispatchModel( model_id="second", base_url="http://localhost:11434/v1", context_window=32768, tier=1, eligible_categories=["file_summarization"], ) monkeypatch.setattr(dispatcher.cfg, "local_dispatch_models", [first, second]) assert _local_dispatch_fallback_entry("file_summarization") == first # --- Todo 8 integration fixtures ------------------------------------------- @pytest.fixture def fallback_router(tmp_path, monkeypatch): """Dispatcher with cloud + local models, classify stubbed to file_summarization.""" db_path = tmp_path / "test.db" conn = sqlite3.connect(db_path) conn.executescript(SCHEMA_SQL) for model_id, completion_price in ((CHEAP, 0.30), (DEAR, 9.00)): conn.execute( """INSERT INTO models (model_id, provider, base_model_id, tier, context_window, effective_context_window, max_output_tokens, cost_per_1m_prompt, cost_per_1m_completion, supports_vision, supports_json_mode, latency_class, reasoning_mode, context_variant, access_level, availability, last_updated) VALUES (?, 'neuralwatt', ?, 2, 262128, 192500, 16384, ?, ?, 1, 1, 'standard', 'default', 'full', 'public', 'active', '2026-08-22T00:00:00+00:00')""", (model_id, model_id, completion_price / 3, completion_price), ) # Seed proficiency so CHEAP is unambiguously the top-ranked cloud candidate # for file_summarization (routing joins proficiency; without rows the join # returns NULL and DEAR, at 0.5 neutral, could otherwise tie or win). conn.execute( """INSERT INTO proficiency (model_id, provider, category, blended_score, last_updated) VALUES (?, 'neuralwatt', 'file_summarization', 0.9, datetime('now'))""", (CHEAP,), ) conn.commit() conn.close() monkeypatch.setattr(dispatcher.cfg.database, "path", str(db_path)) monkeypatch.setattr(dispatcher.cfg.verification, "local_llm_enabled", False) monkeypatch.setattr(dispatcher.cfg.local_vision, "enabled", False) monkeypatch.setattr(dispatcher.cfg.session_cache, "enabled", False) monkeypatch.setattr(dispatcher.cfg.exploration, "epsilon", 0.0) monkeypatch.setattr(dispatcher.cfg.local_energy, "enabled", False) # The circuit breaker is module-global and records failures across tests; # without disabling it, one test's 402/500 opens the breaker and the next # test's routed candidates are all excluded, so routing picks nothing. monkeypatch.setattr(dispatcher.cfg.circuit_breaker, "enabled", False) monkeypatch.setattr(dispatcher.cfg, "local_dispatch_models", [ LocalDispatchModel( model_id="local-m", base_url="http://localhost:11434/v1", context_window=32768, tier=1, eligible_categories=["file_summarization"], ) ]) monkeypatch.setattr( dispatcher, "classify", lambda task, context=None: Classification( task_category="file_summarization", task_tier=1, required_context_tokens=100, confidence=0.9, ), ) monkeypatch.setenv("NEURALWATT_API_KEY", "test-key") calls = [] def fake_post(url, headers=None, json=None, stream=False, timeout=None): calls.append({"url": url, "body": json, "stream": stream}) if "localhost:11434" in url: return FakeResponse(completion("local-m")) # cloud URL — behavior under test varies; tests re-patch this return FakeResponse(completion(json["model"])) monkeypatch.setattr(dispatcher.requests, "post", fake_post) with TestClient(dispatcher.app) as client: yield client, calls, db_path def _decisions(db_path): conn = sqlite3.connect(db_path) conn.row_factory = sqlite3.Row rows = [dict(r) for r in conn.execute( "SELECT * FROM route_decisions ORDER BY id")] conn.close() return rows def _cloud_calls(calls): return [c for c in calls if "localhost:11434" not in c["url"]] # --- Todo 8 integration tests ---------------------------------------------- def _failing_post(monkeypatch, calls, cloud_fn, local_fn): """Re-point requests.post, recording every outbound call into the shared list.""" def fake_post(url, headers=None, json=None, stream=False, timeout=None): calls.append({"url": url, "body": json, "stream": stream}) if "localhost:11434" in url: return local_fn(url, json) return cloud_fn(url, json) monkeypatch.setattr(dispatcher.requests, "post", fake_post) def _cloud_402(): return FakeResponse( status_code=402, reason="Payment Required", payload={"error": {"message": "insufficient credits"}}, ) def test_account_level_refusal_falls_back_to_local_dispatch(fallback_router, monkeypatch): client, calls, db_path = fallback_router _failing_post( monkeypatch, calls, lambda url, body: _cloud_402(), lambda url, body: FakeResponse(completion("local-m")), ) resp = client.post( "/v1/chat/completions", json={"model": "auto", "messages": [{"role": "user", "content": "hi"}]}, ) assert resp.status_code == 200 assert resp.headers["X-Router-Model"] == "local-m" assert resp.json()["model"] == "local-m" # Only CHEAP was tried on the cloud — the 402 short-circuits DEAR. assert len(_cloud_calls(calls)) == 1 rows = _decisions(db_path) fallback_rows = [r for r in rows if r["kind"] == "local_dispatch_fallback"] assert len(fallback_rows) == 1 fb = fallback_rows[0] assert fb["selected_model"] == "local-m" assert fb["selected_provider"] == "ollama-local" assert "cloud_failed:%s:402" % CHEAP in fb["rejected_reason"] # The original 'chat' row is still there with the cloud selection. chat_rows = [r for r in rows if r["kind"] == "chat"] assert len(chat_rows) == 1 assert chat_rows[0]["selected_model"] == CHEAP def test_fallback_respects_eligible_categories(fallback_router, monkeypatch): client, calls, db_path = fallback_router monkeypatch.setattr( dispatcher, "classify", lambda task, context=None: Classification( task_category="coding_general", task_tier=2, required_context_tokens=100, confidence=0.9, ), ) _failing_post( monkeypatch, calls, lambda url, body: _cloud_402(), lambda url, body: FakeResponse(completion("local-m")), ) resp = client.post( "/v1/chat/completions", json={"model": "auto", "messages": [{"role": "user", "content": "hi"}]}, ) # cloud error surfaces loudly; no local call, no fallback row. assert resp.status_code == 402 assert not any("localhost:11434" in c["url"] for c in calls) rows = _decisions(db_path) assert not [r for r in rows if r["kind"] == "local_dispatch_fallback"] def test_fallback_when_all_cloud_candidates_exhausted(fallback_router, monkeypatch): client, calls, db_path = fallback_router # 500 is NOT account-level, so the loop walks both candidates. _failing_post( monkeypatch, calls, lambda url, body: FakeResponse(status_code=500, payload={"error": "boom"}), lambda url, body: FakeResponse(completion("local-m")), ) resp = client.post( "/v1/chat/completions", json={"model": "auto", "messages": [{"role": "user", "content": "hi"}]}, ) assert resp.status_code == 200 assert resp.headers["X-Router-Model"] == "local-m" assert len(_cloud_calls(calls)) == 2 rows = _decisions(db_path) fallback_rows = [r for r in rows if r["kind"] == "local_dispatch_fallback"] assert len(fallback_rows) == 1 assert ":500" in fallback_rows[0]["rejected_reason"] def test_request_attributable_4xx_walks_candidates_not_fallback(fallback_router, monkeypatch): client, calls, db_path = fallback_router def cloud_fn(url, body): if (body or {}).get("model") == CHEAP: return FakeResponse(status_code=400, payload={"error": "bad request"}) return FakeResponse(completion((body or {}).get("model"))) _failing_post( monkeypatch, calls, cloud_fn, lambda url, body: FakeResponse(completion("local-m")), ) resp = client.post( "/v1/chat/completions", json={"model": "auto", "messages": [{"role": "user", "content": "hi"}]}, ) # 400 is request-attributable → walk to DEAR, never local. assert resp.status_code == 200 assert resp.headers["X-Router-Model"] == DEAR assert not any("localhost:11434" in c["url"] for c in calls) rows = _decisions(db_path) assert not [r for r in rows if r["kind"] == "local_dispatch_fallback"] def test_local_failure_falls_through_to_original_cloud_error(fallback_router, monkeypatch, caplog): client, calls, db_path = fallback_router _failing_post( monkeypatch, calls, lambda url, body: _cloud_402(), lambda url, body: FakeResponse(status_code=500, payload={"error": "local boom"}), ) resp = client.post( "/v1/chat/completions", json={"model": "auto", "messages": [{"role": "user", "content": "hi"}]}, ) # The ORIGINAL cloud error surfaces, not the local failure. assert resp.status_code == 402 rows = _decisions(db_path) assert not [r for r in rows if r["kind"] == "local_dispatch_fallback"] assert any( rec.getMessage().startswith("local_dispatch_fallback_skip") for rec in caplog.records ) def test_pinned_model_does_not_fall_back(fallback_router, monkeypatch): client, calls, db_path = fallback_router _failing_post( monkeypatch, calls, lambda url, body: _cloud_402(), lambda url, body: FakeResponse(completion("local-m")), ) resp = client.post( "/v1/chat/completions", json={"model": DEAR, "messages": [{"role": "user", "content": "hi"}]}, ) assert resp.status_code == 402 assert not any("localhost:11434" in c["url"] for c in calls) rows = _decisions(db_path) assert not [r for r in rows if r["kind"] == "local_dispatch_fallback"] class _Headers: def get(self, key, default=None): return default class _SDKFailRequest: url = "https://neuralwatt.example/v1/chat/completions" method = "POST" headers = _Headers() content = b"" class _SDKFailResponse: status_code = 402 reason_phrase = "Payment Required" text = "insufficient credits" request = _SDKFailRequest() headers = _Headers() def json(self): return {"error": {"message": "insufficient credits", "type": "insufficient_quota", "code": "insufficient_quota"}} def test_dispatch_endpoint_does_not_fall_back(fallback_router, monkeypatch): client, calls, db_path = fallback_router def failing_create(**kwargs): raise APIStatusError( "402 Payment Required", response=_SDKFailResponse(), body=None ) class _Raw: def create(self, **kwargs): return failing_create(**kwargs) class _Completions: @property def with_raw_response(self): return _Raw() class _Chat: completions = _Completions() def fake_provider_client(provider): return type("_Client", (), {"chat": _Chat()})() monkeypatch.setattr(dispatcher, "_provider_client", fake_provider_client) resp = client.post( "/dispatch", json={ "task": "summarize this", "task_category": "file_summarization", "task_tier": 1, "required_context_tokens": 100, }, ) # The /dispatch error path is 502, NOT a fallback to local. assert resp.status_code == 502 assert not any("localhost:11434" in c["url"] for c in calls) rows = _decisions(db_path) assert not [r for r in rows if r["kind"] == "local_dispatch_fallback"] # --- Todo 9: streaming path ------------------------------------------------ def test_streaming_account_level_refusal_falls_back_to_local(fallback_router, monkeypatch): """A streamed request that hits an account-level cloud refusal (402) on the first candidate must short-circuit the streaming failover loop and degrade to the local dispatch SSE stream rather than raising the cloud error. """ client, calls, db_path = fallback_router _failing_post( monkeypatch, calls, lambda url, body: _cloud_402(), lambda url, body: FakeResponse(completion("local-m")), ) resp = client.post( "/v1/chat/completions", json={ "model": "auto", "stream": True, "messages": [{"role": "user", "content": "hi"}], }, ) assert resp.status_code == 200 assert resp.headers["content-type"].startswith("text/event-stream") # Parse the bounded SSE frames from the local fallback generator. frames = [] for frame in resp.text.split("\n\n"): frame = frame.strip() if not frame: continue _, _, data = frame.partition("data: ") if data == "[DONE]": continue frames.append(json.loads(data)) assert frames, "expected at least one SSE data frame" # The first chunk carries the assistant role + local model. first = frames[0] assert first["model"] == "local-m" assert first["choices"][0]["delta"]["role"] == "assistant" assert first["choices"][0]["delta"]["content"] == "hello there" # 402 is account-level → the streaming loop short-circuits, so only CHEAP # was ever tried on the cloud (DEAR never attempted). assert len(_cloud_calls(calls)) == 1 rows = _decisions(db_path) fallback_rows = [r for r in rows if r["kind"] == "local_dispatch_fallback"] assert len(fallback_rows) == 1 fb = fallback_rows[0] assert fb["selected_model"] == "local-m" assert fb["selected_provider"] == "ollama-local" assert fb["streamed"] == 1 assert "cloud_failed:%s:402" % CHEAP in fb["rejected_reason"]