Spaces:
Running
Running
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()
|