burtenshaw's picture
burtenshaw HF Staff
feat: publish beam pi study source
5741b22 verified
Raw History Blame Contribute Delete
10.5 kB
#!/usr/bin/env python3
"""Publish sanitized controller events to Trackio from a durable SQLite outbox."""
from __future__ import annotations
import argparse
import hashlib
import json
import math
import os
from pathlib import Path
import re
import signal
import sqlite3
import sys
import time
SPACE_ID = "burtenshaw/beam-pi-programbench"
PROJECT = "beam-pi-programbench-20261009"
LABEL = re.compile(r"^[A-Za-z0-9_.:/-]{1,180}$")
CONFIG_TEXT = {"model", "provider", "reasoning", "task", "condition", "kind", "phase",
"task_id", "benchmark_commit", "test_dataset_commit", "stop_reason"}
KINDS = {"run_started", "progress", "score", "run_finished", "error", "budget_stop"}
STATUSES = {"pending", "running", "completed", "failed", "cancelled", "budget_exhausted",
"wall_time_exhausted", "passed", "blocked", "validated", "initializing",
"calibrating", "calibration_failed", "validation_failed", "interrupted"}
def numeric(value):
return isinstance(value, (int, float)) and not isinstance(value, bool) and math.isfinite(value)
def safe_config(config):
"""Allow scalar numeric settings and explicitly named, short categorical strings."""
result = {}
if not isinstance(config, dict):
return result
for key, value in config.items():
if not isinstance(key, str) or not LABEL.fullmatch(key):
continue
if any(word in key.lower() for word in ("secret", "token_key", "api_key", "password", "credential")):
continue
if isinstance(value, bool) or numeric(value):
result[key] = value
elif key in CONFIG_TEXT and isinstance(value, str) and LABEL.fullmatch(value):
result[key] = value
return result
def sanitize(event):
run_id = event.get("run_id")
if not isinstance(run_id, str) or not LABEL.fullmatch(run_id):
raise ValueError("invalid run_id")
kind = event.get("kind")
if kind not in KINDS:
raise ValueError("invalid event kind")
config = safe_config(event.get("config", {}))
for key in ("task", "condition"):
value = event.get(key)
if isinstance(value, str) and LABEL.fullmatch(value):
config[key] = value
config["phase"] = ("infrastructure_validation" if run_id.startswith("validation-")
else "study_summary" if run_id == "study-summary" else "benchmark")
metrics = {key: value for key, value in event.get("metrics", {}).items()
if isinstance(key, str) and LABEL.fullmatch(key) and numeric(value)}
metrics[f"event/{kind}"] = 1
status = event.get("status")
if status in STATUSES:
for label in sorted(STATUSES):
metrics[f"status/{label}"] = int(label == status)
# Deliberately omit message, traces, prompts, outputs, paths, and arbitrary text.
return {"run_id": run_id, "config": config, "metrics": metrics}
class Tracker:
def __init__(self, state_dir, space_id=SPACE_ID, project=PROJECT, client=None):
self.state_dir = Path(state_dir)
self.state_dir.mkdir(parents=True, exist_ok=True)
os.chmod(self.state_dir, 0o700)
self.db = sqlite3.connect(self.state_dir / "outbox.sqlite")
self.db.execute("PRAGMA journal_mode=DELETE")
self.db.execute("PRAGMA synchronous=FULL")
self.db.executescript("""
CREATE TABLE IF NOT EXISTS events (
id INTEGER PRIMARY KEY, log_id TEXT UNIQUE NOT NULL, payload TEXT NOT NULL,
remote_visible INTEGER NOT NULL DEFAULT 0);
CREATE TABLE IF NOT EXISTS cursor (
path TEXT PRIMARY KEY, identity TEXT NOT NULL, offset INTEGER NOT NULL);
""")
self.db.commit()
self.space_id, self.project, self.client = space_id, project, client
self.last_error = None
def enqueue(self, event):
payload = sanitize(event)
# Hash the source event for stable IDs across replay/restart, without storing it.
log_id = hashlib.sha256(json.dumps(event, sort_keys=True).encode()).hexdigest()
self.db.execute("INSERT OR IGNORE INTO events(log_id,payload) VALUES (?,?)",
(log_id, json.dumps(payload, sort_keys=True)))
self.db.commit()
def ingest(self, events_path):
path = Path(events_path)
if not path.exists():
return
stat = path.stat()
identity = f"{stat.st_dev}:{stat.st_ino}"
saved = self.db.execute("SELECT identity,offset FROM cursor WHERE path=?",
(str(path.resolve()),)).fetchone()
offset = saved[1] if saved and saved[0] == identity and saved[1] <= stat.st_size else 0
with path.open("rb") as source:
source.seek(offset)
while True:
line = source.readline()
if not line or not line.endswith(b"\n"):
break
self.enqueue(json.loads(line))
self.db.execute("INSERT OR REPLACE INTO cursor VALUES (?,?,?)",
(str(path.resolve()), identity, source.tell()))
self.db.commit()
def pending(self):
return self.db.execute("SELECT count(*) FROM events WHERE remote_visible=0").fetchone()[0]
def _connect(self):
if self.client is None:
os.environ.setdefault("TRACKIO_DIR", str(self.state_dir / "trackio"))
from huggingface_hub import get_token
from trackio.remote_client import RemoteClient
self.client = RemoteClient(self.space_id, hf_token=get_token(),
httpx_kwargs={"timeout": 30}, verbose=False)
return self.client
def flush(self):
rows = self.db.execute("SELECT id,log_id,payload FROM events WHERE remote_visible=0 ORDER BY id LIMIT 100").fetchall()
if not rows:
self.write_status()
return True
try:
client = self._connect()
from huggingface_hub import get_token
logs = []
by_run = {}
for seq, log_id, payload_text in rows:
payload = json.loads(payload_text)
run_id = payload["run_id"]
metrics = payload["metrics"] | {"tracking/event_sequence": seq}
logs.append({"project": self.project, "run": run_id, "run_id": run_id,
"metrics": metrics, "step": seq, "log_id": log_id,
"config": payload["config"]})
by_run.setdefault(run_id, set()).add(seq)
client.predict(api_name="/bulk_log", logs=logs, hf_token=get_token())
# An HTTP success can mean queued writes. Only read-back events are acknowledged.
acknowledged = set()
for run_id, sequences in by_run.items():
values = client.predict(api_name="/get_metric_values", project=self.project,
run=run_id, run_id=run_id,
metric_name="tracking/event_sequence", max_points=None)
acknowledged.update(item["value"] for item in values
if item.get("value") in sequences)
self.db.executemany("UPDATE events SET remote_visible=1 WHERE id=?",
[(seq,) for seq in acknowledged])
self.db.commit()
self.last_error = None if len(acknowledged) == len(rows) else "remote_readback_incomplete"
self.write_status()
return self.last_error is None
except Exception as error:
code = getattr(getattr(error, "response", None), "status_code", None)
self.last_error = type(error).__name__ + (f"_HTTP_{code}" if code else "")
print(f"TRACKING_UPLOAD_PENDING: {self.last_error}; {self.pending()} events retained locally",
file=sys.stderr, flush=True)
self.client = None
self.write_status()
return False
def write_status(self):
total, visible = self.db.execute("SELECT count(*),coalesce(sum(remote_visible),0) FROM events").fetchone()
status = {"timestamp": time.time(), "space_id": self.space_id, "project": self.project,
"local_durable_events": total, "remote_readback_events": visible,
"pending_events": total - visible, "last_error": self.last_error,
"remote_acknowledgement": "metric_readback", "local_storage": "sqlite_full_sync"}
tmp = self.state_dir / "status.json.tmp"
tmp.write_text(json.dumps(status, indent=2) + "\n")
tmp.replace(self.state_dir / "status.json")
def close(self):
self.db.close()
def main():
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--events", required=True, type=Path)
parser.add_argument("--state-dir", required=True, type=Path)
parser.add_argument("--space-id", default=SPACE_ID)
parser.add_argument("--project", default=PROJECT)
parser.add_argument("--once", action="store_true")
parser.add_argument("--poll-seconds", type=float, default=15)
args = parser.parse_args()
tracker = Tracker(args.state_dir, args.space_id, args.project)
stopped = False
def stop(*_):
nonlocal stopped
stopped = True
signal.signal(signal.SIGTERM, stop)
signal.signal(signal.SIGINT, stop)
try:
while True:
try:
tracker.ingest(args.events)
except Exception as error:
tracker.last_error = "ingest_" + type(error).__name__
tracker.write_status()
print(f"TRACKING_INGEST_ERROR: {tracker.last_error}; inspect controller event schema",
file=sys.stderr, flush=True)
if args.once or stopped:
return 2
tracker.flush()
if args.once or stopped:
# Drain already-ingested batches once; stop if remote cannot confirm.
while tracker.pending() and tracker.flush():
pass
return 1 if tracker.pending() else 0
deadline = time.monotonic() + args.poll_seconds
while not stopped and time.monotonic() < deadline:
time.sleep(min(1, max(0, deadline - time.monotonic())))
finally:
tracker.close()
if __name__ == "__main__":
raise SystemExit(main())