Files
6krrt/tests/test_seed_local_dispatch.py
adlee-was-taken 1ba4ad3f2c fix: local dispatch metering + seed module LOC ceiling
- _run_local_dispatch now returns metering data it already measured,
  and dispatch_endpoint builds Telemetry from that instead of spawning a
  second local_energy.measure that wrapped nothing and double-logged.
- Split seed_local_dispatch_energy.py pure derivation into
  seed_local_dispatch_core.py so the CLI script stays under 250 LOC.
- Update imports/stubs and assert single measure + single log in the
  /dispatch telemetry test.

Verification: pytest target files 86 passed; full suite 1020 passed.
2026-09-02 02:09:59 -04:00

396 lines
14 KiB
Python

"""Tests for ``src/seed_local_dispatch_energy.py``.
Covers:
- ``derive_token_prices`` on synthetic samples with ground truth
- No-tariff refusal (exit 1, message)
- ``--dry-run`` makes zero HTTP calls and zero DB writes
- Write path updates the temp-DB models row exactly
"""
from __future__ import annotations
import sqlite3
import sys
from pathlib import Path
from unittest.mock import MagicMock, patch
import pytest
import yaml
TEST_DIR = Path(__file__).resolve().parent
ROOT = TEST_DIR.parent
SRC = ROOT / "src"
sys.path.insert(0, str(SRC))
from seed_local_dispatch_core import derive_token_prices
from seed_local_dispatch_energy import _is_loopback
class TestDeriveTokenPrices:
def test_recover_ground_truth(self):
tariff = 8.0
true_a = 1e-7
true_b = 5e-7
expected_prompt = tariff * true_a * 1e6
expected_completion = tariff * true_b * 1e6
samples = [(p, c, true_a * p + true_b * c)
for p in (100, 250, 500, 800, 1000)
for c in (50, 150, 300, 600, 1200)]
cost_p, cost_c, stats = derive_token_prices(samples, tariff)
assert cost_p == pytest.approx(expected_prompt, rel=1e-6)
assert cost_c == pytest.approx(expected_completion, rel=1e-6)
assert stats["r_squared"] == pytest.approx(1.0, abs=1e-10)
assert stats["slope_prompt"] == pytest.approx(true_a, rel=1e-6)
assert stats["slope_completion"] == pytest.approx(true_b, rel=1e-6)
def test_r_squared_near_one_with_noise(self):
tariff = 8.0
samples = []
noise = 1e-9
for p in (100, 500, 1000, 2000):
for c in (50, 200, 500):
energy = 1e-7 * p + 5e-7 * c + noise
samples.append((p, c, energy))
_, _, stats = derive_token_prices(samples, tariff)
assert stats["r_squared"] > 0.99
def test_asymmetric_slopes(self):
tariff = 0.12
samples = [(1000, 100, 0.0012), (500, 500, 0.0036), (200, 1000, 0.0060)]
cost_p, cost_c, _ = derive_token_prices(samples, tariff)
assert cost_c > cost_p
def test_single_sample_raises(self):
with pytest.raises(ValueError, match="need at least 2 samples"):
derive_token_prices([(100, 50, 0.001)], 8.0)
def test_empty_raises(self):
with pytest.raises(ValueError, match="need at least 2 samples"):
derive_token_prices([], 8.0)
def test_stats_contain_n_samples(self):
samples = [(100, 50, 0.0001), (200, 200, 0.0002)]
_, _, stats = derive_token_prices(samples, 8.0)
assert stats["n_samples"] == 2
assert "slope_prompt" in stats
assert "slope_completion" in stats
assert "r_squared" in stats
assert "median_prompt_tokens" in stats
assert "median_completion_tokens" in stats
class TestIsLoopback:
def test_localhost(self):
assert _is_loopback("http://localhost:11434/v1")
assert _is_loopback("http://localhost:11434")
def test_127_0_0_1(self):
assert _is_loopback("http://127.0.0.1:8080/v1")
def test_non_loopback(self):
assert not _is_loopback("http://10.0.0.5:11434")
assert not _is_loopback("https://api.example.com/v1")
def test_trailing_path_after_v1(self):
assert _is_loopback("http://127.0.0.1:11434/v1/test")
def _entry(model_id: str, base_url: str):
return MagicMock(model_id=model_id, base_url=base_url)
class TestCliRefusal:
def test_tariff_none_refuses(self, tmp_path, capsys):
db_path = tmp_path / "test.db"
_init_db(db_path)
_seed_local_model(db_path, "test-model")
import seed_local_dispatch_energy as mod
mock_cfg = MagicMock()
mock_cfg.local_energy.enabled = True
mock_cfg.local_energy.tariff_usd_per_kwh = None
mock_cfg.local_energy.sample_interval_seconds = 0.25
mock_cfg.local_energy.meter = "nvidia_smi"
mock_cfg.local_energy_dispatch_models = frozenset({"test-model"})
mock_cfg.database.path = str(db_path)
mock_cfg.local_dispatch_models = [
_entry("test-model", "http://localhost:11434/v1")
]
with patch.object(mod, "load_config", return_value=mock_cfg), patch.object(
sys, "argv", ["seed_local_dispatch_energy.py"]
):
result = mod.main()
assert result == 1
captured = capsys.readouterr()
assert "tariff" in captured.err.lower()
def test_local_energy_disabled_refuses(self, tmp_path, capsys):
db_path = tmp_path / "x.db"
_init_db(db_path)
_seed_local_model(db_path, "test-model")
import seed_local_dispatch_energy as mod
mock_cfg = MagicMock()
mock_cfg.local_energy.enabled = False
mock_cfg.local_energy.tariff_usd_per_kwh = 8.0
mock_cfg.local_energy.sample_interval_seconds = 0.25
mock_cfg.local_energy.meter = "nvidia_smi"
mock_cfg.local_energy_dispatch_models = frozenset({"test-model"})
mock_cfg.database.path = str(db_path)
mock_cfg.local_dispatch_models = [
_entry("test-model", "http://localhost:11434/v1")
]
with patch.object(mod, "load_config", return_value=mock_cfg), patch.object(
sys, "argv", ["seed_local_dispatch_energy.py"]
):
result = mod.main()
assert result == 1
captured = capsys.readouterr()
assert "enabled" in captured.err.lower()
class TestLoopbackRefusal:
def test_non_loopback_models_refuse(self, tmp_path, capsys):
db_path = tmp_path / "x.db"
_init_db(db_path)
_seed_local_model(db_path, "remote-model")
import seed_local_dispatch_energy as mod
mock_cfg = MagicMock()
mock_cfg.local_energy.enabled = True
mock_cfg.local_energy.tariff_usd_per_kwh = 8.0
mock_cfg.local_energy.sample_interval_seconds = 0.25
mock_cfg.local_energy.meter = "nvidia_smi"
mock_cfg.local_energy_dispatch_models = frozenset({"remote-model"})
mock_cfg.database.path = str(db_path)
mock_cfg.local_dispatch_models = [
_entry("remote-model", "http://10.0.0.5:11434/v1")
]
with patch.object(mod.sqlite3, "connect") as mock_connect, patch.object(
mod, "load_config", return_value=mock_cfg
), patch.object(sys, "argv", ["seed_local_dispatch_energy.py"]):
mock_conn = MagicMock()
mock_conn.execute.return_value.fetchall.return_value = [
{"model_id": "remote-model"}
]
mock_connect.return_value = mock_conn
result = mod.main()
assert result == 1
captured = capsys.readouterr()
assert "loopback" in captured.err.lower()
class TestDryRun:
def test_dry_run_zero_http_calls(self, tmp_path, capsys):
import seed_local_dispatch_energy as mod
db_path = tmp_path / "dry.db"
_init_db(db_path)
_seed_local_model(db_path, "dry-model")
mock_cfg = MagicMock()
mock_cfg.local_energy.enabled = True
mock_cfg.local_energy.tariff_usd_per_kwh = 8.0
mock_cfg.local_energy.sample_interval_seconds = 0.25
mock_cfg.local_energy.meter = "nvidia_smi"
mock_cfg.local_energy_dispatch_models = frozenset({"dry-model"})
mock_cfg.database.path = str(db_path)
mock_cfg.local_dispatch_models = [
_entry("dry-model", "http://localhost:11434/v1")
]
mock_conn = MagicMock()
mock_conn.execute.return_value.fetchall.return_value = [
{"model_id": "dry-model"}
]
mock_conn.row_factory = None
with patch.object(mod.sqlite3, "connect", return_value=mock_conn), patch.object(
mod, "load_config", return_value=mock_cfg
), patch.object(sys, "argv", ["seed_local_dispatch_energy.py", "--dry-run"]):
result = mod.main()
assert result == 0
captured = capsys.readouterr()
assert "sum_small" in captured.out
assert "dry-model" in captured.out
assert not mock_conn.commit.called
def test_dry_run_zero_db_writes(self, tmp_path):
import seed_local_dispatch_energy as mod
db_path = tmp_path / "dry2.db"
_init_db(db_path)
_seed_local_model(db_path, "dry-model2")
mock_cfg = MagicMock()
mock_cfg.local_energy.enabled = True
mock_cfg.local_energy.tariff_usd_per_kwh = 8.0
mock_cfg.local_energy.sample_interval_seconds = 0.25
mock_cfg.local_energy.meter = "nvidia_smi"
mock_cfg.local_energy_dispatch_models = frozenset({"dry-model2"})
mock_cfg.database.path = str(db_path)
mock_cfg.local_dispatch_models = [
_entry("dry-model2", "http://localhost:11434/v1")
]
mock_conn = MagicMock(spec=sqlite3.Connection)
mock_conn.execute.return_value.fetchall.return_value = [
{"model_id": "dry-model2"}
]
mock_conn.row_factory = None
with patch.object(mod.sqlite3, "connect", return_value=mock_conn), patch.object(
mod, "load_config", return_value=mock_cfg
), patch.object(sys, "argv", ["seed_local_dispatch_energy.py", "--dry-run"]):
mod.main()
assert not mock_conn.commit.called
updates = 0
for call in mock_conn.execute.call_args_list:
stmt = call[0][0] if call[0] else ""
if stmt.strip().startswith("UPDATE"):
updates += 1
assert updates == 0
class TestWritePath:
def test_write_path_updates_models(self, tmp_path):
db_path = tmp_path / "write.db"
_init_db(db_path)
_seed_local_model(db_path, "write-model")
tariff = 8.0
true_a, true_b = 1e-7, 5e-7
samples_data = [
{"shape": name, "sample": i + 1, "prompt_tokens": p,
"completion_tokens": c, "gross_kwh": true_a * p + true_b * c,
"avg_watts": 50.0, "duration_s": 1.0}
for name in ("sum_small", "sum_large", "diff_small", "long_answer")
for i, (p, c) in enumerate([
(100, 50), (200, 100), (1000, 500), (500, 250), (300, 300)
])
]
mock_cfg = MagicMock()
mock_cfg.local_energy.enabled = True
mock_cfg.local_energy.tariff_usd_per_kwh = tariff
mock_cfg.local_energy.sample_interval_seconds = 0.25
mock_cfg.local_energy.meter = "nvidia_smi"
mock_cfg.local_energy_dispatch_models = frozenset({"write-model"})
mock_cfg.database.path = str(db_path)
mock_cfg.local_dispatch_models = [
_entry("write-model", "http://localhost:11434/v1")
]
mock_conn = MagicMock()
mock_conn.execute.return_value.fetchall.return_value = [
{"model_id": "write-model"}
]
mock_conn.row_factory = None
mock_conn.commit = MagicMock()
call_idx = [0]
def fake_post(url, headers=None, json=None, timeout=None):
call_idx[0] += 1
resp = MagicMock()
resp.raise_for_status = MagicMock()
sample = samples_data[call_idx[0] - 1]
resp.json.return_value = {
"usage": {
"prompt_tokens": sample["prompt_tokens"],
"completion_tokens": sample["completion_tokens"],
}
}
return resp
import seed_local_dispatch_energy as mod
with patch.object(mod.sqlite3, "connect", return_value=mock_conn), patch.object(
mod, "load_config", return_value=mock_cfg
), patch.object(mod, "measure") as mock_measure, patch.object(
mod.requests, "post", side_effect=fake_post
), patch.object(
sys, "argv", ["seed_local_dispatch_energy.py"]
):
ctx = MagicMock()
ctx.avg_power_watts = 50.0
ctx.duration_seconds = 1.0
mock_measure.return_value.__enter__ = MagicMock(return_value=ctx)
mock_measure.return_value.__exit__ = MagicMock(return_value=False)
result = mod.main()
assert result == 0
update_count = sum(
1 for call in mock_conn.execute.call_args_list
if call[0][0].strip().startswith("UPDATE")
)
assert update_count >= 1
found_model = False
for call in mock_conn.execute.call_args_list:
stmt = call[0][0] if call[0] else ""
if stmt.strip().startswith("UPDATE models SET"):
params = call[0][1]
assert len(params) == 3
assert params[2] == "write-model"
found_model = True
assert isinstance(params[0], float)
assert isinstance(params[1], float)
assert found_model
assert mock_conn.commit.called
def _make_config(local_energy_enabled=False, tariff=None):
raw = yaml.safe_load((ROOT / "config" / "config.yaml").read_text())
raw["dispatch_providers"] = {}
raw["local_energy"]["enabled"] = local_energy_enabled
raw["local_energy"]["tariff_usd_per_kwh"] = tariff
return yaml.safe_dump(raw)
def _init_db(db_path: Path) -> sqlite3.Connection:
conn = sqlite3.connect(str(db_path))
sql = (ROOT / "config" / "schema.sql").read_text()
conn.executescript(sql)
conn.close()
return conn
def _seed_local_model(db_path: Path, model_id: str) -> sqlite3.Connection:
conn = sqlite3.connect(str(db_path))
conn.execute(
"""
INSERT OR IGNORE INTO models (
model_id, provider, base_model_id, display_name,
context_window, effective_context_window, max_output_tokens,
supports_tools, supports_json_mode, supports_vision,
supports_reasoning, reasoning_default_enabled, latency_class,
reasoning_mode, context_variant, access_level,
pricing_tbd, deprecated, availability, last_updated, tier
) VALUES (?, 'ollama-local', ?, ?, 32768, 32768, 2048,
0, 0, 0, 0, 0, 'standard', 'default', 'full', 'public',
0, 0, 'active', '2026-01-01T00:00:00+00:00', 1)
""",
(model_id, model_id, model_id),
)
conn.commit()
conn.close()
return conn