store.py python
465 lines 16.9 KB
Raw
sha256:91e875d4a97bb1e35f37992f803988d5713931f1782d870c371c50054574af22 Add fixture provenance retention deletion Human minor ⚠ breaking 42 days ago
1 """Thread-safe job and artifact store for the in-process fake worker."""
2
3 from __future__ import annotations
4
5 import json
6 import threading
7 from dataclasses import dataclass, field
8 from datetime import UTC, datetime
9 from pathlib import Path
10 from typing import Iterable
11
12 from scooling_lab.contracts import (
13 TrainingJobRequest,
14 TrainingJobStatus,
15 require_artifact_id,
16 require_job_id,
17 )
18 from scooling_lab.errors import ApiError, ErrorCode
19 from scooling_lab.provenance import ProvenanceRecord, validate_provenance_record
20 from scooling_lab.retention import (
21 RetentionPolicy,
22 expires_at,
23 is_expired,
24 retention_policy_from_mapping,
25 )
26 from scooling_lab.state_machine import transition
27
28
29 def utc_now_iso() -> str:
30 """Return a UTC timestamp formatted for public metadata."""
31
32 return datetime.now(UTC).replace(microsecond=0).isoformat().replace("+00:00", "Z")
33
34
35 @dataclass(frozen=True)
36 class ArtifactMetadata:
37 """Placeholder artifact metadata registered by the fake worker."""
38
39 id: str
40 job_id: str
41 dataset_hash: str
42 artifact_hash: str
43 created_at: str
44 provenance_id: str
45 retention_policy: RetentionPolicy
46
47 def to_dict(self) -> dict[str, object]:
48 """Serialize the artifact metadata for API responses and persistence."""
49
50 return {
51 "id": self.id,
52 "jobId": self.job_id,
53 "datasetHash": self.dataset_hash,
54 "artifactHash": self.artifact_hash,
55 "createdAt": self.created_at,
56 "expiresAt": expires_at(self.created_at, self.retention_policy),
57 "provenanceRecordId": self.provenance_id,
58 "retentionPolicy": self.retention_policy.to_public_dict(),
59 }
60
61 @classmethod
62 def from_dict(cls, payload: dict[str, object]) -> "ArtifactMetadata":
63 """Rehydrate persisted artifact metadata."""
64
65 return cls(
66 id=str(payload["id"]),
67 job_id=str(payload["jobId"]),
68 dataset_hash=str(payload["datasetHash"]),
69 artifact_hash=str(payload["artifactHash"]),
70 created_at=str(payload["createdAt"]),
71 provenance_id=str(payload["provenanceRecordId"]),
72 retention_policy=retention_policy_from_mapping(payload["retentionPolicy"]),
73 )
74
75
76 @dataclass(frozen=True)
77 class DeletionReceipt:
78 """Content-free result for idempotent artifact deletion."""
79
80 job_id: str
81 artifact_id: str
82 deleted: bool
83 already_deleted: bool
84 verified: bool
85
86 def to_dict(self) -> dict[str, str | bool]:
87 """Serialize deletion status without hashes, paths, or request content."""
88
89 return {
90 "alreadyDeleted": self.already_deleted,
91 "artifactId": self.artifact_id,
92 "deleted": self.deleted,
93 "jobId": self.job_id,
94 "verified": self.verified,
95 }
96
97
98 @dataclass
99 class TrainingJobRecord:
100 """Internal record for one fixture training job."""
101
102 id: str
103 request: TrainingJobRequest | None
104 status: TrainingJobStatus = TrainingJobStatus.QUEUED
105 created_at: str = field(default_factory=utc_now_iso)
106 updated_at: str = field(default_factory=utc_now_iso)
107 artifacts: list[ArtifactMetadata] = field(default_factory=list)
108 provenance: ProvenanceRecord | None = None
109 deleted_at: str | None = None
110
111 def to_public_dict(self) -> dict[str, object]:
112 """Serialize the safe job status shape for API responses."""
113
114 if self.status == TrainingJobStatus.DELETED:
115 payload: dict[str, object] = {
116 "id": self.id,
117 "status": self.status.value,
118 "createdAt": self.created_at,
119 "updatedAt": self.updated_at,
120 }
121 if self.deleted_at is not None:
122 payload["deletedAt"] = self.deleted_at
123 return payload
124 if self.request is None:
125 raise ApiError(ErrorCode.INTERNAL_ERROR, 500)
126 return {
127 "id": self.id,
128 "status": self.status.value,
129 "createdAt": self.created_at,
130 "updatedAt": self.updated_at,
131 "request": self.request.to_public_dict(),
132 }
133
134 def to_persisted_dict(self) -> dict[str, object]:
135 """Serialize the full non-secret record to the server-controlled store."""
136
137 if self.status == TrainingJobStatus.DELETED:
138 payload: dict[str, object] = {
139 "id": self.id,
140 "status": self.status.value,
141 "createdAt": self.created_at,
142 "updatedAt": self.updated_at,
143 }
144 if self.deleted_at is not None:
145 payload["deletedAt"] = self.deleted_at
146 return payload
147 if self.request is None:
148 raise ApiError(ErrorCode.INTERNAL_ERROR, 500)
149 return {
150 "id": self.id,
151 "status": self.status.value,
152 "createdAt": self.created_at,
153 "updatedAt": self.updated_at,
154 "request": {
155 "idempotencyKey": self.request.idempotency_key,
156 "datasetId": self.request.dataset_id,
157 "modelId": self.request.model_id,
158 "requestedBy": self.request.requested_by,
159 "retentionPolicy": self.request.retention_policy.to_public_dict(),
160 "trainingParameters": dict(self.request.training_parameters),
161 },
162 "artifacts": [artifact.to_dict() for artifact in self.artifacts],
163 "provenance": self.provenance.to_dict()
164 if self.provenance is not None
165 else None,
166 }
167
168 @classmethod
169 def from_persisted_dict(cls, payload: dict[str, object]) -> "TrainingJobRecord":
170 """Rehydrate a job record from server-controlled JSON."""
171
172 status = TrainingJobStatus(str(payload["status"]))
173 if status == TrainingJobStatus.DELETED:
174 deleted_at = payload.get("deletedAt")
175 return cls(
176 id=str(payload["id"]),
177 request=None,
178 status=status,
179 created_at=str(payload["createdAt"]),
180 updated_at=str(payload["updatedAt"]),
181 deleted_at=str(deleted_at) if deleted_at is not None else None,
182 )
183 request_payload = payload["request"]
184 if not isinstance(request_payload, dict):
185 raise ApiError(ErrorCode.INTERNAL_ERROR, 500)
186 artifacts_payload = payload.get("artifacts", [])
187 if not isinstance(artifacts_payload, list):
188 raise ApiError(ErrorCode.INTERNAL_ERROR, 500)
189 provenance_payload = payload.get("provenance")
190 provenance = None
191 if provenance_payload is not None:
192 if not isinstance(provenance_payload, dict):
193 raise ApiError(ErrorCode.INTERNAL_ERROR, 500)
194 provenance = ProvenanceRecord.from_mapping(provenance_payload)
195 return cls(
196 id=str(payload["id"]),
197 request=TrainingJobRequest.from_mapping(request_payload),
198 status=status,
199 created_at=str(payload["createdAt"]),
200 updated_at=str(payload["updatedAt"]),
201 artifacts=[
202 ArtifactMetadata.from_dict(artifact)
203 for artifact in artifacts_payload
204 if isinstance(artifact, dict)
205 ],
206 provenance=provenance,
207 )
208
209
210 class TrainingJobStore:
211 """In-memory job store with optional server-controlled JSON persistence."""
212
213 def __init__(self, persistence_path: Path | None = None, queue_limit: int = 5) -> None:
214 """Create a store with a bounded queued/running fixture capacity."""
215
216 self._lock = threading.RLock()
217 self._jobs: dict[str, TrainingJobRecord] = {}
218 self._persistence_path = persistence_path
219 self._queue_limit = queue_limit
220 if persistence_path is not None and persistence_path.exists():
221 self._load()
222
223 def create(self, request: TrainingJobRequest) -> TrainingJobRecord:
224 """Create or return the deterministic duplicate job for a request."""
225
226 with self._lock:
227 job_id = request.stable_job_id()
228 existing = self._jobs.get(job_id)
229 if existing is not None:
230 return existing
231 if self.active_count() >= self._queue_limit:
232 raise ApiError(ErrorCode.QUEUE_LIMIT_EXCEEDED, 429)
233 record = TrainingJobRecord(id=job_id, request=request)
234 self._jobs[job_id] = record
235 self._save()
236 return record
237
238 def get(self, job_id: str) -> TrainingJobRecord:
239 """Return one job or raise a safe not-found error."""
240
241 with self._lock:
242 record = self._jobs.get(job_id)
243 if record is None:
244 raise ApiError(ErrorCode.NOT_FOUND, 404)
245 return record
246
247 def update_status(
248 self, job_id: str, target_status: TrainingJobStatus
249 ) -> TrainingJobRecord:
250 """Apply the state machine and persist the updated job."""
251
252 with self._lock:
253 record = self.get(job_id)
254 next_status = transition(record.status, target_status)
255 if next_status != record.status:
256 record.status = next_status
257 record.updated_at = utc_now_iso()
258 self._save()
259 return record
260
261 def register_artifact(
262 self, job_id: str, artifact: ArtifactMetadata, provenance: ProvenanceRecord
263 ) -> TrainingJobRecord:
264 """Attach a placeholder artifact once, preserving retry stability."""
265
266 require_job_id(job_id)
267 require_artifact_id(artifact.id)
268 validate_provenance_record(provenance)
269 with self._lock:
270 record = self.get(job_id)
271 if all(existing.id != artifact.id for existing in record.artifacts):
272 record.artifacts.append(artifact)
273 record.provenance = provenance
274 record.updated_at = utc_now_iso()
275 self._save()
276 return record
277
278 def list_artifacts(self, job_id: str) -> list[ArtifactMetadata]:
279 """Return artifacts for one job with no cross-job scan exposure."""
280
281 require_job_id(job_id)
282 with self._lock:
283 record = self.get(job_id)
284 if record.status == TrainingJobStatus.DELETED:
285 return []
286 return list(record.artifacts)
287
288 def get_provenance(self, job_id: str) -> ProvenanceRecord:
289 """Return the validated provenance record for one completed job."""
290
291 require_job_id(job_id)
292 with self._lock:
293 record = self.get(job_id)
294 if record.status == TrainingJobStatus.DELETED or record.provenance is None:
295 raise ApiError(ErrorCode.NOT_FOUND, 404)
296 return record.provenance
297
298 def evaluate_expiry(
299 self, job_id: str, now: datetime | None = None
300 ) -> TrainingJobRecord:
301 """Evaluate a job's artifact expiry and delete it when TTL has elapsed."""
302
303 require_job_id(job_id)
304 with self._lock:
305 record = self.get(job_id)
306 self._delete_expired_artifact_locked(record, now or datetime.now(UTC))
307 return record
308
309 def sweep_expired(self, now: datetime | None = None) -> dict[str, object]:
310 """Delete every expired artifact and return a content-free sweep summary."""
311
312 deleted_job_ids: list[str] = []
313 sweep_time = now or datetime.now(UTC)
314 with self._lock:
315 for record in sorted(self._jobs.values(), key=lambda item: item.id):
316 receipt = self._delete_expired_artifact_locked(record, sweep_time)
317 if receipt is not None and receipt.deleted:
318 deleted_job_ids.append(record.id)
319 return {
320 "deletedJobIds": deleted_job_ids,
321 "deletedCount": len(deleted_job_ids),
322 }
323
324 def delete_artifact(self, job_id: str, artifact_id: str) -> DeletionReceipt:
325 """Delete an artifact, provenance, and job content as an idempotent cascade."""
326
327 require_job_id(job_id)
328 require_artifact_id(artifact_id)
329 with self._lock:
330 record = self.get(job_id)
331 if record.status == TrainingJobStatus.DELETED:
332 return DeletionReceipt(
333 job_id=job_id,
334 artifact_id=artifact_id,
335 deleted=True,
336 already_deleted=True,
337 verified=True,
338 )
339 artifact = self._find_artifact(record, artifact_id)
340 if artifact is None:
341 raise ApiError(ErrorCode.NOT_FOUND, 404)
342 return self._delete_artifact_locked(record, artifact)
343
344 def verify_hash_absence(self, hash_values: Iterable[str]) -> bool:
345 """Verify supplied hashes are absent from every store serialization."""
346
347 hashes = tuple(value for value in hash_values if value)
348 if not hashes:
349 return False
350 with self._lock:
351 output = json.dumps(
352 {
353 "artifacts": {
354 record.id: [artifact.to_dict() for artifact in record.artifacts]
355 for record in self._jobs.values()
356 },
357 "jobs": [
358 record.to_public_dict()
359 for record in sorted(self._jobs.values(), key=lambda item: item.id)
360 ],
361 "persisted": [
362 record.to_persisted_dict()
363 for record in sorted(self._jobs.values(), key=lambda item: item.id)
364 ],
365 "provenance": {
366 record.id: record.provenance.to_dict()
367 for record in self._jobs.values()
368 if record.provenance is not None
369 },
370 },
371 sort_keys=True,
372 )
373 return all(hash_value not in output for hash_value in hashes)
374
375 def active_count(self) -> int:
376 """Count queued and running jobs for quota enforcement."""
377
378 return sum(
379 1
380 for record in self._jobs.values()
381 if record.status
382 in {TrainingJobStatus.QUEUED, TrainingJobStatus.RUNNING}
383 )
384
385 def _delete_expired_artifact_locked(
386 self, record: TrainingJobRecord, now: datetime
387 ) -> DeletionReceipt | None:
388 if record.status != TrainingJobStatus.SUCCEEDED:
389 return None
390 for artifact in record.artifacts:
391 if is_expired(artifact.created_at, artifact.retention_policy, now):
392 return self._delete_artifact_locked(record, artifact, verify=False)
393 return None
394
395 def _delete_artifact_locked(
396 self, record: TrainingJobRecord, artifact: ArtifactMetadata, verify: bool = True
397 ) -> DeletionReceipt:
398 hashes = self._hashes_for_artifact(record, artifact)
399 record.status = transition(record.status, TrainingJobStatus.DELETED)
400 record.request = None
401 record.artifacts = []
402 record.provenance = None
403 record.deleted_at = utc_now_iso()
404 record.updated_at = record.deleted_at
405 self._save()
406 return DeletionReceipt(
407 job_id=record.id,
408 artifact_id=artifact.id,
409 deleted=True,
410 already_deleted=False,
411 verified=self.verify_hash_absence(hashes) if verify else True,
412 )
413
414 def _find_artifact(
415 self, record: TrainingJobRecord, artifact_id: str
416 ) -> ArtifactMetadata | None:
417 for artifact in record.artifacts:
418 if artifact.id == artifact_id:
419 return artifact
420 return None
421
422 def _hashes_for_artifact(
423 self, record: TrainingJobRecord, artifact: ArtifactMetadata
424 ) -> tuple[str, ...]:
425 hashes = [artifact.dataset_hash, artifact.artifact_hash]
426 if record.provenance is not None:
427 hashes.extend(
428 [
429 record.provenance.dataset_hash,
430 record.provenance.artifact_hash,
431 record.provenance.training_config_hash,
432 ]
433 )
434 return tuple(hashes)
435
436 def _save(self) -> None:
437 if self._persistence_path is None:
438 return
439 payload = {
440 "jobs": [
441 record.to_persisted_dict()
442 for record in sorted(self._jobs.values(), key=lambda item: item.id)
443 ]
444 }
445 self._persistence_path.parent.mkdir(parents=True, exist_ok=True)
446 self._persistence_path.write_text(
447 json.dumps(payload, indent=2, sort_keys=True) + "\n",
448 encoding="utf-8",
449 )
450
451 def _load(self) -> None:
452 if self._persistence_path is None:
453 return
454 payload = json.loads(self._persistence_path.read_text(encoding="utf-8"))
455 jobs = payload.get("jobs", [])
456 if not isinstance(jobs, list):
457 raise ApiError(ErrorCode.INTERNAL_ERROR, 500)
458 self._jobs = {
459 record.id: record
460 for record in (
461 TrainingJobRecord.from_persisted_dict(job)
462 for job in jobs
463 if isinstance(job, dict)
464 )
465 }
File History 1 commit
sha256:91e875d4a97bb1e35f37992f803988d5713931f1782d870c371c50054574af22 Add fixture provenance retention deletion Human minor 42 days ago