"""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