331 lines
11 KiB
Python
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
|