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