provenance.py file-level

at sha256:9 · View file ↗ · Intel ↗

History
1 files
1 commits
0 hotspots
0 🧊 dead
0 💥 blast risk
sha256:c Add T4 job cancellation retry and validation · · Jun 11, 2026
1 """Content-free provenance records for completed fixture artifacts."""
2
3 from __future__ import annotations
4
5 import argparse
6 import json
7 import re
8 from dataclasses import dataclass
9 from pathlib import Path
10 from typing import Mapping
11
12 from scooling_lab.errors import ApiError, ErrorCode
13 from scooling_lab.retention import parse_utc_timestamp
14
15
16 PROVENANCE_SCHEMA_VERSION = "scooling-lab.provenance.v1"
17 PROVENANCE_KEYS: frozenset[str] = frozenset(
18 {
19 "jobId",
20 "datasetHash",
21 "artifactHash",
22 "baseModelId",
23 "trainingConfigHash",
24 "createdAt",
25 "schemaVersion",
26 }
27 )
28 HASH_RE = re.compile(r"^[a-f0-9]{64}$")
29 JOB_ID_RE = re.compile(r"^job_[a-f0-9]{24}$")
30 SAFE_MODEL_ID_RE = re.compile(r"^[A-Za-z0-9._:-]{3,96}$")
31 UNSAFE_FREE_TEXT_RE = re.compile(r"(?i)(https?://|file://|ssh://|\.\./|/|\\|[;&|`$<>]|\s)")
32
33
34 @dataclass(frozen=True)
35 class ProvenanceRecord:
36 """Validated, content-free provenance for one completed fixture artifact."""
37
38 job_id: str
39 dataset_hash: str
40 artifact_hash: str
41 base_model_id: str
42 training_config_hash: str
43 created_at: str
44 schema_version: str = PROVENANCE_SCHEMA_VERSION
45
46 def to_dict(self) -> dict[str, str]:
47 """Serialize the provenance record in the public wire shape."""
48
49 return {
50 "artifactHash": self.artifact_hash,
51 "baseModelId": self.base_model_id,
52 "createdAt": self.created_at,
53 "datasetHash": self.dataset_hash,
54 "jobId": self.job_id,
55 "schemaVersion": self.schema_version,
56 "trainingConfigHash": self.training_config_hash,
57 }
58
59 @classmethod
60 def from_mapping(cls, payload: Mapping[str, object]) -> "ProvenanceRecord":
61 """Validate and rehydrate a provenance record from untrusted data."""
62
63 validate_provenance_mapping(payload)
64 return cls(
65 artifact_hash=str(payload["artifactHash"]),
66 base_model_id=str(payload["baseModelId"]),
67 created_at=str(payload["createdAt"]),
68 dataset_hash=str(payload["datasetHash"]),
69 job_id=str(payload["jobId"]),
70 schema_version=str(payload["schemaVersion"]),
71 training_config_hash=str(payload["trainingConfigHash"]),
72 )
73
74
75 def validate_provenance_mapping(payload: Mapping[str, object]) -> None:
76 """Reject provenance records that contain text, paths, URLs, or extra fields."""
77
78 if set(payload) != PROVENANCE_KEYS:
79 raise ApiError(ErrorCode.VALIDATION_ERROR, 400)
80 for key, value in payload.items():
81 if not isinstance(value, str):
82 raise ApiError(ErrorCode.VALIDATION_ERROR, 400)
83 if key != "createdAt" and UNSAFE_FREE_TEXT_RE.search(value):
84 raise ApiError(ErrorCode.VALIDATION_ERROR, 400)
85
86 if not JOB_ID_RE.fullmatch(str(payload["jobId"])):
87 raise ApiError(ErrorCode.VALIDATION_ERROR, 400)
88 for key in ("datasetHash", "artifactHash", "trainingConfigHash"):
89 if not HASH_RE.fullmatch(str(payload[key])):
90 raise ApiError(ErrorCode.VALIDATION_ERROR, 400)
91 if not SAFE_MODEL_ID_RE.fullmatch(str(payload["baseModelId"])):
92 raise ApiError(ErrorCode.VALIDATION_ERROR, 400)
93 if payload["schemaVersion"] != PROVENANCE_SCHEMA_VERSION:
94 raise ApiError(ErrorCode.VALIDATION_ERROR, 400)
95 try:
96 parse_utc_timestamp(str(payload["createdAt"]))
97 except (ApiError, ValueError) as exc:
98 raise ApiError(ErrorCode.VALIDATION_ERROR, 400) from exc
99
100
101 def validate_provenance_record(record: ProvenanceRecord | Mapping[str, object]) -> None:
102 """Validate any provenance record against the content-free schema."""
103
104 if isinstance(record, ProvenanceRecord):
105 validate_provenance_mapping(record.to_dict())
106 return
107 validate_provenance_mapping(record)
108
109
110 def load_provenance_file(path: Path) -> ProvenanceRecord:
111 """Read and validate one JSON provenance record from disk."""
112
113 payload = json.loads(path.read_text(encoding="utf-8"))
114 if not isinstance(payload, dict):
115 raise ApiError(ErrorCode.VALIDATION_ERROR, 400)
116 return ProvenanceRecord.from_mapping(payload)
117
118
119 def self_check() -> None:
120 """Run validator acceptance and rejection checks for CI."""
121
122 valid = {
123 "artifactHash": "b" * 64,
124 "baseModelId": "fixture-tiny-llm",
125 "createdAt": "2026-06-11T00:00:00Z",
126 "datasetHash": "a" * 64,
127 "jobId": "job_" + "1" * 24,
128 "schemaVersion": PROVENANCE_SCHEMA_VERSION,
129 "trainingConfigHash": "c" * 64,
130 }
131 validate_provenance_mapping(valid)
132 rejected = dict(valid)
133 rejected["baseModelId"] = "fixture model with prompt text"
134 try:
135 validate_provenance_mapping(rejected)
136 except ApiError:
137 return
138 raise ApiError(ErrorCode.VALIDATION_ERROR, 400)
139
140
141 def build_parser() -> argparse.ArgumentParser:
142 """Build the offline provenance validator CLI parser."""
143
144 parser = argparse.ArgumentParser(description="Validate Scooling Lab provenance records")
145 parser.add_argument("--record", action="append", type=Path, default=[])
146 parser.add_argument("--self-check", action="store_true")
147 return parser
148
149
150 def main(argv: list[str] | None = None) -> int:
151 """Run offline provenance validation for files and CI self-checks."""
152
153 args = build_parser().parse_args(argv)
154 if args.self_check:
155 self_check()
156 for record_path in args.record:
157 load_provenance_file(record_path)
158 return 0
159
160
161 if __name__ == "__main__":
162 try:
163 raise SystemExit(main())
164 except (ApiError, json.JSONDecodeError) as exc:
165 raise SystemExit(f"provenance validation failed: {exc}") from exc