Files
6krrt/tests/test_local_dispatch_fallback.py
adlee-was-taken 7d9988caf3 test: make the HTTP fakes exercise the boundary instead of modelling it
The mojibake in 8518114 survived a green suite, and not by omission. There
is a test called `test_a_stream_is_proxied_verbatim_including_the_telemetry
_comments` which passed for months against a proxy that was decoding the
stream as ISO-8859-1 and re-encoding it as UTF-8. It asserted three ASCII
substrings, and a test that only ever sees ASCII cannot observe a charset
bug.

Two separate lies, both now fixed.

The fakes. `iter_lines` yielded `self._lines` unchanged in both modes, so a
test holding str lines got those str back even under decode_unicode -- a
model of a stream that had ALREADY been decoded correctly. A fake that
hands back the right answer cannot reproduce a decode bug. They now treat
the wire as UTF-8 bytes whatever the test wrote, derive `encoding` from
Content-Type using requests' own `get_encoding_from_headers`, and default
to the charset-less `text/event-stream` OpenRouter really sends. The
hostile case is the default, so a broken passthrough can no longer look
correct.

The assertions. The verbatim test now compares BYTES against every line the
provider emitted, and the shared STREAM_LINES sample carries a raw UTF-8 em
dash, so the whole streaming surface fails if the proxy ever goes back to
decoding and re-encoding.

Measured, by reverting 8518114 and running the suite:

    before this commit   1 test caught it (the one written for it)
    after                2, incl. the verbatim test that had been lying

The dedicated regression test also gets simpler -- it no longer needs a
bespoke response subclass, because the shared fake is now faithful enough.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01VRQXz5SYZYVWscxS1QqF6U
2026-09-08 14:05:16 -04:00

567 lines
20 KiB
Python

"""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"]