Files
6krrt/tests/test_local_energy.py
2026-09-01 00:20:10 -04:00

209 lines
6.3 KiB
Python

"""Tests for the local energy meter.
The meter has three responsibilities:
- sample the GPU through nvidia-smi (or a stub),
- poll in a background thread while local work runs,
- persist rows to a dedicated local_energy_observations table.
These tests stay offline by stubbing the sampler. The nvidia-smi path is
only exercised as a call-site shape check; no real GPU is required.
"""
import sqlite3
import time
from pathlib import Path
import pytest
import local_energy
from local_energy import (
ensure_local_energy_table,
log_local_energy,
measure,
)
ROOT = Path(__file__).resolve().parent.parent
SCHEMA_SQL = (ROOT / "config" / "schema.sql").read_text()
@pytest.fixture
def tmp_db(tmp_path):
db_path = tmp_path / "local_energy.db"
conn = sqlite3.connect(db_path)
conn.row_factory = sqlite3.Row
try:
yield conn
finally:
conn.close()
# --- nvidia-smi sampler ------------------------------------------------------
def test_sample_nvidia_smi_returns_optional_float():
"""nvidia-smi returns a float or None; on CI/GPU hosts it may succeed."""
value = local_energy.sample_nvidia_smi()
assert value is None or isinstance(value, float)
# --- measure context manager -------------------------------------------------
def test_measure_averages_samples_from_stub_sampler():
"""The context manager collects samples and averages them."""
readings = [50.0, 70.0]
sampler = iter(readings).__next__
with measure(sample_interval_seconds=0.005, sampler=sampler) as result:
time.sleep(0.02)
assert result.result().avg_power_watts == pytest.approx(60.0)
def test_measure_returns_none_avg_when_no_samples_collected():
"""If the sampler always fails, avg_power_watts is None."""
with measure(sample_interval_seconds=0.005, sampler=lambda: None) as result:
time.sleep(0.005)
assert result.result().avg_power_watts is None
assert result.result().duration_seconds >= 0.0
def test_measure_reports_nonzero_duration():
"""Duration covers the full time inside the context manager."""
with measure(sample_interval_seconds=0.005, sampler=lambda: 10.0) as result:
time.sleep(0.05)
assert result.result().duration_seconds >= 0.04
def test_measure_rejects_nonpositive_interval():
"""A zero or negative interval is a programmer error."""
with pytest.raises(ValueError), measure(
sample_interval_seconds=0, sampler=lambda: 1.0
):
pass
def test_measure_returns_result_even_if_sampler_raises():
"""Sampler exceptions stop collection but still yield duration."""
def bad_sampler():
raise RuntimeError("gpu unreachable")
with measure(sample_interval_seconds=0.005, sampler=bad_sampler) as result:
time.sleep(0.005)
assert result.result().avg_power_watts is None
assert result.result().duration_seconds >= 0.0
# --- table migration ---------------------------------------------------------
def test_ensure_local_energy_table_creates_table(tmp_db):
ensure_local_energy_table(tmp_db)
row = tmp_db.execute(
"SELECT name FROM sqlite_master WHERE type='table' AND name=?",
("local_energy_observations",),
).fetchone()
assert row is not None
def test_ensure_local_energy_table_is_idempotent(tmp_db):
ensure_local_energy_table(tmp_db)
ensure_local_energy_table(tmp_db)
cols = {
r[1] for r in tmp_db.execute("PRAGMA table_info(local_energy_observations)")
}
assert "id" in cols
assert "model_id" in cols
assert "call_type" in cols
def test_ensure_local_energy_table_creates_index(tmp_db):
ensure_local_energy_table(tmp_db)
row = tmp_db.execute(
"SELECT name FROM sqlite_master WHERE type='index' AND name=?",
("idx_local_energy_model",),
).fetchone()
assert row is not None
# --- logging -----------------------------------------------------------------
def test_log_local_energy_inserts_row(tmp_db):
ensure_local_energy_table(tmp_db)
log_local_energy(
conn=tmp_db,
model_id="mistral-nemo:12b",
call_type="classification",
avg_power_watts=120.0,
duration_seconds=2.5,
energy_kwh=0.0000833,
cost_usd=0.0001,
carbon_g_co2eq=0.05,
meter="nvidia-smi",
observed_at="2026-08-31T00:00:00+00:00",
)
rows = tmp_db.execute("SELECT * FROM local_energy_observations").fetchall()
assert len(rows) == 1
assert rows[0]["model_id"] == "mistral-nemo:12b"
assert rows[0]["call_type"] == "classification"
assert rows[0]["avg_power_watts"] == pytest.approx(120.0)
assert rows[0]["duration_seconds"] == pytest.approx(2.5)
assert rows[0]["energy_kwh"] == pytest.approx(0.0000833)
assert rows[0]["cost_usd"] == pytest.approx(0.0001)
assert rows[0]["carbon_g_co2eq"] == pytest.approx(0.05)
assert rows[0]["meter"] == "nvidia-smi"
def test_log_local_energy_creates_table_on_demand(tmp_db):
"""The write helper migrates the table so callers don't need to."""
assert (
tmp_db.execute(
"SELECT name FROM sqlite_master WHERE type='table' AND name=?",
("local_energy_observations",),
).fetchone()
is None
)
log_local_energy(
conn=tmp_db,
model_id="qwen3.5",
call_type="verification",
avg_power_watts=None,
duration_seconds=1.0,
energy_kwh=None,
cost_usd=None,
carbon_g_co2eq=None,
meter="stub",
observed_at="2026-08-31T00:00:00+00:00",
)
rows = tmp_db.execute("SELECT * FROM local_energy_observations").fetchall()
assert len(rows) == 1
assert rows[0]["avg_power_watts"] is None
assert rows[0]["energy_kwh"] is None
def test_measure_none_avg_power_admits_null_energy():
"""When the sampler returns only None, the measurement records None power.
Downstream callers compute energy as None, which ``log_local_energy``
stores honestly rather than as zero.
"""
with measure(sample_interval_seconds=0.005, sampler=lambda: None) as result:
time.sleep(0.02)
measurement = result.result()
assert measurement.avg_power_watts is None
assert measurement.duration_seconds >= 0.0
def test_carbon_calculation_from_intensity():
"""Carbon is energy (kWh) * 1000 * grid intensity (g/kWh)."""
energy_kwh = 0.0000025
intensity = 475.0
assert energy_kwh * 1000 * intensity == pytest.approx(1.1875)