Files
6krrt/tests/test_dispatcher_local_energy.py

331 lines
11 KiB
Python

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