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