- _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.
396 lines
14 KiB
Python
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
|