"""LocalDispatchModel config model: parsing, validation, meterable-set property. Tests the LocalDispatchModel class, the eligible_categories validator on RouterConfig, and the local_energy_dispatch_models cached property. """ from __future__ import annotations import copy import sqlite3 from datetime import datetime, timezone from pathlib import Path import pytest import yaml from config import LocalDispatchModel, RouterConfig ROOT = Path(__file__).resolve().parent.parent @pytest.fixture def raw() -> dict: with open(ROOT / "config" / "config.yaml") as fh: return yaml.safe_load(fh) # --- valid entry parses ----------------------------------------------------- def test_valid_entry_parses(): entry = LocalDispatchModel( model_id="test-model", base_url="http://localhost:11434/v1", context_window=16384, tier=1, eligible_categories=["coding_general"], ) assert entry.model_id == "test-model" assert entry.base_url == "http://localhost:11434/v1" assert entry.timeout_seconds == 120.0 assert entry.context_window == 16384 assert entry.max_output_tokens == 2048 assert entry.tier == 1 assert entry.api_key_env is None def test_valid_entry_with_all_fields(): entry = LocalDispatchModel( model_id="test-model", base_url="http://localhost:11434/v1", api_key_env="OLLAMA_API_KEY", timeout_seconds=60.0, context_window=32768, max_output_tokens=4096, tier=2, eligible_categories=["coding_general", "debugging", "summarization"], ) assert entry.api_key_env == "OLLAMA_API_KEY" assert entry.timeout_seconds == 60.0 assert entry.context_window == 32768 assert entry.max_output_tokens == 4096 assert entry.tier == 2 assert entry.eligible_categories == [ "coding_general", "debugging", "summarization", ] # --- valid entry parses in RouterConfig via full config load ----------------- def test_valid_local_dispatch_entry_loads_in_full_config(raw): cfg = copy.deepcopy(raw) cfg["local_dispatch_models"] = [ { "model_id": "qwen2.5-coder-router:14b", "base_url": "http://localhost:11434/v1", "context_window": 16384, "tier": 1, "eligible_categories": ["coding_general"], } ] loaded = RouterConfig(**cfg) assert len(loaded.local_dispatch_models) == 1 assert loaded.local_dispatch_models[0].model_id == "qwen2.5-coder-router:14b" # --- unknown category raises ValueError ------------------------------------- def test_unknown_category_in_eligible_categories_raises(raw): cfg = copy.deepcopy(raw) cfg["local_dispatch_models"] = [ { "model_id": "test-model", "base_url": "http://localhost:11434/v1", "context_window": 16384, "tier": 1, "eligible_categories": ["not_a_real_category"], } ] with pytest.raises(ValueError, match="not_a_real_category"): RouterConfig(**cfg) def test_unknown_category_message_mentions_proficiency_categories(raw): cfg = copy.deepcopy(raw) cfg["local_dispatch_models"] = [ { "model_id": "test-model", "context_window": 16384, "tier": 1, "eligible_categories": ["fake_category"], } ] with pytest.raises(ValueError, match="proficiency.categories"): RouterConfig(**cfg) # --- duplicate model_id raises ---------------------------------------------- def test_duplicate_model_id_raises(raw): cfg = copy.deepcopy(raw) cfg["local_dispatch_models"] = [ { "model_id": "dup-model", "context_window": 16384, "tier": 1, "eligible_categories": ["coding_general"], }, { "model_id": "dup-model", "context_window": 32768, "tier": 2, "eligible_categories": ["debugging"], }, ] with pytest.raises(ValueError, match="duplicate model_id"): RouterConfig(**cfg) # --- empty eligible_categories raises --------------------------------------- def test_empty_eligible_categories_raises(): with pytest.raises(ValueError, match="must contain at least one"): LocalDispatchModel( model_id="test-model", context_window=16384, tier=1, eligible_categories=[], ) # --- duplicate eligible_categories raises ----------------------------------- def test_duplicate_eligible_categories_raises(): with pytest.raises(ValueError, match="duplicate"): LocalDispatchModel( model_id="test-model", context_window=16384, tier=1, eligible_categories=["coding_general", "coding_general"], ) # --- local_energy_dispatch_models when disabled ----------------------------- def test_local_energy_dispatch_models_empty_when_disabled(raw): cfg = copy.deepcopy(raw) cfg["local_energy"]["enabled"] = False loaded = RouterConfig(**cfg) assert loaded.local_energy_dispatch_models == frozenset() def test_local_energy_dispatch_models_empty_when_disabled_with_entries(raw): cfg = copy.deepcopy(raw) cfg["local_energy"]["enabled"] = False cfg["local_dispatch_models"] = [ { "model_id": "test-model", "base_url": "http://localhost:11434/v1", "context_window": 16384, "tier": 1, "eligible_categories": ["coding_general"], } ] loaded = RouterConfig(**cfg) assert loaded.local_energy_dispatch_models == frozenset() # --- local_energy_dispatch_models when enabled with loopback ---------------- def test_local_energy_dispatch_models_contains_loopback_entry_when_enabled(raw): cfg = copy.deepcopy(raw) cfg["local_energy"]["enabled"] = True cfg["local_energy"]["tariff_usd_per_kwh"] = 8.0 cfg["local_dispatch_models"] = [ { "model_id": "local-model", "base_url": "http://localhost:11434/v1", "context_window": 16384, "tier": 1, "eligible_categories": ["coding_general"], } ] loaded = RouterConfig(**cfg) assert loaded.local_energy_dispatch_models == frozenset({"local-model"}) # --- local_energy_dispatch_models excludes non-loopback --------------------- def test_local_energy_dispatch_models_excludes_non_loopback(raw): cfg = copy.deepcopy(raw) cfg["local_energy"]["enabled"] = True cfg["local_energy"]["tariff_usd_per_kwh"] = 8.0 cfg["local_dispatch_models"] = [ { "model_id": "vpn-model", "base_url": "http://192.168.1.100:11434/v1", "context_window": 16384, "tier": 1, "eligible_categories": ["coding_general"], } ] loaded = RouterConfig(**cfg) assert loaded.local_energy_dispatch_models == frozenset() def test_local_energy_dispatch_models_mixed_includes_loopback_only(raw): cfg = copy.deepcopy(raw) cfg["local_energy"]["enabled"] = True cfg["local_energy"]["tariff_usd_per_kwh"] = 8.0 cfg["local_dispatch_models"] = [ { "model_id": "locals", "base_url": "http://localhost:11434/v1", "context_window": 16384, "tier": 1, "eligible_categories": ["coding_general"], }, { "model_id": "vpn", "base_url": "http://10.0.0.5:8080/v1", "context_window": 32768, "tier": 2, "eligible_categories": ["debugging"], }, ] loaded = RouterConfig(**cfg) assert loaded.local_energy_dispatch_models == frozenset({"locals"}) # --- timeout_seconds validation -------------------------------------------- def test_nonpositive_timeout_raises(): with pytest.raises(ValueError, match="timeout_seconds"): LocalDispatchModel( model_id="test-model", context_window=16384, tier=1, timeout_seconds=0, eligible_categories=["coding_general"], ) def test_negative_timeout_raises(): with pytest.raises(ValueError, match="timeout_seconds"): LocalDispatchModel( model_id="test-model", context_window=16384, tier=1, timeout_seconds=-5.0, eligible_categories=["coding_general"], ) # --- tier validation -------------------------------------------------------- def test_tier_below_one_raises(): with pytest.raises(ValueError, match="greater_than_equal"): LocalDispatchModel( model_id="test-model", context_window=16384, tier=0, eligible_categories=["coding_general"], ) def test_tier_above_three_raises(): with pytest.raises(ValueError, match="less_than_equal"): LocalDispatchModel( model_id="test-model", context_window=16384, tier=4, eligible_categories=["coding_general"], ) def test_tier_one_loads(raw): cfg = copy.deepcopy(raw) cfg["local_dispatch_models"] = [ { "model_id": "model", "context_window": 16384, "tier": 1, "eligible_categories": ["coding_general"], } ] loaded = RouterConfig(**cfg) assert loaded.local_dispatch_models[0].tier == 1 def test_tier_three_loads(raw): cfg = copy.deepcopy(raw) cfg["local_dispatch_models"] = [ { "model_id": "model", "context_window": 16384, "tier": 3, "eligible_categories": ["coding_general"], } ] loaded = RouterConfig(**cfg) assert loaded.local_dispatch_models[0].tier == 3 # --- default values -------------------------------------------------------- def test_base_url_defaults_to_localhost(raw): cfg = copy.deepcopy(raw) cfg["local_dispatch_models"] = [ { "model_id": "model", "context_window": 16384, "tier": 1, "eligible_categories": ["coding_general"], } ] loaded = RouterConfig(**cfg) assert loaded.local_dispatch_models[0].base_url == "http://localhost:11434/v1" # --- the shipped config still loads ----------------------------------------- def test_the_shipped_config_has_one_dispatch_model(raw): """Shipped config carries one local dispatch entry for the chosen local model. Pins the model id deliberately: swapping the local dispatch model is a decision that should require updating this test, not something that happens silently. See docs/local-models.md for the swap procedure. """ loaded = RouterConfig(**raw) assert len(loaded.local_dispatch_models) == 1 m = loaded.local_dispatch_models[0] assert m.model_id == "qwen2.5-coder-router:14b" assert m.base_url == "http://localhost:11434/v1" assert m.tier == 1 assert m.eligible_categories == ["file_summarization", "diff_checking"] @pytest.fixture def db_no_column(tmp_path): """Temp DB created from a schema that lacks eligible_categories.""" conn = sqlite3.connect(tmp_path / "test.db") conn.row_factory = sqlite3.Row schema = (ROOT / "config" / "schema.sql").read_text() lines = [l for l in schema.splitlines() if "eligible_categories" not in l] conn.executescript("\n".join(lines)) yield conn conn.close() @pytest.fixture def db_with_column(tmp_path): """Temp DB from the full schema.sql (has eligible_categories).""" conn = sqlite3.connect(tmp_path / "test.db") conn.row_factory = sqlite3.Row conn.executescript((ROOT / "config" / "schema.sql").read_text()) yield conn conn.close() def test_ensure_models_eligible_categories_noop_on_existing_column(db_with_column): """Calling _ensure on an up-to-date DB is a silent no-op.""" from poller import _ensure_models_eligible_categories _ensure_models_eligible_categories(db_with_column) cols = {r[1] for r in db_with_column.execute("PRAGMA table_info(models)")} assert "eligible_categories" in cols def test_ensure_models_eligible_categories_migrates_old_schema(db_no_column): """Adding the column on an old DB works without error.""" from poller import _ensure_models_eligible_categories assert "eligible_categories" not in { r[1] for r in db_no_column.execute("PRAGMA table_info(models)") } _ensure_models_eligible_categories(db_no_column) assert "eligible_categories" in { r[1] for r in db_no_column.execute("PRAGMA table_info(models)") } @pytest.fixture def db_and_cfg(tmp_path): """DB with column + a RouterConfig carrying one local dispatch entry.""" conn = sqlite3.connect(tmp_path / "test.db") conn.row_factory = sqlite3.Row conn.executescript((ROOT / "config" / "schema.sql").read_text()) raw = yaml.safe_load((ROOT / "config" / "config.yaml").read_text()) raw["dispatch_providers"] = {} raw["dispatch_settings"] = {"default_provider": "ollama-local"} raw["dispatch_providers"]["ollama-local"] = { "base_url": "http://localhost:11434/v1", "api_key_env": "OLLAMA_API_KEY", } raw["local_dispatch_models"] = [ { "model_id": "test-model", "context_window": 12288, "tier": 1, "eligible_categories": ["file_summarization", "diff_checking"], } ] cfg = RouterConfig(**raw) yield conn, cfg def test_upsert_creates_row(db_and_cfg): """First upsert inserts a row with correct provider, tier, availability.""" from poller import upsert_local_dispatch_models conn, cfg = db_and_cfg upsert_local_dispatch_models(conn, cfg) row = conn.execute( "SELECT * FROM models WHERE model_id='test-model' AND provider='ollama-local'" ).fetchone() assert row is not None assert row["provider"] == "ollama-local" assert row["tier"] == 1 assert row["availability"] == "active" assert row["supports_tools"] == 0 assert row["supports_json_mode"] == 0 assert row["supports_vision"] == 0 assert row["supports_reasoning"] == 0 assert row["reasoning_default_enabled"] == 0 def test_upsert_eligible_categories_comma_joined(db_and_cfg): """eligible_categories becomes a comma-joined string.""" from poller import upsert_local_dispatch_models conn, cfg = db_and_cfg upsert_local_dispatch_models(conn, cfg) row = conn.execute( "SELECT eligible_categories FROM models WHERE model_id='test-model'" ).fetchone() assert row[0] == "file_summarization,diff_checking" def test_upsert_effective_context_window_math(db_and_cfg): """effective_context_window == int(12288 * 0.75) - 2048.""" from poller import upsert_local_dispatch_models conn, cfg = db_and_cfg upsert_local_dispatch_models(conn, cfg) row = conn.execute( "SELECT effective_context_window FROM models WHERE model_id='test-model'" ).fetchone() assert row[0] == 7168 def test_upsert_cost_columns_null(db_and_cfg): """Fresh rows have NULL cost columns.""" from poller import upsert_local_dispatch_models conn, cfg = db_and_cfg upsert_local_dispatch_models(conn, cfg) row = conn.execute( "SELECT cost_per_1m_prompt, cost_per_1m_completion, cost_per_1m_prompt_cached " "FROM models WHERE model_id='test-model'" ).fetchone() assert row[0] is None assert row[1] is None assert row[2] is None def test_upsert_idempotent_twice(db_and_cfg): """Running upsert twice doesn't error and row still exists.""" from poller import upsert_local_dispatch_models conn, cfg = db_and_cfg upsert_local_dispatch_models(conn, cfg) upsert_local_dispatch_models(conn, cfg) count = conn.execute( "SELECT COUNT(*) FROM models WHERE model_id='test-model' AND provider='ollama-local'" ).fetchone()[0] assert count == 1 def test_upsert_cost_column_survives_re_run(db_and_cfg): """Pre-set cost column survives a re-run (must NOT be in DO UPDATE SET).""" from poller import upsert_local_dispatch_models conn, cfg = db_and_cfg upsert_local_dispatch_models(conn, cfg) sentinel = 0.0123 conn.execute( "UPDATE models SET cost_per_1m_prompt=? WHERE model_id='test-model'", (sentinel,), ) conn.commit() upsert_local_dispatch_models(conn, cfg) row = conn.execute( "SELECT cost_per_1m_prompt FROM models WHERE model_id='test-model'" ).fetchone() assert row[0] == sentinel def test_upsert_updates_changed_config_value(db_and_cfg): """Second upsert with a changed config value updates the row.""" from poller import upsert_local_dispatch_models conn, cfg = db_and_cfg upsert_local_dispatch_models(conn, cfg) raw = yaml.safe_load((ROOT / "config" / "config.yaml").read_text()) raw["dispatch_providers"] = { "ollama-local": { "base_url": "http://localhost:11434/v1", "api_key_env": "OLLAMA_API_KEY", } } raw["dispatch_settings"] = {"default_provider": "ollama-local"} raw["local_dispatch_models"] = [ { "model_id": "test-model", "context_window": 12288, "tier": 2, "eligible_categories": [ "file_summarization", "diff_checking", "coding_general", ], } ] cfg2 = RouterConfig(**raw) upsert_local_dispatch_models(conn, cfg2) row = conn.execute( "SELECT tier, eligible_categories FROM models WHERE model_id='test-model'" ).fetchone() assert row["tier"] == 2 assert row["eligible_categories"] == "file_summarization,diff_checking,coding_general" def test_cloud_row_untouched(db_and_cfg): """A cloud row (provider='neuralwatt', eligible_categories NULL) is not affected.""" from poller import upsert_local_dispatch_models conn, cfg = db_and_cfg conn.execute( """ INSERT INTO models ( model_id, provider, context_window, effective_context_window, last_updated, access_level, eligible_categories ) VALUES (?, 'neuralwatt', 32768, 24576, ?, 'public', NULL) """, ("gpt-5.1", datetime.now(timezone.utc).isoformat()), ) conn.commit() upsert_local_dispatch_models(conn, cfg) cloud = conn.execute( "SELECT eligible_categories FROM models WHERE provider='neuralwatt'" ).fetchone() assert cloud[0] is None # --- _local_dispatch_config_for ------------------------------------------- def test_local_dispatch_config_for_found(): """Linear scan finds an entry by model_id.""" entry = LocalDispatchModel( model_id="qwen2.5-coder-router:14b", base_url="http://localhost:11434/v1", context_window=16384, tier=1, eligible_categories=["file_summarization", "diff_checking"], ) result = LocalDispatchModel( model_id="other-model", base_url="http://localhost:8080/v1", context_window=8192, tier=2, eligible_categories=["debugging"], ) entries = [entry, result] found = None for e in entries: if e.model_id == "qwen2.5-coder-router:14b": found = e assert found is entry def test_local_dispatch_config_for_missing(): """Linear scan returns None when model_id is not present.""" entries = [ LocalDispatchModel( model_id="model-a", context_window=8192, tier=1, eligible_categories=["coding_general"], ), ] found = None for e in entries: if e.model_id == "absent": found = e assert found is None # --- _run_local_dispatch + _local_dispatch_response ------------------------- import json from unittest import mock import pytest import dispatcher from dispatcher import _local_dispatch_response, _run_local_dispatch ROOT = Path(__file__).resolve().parent.parent @pytest.fixture def entry(): return LocalDispatchModel( model_id="qwen2.5-coder-router:14b", base_url="http://localhost:11434/v1", context_window=16384, max_output_tokens=2048, tier=1, eligible_categories=["file_summarization"], ) @pytest.fixture def messages(): return [ {"role": "user", "content": "Summarize this file"}, ] def _good_ollama_payload(): return { "id": "cmpl-ollama-99", "object": "chat.completion", "created": 1700000000, "model": "qwen2.5-coder-router:14b", "choices": [ { "index": 0, "message": { "role": "assistant", "content": "The file refactors the auth module.", }, "finish_reason": "stop", } ], "usage": {"prompt_tokens": 30, "completion_tokens": 12}, } def _make_cfg(dispatch_models=None, local_energy_enabled=False): raw = yaml.safe_load((ROOT / "config" / "config.yaml").read_text()) raw["dispatch_providers"] = {} raw["dispatch_settings"] = {"default_provider": "ollama-local"} if dispatch_models: raw["local_dispatch_models"] = dispatch_models raw["local_energy"]["enabled"] = bool(local_energy_enabled) if local_energy_enabled: raw["local_energy"]["tariff_usd_per_kwh"] = 8.0 raw["dispatch_providers"]["ollama-local"] = { "base_url": "http://localhost:11434/v1", "api_key_env": "OLLAMA_API_KEY", } return RouterConfig(**raw) # --- happy path: non-streaming 200 JSONResponse ---------------------------- def test_run_local_dispatch_happy_non_streaming(entry, messages, monkeypatch): """POST to Ollama returns 200, payload passthrough, request_id generated.""" payload = _good_ollama_payload() call_capture = [] def fake_post(url, headers=None, json=None, timeout=None): call_capture.append({"url": url, "headers": headers, "json": json}) class R: status_code = 200 def json(self): return payload return R() monkeypatch.setattr(dispatcher.requests, "post", fake_post) # Patch metering off so we don't need local_energy context cfg = _make_cfg( dispatch_models=[ { "model_id": entry.model_id, "base_url": entry.base_url, "context_window": entry.context_window, "max_output_tokens": entry.max_output_tokens, "tier": entry.tier, "eligible_categories": entry.eligible_categories, } ], local_energy_enabled=False, ) monkeypatch.setattr(dispatcher, "cfg", cfg) # Seed the effective_context_window so PIN case guard passes result = _run_local_dispatch( entry, messages, category="file_summarization", body={} ) assert "payload" in result assert result["payload"] is payload assert result["request_id"].startswith("local-dispatch-") assert len(call_capture) == 1 assert call_capture[0]["json"]["stream"] is False assert call_capture[0]["json"]["model"] == entry.model_id assert "temperature" not in call_capture[0]["json"] assert "tools" not in call_capture[0]["json"] assert "tool_choice" not in call_capture[0]["json"] def test_local_dispatch_response_non_streaming(entry): """Response shape: id=request_id, X-Router-Model, X-Router-Verification, right model id.""" payload = _good_ollama_payload() rid = "local-dispatch-aabbcc" resp = _local_dispatch_response(payload, entry, streaming=False, request_id=rid) assert resp.status_code == 200 body = resp.body.decode() data = json.loads(body) assert data["id"] == rid assert data["model"] == entry.model_id assert data["choices"][0]["message"]["content"] == "The file refactors the auth module." assert data["choices"][0]["finish_reason"] == "stop" assert data["usage"]["prompt_tokens"] == 30 assert data["usage"]["completion_tokens"] == 12 assert resp.headers["X-Router-Model"] == entry.model_id assert "X-Router-Verification" in resp.headers @pytest.mark.anyio async def test_local_dispatch_response_streaming(entry): """SSE emits content chunk, stop+usage chunk, [DONE].""" payload = _good_ollama_payload() rid = "local-dispatch-ff00ff" resp = _local_dispatch_response(payload, entry, streaming=True, request_id=rid) lines = [] async for chunk in resp.body_iterator: lines.append(chunk) # content chunk chunk0 = json.loads(lines[0].decode().removeprefix("data: ")) assert chunk0["id"] == rid assert chunk0["choices"][0]["delta"]["content"] == "The file refactors the auth module." assert chunk0["choices"][0]["finish_reason"] is None # stop+usage chunk chunk1 = json.loads(lines[1].decode().removeprefix("data: ")) assert chunk1["model"] == entry.model_id assert chunk1["choices"][0]["finish_reason"] == "stop" assert chunk1["usage"]["prompt_tokens"] == 30 assert chunk1["usage"]["completion_tokens"] == 12 # [DONE] assert lines[2] == b"data: [DONE]\n\n" def test_temperature_forwarded_when_sent(entry, messages, monkeypatch): """Temperature in body → forwarded to Ollama.""" captured = {} def fake_post(url, headers=None, json=None, timeout=None): captured["json"] = json class R: status_code = 200 def json(self): return _good_ollama_payload() return R() monkeypatch.setattr(dispatcher.requests, "post", fake_post) cfg = _make_cfg( dispatch_models=[ { "model_id": entry.model_id, "base_url": entry.base_url, "context_window": entry.context_window, "max_output_tokens": entry.max_output_tokens, "tier": entry.tier, "eligible_categories": ["coding_general"], } ], local_energy_enabled=False, ) monkeypatch.setattr(dispatcher, "cfg", cfg) _run_local_dispatch(entry, messages, category=None, body={"temperature": 0.7}) assert captured["json"]["temperature"] == 0.7 def test_temperature_not_sent_when_absent(entry, messages, monkeypatch): """No temperature in body → NOT forwarded.""" captured = {} def fake_post(url, headers=None, json=None, timeout=None): captured["json"] = json class R: status_code = 200 def json(self): return _good_ollama_payload() return R() monkeypatch.setattr(dispatcher.requests, "post", fake_post) cfg = _make_cfg( dispatch_models=[ { "model_id": entry.model_id, "base_url": entry.base_url, "context_window": entry.context_window, "max_output_tokens": entry.max_output_tokens, "tier": entry.tier, "eligible_categories": ["coding_general"], } ], local_energy_enabled=False, ) monkeypatch.setattr(dispatcher, "cfg", cfg) _run_local_dispatch(entry, messages, category=None, body={}) assert "temperature" not in captured.get("json", {}) # --- request body: max_tokens clamp ---------------------------------------- def test_max_tokens_clamped_to_entry_max(entry, messages, monkeypatch): """Client max_tokens > entry.max_output_tokens → clamped.""" captured = {} def fake_post(url, headers=None, json=None, timeout=None): captured["json"] = json class R: status_code = 200 def json(self): return _good_ollama_payload() return R() monkeypatch.setattr(dispatcher.requests, "post", fake_post) cfg = _make_cfg( dispatch_models=[ { "model_id": entry.model_id, "base_url": entry.base_url, "context_window": entry.context_window, "max_output_tokens": 1024, # lower than client max "tier": entry.tier, "eligible_categories": ["coding_general"], } ], local_energy_enabled=False, ) monkeypatch.setattr(dispatcher, "cfg", cfg) entry_with_small_max = LocalDispatchModel( model_id=entry.model_id, base_url=entry.base_url, context_window=entry.context_window, max_output_tokens=1024, tier=entry.tier, eligible_categories=entry.eligible_categories, ) _run_local_dispatch( entry_with_small_max, messages, category=None, body={"max_tokens": 4096} ) assert captured["json"]["max_tokens"] == 1024 def test_max_tokens_uses_entry_max_when_client_absent(entry, messages, monkeypatch): """No client max_tokens → uses entry.max_output_tokens.""" captured = {} def fake_post(url, headers=None, json=None, timeout=None): captured["json"] = json class R: status_code = 200 def json(self): return _good_ollama_payload() return R() monkeypatch.setattr(dispatcher.requests, "post", fake_post) cfg = _make_cfg( dispatch_models=[ { "model_id": entry.model_id, "base_url": entry.base_url, "context_window": entry.context_window, "max_output_tokens": 2048, "tier": entry.tier, "eligible_categories": ["coding_general"], } ], local_energy_enabled=False, ) monkeypatch.setattr(dispatcher, "cfg", cfg) _run_local_dispatch(entry, messages, category=None, body={}) assert captured["json"]["max_tokens"] == entry.max_output_tokens # --- tools/tool_choice stripping ------------------------------------------ def test_tools_stripped_from_body(entry, messages, monkeypatch): """tools/tool_choice removed from request body.""" captured = {} def fake_post(url, headers=None, json=None, timeout=None): captured["json"] = json class R: status_code = 200 def json(self): return _good_ollama_payload() return R() monkeypatch.setattr(dispatcher.requests, "post", fake_post) cfg = _make_cfg( dispatch_models=[ { "model_id": entry.model_id, "base_url": entry.base_url, "context_window": entry.context_window, "max_output_tokens": entry.max_output_tokens, "tier": entry.tier, "eligible_categories": ["coding_general"], } ], local_energy_enabled=False, ) monkeypatch.setattr(dispatcher, "cfg", cfg) body_with_tools = { "tools": [{"type": "function", "function": {"name": "echo"}}], "tool_choice": "auto", } _run_local_dispatch(entry, messages, category=None, body=body_with_tools) assert "tools" not in captured["json"] assert "tool_choice" not in captured["json"] def test_tools_stripped_debug_log(entry, messages, monkeypatch): """logs.debug('local_dispatch_tools_stripped') called once per call.""" logs_debug_calls = [] def fake_post(url, headers=None, json=None, timeout=None): class R: status_code = 200 def json(self): return _good_ollama_payload() return R() monkeypatch.setattr(dispatcher.requests, "post", fake_post) with mock.patch.object( dispatcher.logs, "debug", side_effect=lambda name, **kw: logs_debug_calls.append((name, kw)), ): cfg = _make_cfg( dispatch_models=[ { "model_id": entry.model_id, "base_url": entry.base_url, "context_window": entry.context_window, "max_output_tokens": entry.max_output_tokens, "tier": entry.tier, "eligible_categories": ["coding_general"], } ], local_energy_enabled=False, ) monkeypatch.setattr(dispatcher, "cfg", cfg) _run_local_dispatch(entry, messages, category=None, body={}) assert "local_dispatch_tools_stripped" in {name for name, _ in logs_debug_calls} # --- failure paths: RequestException → HTTPException 502 ------------------- def test_request_exception_raises_502(entry, messages, monkeypatch): """requests.RequestException → HTTPException(502).""" def fake_post(url, headers=None, json=None, timeout=None): raise dispatcher.requests.RequestException("connection refused") monkeypatch.setattr(dispatcher.requests, "post", fake_post) cfg = _make_cfg( dispatch_models=[ { "model_id": entry.model_id, "base_url": entry.base_url, "context_window": entry.context_window, "max_output_tokens": entry.max_output_tokens, "tier": entry.tier, "eligible_categories": ["coding_general"], } ], local_energy_enabled=False, ) monkeypatch.setattr(dispatcher, "cfg", cfg) from fastapi import HTTPException with pytest.raises(HTTPException) as exc: _run_local_dispatch(entry, messages, category=None, body={}) assert exc.value.status_code == 502 assert entry.model_id in str(exc.value.detail) # --- failure paths: non-200 status → HTTPException 502 --------------------- def test_non_200_status_raises_502(entry, messages, monkeypatch): """Ollama returns 500 → HTTPException(502).""" def fake_post(url, headers=None, json=None, timeout=None): class R: status_code = 500 return R() monkeypatch.setattr(dispatcher.requests, "post", fake_post) cfg = _make_cfg( dispatch_models=[ { "model_id": entry.model_id, "base_url": entry.base_url, "context_window": entry.context_window, "max_output_tokens": entry.max_output_tokens, "tier": entry.tier, "eligible_categories": ["coding_general"], } ], local_energy_enabled=False, ) monkeypatch.setattr(dispatcher, "cfg", cfg) from fastapi import HTTPException with pytest.raises(HTTPException) as exc: _run_local_dispatch(entry, messages, category=None, body={}) assert exc.value.status_code == 502 # --- failure paths: unparseable JSON → HTTPException 502 ------------------- def test_unparseable_json_raises_502(entry, messages, monkeypatch): """Ollama returns non-JSON body → HTTPException(502).""" def fake_post(url, headers=None, json=None, timeout=None): class R: status_code = 200 def json(self): raise ValueError("no JSON") return R() monkeypatch.setattr(dispatcher.requests, "post", fake_post) cfg = _make_cfg( dispatch_models=[ { "model_id": entry.model_id, "base_url": entry.base_url, "context_window": entry.context_window, "max_output_tokens": entry.max_output_tokens, "tier": entry.tier, "eligible_categories": ["coding_general"], } ], local_energy_enabled=False, ) monkeypatch.setattr(dispatcher, "cfg", cfg) from fastapi import HTTPException with pytest.raises(HTTPException) as exc: _run_local_dispatch(entry, messages, category=None, body={}) assert exc.value.status_code == 502 # --- failure paths: empty content → HTTPException 502 ---------------------- def test_empty_content_raises_502(entry, messages, monkeypatch): """Ollama returns choices with empty content → HTTPException(502).""" empty_payload = { "id": "cmpl-0", "choices": [{"message": {"role": "assistant", "content": ""}, "finish_reason": "stop"}], } def fake_post(url, headers=None, json=None, timeout=None): class R: status_code = 200 def json(self): return empty_payload return R() monkeypatch.setattr(dispatcher.requests, "post", fake_post) cfg = _make_cfg( dispatch_models=[ { "model_id": entry.model_id, "base_url": entry.base_url, "context_window": entry.context_window, "max_output_tokens": entry.max_output_tokens, "tier": entry.tier, "eligible_categories": ["coding_general"], } ], local_energy_enabled=False, ) monkeypatch.setattr(dispatcher, "cfg", cfg) from fastapi import HTTPException with pytest.raises(HTTPException) as exc: _run_local_dispatch(entry, messages, category=None, body={}) assert exc.value.status_code == 502 # --- metering: local_energy.measure called and _log_local_energy captured --- def test_metering_calls_measure_and_log_local_energy(entry, messages, monkeypatch): """When metered: measure context manager enters/exits, _log_local_energy called.""" captured_log = [] def fake_measure(**kw): class MCtx: def __enter__(s): return s def __exit__(s, *a): pass avg_power_watts = 500.0 duration_seconds = 1.5 return MCtx() def fake_log(model_id=None, call_type=None, measurement=None, **kw): captured_log.append({"model_id": model_id, "call_type": call_type}) def fake_post(url, headers=None, json=None, timeout=None): class R: status_code = 200 def json(self): return _good_ollama_payload() return R() monkeypatch.setattr(dispatcher.local_energy, "measure", fake_measure) monkeypatch.setattr(dispatcher, "_log_local_energy", fake_log) monkeypatch.setattr(dispatcher.requests, "post", fake_post) cfg = _make_cfg( dispatch_models=[ { "model_id": entry.model_id, "base_url": "http://localhost:11434/v1", "context_window": entry.context_window, "max_output_tokens": entry.max_output_tokens, "tier": entry.tier, "eligible_categories": ["coding_general"], } ], local_energy_enabled=True, ) monkeypatch.setattr(dispatcher, "cfg", cfg) assert entry.model_id in cfg.local_energy_dispatch_models _run_local_dispatch(entry, messages, category="file_summarization", body={}) assert len(captured_log) == 1 assert captured_log[0]["model_id"] == entry.model_id assert captured_log[0]["call_type"] == "file_summarization" def test_metering_not_when_disabled(entry, messages, monkeypatch): """When local_energy.enabled is False: measure not called, _log_local_energy not called.""" measured = [] logged = [] def fake_measure(**kw): measured.append(True) class MCtx: def __enter__(s): return s def __exit__(s, *a): pass avg_power_watts = None duration_seconds = 0.0 return MCtx() def fake_log(**kw): logged.append(True) def fake_post(url, headers=None, json=None, timeout=None): class R: status_code = 200 def json(self): return _good_ollama_payload() return R() monkeypatch.setattr(dispatcher.local_energy, "measure", fake_measure) monkeypatch.setattr(dispatcher, "_log_local_energy", fake_log) monkeypatch.setattr(dispatcher.requests, "post", fake_post) cfg = _make_cfg( dispatch_models=[ { "model_id": entry.model_id, "base_url": entry.base_url, "context_window": entry.context_window, "max_output_tokens": entry.max_output_tokens, "tier": entry.tier, "eligible_categories": ["coding_general"], } ], local_energy_enabled=False, ) monkeypatch.setattr(dispatcher, "cfg", cfg) _run_local_dispatch(entry, messages, category=None, body={}) assert measured == [] assert logged == [] # --- circuit breaker: record_success on happy path -------------------------- def test_circuit_breaker_record_success_on_happy(entry, messages, monkeypatch): """Success path calls record_success with (model_id, 'ollama-local').""" calls = [] def fake_post(url, headers=None, json=None, timeout=None): class R: status_code = 200 def json(self): return _good_ollama_payload() return R() def fake_record_failure(model_id, provider, *a, **kw): calls.append(("failure", model_id, provider)) def fake_record_success(model_id, provider): calls.append(("success", model_id, provider)) monkeypatch.setattr(dispatcher.requests, "post", fake_post) monkeypatch.setattr(dispatcher.circuit_breaker, "record_failure", fake_record_failure) monkeypatch.setattr(dispatcher.circuit_breaker, "record_success", fake_record_success) cfg = _make_cfg( dispatch_models=[ { "model_id": entry.model_id, "base_url": entry.base_url, "context_window": entry.context_window, "max_output_tokens": entry.max_output_tokens, "tier": entry.tier, "eligible_categories": ["coding_general"], } ], local_energy_enabled=False, ) monkeypatch.setattr(dispatcher, "cfg", cfg) monkeypatch.setattr(cfg.circuit_breaker, "enabled", True) _run_local_dispatch(entry, messages, category="diff_checking", body={}) assert calls == [("success", entry.model_id, "ollama-local")] # --- circuit breaker: record_failure before each 502 path ------------------- def test_circuit_breaker_record_failure_on_request_exception(entry, messages, monkeypatch): """Connection error records circuit_breaker record_failure before raising 502.""" calls = [] def fake_post(url, headers=None, json=None, timeout=None): raise dispatcher.requests.RequestException("connection") def fake_record_failure(model_id, provider, *a, **kw): calls.append(("failure", model_id, provider)) def fake_record_success(model_id, provider): calls.append(("success", model_id, provider)) monkeypatch.setattr(dispatcher.requests, "post", fake_post) monkeypatch.setattr(dispatcher.circuit_breaker, "record_failure", fake_record_failure) monkeypatch.setattr(dispatcher.circuit_breaker, "record_success", fake_record_success) cfg = _make_cfg( dispatch_models=[ { "model_id": entry.model_id, "base_url": entry.base_url, "context_window": entry.context_window, "max_output_tokens": entry.max_output_tokens, "tier": entry.tier, "eligible_categories": ["coding_general"], } ], local_energy_enabled=False, ) monkeypatch.setattr(dispatcher, "cfg", cfg) monkeypatch.setattr(cfg.circuit_breaker, "enabled", True) from fastapi import HTTPException with pytest.raises(HTTPException) as exc: _run_local_dispatch(entry, messages, category="diff_checking", body={}) assert exc.value.status_code == 502 assert calls == [("failure", entry.model_id, "ollama-local")] def test_circuit_breaker_record_failure_on_500_status(entry, messages, monkeypatch): """Non-200 status records circuit_breaker record_failure before raising 502.""" calls = [] def fake_post(url, headers=None, json=None, timeout=None): class R: status_code = 500 return R() def fake_record_failure(model_id, provider, *a, **kw): calls.append(("failure", model_id, provider)) monkeypatch.setattr(dispatcher.requests, "post", fake_post) monkeypatch.setattr(dispatcher.circuit_breaker, "record_failure", fake_record_failure) cfg = _make_cfg( dispatch_models=[ { "model_id": entry.model_id, "base_url": entry.base_url, "context_window": entry.context_window, "max_output_tokens": entry.max_output_tokens, "tier": entry.tier, "eligible_categories": ["coding_general"], } ], local_energy_enabled=False, ) monkeypatch.setattr(dispatcher, "cfg", cfg) monkeypatch.setattr(cfg.circuit_breaker, "enabled", True) from fastapi import HTTPException with pytest.raises(HTTPException) as exc: _run_local_dispatch(entry, messages, category=None, body={}) assert exc.value.status_code == 502 assert calls == [("failure", entry.model_id, "ollama-local")] def test_circuit_breaker_record_failure_on_unparseable_json(entry, messages, monkeypatch): """Unparseable JSON records circuit_breaker record_failure before raising 502.""" calls = [] def fake_post(url, headers=None, json=None, timeout=None): class R: status_code = 200 def json(self): raise ValueError("no json") return R() def fake_record_failure(model_id, provider, *a, **kw): calls.append(("failure", model_id, provider)) monkeypatch.setattr(dispatcher.requests, "post", fake_post) monkeypatch.setattr(dispatcher.circuit_breaker, "record_failure", fake_record_failure) cfg = _make_cfg( dispatch_models=[ { "model_id": entry.model_id, "base_url": entry.base_url, "context_window": entry.context_window, "max_output_tokens": entry.max_output_tokens, "tier": entry.tier, "eligible_categories": ["coding_general"], } ], local_energy_enabled=False, ) monkeypatch.setattr(dispatcher, "cfg", cfg) monkeypatch.setattr(cfg.circuit_breaker, "enabled", True) from fastapi import HTTPException with pytest.raises(HTTPException) as exc: _run_local_dispatch(entry, messages, category=None, body={}) assert exc.value.status_code == 502 assert calls == [("failure", entry.model_id, "ollama-local")] def test_circuit_breaker_record_failure_on_empty_content(entry, messages, monkeypatch): """Empty content records circuit_breaker record_failure before raising 502.""" calls = [] def fake_post(url, headers=None, json=None, timeout=None): class R: status_code = 200 def json(self): return { "choices": [{"message": {"content": ""}, "finish_reason": "stop"}] } return R() def fake_record_failure(model_id, provider, *a, **kw): calls.append(("failure", model_id, provider)) monkeypatch.setattr(dispatcher.requests, "post", fake_post) monkeypatch.setattr(dispatcher.circuit_breaker, "record_failure", fake_record_failure) cfg = _make_cfg( dispatch_models=[ { "model_id": entry.model_id, "base_url": entry.base_url, "context_window": entry.context_window, "max_output_tokens": entry.max_output_tokens, "tier": entry.tier, "eligible_categories": ["coding_general"], } ], local_energy_enabled=False, ) monkeypatch.setattr(dispatcher, "cfg", cfg) monkeypatch.setattr(cfg.circuit_breaker, "enabled", True) from fastapi import HTTPException with pytest.raises(HTTPException) as exc: _run_local_dispatch(entry, messages, category=None, body={}) assert exc.value.status_code == 502 assert calls == [("failure", entry.model_id, "ollama-local")] import sqlite3 from unittest import mock import yaml from starlette.testclient import TestClient from dispatcher import Classification, app SCHEMA = (ROOT / "config" / "schema.sql").read_text() CLOUD_MODEL = "cloud-gpt" LOCAL_MODEL = "qwen2.5-coder-router:14b" @pytest.fixture def router_local(tmp_path, monkeypatch): """TestClient fixture with a throwaway catalog and one local dispatch entry.""" db_path = tmp_path / "test.db" conn = sqlite3.connect(db_path) conn.executescript(SCHEMA) conn.commit() conn.close() raw = yaml.safe_load((ROOT / "config" / "config.yaml").read_text()) raw["dispatch_providers"] = { "neuralwatt": {"base_url": "https://api.neuralwatt.com/v1", "api_key_env": "NEURALWATT_API_KEY"}, } raw["local_dispatch_models"] = [ { "model_id": LOCAL_MODEL, "base_url": "http://localhost:11434/v1", "context_window": 16384, "max_output_tokens": 2048, "tier": 1, "eligible_categories": ["file_summarization", "diff_checking"], } ] cfg = RouterConfig(**raw) monkeypatch.setattr(dispatcher, "cfg", cfg) 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) yield db_path def _seed_local_row(db_path): conn = sqlite3.connect(db_path) conn.row_factory = sqlite3.Row 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, eligible_categories ) VALUES (?, 'ollama-local', ?, 1, 16384, 7168, 2048, NULL, NULL, 0, 0, 'standard', 'default', 'full', 'public', 'active', '2026-08-22T00:00:00+00:00', 'file_summarization,diff_checking') """, (LOCAL_MODEL, LOCAL_MODEL), ) conn.commit() conn.close() def _seed_cloud_row(db_path, model_id=CLOUD_MODEL, tier=2, cost_prompt=0.10, cost_completion=0.30): 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 (?, 'neuralwatt', ?, ?, 262128, 192500, 16384, ?, ?, 1, 1, 'standard', 'default', 'full', 'public', 'active', '2026-08-22T00:00:00+00:00') """, (model_id, model_id, tier, cost_prompt, cost_completion), ) conn.commit() conn.close() def test_chat_auto_routes_to_local(router_local, monkeypatch): """model:auto selecting local row -> 200, X-Router-Model, decision row.""" _seed_local_row(router_local) monkeypatch.setattr( dispatcher, "classify", lambda task, context: Classification( task_category="file_summarization", task_tier=1, required_context_tokens=100, confidence=0.9, ), ) calls = [] def fake_post(url, headers=None, json=None, timeout=None): calls.append({"url": url, "json": json, "stream": False}) class R: status_code = 200 def json(self): return { "id": "local-1", "choices": [ { "message": {"content": "summary"}, "finish_reason": "stop", } ], "usage": {"prompt_tokens": 5, "completion_tokens": 2}, } return R() monkeypatch.setattr(dispatcher.requests, "post", fake_post) monkeypatch.setenv("NEURALWATT_API_KEY", "x") with TestClient(app) as client: resp = client.post( "/v1/chat/completions", json={"model": "auto", "messages": [{"role": "user", "content": "Summarize"}]}, ) assert resp.status_code == 200 data = resp.json() assert data["model"] == LOCAL_MODEL assert resp.headers["X-Router-Model"] == LOCAL_MODEL assert calls assert "localhost:11434" in calls[0]["url"] conn = sqlite3.connect(router_local) conn.row_factory = sqlite3.Row row = conn.execute( "SELECT selected_provider FROM route_decisions ORDER BY id DESC LIMIT 1" ).fetchone() conn.close() assert row["selected_provider"] == "ollama-local" def test_chat_auto_does_not_select_local_for_unsupported_category( router_local, monkeypatch ): """Category outside local row's eligible list -> cloud row chosen.""" _seed_local_row(router_local) _seed_cloud_row(router_local) monkeypatch.setattr( dispatcher, "classify", lambda task, context: Classification( task_category="coding_general", task_tier=1, required_context_tokens=100, confidence=0.9, ), ) calls = [] def fake_post(url, headers=None, json=None, timeout=None): calls.append({"url": url, "json": json}) class R: status_code = 200 def json(self): return { "id": "cloud-1", "model": CLOUD_MODEL, "choices": [ { "message": {"content": "hello"}, "finish_reason": "stop", } ], "usage": {"prompt_tokens": 3, "completion_tokens": 1}, } return R() monkeypatch.setattr(dispatcher.requests, "post", fake_post) monkeypatch.setenv("NEURALWATT_API_KEY", "x") with TestClient(app) as client: 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"] == CLOUD_MODEL assert calls[0]["json"]["model"] == CLOUD_MODEL def test_chat_pinned_local_dispatches_locally(router_local, monkeypatch): """Pinned local model id -> POST to local URL, not NeuralWatt.""" _seed_local_row(router_local) calls = [] def fake_post(url, headers=None, json=None, timeout=None): calls.append({"url": url, "json": json}) class R: status_code = 200 def json(self): return { "id": "local-2", "choices": [ { "message": {"content": "pinned answer"}, "finish_reason": "stop", } ], "usage": {"prompt_tokens": 4, "completion_tokens": 3}, } return R() monkeypatch.setattr(dispatcher.requests, "post", fake_post) monkeypatch.setenv("NEURALWATT_API_KEY", "x") with TestClient(app) as client: resp = client.post( "/v1/chat/completions", json={ "model": LOCAL_MODEL, "messages": [{"role": "user", "content": "Hi"}], }, ) assert resp.status_code == 200 assert resp.json()["model"] == LOCAL_MODEL assert "localhost:11434" in calls[0]["url"] def test_list_models_includes_local_row(router_local, monkeypatch): """GET /v1/models lists the local row with owned_by 'ollama-local'.""" _seed_local_row(router_local) monkeypatch.setenv("NEURALWATT_API_KEY", "x") with TestClient(app) as client: resp = client.get("/v1/models") assert resp.status_code == 200 data = resp.json()["data"] local_entries = [m for m in data if m["id"] == LOCAL_MODEL] assert len(local_entries) == 1 assert local_entries[0]["owned_by"] == "ollama-local" def test_dispatch_endpoint_local_with_telemetry(router_local, monkeypatch): _seed_local_row(router_local) def fake_post(url, headers=None, json=None, timeout=None): class R: status_code = 200 def json(self): return { "id": "local-3", "choices": [ { "message": {"content": "dispatched"}, "finish_reason": "stop", } ], "usage": {"prompt_tokens": 7, "completion_tokens": 4}, } return R() monkeypatch.setattr(dispatcher.requests, "post", fake_post) calls = {"measure": 0, "log": []} class FakeMeasurement: def __enter__(s): return s def __exit__(s, *a): return None avg_power_watts = 250.0 duration_seconds = 2.0 def fake_measure(**kw): calls["measure"] += 1 return FakeMeasurement() def fake_log(**kw): calls["log"].append(kw) monkeypatch.setattr(dispatcher.local_energy, "measure", fake_measure) monkeypatch.setattr(dispatcher, "_log_local_energy", fake_log) raw = yaml.safe_load((ROOT / "config" / "config.yaml").read_text()) raw["dispatch_providers"] = { "ollama-local": { "base_url": "http://localhost:11434/v1", "api_key_env": "OLLAMA_API_KEY", } } raw["dispatch_settings"] = {"default_provider": "ollama-local"} raw["local_dispatch_models"] = [ { "model_id": LOCAL_MODEL, "base_url": "http://localhost:11434/v1", "context_window": 16384, "max_output_tokens": 2048, "tier": 1, "eligible_categories": ["file_summarization"], } ] raw["local_energy"]["enabled"] = True raw["local_energy"]["tariff_usd_per_kwh"] = 8.0 cfg = RouterConfig(**raw) monkeypatch.setattr(dispatcher, "cfg", cfg) monkeypatch.setattr(dispatcher.cfg.database, "path", str(router_local)) with TestClient(app) as client: resp = client.post( "/dispatch", json={ "task": "Summarize this", "task_category": "file_summarization", "task_tier": 1, "required_context_tokens": 50, }, ) assert resp.status_code == 200 body = resp.json() assert body["content"] == "dispatched" assert body["prompt_tokens"] == 7 assert body["completion_tokens"] == 4 assert body["telemetry"]["avg_power_watts"] == 250.0 assert body["telemetry"]["duration_seconds"] == 2.0 assert body["telemetry"]["energy_kwh"] > 0 assert calls["measure"] == 1, "local_energy.measure must be called once per dispatch" assert len(calls["log"]) == 1, "exactly one local_energy_observations row per dispatch" assert calls["log"][0]["model_id"] == LOCAL_MODEL assert calls["log"][0]["call_type"] == "file_summarization" def test_chat_auto_tools_stripped_for_file_summarization(router_local, monkeypatch): """Tools-carrying auto request classified file_summarization -> no 'tools' key forwarded.""" _seed_local_row(router_local) monkeypatch.setattr( dispatcher, "classify", lambda task, context: Classification( task_category="file_summarization", task_tier=1, required_context_tokens=100, confidence=0.9, ), ) calls = [] def fake_post(url, headers=None, json=None, timeout=None): calls.append({"url": url, "json": json}) class R: status_code = 200 def json(self): return { "id": "local-tool-1", "choices": [ {"message": {"content": "ok"}, "finish_reason": "stop"} ], "usage": {"prompt_tokens": 1, "completion_tokens": 1}, } return R() monkeypatch.setattr(dispatcher.requests, "post", fake_post) monkeypatch.setenv("NEURALWATT_API_KEY", "x") with TestClient(app) as client: resp = client.post( "/v1/chat/completions", json={ "model": "auto", "messages": [{"role": "user", "content": "Summarize"}], "tools": [{"type": "function", "function": {"name": "x"}}], }, ) assert resp.status_code == 200 assert "tools" not in calls[0]["json"] def test_pinned_unknown_model_still_passthroughs_to_neuralwatt( router_local, monkeypatch ): """Pinning an unknown model keeps llm-router/ alias strip and hits cloud.""" conn = sqlite3.connect(router_local) 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 ('known-cloud', 'neuralwatt', 'known-cloud', 2, 262128, 192500, 16384, 0.10, 0.30, 1, 1, 'standard', 'default', 'full', 'public', 'active', '2026-08-22T00:00:00+00:00') """ ) conn.commit() conn.close() calls = [] def fake_post(url, headers=None, json=None, timeout=None): calls.append({"url": url, "json": json}) class R: status_code = 200 def json(self): return { "id": "cloud-2", "model": "known-cloud", "choices": [ {"message": {"content": "cloud"}, "finish_reason": "stop"} ], "usage": {"prompt_tokens": 2, "completion_tokens": 1}, } return R() monkeypatch.setattr(dispatcher.requests, "post", fake_post) monkeypatch.setenv("NEURALWATT_API_KEY", "x") with TestClient(app) as client: resp = client.post( "/v1/chat/completions", json={ "model": "llm-router/known-cloud", "messages": [{"role": "user", "content": "Hi"}], }, ) assert resp.status_code == 200 assert calls[0]["json"]["model"] == "known-cloud" from eval_proficiency import _endpoint_for, add_self_eval def _make_eval_cfg(): raw = yaml.safe_load((ROOT / "config" / "config.yaml").read_text()) raw["dispatch_providers"] = { "neuralwatt": { "base_url": "https://api.neuralwatt.com/v1", "api_key_env": "NEURALWATT_API_KEY", } } raw["local_dispatch_models"] = [ { "model_id": LOCAL_MODEL, "base_url": "http://localhost:11434/v1", "context_window": 16384, "max_output_tokens": 2048, "tier": 1, "eligible_categories": ["file_summarization", "diff_checking"], } ] return RouterConfig(**raw) def test_endpoint_for_ollama_local(): """Local identity resolves to local base_url, no auth header, custom timeout.""" identity = { "model_id": LOCAL_MODEL, "provider": "ollama-local", } cfg = _make_eval_cfg() base_url, api_key, kwargs = _endpoint_for(identity, cfg) assert base_url == "http://localhost:11434/v1" assert api_key is None assert kwargs == {"timeout": 120.0} def test_endpoint_for_ollama_local_with_api_key_env(monkeypatch): """Local entry with api_key_env forwards the env value when set.""" monkeypatch.setenv("OLLAMA_API_KEY", "local-secret") raw = yaml.safe_load((ROOT / "config" / "config.yaml").read_text()) raw["dispatch_providers"] = { "ollama-local": { "base_url": "http://localhost:11434/v1", "api_key_env": "OLLAMA_API_KEY", } } raw["dispatch_settings"] = {"default_provider": "ollama-local"} raw["local_dispatch_models"] = [ { "model_id": LOCAL_MODEL, "base_url": "http://localhost:11434/v1", "api_key_env": "OLLAMA_API_KEY", "timeout_seconds": 60.0, "context_window": 16384, "tier": 1, "eligible_categories": ["file_summarization"], } ] cfg = RouterConfig(**raw) base_url, api_key, kwargs = _endpoint_for( {"model_id": LOCAL_MODEL, "provider": "ollama-local"}, cfg ) assert base_url == "http://localhost:11434/v1" assert api_key == "local-secret" assert kwargs == {"timeout": 60.0} def test_endpoint_for_neuralwatt(): """Cloud identity resolves to neuralwatt provider settings.""" cfg = _make_eval_cfg() identity = {"model_id": "kimi-k3", "provider": "neuralwatt"} base_url, _api_key, kwargs = _endpoint_for(identity, cfg) assert base_url == cfg.dispatch_providers["neuralwatt"].base_url assert kwargs == {"timeout": 300.0} def test_add_self_eval_ollama_local_writes_proficiency_row(tmp_path): """Writing self-eval for provider='ollama-local' keys by that provider.""" db_path = tmp_path / "test.db" conn = sqlite3.connect(str(db_path)) conn.executescript((ROOT / "config" / "schema.sql").read_text()) conn.row_factory = sqlite3.Row 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, eligible_categories ) VALUES (?, 'ollama-local', ?, 1, 16384, 7168, 2048, NULL, NULL, 0, 0, 'standard', 'default', 'full', 'public', 'active', '2026-08-22T00:00:00+00:00', 'file_summarization,diff_checking') """, ("qwen2.5-coder-router:14b", "qwen2.5-coder-router:14b"), ) conn.commit() cfg = _make_eval_cfg() add_self_eval( conn, cfg, "qwen2.5-coder-router:14b", "ollama-local", "diff_checking", [1.0, 0.5] ) row = conn.execute( """ SELECT model_id, provider, category, self_eval_score, self_eval_samples FROM proficiency WHERE model_id = ? AND provider = ? AND category = ? """, ("qwen2.5-coder-router:14b", "ollama-local", "diff_checking"), ).fetchone() assert row is not None assert row["provider"] == "ollama-local" assert row["category"] == "diff_checking" assert row["self_eval_samples"] == 2 assert round(row["self_eval_score"], 3) == 0.75 conn.close() # --- POST /outcome attribution for local-dispatch answers ------------------- @pytest.fixture def router_local_outcome(router_local, monkeypatch): """Same throwaway DB, but local_energy enabled and no verify/classify.""" raw = yaml.safe_load((ROOT / "config" / "config.yaml").read_text()) raw["dispatch_providers"] = {} raw["dispatch_settings"] = {"default_provider": "ollama-local"} raw["local_dispatch_models"] = [ { "model_id": LOCAL_MODEL, "base_url": "http://localhost:11434/v1", "context_window": 16384, "max_output_tokens": 2048, "tier": 1, "eligible_categories": ["file_summarization", "diff_checking"], } ] raw["local_energy"]["enabled"] = True raw["local_energy"]["tariff_usd_per_kwh"] = 8.0 raw["dispatch_providers"]["neuralwatt"] = { "base_url": "https://api.neuralwatt.com/v1", "api_key_env": "NEURALWATT_API_KEY", } raw["dispatch_settings"] = {"default_provider": "neuralwatt"} cfg = RouterConfig(**raw) monkeypatch.setattr(dispatcher, "cfg", cfg) monkeypatch.setattr(dispatcher.cfg.database, "path", str(router_local)) 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) return router_local def _seed_local_energy_row( db_path, *, request_id, session_dir=None, model_id=LOCAL_MODEL, call_type="file_summarization", observed_at=None, with_request_id=True, ): """Insert a local_energy_observations row; can omit request_id/session_dir columns.""" if observed_at is None: observed_at = datetime.now(timezone.utc).isoformat() conn = sqlite3.connect(db_path) cols = "observed_at, model_id, call_type, avg_power_watts, duration_seconds" vals = [observed_at, model_id, call_type, 250.0, 2.0] if with_request_id: cols += ", request_id" vals.append(request_id) if session_dir is not None: cols += ", session_dir" vals.append(session_dir) placeholders = ", ".join("?" for _ in vals) conn.execute( f"INSERT INTO local_energy_observations ({cols}) VALUES ({placeholders})", vals, ) conn.commit() conn.close() def test_outcome_local_dispatch_request_id_records_verification( router_local_outcome, monkeypatch ): """(1) local row with request_id + POST /outcome -> 200, ollama-local client_outcome.""" _seed_local_row(router_local_outcome) rid = "local-dispatch-deadbeef" _seed_local_energy_row(router_local_outcome, request_id=rid) monkeypatch.setenv("NEURALWATT_API_KEY", "x") with TestClient(app) as client: resp = client.post( "/outcome", json={"request_id": rid, "ok": True, "detail": "tests passed"}, ) assert resp.status_code == 200 data = resp.json() assert data["recorded"] is True assert data["model_id"] == LOCAL_MODEL assert data["task_category"] == "file_summarization" conn = sqlite3.connect(router_local_outcome) conn.row_factory = sqlite3.Row row = conn.execute( "SELECT model_id, provider, task_category, kind, verdict " "FROM verifications ORDER BY id DESC LIMIT 1" ).fetchone() conn.close() assert row is not None assert row["model_id"] == LOCAL_MODEL assert row["provider"] == "ollama-local" assert row["task_category"] == "file_summarization" assert row["kind"] == "client_outcome" assert row["verdict"] == "succeeded" def test_outcome_cloud_wins_when_request_id_collides(router_local_outcome, monkeypatch): """(2) same request_id exists in cloud energy_observations -> cloud wins.""" _seed_local_row(router_local_outcome) rid = "same-rid-cloud-and-local" _seed_local_energy_row(router_local_outcome, request_id=rid, model_id=LOCAL_MODEL) conn = sqlite3.connect(router_local_outcome) conn.row_factory = sqlite3.Row conn.execute( "INSERT INTO energy_observations (model_id, provider, request_id, task_category, observed_at) " "VALUES (?, 'neuralwatt', ?, 'coding_general', ?)", (CLOUD_MODEL, rid, datetime.now(timezone.utc).isoformat()), ) conn.commit() conn.close() monkeypatch.setenv("NEURALWATT_API_KEY", "x") with TestClient(app) as client: resp = client.post( "/outcome", json={"request_id": rid, "ok": False, "detail": "cloud failed"}, ) assert resp.status_code == 200 data = resp.json() assert data["model_id"] == CLOUD_MODEL assert data["task_category"] == "coding_general" def test_outcome_session_dir_local_attribution(router_local_outcome, monkeypatch): """(3) source matches session_dir on a local row -> resolves.""" _seed_local_row(router_local_outcome) rid = "local-dispatch-sourcedir" project_dir = "/home/alee/Sources/6krrt" _seed_local_energy_row( router_local_outcome, request_id=rid, call_type="diff_checking", session_dir=project_dir, ) monkeypatch.setenv("NEURALWATT_API_KEY", "x") with TestClient(app) as client: resp = client.post( "/outcome", json={"source": project_dir, "ok": True, "detail": "diff matched"}, ) assert resp.status_code == 200 data = resp.json() assert data["model_id"] == LOCAL_MODEL assert data["request_id"] == rid assert data["task_category"] == "diff_checking" def test_outcome_unknown_request_id_still_404(router_local_outcome, monkeypatch): """(4) request_id in neither table -> 404.""" _seed_local_row(router_local_outcome) monkeypatch.setenv("NEURALWATT_API_KEY", "x") with TestClient(app) as client: resp = client.post( "/outcome", json={"request_id": "does-not-exist", "ok": True}, ) assert resp.status_code == 404 def test_outcome_two_local_sessions_still_ambiguous( router_local_outcome, monkeypatch ): """(5) two distinct local sessions inside window -> 409, no weakening.""" _seed_local_row(router_local_outcome) now = datetime.now(timezone.utc) _seed_local_energy_row( router_local_outcome, request_id="local-session-a", session_dir="/tmp/projA", observed_at=now.isoformat(), ) _seed_local_energy_row( router_local_outcome, request_id="local-session-b", session_dir="/tmp/projB", observed_at=now.isoformat(), ) monkeypatch.setattr( dispatcher.cfg.verification, "outcome_attribution_window_seconds", 60, ) monkeypatch.setenv("NEURALWATT_API_KEY", "x") with TestClient(app) as client: resp = client.post( "/outcome", json={"ok": True, "detail": "no request id"}, ) assert resp.status_code == 409 assert "More than one conversation" in resp.text def test_outcome_pre_column_local_row_does_not_error(router_local_outcome, monkeypatch): """(6) writing a new-style local row into a DB missing request_id/session_dir doesn't error.""" # Simulate pre-Todo-11 state: drop the index first, then the columns. conn = sqlite3.connect(router_local_outcome) conn.execute("DROP INDEX IF EXISTS idx_local_energy_request") conn.execute("ALTER TABLE local_energy_observations DROP COLUMN request_id") conn.execute("ALTER TABLE local_energy_observations DROP COLUMN session_dir") conn.commit() conn.close() _seed_local_row(router_local_outcome) fake_measure = mock.MagicMock() fake_measure.__enter__.return_value.avg_power_watts = 250.0 fake_measure.__enter__.return_value.duration_seconds = 2.0 monkeypatch.setattr(dispatcher.local_energy, "measure", lambda **kw: fake_measure) def fake_post(url, headers=None, json=None, timeout=None): class R: status_code = 200 def json(self): return { "id": "local-compat", "choices": [ { "message": {"content": "compat"}, "finish_reason": "stop", } ], "usage": {"prompt_tokens": 5, "completion_tokens": 1}, } return R() monkeypatch.setattr(dispatcher.requests, "post", fake_post) monkeypatch.setenv("NEURALWATT_API_KEY", "x") with TestClient(app) as client: resp = client.post( "/v1/chat/completions", json={ "model": LOCAL_MODEL, "messages": [{"role": "user", "content": "Summarize"}], }, ) assert resp.status_code == 200 # The most important assertion: the log succeeded despite missing columns.