File size: 3,708 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
import importlib.util
import json
from pathlib import Path
import tempfile
import unittest
from unittest import mock
from types import SimpleNamespace


spec = importlib.util.spec_from_file_location("study_tracking", Path(__file__).parents[1] / "tracking.py")
tracking = importlib.util.module_from_spec(spec)
spec.loader.exec_module(tracking)


class FakeRemote:
    def __init__(self):
        self.fail = False
        self.visible = True
        self.logs = {}

    def predict(self, api_name, **kwargs):
        if self.fail:
            raise ConnectionError("secret must not appear in stored error")
        if api_name == "/bulk_log":
            for row in kwargs["logs"]:
                self.logs[row["log_id"]] = row
            return None
        return [{"value": row["metrics"]["tracking/event_sequence"]}
                for row in self.logs.values() if self.visible and row["run_id"] == kwargs["run_id"]]


class TrackingTests(unittest.TestCase):
    def setUp(self):
        self.auth = mock.patch.dict("sys.modules", {"huggingface_hub": SimpleNamespace(get_token=lambda: None)})
        self.auth.start()
        self.temp = tempfile.TemporaryDirectory()
        self.client = FakeRemote()
        self.tracker = tracking.Tracker(self.temp.name, client=self.client)
        self.event = {"timestamp": "2026-10-09T00:00:00Z", "run_id": "jq-single-r1",
                      "task": "jqlang__jq.b33a763", "condition": "single", "kind": "progress",
                      "metrics": {"api_input_tokens": 123, "score": 0.2}}

    def tearDown(self):
        self.tracker.close()
        self.temp.cleanup()
        self.auth.stop()

    def test_replay_is_idempotent(self):
        self.tracker.enqueue(self.event)
        self.tracker.enqueue(self.event)
        self.assertEqual(self.tracker.pending(), 1)
        self.assertTrue(self.tracker.flush())
        self.assertEqual(self.tracker.pending(), 0)
        self.assertEqual(len(self.client.logs), 1)

    def test_failure_preserves_outbox_and_redacts_error(self):
        self.tracker.enqueue(self.event)
        self.client.fail = True
        self.assertFalse(self.tracker.flush())
        self.assertEqual(self.tracker.pending(), 1)
        status = json.loads((Path(self.temp.name) / "status.json").read_text())
        self.assertEqual(status["last_error"], "ConnectionError")
        self.tracker.client = self.client
        self.client.fail = False
        self.assertTrue(self.tracker.flush())

    def test_http_success_without_readback_is_not_acknowledged(self):
        self.tracker.enqueue(self.event)
        self.client.visible = False
        self.assertFalse(self.tracker.flush())
        self.assertEqual(self.tracker.pending(), 1)

    def test_partial_line_waits_and_restart_keeps_cursor(self):
        path = Path(self.temp.name) / "input.jsonl"
        line = json.dumps(self.event)
        path.write_text(line)
        self.tracker.ingest(path)
        self.assertEqual(self.tracker.pending(), 0)
        path.write_text(line + "\n")
        self.tracker.ingest(path)
        self.tracker.ingest(path)
        self.assertEqual(self.tracker.pending(), 1)

    def test_raw_text_and_credentials_are_never_published(self):
        event = self.event | {"message": "secret", "config": {"api_key": "secret", "model": "Beam-501B-A23B",
                                                                "prompt": "private prompt", "max_tokens": 8}}
        payload = tracking.sanitize(event)
        self.assertNotIn("secret", json.dumps(payload))
        self.assertNotIn("private prompt", json.dumps(payload))
        self.assertEqual(payload["config"]["max_tokens"], 8)


if __name__ == "__main__":
    unittest.main()