contracts.py python
225 lines 7.7 KB
Raw
sha256:91e875d4a97bb1e35f37992f803988d5713931f1782d870c371c50054574af22 Add fixture provenance retention deletion Human minor ⚠ breaking 42 days ago
1 """Training API request contracts and schema-layer validation."""
2
3 from __future__ import annotations
4
5 import hashlib
6 import json
7 import re
8 from dataclasses import dataclass, field
9 from enum import Enum
10 from types import MappingProxyType
11 from typing import Mapping
12
13 from scooling_lab.errors import ApiError, ErrorCode
14 from scooling_lab.retention import (
15 RetentionPolicy,
16 default_retention_policy,
17 retention_policy_from_mapping,
18 )
19
20
21 class TrainingJobStatus(str, Enum):
22 """States allowed by the T2 training job state machine."""
23
24 QUEUED = "queued"
25 RUNNING = "running"
26 SUCCEEDED = "succeeded"
27 FAILED = "failed"
28 CANCELLED = "cancelled"
29 DELETED = "deleted"
30
31
32 ALLOWED_MODEL_IDS: frozenset[str] = frozenset({"fixture-tiny-llm"})
33 ALLOWED_DATASET_IDS: frozenset[str] = frozenset({"fixture:synthetic-tiny-v1"})
34 ALLOWED_REQUEST_KEYS: frozenset[str] = frozenset(
35 {
36 "idempotencyKey",
37 "datasetId",
38 "modelId",
39 "requestedBy",
40 "retentionPolicy",
41 "trainingParameters",
42 }
43 )
44 FORBIDDEN_KEY_TERMS: tuple[str, ...] = (
45 "url",
46 "uri",
47 "path",
48 "file",
49 "shell",
50 "command",
51 "callback",
52 "webhook",
53 "worker",
54 )
55 SAFE_IDENTIFIER_RE = re.compile(r"^[A-Za-z0-9._:-]{3,96}$")
56 REQUESTER_RE = re.compile(r"^[A-Za-z0-9._:-]{3,80}$")
57 JOB_ID_RE = re.compile(r"^job_[a-f0-9]{24}$")
58 ARTIFACT_ID_RE = re.compile(r"^artifact_[a-f0-9]{24}$")
59 FORBIDDEN_STRING_RE = re.compile(
60 r"(?i)(https?://|file://|ssh://|[;&|`$<>]|\.\./|/\w|[A-Za-z]:\\)"
61 )
62
63
64 @dataclass(frozen=True)
65 class TrainingJobRequest:
66 """Validated createTrainingJob payload.
67
68 The schema only accepts inert fixture identifiers, bounded numeric training
69 parameters, and a caller-provided idempotency key. It rejects unknown keys so
70 browser-supplied worker URLs, callback URLs, file paths, and shell strings
71 fail before reaching the service layer.
72 """
73
74 idempotency_key: str
75 dataset_id: str
76 model_id: str
77 requested_by: str
78 retention_policy: RetentionPolicy = field(default_factory=default_retention_policy)
79 training_parameters: Mapping[str, int | float | bool] = field(
80 default_factory=lambda: MappingProxyType({})
81 )
82
83 @classmethod
84 def from_mapping(cls, payload: Mapping[str, object]) -> "TrainingJobRequest":
85 """Validate and convert an untrusted JSON object into a request."""
86
87 reject_forbidden_keys(payload)
88 unknown_keys = set(payload).difference(ALLOWED_REQUEST_KEYS)
89 if unknown_keys:
90 raise ApiError(ErrorCode.VALIDATION_ERROR, 400)
91
92 idempotency_key = require_safe_identifier(payload.get("idempotencyKey"))
93 dataset_id = require_safe_identifier(payload.get("datasetId"))
94 model_id = require_safe_identifier(payload.get("modelId"))
95 requested_by = require_requester(payload.get("requestedBy"))
96 retention_policy = retention_policy_from_mapping(payload.get("retentionPolicy"))
97 training_parameters = validate_training_parameters(
98 payload.get("trainingParameters", {})
99 )
100
101 if dataset_id not in ALLOWED_DATASET_IDS:
102 raise ApiError(ErrorCode.VALIDATION_ERROR, 400)
103 if model_id not in ALLOWED_MODEL_IDS:
104 raise ApiError(ErrorCode.VALIDATION_ERROR, 400)
105
106 return cls(
107 idempotency_key=idempotency_key,
108 dataset_id=dataset_id,
109 model_id=model_id,
110 requested_by=requested_by,
111 retention_policy=retention_policy,
112 training_parameters=MappingProxyType(training_parameters),
113 )
114
115 def fingerprint(self) -> str:
116 """Return a stable hash input for idempotent job creation."""
117
118 return json.dumps(
119 {
120 "datasetId": self.dataset_id,
121 "idempotencyKey": self.idempotency_key,
122 "modelId": self.model_id,
123 "requestedBy": self.requested_by,
124 "retentionPolicy": self.retention_policy.to_public_dict(),
125 "trainingParameters": dict(sorted(self.training_parameters.items())),
126 },
127 separators=(",", ":"),
128 sort_keys=True,
129 )
130
131 def stable_job_id(self) -> str:
132 """Return the deterministic job id for this request."""
133
134 digest = hashlib.sha256(self.fingerprint().encode("utf-8")).hexdigest()[:24]
135 return f"job_{digest}"
136
137 def to_public_dict(self) -> dict[str, object]:
138 """Serialize only safe, non-private request fields."""
139
140 return {
141 "datasetId": self.dataset_id,
142 "modelId": self.model_id,
143 "requestedBy": self.requested_by,
144 "retentionPolicy": self.retention_policy.to_public_dict(),
145 "trainingParameters": dict(sorted(self.training_parameters.items())),
146 }
147
148
149 def reject_forbidden_keys(payload: Mapping[str, object]) -> None:
150 """Reject explicitly dangerous key shapes anywhere in a request payload."""
151
152 for key, value in payload.items():
153 lowered = key.lower()
154 if any(term in lowered for term in FORBIDDEN_KEY_TERMS):
155 raise ApiError(ErrorCode.VALIDATION_ERROR, 400)
156 if isinstance(value, Mapping):
157 reject_forbidden_keys(value)
158
159
160 def require_safe_identifier(value: object) -> str:
161 """Validate compact fixture identifiers and reject path or URL strings."""
162
163 if not isinstance(value, str) or not SAFE_IDENTIFIER_RE.fullmatch(value):
164 raise ApiError(ErrorCode.VALIDATION_ERROR, 400)
165 if FORBIDDEN_STRING_RE.search(value):
166 raise ApiError(ErrorCode.VALIDATION_ERROR, 400)
167 return value
168
169
170 def require_requester(value: object) -> str:
171 """Validate the non-secret caller label used for fixture audit context."""
172
173 if not isinstance(value, str) or not REQUESTER_RE.fullmatch(value):
174 raise ApiError(ErrorCode.VALIDATION_ERROR, 400)
175 if FORBIDDEN_STRING_RE.search(value):
176 raise ApiError(ErrorCode.VALIDATION_ERROR, 400)
177 return value
178
179
180 def require_job_id(value: object) -> str:
181 """Validate a server-generated job id before any store lookup."""
182
183 if not isinstance(value, str) or not JOB_ID_RE.fullmatch(value):
184 raise ApiError(ErrorCode.VALIDATION_ERROR, 400)
185 return value
186
187
188 def require_artifact_id(value: object) -> str:
189 """Validate a server-generated artifact id before any mutation."""
190
191 if not isinstance(value, str) or not ARTIFACT_ID_RE.fullmatch(value):
192 raise ApiError(ErrorCode.VALIDATION_ERROR, 400)
193 return value
194
195
196 def validate_training_parameters(value: object) -> dict[str, int | float | bool]:
197 """Validate the bounded dry-run parameter set for the fake worker."""
198
199 if not isinstance(value, Mapping):
200 raise ApiError(ErrorCode.VALIDATION_ERROR, 400)
201 allowed_keys = {"epochs", "learningRate", "dryRun"}
202 if set(value).difference(allowed_keys):
203 raise ApiError(ErrorCode.VALIDATION_ERROR, 400)
204
205 parameters: dict[str, int | float | bool] = {}
206 if "epochs" in value:
207 epochs = value["epochs"]
208 if not isinstance(epochs, int) or isinstance(epochs, bool) or not 1 <= epochs <= 3:
209 raise ApiError(ErrorCode.VALIDATION_ERROR, 400)
210 parameters["epochs"] = epochs
211 if "learningRate" in value:
212 learning_rate = value["learningRate"]
213 if (
214 not isinstance(learning_rate, (int, float))
215 or isinstance(learning_rate, bool)
216 or not 0 < float(learning_rate) <= 1
217 ):
218 raise ApiError(ErrorCode.VALIDATION_ERROR, 400)
219 parameters["learningRate"] = float(learning_rate)
220 if "dryRun" in value:
221 dry_run = value["dryRun"]
222 if not isinstance(dry_run, bool) or dry_run is not True:
223 raise ApiError(ErrorCode.VALIDATION_ERROR, 400)
224 parameters["dryRun"] = dry_run
225 return parameters
File History 1 commit
sha256:91e875d4a97bb1e35f37992f803988d5713931f1782d870c371c50054574af22 Add fixture provenance retention deletion Human minor 42 days ago