from __future__ import annotations import hashlib import io import json import threading import uuid from collections import defaultdict from datetime import datetime, timezone from pathlib import Path from typing import Any, Iterator, Literal InputKind = Literal["image", "text"] def hash_text(text: str) -> tuple[str, dict[str, Any]]: raw = text.encode("utf-8") digest = hashlib.sha256(raw).hexdigest() return f"sha256:{digest}", {"len": len(text)} def hash_image(image) -> tuple[str, dict[str, Any]]: buf = io.BytesIO() image.save(buf, format="PNG") data = buf.getvalue() digest = hashlib.sha256(data).hexdigest() width, height = image.size channels = len(image.getbands()) return f"sha256:{digest}", { "shape": [height, width, channels], "bytes": len(data), } def _percentile(values: list[float], pct: float) -> float | None: if not values: return None values = sorted(values) if len(values) == 1: return values[0] rank = (pct / 100.0) * (len(values) - 1) lo = int(rank) hi = min(lo + 1, len(values) - 1) frac = rank - lo return values[lo] + (values[hi] - values[lo]) * frac class RunTracker: def __init__(self, log_path: Path): self._log_path = Path(log_path) self._lock = threading.Lock() self._log_path.parent.mkdir(parents=True, exist_ok=True) @property def log_path(self) -> Path: return self._log_path def record( self, *, model_version: str, input_kind: InputKind, input_hash: str, input_meta: dict, output: object, latency_ms: float, client_meta: dict | None = None, ) -> str: run_id = str(uuid.uuid4()) row = { "run_id": run_id, "timestamp": datetime.now(timezone.utc).isoformat().replace("+00:00", "Z"), "model_version": model_version, "input_kind": input_kind, "input_hash": input_hash, "input_meta": input_meta, "output": output, "latency_ms": float(latency_ms), "client_meta": client_meta or {}, } line = json.dumps(row, default=str, ensure_ascii=False) + "\n" with self._lock: with self._log_path.open("a", encoding="utf-8") as fh: fh.write(line) return run_id def _read_all(self) -> list[dict]: if not self._log_path.exists(): return [] rows: list[dict] = [] with self._log_path.open("r", encoding="utf-8") as fh: for line in fh: line = line.strip() if not line: continue try: rows.append(json.loads(line)) except json.JSONDecodeError: continue return rows def iter_runs( self, *, model_version: str | None = None, since: datetime | None = None, until: datetime | None = None, limit: int | None = None, ) -> Iterator[dict]: rows = self._read_all() rows.sort(key=lambda r: r.get("timestamp", ""), reverse=True) count = 0 for row in rows: if model_version and row.get("model_version") != model_version: continue ts_str = row.get("timestamp", "") ts: datetime | None = None if ts_str: try: ts = datetime.fromisoformat(ts_str.replace("Z", "+00:00")) except ValueError: ts = None if since and ts and ts < since: continue if until and ts and ts > until: continue yield row count += 1 if limit is not None and count >= limit: break def stats(self) -> dict: rows = self._read_all() per_model_latencies: dict[str, list[float]] = defaultdict(list) per_day: dict[str, int] = defaultdict(int) for row in rows: model = row.get("model_version", "unknown") latency = row.get("latency_ms") if isinstance(latency, (int, float)): per_model_latencies[model].append(float(latency)) ts_str = row.get("timestamp", "") if ts_str: day = ts_str[:10] per_day[day] += 1 per_model = { model: { "count": len(latencies), "p50_ms": _percentile(latencies, 50), "p95_ms": _percentile(latencies, 95), } for model, latencies in per_model_latencies.items() } return {"per_model": per_model, "per_day": dict(sorted(per_day.items()))} def clear(self) -> int: with self._lock: if not self._log_path.exists(): return 0 count = sum(1 for _ in self._log_path.open("r", encoding="utf-8")) self._log_path.write_text("", encoding="utf-8") return count