Files
6krrt/tests/test_chat_completions.py
2026-09-28 22:57:28 -04:00

1945 lines
71 KiB
Python
Raw Permalink Blame History

"""Tests for the OpenAI-compatible surface — the endpoint every client uses.
This file exists because it did not. `/v1/chat/completions` is how opencode,
an SDK and plain curl all reach the router, and it had zero tests, which is
how a NameError survived on the pass-through path: `alternatives` read
`decision.runners_up`, but `decision` is only bound inside `if wants_routing`,
so every non-streaming request naming a real model id returned 500 before the
provider was ever called.
Nothing here touches the network. The classifier, the provider and the local
verifier are all stubbed, so the tests run in milliseconds and pin behaviour
rather than reachability.
"""
import io
import json
import sqlite3
from pathlib import Path
from unittest.mock import MagicMock
import pytest
from starlette.testclient import TestClient
import dispatcher
import session_cache
from config import RoutingProfile
from dispatcher import Classification, app
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},
}
@pytest.fixture
def router(tmp_path, monkeypatch):
"""The dispatcher pointed at a throwaway catalog, with nothing dialled out."""
db_path = tmp_path / "test.db"
conn = sqlite3.connect(db_path)
conn.executescript(SCHEMA_SQL)
for model_id, completion_price, vision in (
(CHEAP, 0.30, 1),
(DEAR, 9.00, 0),
):
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, 'standard', 'default', 'full', 'public', 'active',
'2026-08-22T00:00:00+00:00')
""",
(model_id, model_id, completion_price / 3, completion_price, vision),
)
conn.commit()
conn.close()
monkeypatch.setattr(dispatcher.cfg.database, "path", str(db_path))
# No Ollama in the test environment, and the local check would otherwise
# fire as a background task and try to reach localhost:11434.
monkeypatch.setattr(dispatcher.cfg.verification, "local_llm_enabled", False)
# The local vision fallback is off unless a test opts in; without this it
# would fire for every image request and try to reach localhost.
monkeypatch.setattr(dispatcher.cfg.local_vision, "enabled", False)
# Isolate the shared module-level session cache: a real cache keyed on the
# same fingerprint leaks a classification from one test into the next.
monkeypatch.setattr(dispatcher.cfg.session_cache, "enabled", False)
# Exploration is on in production, but this fixture asserts deterministic
# tier-2 winners; pin epsilon to 0 so exploration never fires.
monkeypatch.setattr(dispatcher.cfg.exploration, "epsilon", 0.0)
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 stream:
return FakeResponse(lines=STREAM_LINES)
return FakeResponse(completion(json["model"]))
monkeypatch.setattr(dispatcher.requests, "post", fake_post)
monkeypatch.setattr(
dispatcher, "classify",
lambda task, context: Classification(
task_category="coding_general", task_tier=2,
required_context_tokens=100, confidence=0.9,
),
)
yield TestClient(app), calls, db_path
# The content carries a raw UTF-8 em dash on purpose. An OpenAI-compatible
# SSE stream sends non-ASCII unescaped, and this sample is what every
# streaming test in this file exercises -- so the whole streaming surface now
# fails if the proxy ever goes back to decoding and re-encoding the stream
# rather than forwarding the provider's bytes. All-ASCII fixtures are how a
# charset bug survived a test literally named "proxied verbatim".
STREAM_LINES = [
'data: {"id":"chatcmpl-stream-1","choices":[{"delta":{"content":"hel—"}}]}',
"",
'data: {"id":"chatcmpl-stream-1","choices":[{"delta":{"content":"lo"},'
'"finish_reason":"stop"}],"usage":{"prompt_tokens":31,'
'"completion_tokens":9}}',
"",
': energy {"energy_kwh": 5e-05, "carbon_g_co2eq": 0.0024, '
'"carbon_source": "agent_cache"}',
': cost {"request_cost_usd": 0.0004}',
"data: [DONE]",
"",
]
def _messages(text="write me a function"):
return [{"role": "user", "content": text}]
# --- the regression -------------------------------------------------------
def test_a_real_model_id_is_dispatched_rather_than_routed(router):
"""The documented pass-through: 'any real model id — dispatched as asked'.
This raised NameError on `decision` before the provider was ever called,
so the whole path 500'd while streaming clients never noticed.
"""
client, calls, _ = router
resp = client.post(
"/v1/chat/completions",
json={"model": DEAR, "messages": _messages()},
)
assert resp.status_code == 200
assert calls[0]["body"]["model"] == DEAR, "the caller's choice must survive"
assert resp.headers["X-Router-Model"] == DEAR
def test_a_real_model_id_does_not_pay_for_classification(router, monkeypatch):
"""A named model has nothing to classify; the ~2s round-trip is skipped."""
client, _, _ = router
def explode(task, context):
raise AssertionError("classifier consulted for an explicit model id")
monkeypatch.setattr(dispatcher, "classify", explode)
assert client.post(
"/v1/chat/completions", json={"model": CHEAP, "messages": _messages()}
).status_code == 200
# --- routing --------------------------------------------------------------
def test_auto_routes_and_reports_the_model_it_actually_used(router):
"""Quality ties, so cost breaks it — and the client is told what ran."""
client, calls, _ = router
resp = client.post(
"/v1/chat/completions",
json={"model": "auto", "messages": _messages()},
)
assert resp.status_code == 200
assert calls[0]["body"]["model"] == CHEAP
assert resp.headers["X-Router-Model"] == CHEAP
# The body must stay a valid OpenAI response naming the real model, not 'auto'.
assert resp.json()["model"] == CHEAP
def test_auto_locality_profile_routes_to_locality_provider(router):
client, _, _ = router
resp = client.post(
"/v1/chat/completions",
json={"model": "auto:locality", "messages": _messages()},
)
assert resp.status_code == 422, "no ollama-local rows in fixture"
def test_routed_chat_named_profile_persists_profile(router):
client, calls, _ = router
resp = client.post(
"/v1/chat/completions",
json={"model": "auto:batch", "messages": _messages()},
)
assert resp.status_code == 200
assert calls[0]["body"]["model"] == CHEAP
db_path = dispatcher.cfg.database.path
conn = sqlite3.connect(db_path)
conn.row_factory = sqlite3.Row
rows = conn.execute(
"SELECT * FROM route_decisions ORDER BY id"
).fetchall()
conn.close()
assert len(rows) == 1
assert rows[0]["profile"] == "batch"
assert rows[0]["kind"] == "chat"
def test_a_provider_prefixed_router_name_still_routes(router):
"""opencode sends `llm-router/auto`; only the virtual names are stripped."""
client, calls, _ = router
resp = client.post(
"/v1/chat/completions",
json={"model": "llm-router/auto", "messages": _messages()},
)
assert resp.status_code == 200
assert calls[0]["body"]["model"] == CHEAP
def test_no_eligible_model_is_a_422_naming_the_filters(router, monkeypatch):
"""A dead end must say which constraint killed it, not just 'no model'."""
client, calls, _ = router
monkeypatch.setattr(
dispatcher, "classify",
lambda task, context: Classification(
task_category="coding_general", task_tier=3,
required_context_tokens=100, confidence=0.9,
),
)
resp = client.post(
"/v1/chat/completions", json={"model": "auto", "messages": _messages()}
)
assert resp.status_code == 422
assert "tier >= 3" in resp.json()["detail"]
assert not calls, "nothing should be dispatched when nothing qualifies"
# --- streaming ------------------------------------------------------------
def test_a_stream_is_proxied_verbatim_including_the_telemetry_comments(router):
"""NeuralWatt's energy/cost blocks are SSE comments; clients ignore them.
They must still reach the client untouched — the router reads them on the
way past rather than buffering the stream to strip them.
"Verbatim" is asserted on BYTES, against every line the provider emitted.
The previous version of this test checked three ASCII substrings, and it
passed for months against a proxy that was decoding the stream as
ISO-8859-1 and re-encoding it as UTF-8 -- i.e. against a proxy that was
not verbatim at all. A test that only ever sees ASCII cannot observe a
charset bug, so the sample below carries non-ASCII and the assertion is
equality rather than containment.
"""
client, calls, _ = router
resp = client.post(
"/v1/chat/completions",
json={"model": DEAR, "messages": _messages(), "stream": True},
)
assert resp.status_code == 200
assert calls[0]["stream"] is True
body = resp.content
for line in STREAM_LINES:
expected = line.encode("utf-8") if isinstance(line, str) else line
assert expected in body, (expected, body)
assert body.decode("utf-8") == "".join(
(line if isinstance(line, str) else line.decode("utf-8")) + "\n"
for line in STREAM_LINES
)
def test_a_stream_does_not_re_encode_non_ascii(router, monkeypatch):
"""Non-ASCII must reach the client byte-for-byte as the provider sent it.
This is a regression test for real mojibake, not a hypothetical. The
proxy used to read the stream with ``iter_lines(decode_unicode=True)``
and write it back with ``f"{raw}\\n".encode()``. ``decode_unicode`` uses
``r.encoding``, which requests derives from Content-Type, and
``get_encoding_from_headers`` falls back to ISO-8859-1 for any ``text/*``
without a charset. OpenRouter returns exactly that -- verified live:
``Content-Type: text/event-stream`` with no charset, giving
``r.encoding == 'ISO-8859-1'``.
So an em dash left the provider as UTF-8, was decoded as latin-1 into
'a\\x80\\x94', and was re-encoded as UTF-8 on the way out: a double
encode. Every non-ASCII character in a streamed completion reached the
client mangled. The buffered path was unaffected because ``resp.json()``
decodes as UTF-8, which is why this only ever showed up in streaming --
and every agent client streams.
``FakeResponse`` reproduces the failure exactly: it holds UTF-8 body
bytes and derives its charset with requests' own
``get_encoding_from_headers``, defaulting to the charset-less
``text/event-stream`` OpenRouter really sends. So this test tracks
requests' behaviour rather than restating a belief about it.
"""
payload = "an em dash — and e-acute é"
# ensure_ascii=False on purpose: an OpenAI-compatible SSE stream carries
# raw UTF-8 in the delta rather than \uXXXX escapes, and it is only the
# raw form that the latin-1 decode could corrupt. Escaped, the body is
# pure ASCII and the test would pass against the bug.
sse = [
('data: {"id":"chatcmpl-utf8","choices":[{"delta":{"content":'
+ json.dumps(payload, ensure_ascii=False) + '}}]}').encode("utf-8"),
b"",
('data: {"id":"chatcmpl-utf8","choices":[{"delta":{},'
'"finish_reason":"stop"}],"usage":{"prompt_tokens":31,'
'"completion_tokens":9}}').encode("utf-8"),
b"",
b"data: [DONE]",
]
# The precondition the whole bug depended on. If requests ever starts
# defaulting a charset-less SSE response to UTF-8, this fails loudly
# rather than the test quietly ceasing to cover anything.
assert FakeResponse(lines=sse).encoding == "ISO-8859-1"
def fake_post(url, headers=None, json=None, stream=False, timeout=None):
if stream:
return FakeResponse(lines=sse)
return FakeResponse(completion(json["model"]))
monkeypatch.setattr(dispatcher.requests, "post", fake_post)
client, _, _ = router
resp = client.post(
"/v1/chat/completions",
json={"model": DEAR, "messages": _messages(), "stream": True},
)
assert resp.status_code == 200
body = resp.content
assert payload.encode("utf-8") in body, body
# The specific corruption, named so a regression is unambiguous: an em
# dash double-encoded is c3 a2 c2 80 c2 94.
assert b"\xc3\xa2\xc2\x80\xc2\x94" not in body
assert "<EFBFBD>" not in body.decode("utf-8")
def test_a_streamed_call_still_logs_its_energy(router):
"""Streaming is how every agent client talks; unlogged, it is most of the traffic."""
client, _, db_path = router
client.post(
"/v1/chat/completions",
json={"model": DEAR, "messages": _messages(), "stream": True},
)
conn = sqlite3.connect(db_path)
conn.row_factory = sqlite3.Row
row = conn.execute(
"SELECT model_id, request_id, energy_kwh, cost_usd, completion_tokens "
"FROM energy_observations ORDER BY id DESC LIMIT 1"
).fetchone()
conn.close()
assert row["model_id"] == DEAR
assert row["request_id"] == "chatcmpl-stream-1"
assert row["energy_kwh"] == pytest.approx(5e-05)
assert row["cost_usd"] == pytest.approx(4e-04)
assert row["completion_tokens"] == 9
def test_a_dropped_upstream_connection_ends_the_stream_cleanly(router, monkeypatch):
"""The upstream connection can die mid-stream (NeuralWatt closing early, a
network blip). Left uncaught, requests.exceptions.ChunkedEncodingError
propagated straight out of the generator, which Starlette surfaced as an
unhandled ASGI exception -- a full traceback in the log, and the client's
connection cut dead with no error payload or [DONE].
"""
import requests as requests_module
client, _, _ = router
def broken_lines():
yield 'data: {"id":"chatcmpl-stream-1","choices":[{"delta":{"content":"hel"}}]}'
raise requests_module.exceptions.ChunkedEncodingError("Response ended prematurely")
def fake_post(url, headers=None, json=None, stream=False, timeout=None):
return FakeResponse(lines=broken_lines())
monkeypatch.setattr(dispatcher.requests, "post", fake_post)
resp = client.post(
"/v1/chat/completions",
json={"model": DEAR, "messages": _messages(), "stream": True},
)
assert resp.status_code == 200
body = resp.text
assert "hel" in body
assert "upstream stream interrupted" in body
assert "data: [DONE]" in body
# --- what a request leaves in the journal ---------------------------------
#
# The reason this exists: watching a live session showed only uvicorn's access
# line, which does not even name the model that served it.
@pytest.fixture
def logbuf():
"""Capture what would reach the journal, through the real formatter."""
import logs
buf = io.StringIO()
logs.configure("debug", stream=buf, journald=False)
yield buf
for handler in list(logs.log.handlers):
logs.log.removeHandler(handler)
def _lines(logbuf, event):
return [ln for ln in logbuf.getvalue().splitlines() if ln.startswith(f"{event} ")]
def _fields(line):
out = {}
for token in line.split(" ")[1:]:
if "=" in token:
key, _, value = token.partition("=")
out[key] = value
return out
def test_a_routed_request_logs_exactly_one_decision_line(router, logbuf):
client, _, _ = router
client.post("/v1/chat/completions",
json={"model": "auto", "messages": _messages()})
routes = _lines(logbuf, "route")
assert len(routes) == 1, "routing twice for context must not log twice"
fields = _fields(routes[0])
assert fields["pick"] == CHEAP
assert fields["cat"] == "coding_general"
assert fields["tier"] == "2"
assert int(fields["ms"]) >= 0
def test_the_dispatch_line_carries_the_join_key(router, logbuf):
"""rid is what pivots a journal line to its energy_observations row."""
client, _, _ = router
client.post("/v1/chat/completions",
json={"model": "auto", "messages": _messages()})
fields = _fields(_lines(logbuf, "dispatch")[0])
assert fields["rid"] == "chatcmpl-test-1"
assert fields["model"] == CHEAP
assert fields["c_tok"] == "12"
def test_every_line_of_one_request_shares_a_trace_id(router, logbuf):
client, _, _ = router
client.post("/v1/chat/completions",
json={"model": "auto", "messages": _messages()})
ids = {_fields(ln)["id"] for ln in logbuf.getvalue().splitlines()}
assert len(ids) == 1
assert ids != {"-"}
def test_a_streamed_request_still_has_a_trace_id(router, logbuf):
"""The trap: a StreamingResponse generator runs in a different context.
Read from the ContextVar inside the generator it comes back empty, and
every streamed request -- which is all agent traffic -- logs a blank id.
"""
client, _, _ = router
client.post("/v1/chat/completions",
json={"model": DEAR, "messages": _messages(), "stream": True})
dispatched = _fields(_lines(logbuf, "dispatch")[0])
assert dispatched["id"] != "-"
assert dispatched["stream"] == "1"
assert dispatched["rid"] == "chatcmpl-stream-1"
def test_debug_says_which_filter_dropped_a_model(router, logbuf):
"""'No model satisfies the hard filters' is otherwise a dead end."""
client, _, db_path = router
conn = sqlite3.connect(db_path)
conn.execute(
"""
INSERT INTO models (
model_id, provider, tier, context_window, effective_context_window,
latency_class, access_level, availability, last_updated
) VALUES ('held-model', 'neuralwatt', 2, 262128, 192500, 'flex',
'public', 'active', '2026-08-22T00:00:00+00:00')
"""
)
conn.commit()
conn.close()
client.post("/v1/chat/completions",
json={"model": "auto", "messages": _messages()})
dropped = {_fields(ln)["model"]: _fields(ln)["reason"]
for ln in _lines(logbuf, "filter")}
assert dropped["held-model"] == "latency_class(flex)"
def test_a_pass_through_is_labelled_as_one(router, logbuf):
client, _, _ = router
client.post("/v1/chat/completions",
json={"model": DEAR, "messages": _messages()})
assert not _lines(logbuf, "route"), "nothing was routed"
assert _fields(_lines(logbuf, "passthrough")[0])["model"] == DEAR
def test_no_conversation_text_ever_reaches_the_log(router, logbuf):
"""Prompts here run 60k-150k tokens and the journal is on disk.
Asserted at DEBUG, the most verbose level, on both the request and the
answer.
"""
secret = "SQUAMOUS-EPHEMERAL-9137"
client, _, _ = router
client.post("/v1/chat/completions",
json={"model": "auto", "messages": _messages(f"refactor {secret}")})
assert secret not in logbuf.getvalue()
# The stubbed provider answers "hello there"; that must not appear either.
assert "hello there" not in logbuf.getvalue()
def test_the_decision_line_reports_the_real_provenance(router, logbuf):
"""The re-route for measured context must not read as a client override.
chat_completions routes a second time when the conversation measures larger
than the classifier guessed, passing category and tier back in — which
makes the resulting Classification say source='override'. In a log line
that means "the client chose this", and the client chose nothing: the
classifier ran and only the context figure was replaced.
"""
client, _, _ = router
# ~1200 chars / 3 = 400 tokens measured, above the stubbed estimate of 100.
client.post("/v1/chat/completions",
json={"model": "auto", "messages": _messages("refactor " + "x " * 600)})
fields = _fields(_lines(logbuf, "route")[0])
assert fields["src"] == "classifier", "the classifier did decide this"
assert fields["ctx_src"] == "measured", "but the context came from measurement"
assert int(fields["ctx"]) > 100
def test_a_caller_supplied_context_says_so(router, logbuf):
client, _, _ = router
client.post("/route", json={"task": "refactor this",
"task_category": "coding_general",
"task_tier": 2,
"required_context_tokens": 5000})
fields = _fields(_lines(logbuf, "route")[0])
assert fields["src"] == "override"
assert fields["ctx_src"] == "caller"
# --- capability gates: vision and JSON mode -------------------------------
def _image_messages(text="what is in this image?"):
"""A user turn carrying an inline image_url part, OpenAI multimodal shape."""
return [
{
"role": "user",
"content": [
{"type": "text", "text": text},
{
"type": "image_url",
"image_url": {"url": "data:image/png;base64,AAAA"},
},
],
}
]
def _drop_cheap_by_tier(client_type, db_path):
"""Make CHEAP ineligible so DEAR is the only remaining candidate."""
conn = sqlite3.connect(db_path)
conn.execute("UPDATE models SET tier = 1 WHERE model_id = ?", (CHEAP,))
conn.commit()
conn.close()
def _local_vision_fake(monkeypatch, router_calls, content="local caption", status=200):
"""Point the local vision fallback at a fake that returns `content`."""
def fake_post(url, headers=None, json=None, stream=False, timeout=None):
if "/chat/completions" in url:
router_calls.append({"url": url, "body": json, "stream": stream, "local": True})
return FakeResponse(
{
"choices": [
{"message": {"role": "assistant", "content": content},
"finish_reason": "stop"}
]
},
status_code=status,
)
router_calls.append({"url": url, "body": json, "stream": stream})
return FakeResponse(completion(json["model"]))
monkeypatch.setattr(dispatcher.requests, "post", fake_post)
def test_an_image_request_is_routed_to_a_vision_model(router):
"""Image parts hard-restrict to vision-capable rows; CHEAP has vision."""
client, calls, _ = router
resp = client.post(
"/v1/chat/completions",
json={"model": "auto", "messages": _image_messages()},
)
assert resp.status_code == 200
assert calls[0]["body"]["model"] == CHEAP
assert resp.headers["X-Router-Model"] == CHEAP
def test_an_image_request_excludes_a_non_vision_candidate(router):
"""DEAR is eligible-tiered but lacks vision, so no candidate meets the ask."""
client, calls, db_path = router
_drop_cheap_by_tier(dispatcher, db_path)
resp = client.post(
"/v1/chat/completions",
json={"model": "auto", "messages": _image_messages()},
)
assert resp.status_code == 422
assert "vision" in resp.json()["detail"]
assert not calls, "a non-vision model must never be dispatched for an image"
def test_no_vision_model_uses_local_fallback_when_enabled(router, monkeypatch):
"""No cloud vision candidate, local vision on: the local model answers."""
client, calls, db_path = router
_drop_cheap_by_tier(dispatcher, db_path)
monkeypatch.setattr(dispatcher.cfg.local_vision, "enabled", True)
_local_vision_fake(monkeypatch, calls, content="local caption")
resp = client.post(
"/v1/chat/completions",
json={"model": "auto", "messages": _image_messages()},
)
assert resp.status_code == 200
assert resp.json()["choices"][0]["message"]["content"] == "local caption"
local = [c for c in calls if c["url"].endswith("/chat/completions")]
assert local, "the local fallback should POST to its own /chat/completions"
parts = local[0]["body"]["messages"][0]["content"]
assert any(
isinstance(p, dict) and p.get("type") == "image_url" for p in parts
), "the image_url part must survive into the local call"
def test_no_vision_model_with_local_fallback_off_is_a_422(router, monkeypatch):
"""local_vision.enabled off: no cloud vision candidate 422s, no local call."""
client, calls, db_path = router
_drop_cheap_by_tier(dispatcher, db_path)
monkeypatch.setattr(dispatcher.cfg.local_vision, "enabled", False)
_local_vision_fake(monkeypatch, calls, content="must not be used")
resp = client.post(
"/v1/chat/completions",
json={"model": "auto", "messages": _image_messages()},
)
assert resp.status_code == 422
assert "vision" in resp.json()["detail"]
local = [c for c in calls if c.get("local")]
assert not local, "the local vision fallback must not run when disabled"
def test_local_fallback_skips_when_json_mode_is_requested(router, monkeypatch):
"""Local vision answers in prose; a json_object request must not use it.
JSON mode -- not vision -- is what excluded every cloud candidate here, so
the fallback must not fire or it would silently break the json_object
contract with free-form text instead of the 422 naming the missing
capability.
"""
client, calls, db_path = router
_drop_cheap_by_tier(dispatcher, db_path)
monkeypatch.setattr(dispatcher.cfg.local_vision, "enabled", True)
_local_vision_fake(monkeypatch, calls, content="should not be used")
resp = client.post(
"/v1/chat/completions",
json={
"model": "auto",
"messages": _image_messages(),
"response_format": {"type": "json_object"},
},
)
assert resp.status_code == 422
assert "json" in resp.json()["detail"]
local = [c for c in calls if c.get("local")]
assert not local, "the local vision fallback must not run for a json_object request"
def test_local_fallback_respects_streaming(router, monkeypatch):
client, calls, db_path = router
_drop_cheap_by_tier(dispatcher, db_path)
monkeypatch.setattr(dispatcher.cfg.local_vision, "enabled", True)
_local_vision_fake(monkeypatch, calls, content="streamed caption")
resp = client.post(
"/v1/chat/completions",
json={"model": "auto", "messages": _image_messages(), "stream": True},
)
assert resp.status_code == 200
assert "streamed caption" in resp.text
assert "data: [DONE]" in resp.text
def test_local_fallback_failure_still_422s(router, monkeypatch):
"""A local vision failure must be visible, not swallowed into a 200."""
client, calls, db_path = router
_drop_cheap_by_tier(dispatcher, db_path)
monkeypatch.setattr(dispatcher.cfg.local_vision, "enabled", True)
_local_vision_fake(monkeypatch, calls, status=500)
resp = client.post(
"/v1/chat/completions",
json={"model": "auto", "messages": _image_messages()},
)
assert resp.status_code == 422
assert "vision" in resp.json()["detail"]
def test_a_json_object_response_format_routes_to_a_json_capable_model(router):
"""response_format json_object hard-restricts to JSON-mode-capable rows."""
client, calls, _ = router
resp = client.post(
"/v1/chat/completions",
json={
"model": "auto",
"messages": _messages(),
"response_format": {"type": "json_object"},
},
)
assert resp.status_code == 200
assert calls[0]["body"]["model"] == CHEAP
def test_a_json_object_request_without_a_json_capable_model_422s(router):
client, calls, db_path = router
conn = sqlite3.connect(db_path)
conn.execute("UPDATE models SET supports_json_mode = 0")
conn.commit()
conn.close()
resp = client.post(
"/v1/chat/completions",
json={
"model": "auto",
"messages": _messages(),
"response_format": {"type": "json_object"},
},
)
assert resp.status_code == 422
assert "json" in resp.json()["detail"]
assert not calls
def test_auto_batch_sets_latency_tolerance_to_batch(router):
"""auto:batch must resolve to the built-in batch profile."""
client, calls, _ = router
resp = client.post(
"/v1/chat/completions",
json={"model": "auto:batch", "messages": _messages()},
)
assert resp.status_code == 200
assert resp.headers["X-Router-Model"] == CHEAP
def test_unknown_profile_returns_422_naming_valid_profiles(router):
"""A misspelled profile must not silently fall back to default."""
client, calls, _ = router
resp = client.post(
"/v1/chat/completions",
json={"model": "auto:nosuchprofile", "messages": _messages()},
)
assert resp.status_code == 422
detail = resp.json()["detail"]
assert "nosuchprofile" in detail
assert "default" in detail
assert "batch" in detail
assert not calls
def test_auto_routes_through_configured_default_profile(router, monkeypatch):
"""Bare ``auto`` uses routing.default_profile, not the hard-coded default.
CHEAP wins under both interactive and batch tolerance (every fixture row
is latency_class='standard'), so the model pick alone cannot prove the
batch-local profile ran; the persisted decision row carries the actual
latency_tolerance='batch' and profile='batch-local' under a bare auto.
"""
client, calls, db_path = router
monkeypatch.setattr(dispatcher.cfg.routing, "default_profile", "batch-local")
monkeypatch.setitem(
dispatcher.cfg.profiles,
"batch-local",
RoutingProfile(latency_tolerance="batch"),
)
resp = client.post(
"/v1/chat/completions",
json={"model": "auto", "messages": _messages()},
)
assert resp.status_code == 200
assert resp.headers["X-Router-Model"] == CHEAP
assert calls[0]["body"]["model"] == CHEAP
conn = sqlite3.connect(db_path)
conn.row_factory = sqlite3.Row
row = conn.execute(
"SELECT profile, latency_tolerance FROM route_decisions "
"ORDER BY id DESC LIMIT 1"
).fetchone()
conn.close()
assert row is not None
assert row["latency_tolerance"] == "batch"
assert row["profile"] == "batch-local"
def test_auto_named_profile_overrides_configured_default(router, monkeypatch):
"""`auto:<name>` still overrides routing.default_profile."""
client, calls, _ = router
monkeypatch.setattr(dispatcher.cfg.routing, "default_profile", "batch-local")
monkeypatch.setitem(
dispatcher.cfg.profiles,
"batch-local",
RoutingProfile(latency_tolerance="batch"),
)
resp = client.post(
"/v1/chat/completions",
json={"model": "auto:batch", "messages": _messages()},
)
assert resp.status_code == 200
assert resp.headers["X-Router-Model"] == CHEAP
assert calls[0]["body"]["model"] == CHEAP
def test_v1_models_lists_all_profiles(router):
"""Every valid profile appears as auto:<name> in the models list."""
client, _, _ = router
resp = client.get("/v1/models")
assert resp.status_code == 200
ids = {m["id"] for m in resp.json()["data"]}
assert "auto" in ids
assert "auto:default" in ids
assert "auto:batch" in ids
assert "auto:locality" in ids
assert "auto:bigboybritches" in ids
assert "auto:onlycheaps" in ids
def test_a_plain_text_request_is_unaffected_by_the_gates(router):
"""No images, no response_format: routing is exactly as before."""
client, calls, _ = router
resp = client.post(
"/v1/chat/completions",
json={"model": "auto", "messages": _messages()},
)
assert resp.status_code == 200
assert calls[0]["body"]["model"] == CHEAP
assert resp.headers["X-Router-Model"] == CHEAP
def test_a_pinned_non_vision_model_with_an_image_is_a_clear_422(router):
"""A pin that cannot satisfy the request fails early with a named reason."""
client, calls, _ = router
resp = client.post(
"/v1/chat/completions",
json={"model": DEAR, "messages": _image_messages()},
)
assert resp.status_code == 422
assert "vision" in resp.json()["detail"]
assert not calls, "no provider call should happen for an impossible pin"
def test_a_pinned_non_vision_model_dispatches_when_vision_gate_is_off(router, monkeypatch):
"""The pin check honors `routing.require_vision`, like routed traffic.
`require_vision: false` accepts the occasional provider 400 in exchange
for not fail-closing on an unconfirmed flag; a pinned model must face the
same config switch as `auto` routing, not an unconditional 422.
"""
client, calls, _ = router
monkeypatch.setattr(dispatcher.cfg.routing, "require_vision", False)
resp = client.post(
"/v1/chat/completions",
json={"model": DEAR, "messages": _image_messages()},
)
assert resp.status_code == 200
assert calls[0]["body"]["model"] == DEAR, "the pin must dispatch as asked"
def test_a_provider_prefixed_pin_is_resolved_for_the_capability_check(router):
"""opencode sends a pin as `provider/model` too, not just for `auto`.
Before the fix, `requested` kept its `llm-router/` prefix past this
point: the capability pre-check looked up a model_id the catalog has
never heard of, found no row, and fail-closed a genuinely capable pin
into a false 422. It must resolve to the bare id and go through, and the
provider must see the bare id too -- NeuralWatt has never heard of
`llm-router/...` either.
"""
client, calls, _ = router
resp = client.post(
"/v1/chat/completions",
json={"model": f"llm-router/{CHEAP}", "messages": _image_messages()},
)
assert resp.status_code == 200
assert calls[0]["body"]["model"] == CHEAP
assert resp.headers["X-Router-Model"] == CHEAP
def test_a_provider_prefixed_pin_still_fails_closed_when_incapable(router):
"""The prefix-resolution fix must not turn into a bypass."""
client, calls, _ = router
resp = client.post(
"/v1/chat/completions",
json={"model": f"llm-router/{DEAR}", "messages": _image_messages()},
)
assert resp.status_code == 422
assert "vision" in resp.json()["detail"]
assert not calls, "no provider call should happen for an impossible pin"
def test_a_non_neuralwatt_vendor_pin_is_not_stripped(router, monkeypatch):
"""A real second-provider id containing a slash must not be stripped.
The router strips a `provider/` prefix only when the prefix is not a
configured real provider. With a non-neuralwatt provider in the config,
`vendor/model_id` must stay intact and be dispatched as asked.
"""
client, calls, db_path = router
from config import DispatchProvider
conn = sqlite3.connect(db_path)
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 (?, 'other-vendor', ?, 2, 262128, 192500, 16384, 0.10, 0.30,
1, 1, 'standard', 'default', 'full', 'public', 'active',
'2026-08-22T00:00:00+00:00')
""",
("other-vendor/model_id", "other-vendor/model_id"),
)
conn.commit()
conn.close()
monkeypatch.setattr(
dispatcher.cfg.dispatch_settings, "default_provider", "other-vendor"
)
monkeypatch.setitem(
dispatcher.cfg.dispatch_providers,
"other-vendor",
DispatchProvider(base_url="https://api.other-vendor.example/v1", api_key_env="OTHER_VENDOR_KEY"),
)
monkeypatch.setenv("OTHER_VENDOR_KEY", "other-key")
resp = client.post(
"/v1/chat/completions",
json={"model": "other-vendor/model_id", "messages": _messages()},
)
assert resp.status_code == 200
assert calls[0]["body"]["model"] == "other-vendor/model_id"
def test_pinned_openrouter_model_uses_openrouter_provider(router, monkeypatch):
"""A pin to an OpenRouter-only model must resolve to the openrouter provider.
The old code always used ``cfg.dispatch_settings.default_provider``
(neuralwatt) on the passthrough path, so a pin targeting an OpenRouter
id would end up calling **neuralwatt** with an unrecognised model name,
getting a 400. Now ``_resolve_pinned_provider`` looks up the owning
provider in the catalog and returns it.
"""
client, calls, db_path = router
conn = sqlite3.connect(db_path)
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 (?, 'openrouter', ?, 3, 262128, 192500, 16384, 0.10, 20.0,
1, 1, 'standard', 'default', 'full', 'public', 'active',
'2026-08-22T00:00:00+00:00')
""",
("openai/gpt-6-astra", "openai/gpt-6-astra"),
)
conn.commit()
conn.close()
# Wire up the openrouter provider config
from config import DispatchProvider
monkeypatch.setitem(
dispatcher.cfg.dispatch_providers,
"openrouter",
DispatchProvider(
base_url="https://openrouter.ai/api/v1",
api_key_env="OPENROUTER_API_KEY",
),
)
monkeypatch.setenv("OPENROUTER_API_KEY", "or-key")
resp = client.post(
"/v1/chat/completions",
json={"model": "openai/gpt-6-astra", "messages": _messages()},
)
assert resp.status_code == 200
assert calls[0]["body"]["model"] == "openai/gpt-6-astra"
def test_llm_router_openrouter_prefix_strips_correctly(router, monkeypatch):
"""``llm-router/openai/gpt-6-astra`` strips to ``openai/gpt-6-astra``
and resolves to the ``openrouter`` provider.
The old rsplit heuristic stripped to just ``gpt-6-astra``, which didn't
exist in any provider's table. The new prefix-first logic strips the
``llm-router/`` prefix and checks if the remainder exists — for
OpenRouter vendor/model ids this resolves correctly.
"""
client, calls, db_path = router
conn = sqlite3.connect(db_path)
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 (?, 'openrouter', ?, 3, 262128, 192500, 16384, 0.10, 20.0,
1, 1, 'standard', 'default', 'full', 'public', 'active',
'2026-08-22T00:00:00+00:00')
""",
("openai/gpt-6-astra", "openai/gpt-6-astra"),
)
conn.commit()
conn.close()
from config import DispatchProvider
monkeypatch.setitem(
dispatcher.cfg.dispatch_providers,
"openrouter",
DispatchProvider(
base_url="https://openrouter.ai/api/v1",
api_key_env="OPENROUTER_API_KEY",
),
)
monkeypatch.setenv("OPENROUTER_API_KEY", "or-key")
resp = client.post(
"/v1/chat/completions",
json={"model": "llm-router/openai/gpt-6-astra", "messages": _messages()},
)
assert resp.status_code == 200
assert calls[0]["body"]["model"] == "openai/gpt-6-astra"
assert resp.headers["X-Router-Model"] == "openai/gpt-6-astra"
def test_neuralwatt_only_pin_uses_default_provider(router):
"""Models seeded only on neuralwatt resolve back to the default provider.
The existing fixture seeds CHEAP and DEAR on the neuralwatt provider.
``_resolve_pinned_provider`` finds exactly one row and returns
``neuralwatt``, which matches ``cfg.dispatch_settings.default_provider``
— but the critical thing is the resolution path goes through the
catalog, not a hardcoded default.
"""
client, calls, _ = router
resp = client.post(
"/v1/chat/completions",
json={"model": CHEAP, "messages": _messages()},
)
assert resp.status_code == 200
assert calls[0]["body"]["model"] == CHEAP
assert resp.headers["X-Router-Model"] == CHEAP
def test_the_two_pass_reroute_passes_image_capability(router):
"""The measured-context reroute must carry the image gate, not lose it."""
client, calls, _ = router
# ~1200 chars / 3 = 400 measured tokens > the stubbed estimate of 100, so
# chat_completions reroutes — and the reroute must still require vision.
resp = client.post(
"/v1/chat/completions",
json={"model": "auto", "messages": _image_messages("refactor " + "x " * 600)},
)
assert resp.status_code == 200
assert resp.headers["X-Router-Model"] == CHEAP
assert calls[0]["body"]["model"] == CHEAP
def test_local_fallback_refuses_a_remote_image_url(router, monkeypatch):
"""A remote http(s) URL is an SSRF vector by indirection and must not be
forwarded to the local model, even when the fallback is enabled."""
client, calls, db_path = router
_drop_cheap_by_tier(dispatcher, db_path)
monkeypatch.setattr(dispatcher.cfg.local_vision, "enabled", True)
_local_vision_fake(monkeypatch, calls, content="should not be used")
messages = [
{
"role": "user",
"content": [
{"type": "text", "text": "describe this"},
{
"type": "image_url",
"image_url": {"url": "http://internal/private.png"},
},
],
}
]
resp = client.post(
"/v1/chat/completions",
json={"model": "auto", "messages": messages},
)
assert resp.status_code == 422
assert "vision" in resp.json()["detail"]
local = [c for c in calls if c.get("local")]
assert not local, "a remote image URL must never reach the local vision endpoint"
def test_local_fallback_refuses_a_degenerate_image_url_part(router, monkeypatch):
"""A part that declares itself an image but carries no readable value at
all must still fail the "every part is a verified data: URI" check --
silently skipping it would let it slip past the SSRF guard the same way
a remote URL is refused above."""
client, calls, db_path = router
_drop_cheap_by_tier(dispatcher, db_path)
monkeypatch.setattr(dispatcher.cfg.local_vision, "enabled", True)
_local_vision_fake(monkeypatch, calls, content="should not be used")
messages = [
{
"role": "user",
"content": [
{"type": "text", "text": "describe this"},
{"type": "image_url"},
],
}
]
resp = client.post(
"/v1/chat/completions",
json={"model": "auto", "messages": messages},
)
assert resp.status_code == 422
assert "vision" in resp.json()["detail"]
local = [c for c in calls if c.get("local")]
assert not local, "an unverifiable image part must never reach the local vision endpoint"
def test_local_fallback_wont_masquerade_an_empty_answer_as_200(router, monkeypatch):
"""A 200 with empty content is a failed answer, and must fall through to 422
rather than return an empty 200 - a hidden failure must not look like a win."""
client, calls, db_path = router
_drop_cheap_by_tier(dispatcher, db_path)
monkeypatch.setattr(dispatcher.cfg.local_vision, "enabled", True)
_local_vision_fake(monkeypatch, calls, content=" ")
resp = client.post(
"/v1/chat/completions",
json={"model": "auto", "messages": _image_messages()},
)
assert resp.status_code == 422
assert "vision" in resp.json()["detail"]
def test_local_fallback_refuses_too_many_images(router, monkeypatch):
"""An image count above the budget refuses the local call before spending it."""
client, calls, db_path = router
_drop_cheap_by_tier(dispatcher, db_path)
monkeypatch.setattr(dispatcher.cfg.local_vision, "enabled", True)
_local_vision_fake(monkeypatch, calls, content="caption")
many = _image_messages()
for _ in range(dispatcher.cfg.local_vision.max_images + 1):
many[0]["content"].append({"type": "image_url", "image_url": {"url": "data:image/png;base64,AAAA"}})
resp = client.post(
"/v1/chat/completions",
json={"model": "auto", "messages": many},
)
assert resp.status_code == 422
local = [c for c in calls if c.get("local")]
assert not local, "over-budget image count must skip the local call"
def test_local_fallback_refuses_when_api_key_env_is_missing(router, monkeypatch):
"""A configured local vision api_key_env with no env var declines the call."""
client, calls, db_path = router
_drop_cheap_by_tier(dispatcher, db_path)
monkeypatch.setattr(dispatcher.cfg.local_vision, "enabled", True)
monkeypatch.setattr(dispatcher.cfg.local_vision, "api_key_env", "LOCAL_VISION_KEY")
monkeypatch.delenv("LOCAL_VISION_KEY", raising=False)
_local_vision_fake(monkeypatch, calls, content="caption")
resp = client.post(
"/v1/chat/completions",
json={"model": "auto", "messages": _image_messages()},
)
assert resp.status_code == 422
local = [c for c in calls if c.get("local")]
assert not local, "a missing vision key must skip the local call"
def test_local_fallback_refuses_an_unparseable_response(router, monkeypatch):
"""An unparseable local response body must not yield a 200."""
client, calls, db_path = router
_drop_cheap_by_tier(dispatcher, db_path)
monkeypatch.setattr(dispatcher.cfg.local_vision, "enabled", True)
def fake_post(url, headers=None, json=None, stream=False, timeout=None):
calls.append({"url": url, "body": json, "stream": stream, "local": True})
return FakeResponse({}, status_code=200)
monkeypatch.setattr(dispatcher.requests, "post", fake_post)
resp = client.post(
"/v1/chat/completions",
json={"model": "auto", "messages": _image_messages()},
)
assert resp.status_code == 422
assert "vision" in resp.json()["detail"]
def test_a_pinned_non_json_model_with_json_response_format_422s(router):
"""The JSON-mode arm of the pass-through check mirrors the vision one: a pin
whose model lacks JSON mode gets a clear 422, not a provider 400."""
client, calls, db_path = router
conn = sqlite3.connect(db_path)
conn.execute("UPDATE models SET supports_json_mode = 0 WHERE model_id = ?", (DEAR,))
conn.commit()
conn.close()
resp = client.post(
"/v1/chat/completions",
json={
"model": DEAR,
"messages": _messages(),
"response_format": {"type": "json_object"},
},
)
assert resp.status_code == 422
assert "json" in resp.json()["detail"]
assert not calls, "no provider call should happen for an impossible json pin"
# --- session classification cache -----------------------------------------
def _session_messages(text="write me a function", system="You are a coding agent"):
"""Messages with a stable opening (system) message, so the session
fingerprint is constant across turns - the condition a cache hit needs."""
return [
{"role": "system", "content": system},
{"role": "user", "content": text},
]
def _session_key(messages=None):
return dispatcher.session_fingerprint(messages or _session_messages())
@pytest.fixture
def session_cache_on(monkeypatch):
"""Enable the cache and isolate the module-level store between tests."""
session_cache.clear()
monkeypatch.setattr(dispatcher.cfg.session_cache, "enabled", True)
yield
session_cache.clear()
def _count_classify(monkeypatch, result=None):
"""Replace classify with a counting stub returning ``result`` (or the
default real-classifier shape)."""
result = result or Classification(
task_category="coding_general", task_tier=2,
required_context_tokens=100, confidence=0.9,
)
counts = {"calls": 0}
def classify(task, context):
counts["calls"] += 1
return result
monkeypatch.setattr(dispatcher, "classify", classify)
return counts
def test_cache_hit_skips_classify_and_reports_cached_source(router, logbuf, session_cache_on, monkeypatch):
"""A second request in the same session reuses the cached category/tier and
does not pay another classifier round-trip."""
client, _, _ = router
counts = _count_classify(monkeypatch)
client.post("/v1/chat/completions",
json={"model": "auto", "messages": _session_messages()})
assert counts["calls"] == 1, "first turn classifies"
client.post("/v1/chat/completions",
json={"model": "auto", "messages": _session_messages("follow up")})
assert counts["calls"] == 1, "second turn must not classify again"
routes = _lines(logbuf, "route")
assert _fields(routes[0])["src"] == "classifier", "first turn came from the classifier"
assert _fields(routes[1])["src"] == "cached", "second turn came from the cache"
def test_cache_miss_writes_the_successful_classification(router, session_cache_on):
"""After a real (non-fallback) classification, the session's category/tier
land in the cache under its fingerprint."""
client, _, _ = router
client.post("/v1/chat/completions",
json={"model": "auto", "messages": _session_messages()})
cached = session_cache.get(_session_key(), staleness_seconds=60)
assert cached is not None
assert cached.task_category == "coding_general"
assert cached.task_tier == 2
def test_cache_expiry_reclassifies(router, logbuf, session_cache_on, monkeypatch):
"""Once stale, a session's next turn reclassifies from scratch."""
client, _, _ = router
counts = _count_classify(monkeypatch)
client.post("/v1/chat/completions",
json={"model": "auto", "messages": _session_messages()})
assert counts["calls"] == 1
# Advance beyond the 1200-second TTL and send another turn.
now = [session_cache.time.time()]
monkeypatch.setattr(session_cache.time, "time", lambda: now[0] + 20 * 60 + 1)
client.post("/v1/chat/completions",
json={"model": "auto", "messages": _session_messages("again")})
assert counts["calls"] == 2, "expired entry must reclassify"
assert _fields(_lines(logbuf, "route")[1])["src"] == "classifier"
def test_fallback_classification_is_never_cached(router, session_cache_on, monkeypatch):
"""A classifier fallback must not populate the cache: the next turn still
pays the round-trip."""
client, _, _ = router
fallback = Classification(
task_category="general_chat", task_tier=2,
required_context_tokens=0, confidence=0.0, source="fallback",
)
counts = _count_classify(monkeypatch, result=fallback)
client.post("/v1/chat/completions",
json={"model": "auto", "messages": _session_messages()})
assert session_cache.get(_session_key(), staleness_seconds=60) is None, \
"fallback must never be cached"
client.post("/v1/chat/completions",
json={"model": "auto", "messages": _session_messages("again")})
assert counts["calls"] == 2, "a fallback must not be served from cache"
def test_capability_flags_are_still_read_fresh_on_a_cache_hit(router, session_cache_on, monkeypatch):
"""A cache hit reuses only category/tier; a request that suddenly carries an
image is still gated by its (freshly read) capability flags, not served
blindly from the cached decision."""
client, _, db_path = router
_count_classify(monkeypatch)
client.post("/v1/chat/completions",
json={"model": "auto", "messages": _session_messages()})
# The SAME session (same system prompt) now turns into an image request:
# it must route to the vision-capable model rather than being treated as
# the cached text turn.
image_turn = [{
"role": "user",
"content": [
{"type": "text", "text": "what is in this image?"},
{"type": "image_url",
"image_url": {"url": "data:image/png;base64,AAAA"}},
],
}]
resp = client.post(
"/v1/chat/completions",
json={"model": "auto",
"messages": _session_messages()[:1] + image_turn},
)
assert resp.status_code == 200
assert resp.json()["model"].startswith(CHEAP), "vision-capable model picked"
conn = sqlite3.connect(db_path)
srcs = [r[0] for r in conn.execute(
"SELECT classification_source FROM route_decisions ORDER BY id"
)]
conn.close()
assert srcs == ["classifier", "cached"], "cached decision persisted with source=cached"
# --- circuit breaker failover ---------------------------------------------
def test_breaker_records_failure_then_success_across_availability_failover(
router, monkeypatch
):
"""A 503 on the routed target must record a failure, then a 200 on the
runner-up must record a success — the passive-recovery trust rebuild for
the non-streaming loop (circuit_breaker.record_success was previously
never called from the dispatch path).
"""
import circuit_breaker
client, calls, _ = router
monkeypatch.setattr(dispatcher.cfg.circuit_breaker, "enabled", True)
failures = []
successes = []
monkeypatch.setattr(
dispatcher.circuit_breaker, "record_failure",
lambda model_id, provider, *a, **kw: failures.append((model_id, provider)),
)
monkeypatch.setattr(
dispatcher.circuit_breaker, "record_success",
lambda model_id, provider="neuralwatt": successes.append((model_id, provider)),
)
statuses = iter([503, 200])
def fake_post(url, headers=None, json=None, stream=False, timeout=None):
status = next(statuses)
calls.append({"url": url, "body": json, "stream": stream})
if status >= 400:
return FakeResponse({"error": "down"}, status_code=status)
return FakeResponse(completion(json["model"]))
monkeypatch.setattr(dispatcher.requests, "post", fake_post)
try:
circuit_breaker.clear()
resp = client.post(
"/v1/chat/completions",
json={"model": "auto", "messages": _messages()},
)
finally:
circuit_breaker.clear()
assert resp.status_code == 200
# Routing picks the cheapest (CHEAP) as target; DEAR is the runner-up.
assert calls[0]["body"]["model"] == CHEAP
assert calls[1]["body"]["model"] == DEAR
assert resp.headers["X-Router-Model"] == DEAR
assert resp.json()["model"] == DEAR
# The spies prove the wiring: a failure was recorded for the 503 target
# and a success for the 200 runner-up, without depending on is_down timing.
assert failures == [(CHEAP, "neuralwatt")]
assert successes == [(DEAR, "neuralwatt")]
def test_upstream_402_is_logged_with_full_status_line_and_body(router, caplog, monkeypatch):
client, _, _ = router
caplog.set_level("ERROR", logger="llm_router")
def fake_402(url, headers=None, json=None, stream=False, timeout=None):
return FakeResponse(
payload={
"error": "INSUFFICIENT_CREDITS_MARKER_7f3a",
"pad": "x" * 200,
},
status_code=402,
reason="Payment Required",
)
monkeypatch.setattr(dispatcher.cfg.circuit_breaker, "enabled", False)
monkeypatch.setattr(dispatcher.requests, "post", fake_402)
resp = client.post(
"/v1/chat/completions",
json={"model": CHEAP, "messages": _messages()},
)
assert resp.status_code == 402
upstreams = [r for r in caplog.records if r.getMessage().startswith("upstream ")]
assert len(upstreams) == 1
m = upstreams[0].getMessage()
assert "status=402" in m
assert 'reason="Payment Required"' in m
assert "INSUFFICIENT_CREDITS_MARKER_7f3a" in m
assert "x" * 200 in m
def test_upstream_402_streaming_failover_logs_full_body(router, caplog, monkeypatch):
client, _, _ = router
caplog.set_level("ERROR", logger="llm_router")
upstream_response = FakeResponse(
payload={"error": "SOME_MARKER_xf29b", "pad": "x" * 200},
status_code=402,
reason="Payment Required",
)
monkeypatch.setattr(
dispatcher, "_open_upstream",
lambda *a, **kw: upstream_response,
)
resp = client.post(
"/v1/chat/completions",
json={"model": "auto", "messages": _messages(), "stream": True},
)
assert resp.status_code == 402
upstreams = [r for r in caplog.records if r.getMessage().startswith("upstream ")]
streaming = [r for r in upstreams if "stream=1" in r.getMessage()]
# 402 is account-level → short-circuits after first candidate.
assert len(streaming) == 1
rec = streaming[0]
m = rec.getMessage()
assert m.count("status=402") == 1
assert m.count('reason="Payment Required"') == 1
assert "SOME_MARKER_xf29b" in m
assert "x" * 200 in m
def test_upstream_sdk_status_error_logs_status_and_body(router, caplog, monkeypatch):
from openai import APIStatusError
client, _, _ = router
caplog.set_level("ERROR", logger="llm_router")
marker = "INSUFFICIENT_CREDITS_SDK_DEAD"
body_text = json.dumps({
"error": marker,
"pad": "x" * 200,
})
# Build a FakeResponse that mimics the HTTP response inside APIStatusError
fake_resp = FakeResponse(
payload=json.loads(body_text),
status_code=402,
reason="Payment Required",
request=object(),
)
fake_resp.reason_phrase = "Payment Required"
def fake_provider_client(provider):
"""Return a client whose raw chat completions raises APIStatusError."""
mock_client = MagicMock()
mock_client.chat.completions.with_raw_response.create.side_effect = APIStatusError(
message="payment required",
response=fake_resp,
body={"error": marker, "pad": "x" * 200},
)
return mock_client
monkeypatch.setattr(dispatcher.cfg.circuit_breaker, "enabled", False)
monkeypatch.setattr(dispatcher, "_provider_client", fake_provider_client)
resp = client.post(
"/dispatch",
json={
"task": "write me a function",
"task_category": "general_chat",
"task_tier": 2,
},
)
assert resp.status_code == 502
assert f"Dispatch to {CHEAP} failed" in resp.text
upstreams = [r for r in caplog.records if r.getMessage().startswith("upstream ")]
assert len(upstreams) == 1
m = upstreams[0].getMessage()
assert "status=402" in m
assert 'reason="Payment Required"' in m
assert marker in m
assert "x" * 200 in m
def test_upstream_proxy_generator_failure_logs_full_body(router, caplog, monkeypatch):
import circuit_breaker
circuit_breaker.clear()
client, _, _ = router
caplog.set_level("ERROR", logger="llm_router")
monkeypatch.setattr(dispatcher.cfg.circuit_breaker, "enabled", False)
marker = "PROXY_MARKER_9d4e"
body_text = json.dumps({
"error": marker,
"pad": "x" * 200,
})
upstream_response = FakeResponse(
payload=json.loads(body_text),
status_code=503,
reason="Service Unavailable",
)
monkeypatch.setattr(
dispatcher, "_open_upstream",
lambda *a, **kw: upstream_response,
)
resp = client.post(
"/v1/chat/completions",
json={"model": "auto", "messages": _messages(), "stream": True},
)
assert resp.status_code == 503
upstreams = [r for r in caplog.records if r.getMessage().startswith("upstream ")]
streaming = [r for r in upstreams if "stream=1" in r.getMessage()]
assert len(streaming) >= 1
rec = streaming[0]
assert "status=503" in rec.getMessage()
assert 'reason="Service Unavailable"' in rec.getMessage()
# body must be present with full content, not just [:200]
assert marker in rec.getMessage()
assert "x" * 200 in rec.getMessage()
# --- cross-provider failover ----------------------------------------------
#
# Failover was written when every cloud row shared one NeuralWatt account, and
# two things in it stayed true to that premise after OpenRouter arrived:
#
# 1. `alternatives` carried only (model_id, ceiling) and reused the SELECTED
# model's provider for every candidate, so a cross-provider failover
# posted one provider's model id to another's endpoint with the wrong key.
# 2. `_account_level_refusal` (any 4xx except 400/404/422, so 429 included)
# aborted failover outright. A free-tier rate limit on OpenRouter threw
# away a NeuralWatt candidate that would have worked.
#
# 78 decisions in the 7 days before this was written had a cross-provider
# runner-up, so neither was hypothetical.
REMOTE = "othervendor/remote-model"
def _add_openrouter_runner_up(db_path, monkeypatch, *, with_key=True):
"""Insert an openrouter row priced BETWEEN cheap and dear.
That ordering is what makes it the first runner-up behind CHEAP while
DEAR sits behind it, so a test can tell "skipped this provider's rows"
apart from "gave up on the whole list".
"""
conn = sqlite3.connect(db_path)
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 (?, 'openrouter', ?, 2, 262128, 192500, 16384, 0.33, 1.00,
1, 1, 'standard', 'default', 'full', 'public', 'active',
'2026-08-22T00:00:00+00:00')
""",
(REMOTE, REMOTE),
)
conn.commit()
conn.close()
from config import DispatchProvider
monkeypatch.setitem(
dispatcher.cfg.dispatch_providers,
"openrouter",
DispatchProvider(
base_url="https://openrouter.example/api/v1",
api_key_env="OPENROUTER_API_KEY",
),
)
if with_key:
monkeypatch.setenv("OPENROUTER_API_KEY", "or-key")
else:
monkeypatch.delenv("OPENROUTER_API_KEY", raising=False)
def test_failover_to_another_provider_uses_that_providers_endpoint(
router, monkeypatch
):
"""A 503 on a NeuralWatt target must reach OpenRouter's URL with
OpenRouter's key -- not NeuralWatt's URL carrying an OpenRouter model id."""
client, calls, db_path = router
_add_openrouter_runner_up(db_path, monkeypatch)
monkeypatch.setattr(dispatcher.cfg.circuit_breaker, "enabled", False)
statuses = iter([503, 200])
def fake_post(url, headers=None, json=None, stream=False, timeout=None):
calls.append({"url": url, "headers": headers, "body": json})
if next(statuses) >= 400:
return FakeResponse({"error": "down"}, status_code=503)
return FakeResponse(completion(json["model"]))
monkeypatch.setattr(dispatcher.requests, "post", fake_post)
resp = client.post(
"/v1/chat/completions",
json={"model": "auto", "messages": _messages()},
)
assert resp.status_code == 200
assert calls[0]["body"]["model"] == CHEAP
assert calls[1]["body"]["model"] == REMOTE
assert "openrouter.example" in calls[1]["url"]
assert calls[1]["headers"]["authorization"] == "Bearer or-key"
assert resp.headers["X-Router-Model"] == REMOTE
def test_a_429_skips_that_providers_rows_but_not_the_other_providers(
router, monkeypatch
):
"""The fix. A NeuralWatt 429 is an account-level refusal, so DEAR (also
NeuralWatt) is skipped -- but REMOTE is on a different account and is
tried, where the old code raised 429 without trying anything."""
client, calls, db_path = router
_add_openrouter_runner_up(db_path, monkeypatch)
monkeypatch.setattr(dispatcher.cfg.circuit_breaker, "enabled", False)
def fake_post(url, headers=None, json=None, stream=False, timeout=None):
calls.append({"url": url, "body": json})
if json["model"] == CHEAP:
return FakeResponse({"error": "rate limited"}, status_code=429)
return FakeResponse(completion(json["model"]))
monkeypatch.setattr(dispatcher.requests, "post", fake_post)
resp = client.post(
"/v1/chat/completions",
json={"model": "auto", "messages": _messages()},
)
assert resp.status_code == 200
attempted = [c["body"]["model"] for c in calls]
assert attempted == [CHEAP, REMOTE], (
"DEAR shares the refusing NeuralWatt account and must be skipped; "
"REMOTE does not and must be tried"
)
def test_a_429_with_no_other_provider_still_gives_up_immediately(
router, monkeypatch
):
"""Scoping must not become 'walk the whole list'. With only NeuralWatt
rows, a 429 still refuses without trying DEAR."""
client, calls, _ = router
monkeypatch.setattr(dispatcher.cfg.circuit_breaker, "enabled", False)
def fake_post(url, headers=None, json=None, stream=False, timeout=None):
calls.append({"body": json})
return FakeResponse({"error": "rate limited"}, status_code=429)
monkeypatch.setattr(dispatcher.requests, "post", fake_post)
resp = client.post(
"/v1/chat/completions",
json={"model": "auto", "messages": _messages()},
)
assert resp.status_code == 429
assert [c["body"]["model"] for c in calls] == [CHEAP]
def test_a_provider_with_no_key_is_skipped_not_posted_to_the_wrong_url(
router, monkeypatch
):
"""A missing key for an ALTERNATIVE must skip that candidate, not fail the
request and not send its model id to the previous provider's endpoint."""
client, calls, db_path = router
_add_openrouter_runner_up(db_path, monkeypatch, with_key=False)
monkeypatch.setattr(dispatcher.cfg.circuit_breaker, "enabled", False)
def fake_post(url, headers=None, json=None, stream=False, timeout=None):
calls.append({"url": url, "body": json})
if json["model"] == CHEAP:
return FakeResponse({"error": "down"}, status_code=503)
return FakeResponse(completion(json["model"]))
monkeypatch.setattr(dispatcher.requests, "post", fake_post)
resp = client.post(
"/v1/chat/completions",
json={"model": "auto", "messages": _messages()},
)
assert resp.status_code == 200
attempted = [c["body"]["model"] for c in calls]
assert REMOTE not in attempted, "no key means skip, never post"
assert attempted == [CHEAP, DEAR]
def test_streaming_failover_crosses_providers_too(router, monkeypatch):
"""The streaming path had both defects independently: it built
`stream_candidates` from model ids alone and broke out of the loop on any
account-level refusal."""
client, _, db_path = router
_add_openrouter_runner_up(db_path, monkeypatch)
monkeypatch.setattr(dispatcher.cfg.circuit_breaker, "enabled", False)
opened = []
def fake_open(model, url, headers, body):
opened.append({"model": model, "url": url, "headers": headers})
if model == CHEAP:
return FakeResponse({"error": "rate limited"}, status_code=429)
return FakeResponse(lines=STREAM_LINES)
monkeypatch.setattr(dispatcher, "_open_upstream", fake_open)
resp = client.post(
"/v1/chat/completions",
json={"model": "auto", "messages": _messages(), "stream": True},
)
assert resp.status_code == 200
assert [o["model"] for o in opened] == [CHEAP, REMOTE]
assert "openrouter.example" in opened[1]["url"]
assert opened[1]["headers"]["authorization"] == "Bearer or-key"