Ultraworked with [Sisyphus](https://github.com/code-yeongyu/oh-my-openagent) Co-authored-by: Sisyphus <clio-agent@sisyphuslabs.ai>
2298 lines
75 KiB
Python
2298 lines
75 KiB
Python
"""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.
|
|
|