test / tests /test_tracker.py
Nanny7's picture
Claude Opus 4.7 (1M context)
Add model prediction tracker with 4-tab Gradio UI
6933ba4
Raw History Blame Contribute Delete
3.62 kB
from __future__ import annotations
import json
import sys
from concurrent.futures import ThreadPoolExecutor
from datetime import datetime, timedelta, timezone
from pathlib import Path
import pytest
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from tracker import RunTracker, hash_text # noqa: E402
@pytest.fixture
def tracker(tmp_path: Path) -> RunTracker:
return RunTracker(tmp_path / "runs.jsonl")
def _record_text(tracker: RunTracker, text: str, model: str = "mock-fast@1.0.0", latency: float = 1.0) -> str:
input_hash, input_meta = hash_text(text)
return tracker.record(
model_version=model,
input_kind="text",
input_hash=input_hash,
input_meta=input_meta,
output={"length": len(text)},
latency_ms=latency,
)
def test_hash_text_is_stable():
h1, _ = hash_text("hello")
h2, _ = hash_text("hello")
h3, _ = hash_text("goodbye")
assert h1 == h2
assert h1.startswith("sha256:")
assert h1 != h3
def test_record_and_iter_roundtrip(tracker: RunTracker):
run_id = _record_text(tracker, "hello world")
rows = list(tracker.iter_runs())
assert len(rows) == 1
assert rows[0]["run_id"] == run_id
assert rows[0]["model_version"] == "mock-fast@1.0.0"
assert rows[0]["input_hash"].startswith("sha256:")
def test_iter_runs_newest_first(tracker: RunTracker):
for i in range(3):
_record_text(tracker, f"msg-{i}")
rows = list(tracker.iter_runs())
timestamps = [r["timestamp"] for r in rows]
assert timestamps == sorted(timestamps, reverse=True)
def test_filter_by_model(tracker: RunTracker):
_record_text(tracker, "a", model="mock-fast@1.0.0")
_record_text(tracker, "b", model="mock-accurate@1.0.0")
_record_text(tracker, "c", model="mock-fast@1.0.0")
fast_rows = list(tracker.iter_runs(model_version="mock-fast@1.0.0"))
assert len(fast_rows) == 2
assert all(r["model_version"] == "mock-fast@1.0.0" for r in fast_rows)
def test_filter_by_date(tracker: RunTracker):
_record_text(tracker, "now")
future = datetime.now(timezone.utc) + timedelta(days=1)
rows_future = list(tracker.iter_runs(since=future))
assert rows_future == []
past = datetime.now(timezone.utc) - timedelta(days=1)
rows_past = list(tracker.iter_runs(since=past))
assert len(rows_past) == 1
def test_clear_empties_log(tracker: RunTracker):
for i in range(3):
_record_text(tracker, f"x-{i}")
dropped = tracker.clear()
assert dropped == 3
assert list(tracker.iter_runs()) == []
def test_stats_per_model_and_per_day(tracker: RunTracker):
_record_text(tracker, "a", model="mock-fast@1.0.0", latency=10.0)
_record_text(tracker, "b", model="mock-fast@1.0.0", latency=20.0)
_record_text(tracker, "c", model="mock-accurate@1.0.0", latency=200.0)
stats = tracker.stats()
assert stats["per_model"]["mock-fast@1.0.0"]["count"] == 2
assert stats["per_model"]["mock-accurate@1.0.0"]["count"] == 1
assert sum(stats["per_day"].values()) == 3
def test_concurrent_record_no_corruption(tracker: RunTracker):
def worker(i: int) -> str:
return _record_text(tracker, f"concurrent-{i}", latency=float(i))
with ThreadPoolExecutor(max_workers=10) as pool:
run_ids = list(pool.map(worker, range(100)))
assert len(set(run_ids)) == 100
lines = tracker.log_path.read_text(encoding="utf-8").splitlines()
assert len(lines) == 100
for line in lines:
parsed = json.loads(line)
assert "run_id" in parsed
assert parsed["input_hash"].startswith("sha256:")