File size: 10,512 Bytes
5741b22
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
#!/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())