"""Tests that dispatcher-level local energy metering is wired correctly. Stays offline: stubs ``local_energy.measure`` and ``local_energy.log_local_energy`` so no real Ollama or GPU is contacted. Verifies the three call sites (classify, verify, local_vision) and the gate logic. """ from __future__ import annotations import json from unittest.mock import MagicMock import pytest import dispatcher import local_energy import local_encoder class _FakeMeasurement: """A completed measurement for the dispatcher helper.""" def __init__(self, avg_power_watts, duration_seconds): self.avg_power_watts = avg_power_watts self.duration_seconds = duration_seconds def __enter__(self): return self def __exit__(self, *args): return False def _energy_cfg(enabled: bool, tariff: float = 0.12, intensity: float = 475.0): """Return a bare config object with the local_energy fields we need.""" cfg = MagicMock() cfg.local_energy = MagicMock() cfg.local_energy.enabled = enabled cfg.local_energy.tariff_usd_per_kwh = tariff cfg.local_energy.grid_intensity_g_per_kwh = intensity cfg.local_energy.meter = "nvidia_smi" cfg.local_energy.sample_interval_seconds = 0.25 cfg.local_energy_call_sites = { "classify": True, "verify": True, "local_vision": True, } cfg.classifier = MagicMock() cfg.classifier.model = "classifier-model" cfg.classifier.max_input_chars = 8000 cfg.classifier.fallback_tier = 2 cfg.classifier.encoder = MagicMock() cfg.classifier.encoder.model = "encoder-model" cfg.classifier.encoder.device = "cpu" cfg.classifier.encoder.confidence_min = 0.5 cfg.verification = MagicMock() cfg.verification.model = "verifier-model" cfg.verification.base_url = "http://localhost:11434" cfg.verification.timeout_seconds = 60 cfg.verification.max_output_tokens = 1024 cfg.local_vision = MagicMock() cfg.local_vision.model = "vision-model" cfg.local_vision.base_url = "http://localhost:11434/v1" cfg.local_vision.max_images = 4 cfg.local_vision.max_image_bytes = 9 * 1024 * 1024 cfg.local_vision.api_key_env = None cfg.local_vision.timeout_seconds = 60 return cfg @pytest.fixture def capture_log(tmp_path, monkeypatch): """Capture calls to ``local_energy.log_local_energy`` and use temp DB.""" conn = None def fake_log_local_energy(*, conn=None, **kwargs): calls.append(kwargs) calls = [] monkeypatch.setattr(local_energy, "log_local_energy", fake_log_local_energy) def make_measurement(avg_power_watts=100.0, duration_seconds=2.0): return _FakeMeasurement(avg_power_watts, duration_seconds) monkeypatch.setattr( local_energy, "measure", lambda sample_interval_seconds, sampler: make_measurement(), ) def fake_db(): nonlocal conn if conn is None: db_path = tmp_path / "router.db" import sqlite3 conn = sqlite3.connect(str(db_path)) conn.row_factory = sqlite3.Row return conn monkeypatch.setattr(dispatcher, "_db", fake_db) return calls def test_disabled_local_energy_skips_logging(monkeypatch, capture_log): """When local_energy.enabled is false, classify does NOT log a row.""" cfg = _energy_cfg(enabled=False) monkeypatch.setattr(dispatcher, "cfg", cfg) def fake_classify_once(client, system_prompt, user_content): return dispatcher.Classification( task_category="general_chat", task_tier=1, required_context_tokens=10, confidence=0.9, ) monkeypatch.setattr(dispatcher, "_classify_once", fake_classify_once) monkeypatch.setattr(dispatcher, "_classifier_client", lambda: MagicMock()) dispatcher.classify("hello", None) assert capture_log == [] def test_enabled_local_energy_logs_classify(monkeypatch, capture_log): """When enabled and loopback, classify records one local_energy row.""" cfg = _energy_cfg(enabled=True) monkeypatch.setattr(dispatcher, "cfg", cfg) def fake_classify_once(client, system_prompt, user_content): return dispatcher.Classification( task_category="general_chat", task_tier=1, required_context_tokens=10, confidence=0.9, ) monkeypatch.setattr(dispatcher, "_classify_once", fake_classify_once) monkeypatch.setattr(dispatcher, "_classifier_client", lambda: MagicMock()) dispatcher.classify("hello", None) assert len(capture_log) == 1 entry = capture_log[0] assert entry["model_id"] == "classifier-model" assert entry["call_type"] == "classify" assert entry["avg_power_watts"] == 100.0 assert entry["duration_seconds"] == 2.0 assert entry["energy_kwh"] == pytest.approx(100.0 * 2.0 / 3_600_000) assert entry["cost_usd"] == pytest.approx(entry["energy_kwh"] * 0.12) assert entry["carbon_g_co2eq"] == pytest.approx(entry["energy_kwh"] * 1000 * 475.0) assert entry["meter"] == "nvidia_smi" def test_disabled_call_site_skips_logging(monkeypatch, capture_log): """If local_energy_call_sites["classify"] is false, classify does NOT log.""" cfg = _energy_cfg(enabled=True) cfg.local_energy_call_sites["classify"] = False monkeypatch.setattr(dispatcher, "cfg", cfg) def fake_classify_once(client, system_prompt, user_content): return dispatcher.Classification( task_category="general_chat", task_tier=1, required_context_tokens=10, confidence=0.9, ) monkeypatch.setattr(dispatcher, "_classify_once", fake_classify_once) monkeypatch.setattr(dispatcher, "_classifier_client", lambda: MagicMock()) dispatcher.classify("hello", None) assert capture_log == [] # --- local_encoder mode ----------------------------------------------------- def test_enabled_local_energy_logs_local_encoder(monkeypatch, capture_log): """local_encoder mode: a GPU draw is metered the same as the LLM path.""" cfg = _energy_cfg(enabled=True) cfg.classifier.mode = "local_encoder" monkeypatch.setattr(dispatcher, "cfg", cfg) monkeypatch.setattr( local_encoder, "classify_zero_shot", lambda task, categories, **k: ("general_chat", 0.9), ) dispatcher.classify("hello", None) assert len(capture_log) == 1 entry = capture_log[0] assert entry["model_id"] == "encoder-model" assert entry["call_type"] == "classify" assert entry["avg_power_watts"] == 100.0 assert entry["duration_seconds"] == 2.0 def test_enabled_local_energy_logs_local_encoder_on_below_threshold( monkeypatch, capture_log ): """A below-threshold encoder answer still logs the machine's real draw. The measurement is captured while the GPU is hot; the low-confidence RuntimeError must not drop that observation. """ cfg = _energy_cfg(enabled=True) cfg.classifier.mode = "local_encoder" cfg.classifier.encoder.confidence_min = 0.6 monkeypatch.setattr(dispatcher, "cfg", cfg) monkeypatch.setattr( local_encoder, "classify_zero_shot", lambda task, categories, **k: ("general_chat", 0.2), ) with pytest.raises(RuntimeError): # Drive the encoder path directly: classify() would cascade, but here # we want the raise that happens inside the metered block itself. dispatcher._classify_via_local_encoder("hello") assert len(capture_log) == 1 assert capture_log[0]["model_id"] == "encoder-model" assert capture_log[0]["call_type"] == "classify" def test_disabled_local_energy_skips_logging_for_local_encoder(monkeypatch, capture_log): """local_encoder mode with local_energy.enabled=false logs nothing.""" cfg = _energy_cfg(enabled=False) cfg.classifier.mode = "local_encoder" monkeypatch.setattr(dispatcher, "cfg", cfg) monkeypatch.setattr( local_encoder, "classify_zero_shot", lambda task, categories, **k: ("general_chat", 0.9), ) got = dispatcher.classify("hello", None) assert capture_log == [] assert got.source == "classifier" def test_enabled_local_energy_logs_verify(monkeypatch, capture_log): """When enabled, the local verification POST records one row.""" cfg = _energy_cfg(enabled=True) monkeypatch.setattr(dispatcher, "cfg", cfg) def fake_post(url, **kwargs): resp = MagicMock() resp.status_code = 200 resp.json.return_value = { "message": {"content": json.dumps({"verdict": "ok"})} } resp.raise_for_status = lambda: None return resp monkeypatch.setattr(dispatcher.requests, "post", fake_post) monkeypatch.setattr( dispatcher, "interpret_local_verdict", lambda v: MagicMock(verdict="ok", detail="looks fine"), ) monkeypatch.setattr(dispatcher, "log_verification", lambda *a, **k: None) dispatcher.run_local_verification( model_id="cloud-model", provider="neuralwatt", task_category="general_chat", request_text="what is 2+2?", answer="4", completion_tokens=700, ) assert len(capture_log) == 1 assert capture_log[0]["model_id"] == "verifier-model" assert capture_log[0]["call_type"] == "verify" def test_enabled_local_energy_logs_local_vision(monkeypatch, capture_log): """When enabled, a local-vision fallback POST records one row.""" cfg = _energy_cfg(enabled=True) monkeypatch.setattr(dispatcher, "cfg", cfg) def fake_post(url, **kwargs): resp = MagicMock() resp.status_code = 200 resp.json.return_value = {"choices": [{"message": {"content": "a cat"}}]} return resp monkeypatch.setattr(dispatcher.requests, "post", fake_post) result = dispatcher._run_local_vision( [ { "role": "user", "content": [ {"type": "text", "text": "describe"}, {"type": "image_url", "image_url": {"url": "data:image/png;base64,AAAA"}}, ], } ], cfg.local_vision, ) assert result == "a cat" assert len(capture_log) == 1 assert capture_log[0]["model_id"] == "vision-model" assert capture_log[0]["call_type"] == "local_vision" def test_verify_logs_even_on_post_failure(monkeypatch, capture_log): """A failed verification POST still records duration/power.""" cfg = _energy_cfg(enabled=True) monkeypatch.setattr(dispatcher, "cfg", cfg) def fake_post(url, **kwargs): raise dispatcher.requests.RequestException("timeout") monkeypatch.setattr(dispatcher.requests, "post", fake_post) dispatcher.run_local_verification( model_id="cloud-model", provider="neuralwatt", task_category="general_chat", request_text="what is 2+2?", answer="4", completion_tokens=700, ) assert len(capture_log) == 1 assert capture_log[0]["call_type"] == "verify" assert capture_log[0]["avg_power_watts"] == 100.0