api.py python
209 lines 8.0 KB
Raw
sha256:fc4c9ad652d1fff3dc508cb6ea02ee710ee6dfc4cb3761291d9900b5e029ea8a feat(slice-7): T3 dataset review lifecycle, job queue, prov… Human minor ⚠ breaking 41 days ago
1 """Dependency-free HTTP API for the Scooling Lab T2/T3 contract."""
2
3 from __future__ import annotations
4
5 import json
6 from http import HTTPStatus
7 from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
8 from pathlib import Path
9 import re
10 from typing import Callable
11 from urllib.parse import urlparse
12
13 from scooling_lab.errors import ApiError, ErrorCode, error_payload
14 from scooling_lab.service import TrainingApiService
15 from scooling_lab.store import TrainingJobStore
16
17
18 MAX_BODY_BYTES = 16_384
19 JOB_ID_RE = re.compile(r"^job_[a-f0-9]{24}$")
20 ARTIFACT_ID_RE = re.compile(r"^artifact_[a-f0-9]{24}$")
21 DATASET_ID_RE = re.compile(r"^[A-Za-z0-9._:-]{3,96}$")
22
23
24 def parse_job_route(path: str, suffix: str = "") -> str | None:
25 """Extract a safe job id from supported job subresource routes."""
26
27 prefix = "/training/jobs/"
28 if not path.startswith(prefix):
29 return None
30 remainder = path.removeprefix(prefix)
31 if suffix:
32 ending = f"/{suffix}"
33 if not remainder.endswith(ending):
34 return None
35 remainder = remainder[: -len(ending)]
36 if "/" in remainder or not JOB_ID_RE.fullmatch(remainder):
37 return None
38 return remainder
39
40
41 def parse_artifact_route(path: str) -> tuple[str, str] | None:
42 """Extract safe job and artifact ids from artifact deletion routes."""
43
44 prefix = "/training/jobs/"
45 marker = "/artifacts/"
46 if not path.startswith(prefix) or marker not in path:
47 return None
48 remainder = path.removeprefix(prefix)
49 job_id, separator, artifact_id = remainder.partition(marker)
50 if separator != marker:
51 return None
52 if "/" in artifact_id:
53 return None
54 if not JOB_ID_RE.fullmatch(job_id) or not ARTIFACT_ID_RE.fullmatch(artifact_id):
55 return None
56 return job_id, artifact_id
57
58
59 def parse_dataset_route(path: str, suffix: str = "") -> str | None:
60 """Extract a safe dataset id from supported dataset subresource routes."""
61
62 prefix = "/datasets/"
63 if not path.startswith(prefix):
64 return None
65 remainder = path.removeprefix(prefix)
66 if suffix:
67 ending = f"/{suffix}"
68 if not remainder.endswith(ending):
69 return None
70 remainder = remainder[: -len(ending)]
71 if "/" in remainder or not DATASET_ID_RE.fullmatch(remainder):
72 return None
73 return remainder
74
75
76 def make_handler(service: TrainingApiService) -> type[BaseHTTPRequestHandler]:
77 """Create a request handler bound to the supplied service."""
78
79 class ScoolingLabRequestHandler(BaseHTTPRequestHandler):
80 """HTTP handler exposing create/get/cancel/list artifact endpoints."""
81
82 server_version = "ScoolingLab/0.1"
83
84 def do_POST(self) -> None:
85 """Handle createTrainingJob, cancelTrainingJob, and dataset routes."""
86
87 path = urlparse(self.path).path
88 if path == "/training/jobs":
89 self._handle_json(lambda: service.create_training_job(self._read_json()))
90 return
91 cancel_job_id = parse_job_route(path, "cancel")
92 if cancel_job_id is not None:
93 self._handle_json(lambda: service.cancel_training_job(cancel_job_id))
94 return
95 if path == "/datasets":
96 self._handle_json(lambda: service.register_dataset(self._read_json()))
97 return
98 review_dataset_id = parse_dataset_route(path, "review")
99 if review_dataset_id is not None:
100 _id = review_dataset_id
101 self._handle_json(
102 lambda: service.review_dataset(_id, self._read_json())
103 )
104 return
105 submit_dataset_id = parse_dataset_route(path, "submit")
106 if submit_dataset_id is not None:
107 _sid = submit_dataset_id
108 self._handle_json(lambda: service.submit_dataset_for_review(_sid))
109 return
110 self._send_error(ApiError(ErrorCode.NOT_FOUND, 404))
111
112 def do_GET(self) -> None:
113 """Handle getTrainingJob, listArtifacts, queue state, and dataset routes."""
114
115 path = urlparse(self.path).path
116 artifacts_job_id = parse_job_route(path, "artifacts")
117 if artifacts_job_id is not None:
118 self._handle_json(lambda: service.list_artifacts(artifacts_job_id))
119 return
120 provenance_job_id = parse_job_route(path, "provenance")
121 if provenance_job_id is not None:
122 self._handle_json(lambda: service.get_provenance(provenance_job_id))
123 return
124 job_id = parse_job_route(path)
125 if job_id is not None:
126 self._handle_json(lambda: service.get_training_job(job_id))
127 return
128 if path == "/training/queue":
129 self._handle_json(service.get_queue_state)
130 return
131 dataset_id = parse_dataset_route(path)
132 if dataset_id is not None:
133 _did = dataset_id
134 self._handle_json(lambda: service.get_dataset(_did))
135 return
136 self._send_error(ApiError(ErrorCode.NOT_FOUND, 404))
137
138 def do_PUT(self) -> None:
139 """Reject unsupported mutation routes with a stable error."""
140
141 self._send_error(ApiError(ErrorCode.METHOD_NOT_ALLOWED, 405))
142
143 def do_DELETE(self) -> None:
144 """Handle idempotent deleteArtifact routes."""
145
146 path = urlparse(self.path).path
147 artifact_route = parse_artifact_route(path)
148 if artifact_route is not None:
149 job_id, artifact_id = artifact_route
150 self._handle_json(lambda: service.delete_artifact(job_id, artifact_id))
151 return
152 self._send_error(ApiError(ErrorCode.METHOD_NOT_ALLOWED, 405))
153
154 def log_message(self, format: str, *args: object) -> None:
155 """Suppress default request logging to avoid payload/path leakage."""
156
157 return
158
159 def _read_json(self) -> dict[str, object]:
160 content_length = self.headers.get("Content-Length")
161 if content_length is None:
162 raise ApiError(ErrorCode.MALFORMED_JSON, 400)
163 length = int(content_length)
164 if length <= 0 or length > MAX_BODY_BYTES:
165 raise ApiError(ErrorCode.VALIDATION_ERROR, 400)
166 try:
167 payload = json.loads(self.rfile.read(length).decode("utf-8"))
168 except (UnicodeDecodeError, json.JSONDecodeError) as exc:
169 raise ApiError(ErrorCode.MALFORMED_JSON, 400) from exc
170 if not isinstance(payload, dict):
171 raise ApiError(ErrorCode.VALIDATION_ERROR, 400)
172 return payload
173
174 def _handle_json(self, action: Callable[[], object]) -> None:
175 try:
176 result = action()
177 except ApiError as error:
178 self._send_error(error)
179 return
180 except Exception:
181 self._send_error(ApiError(ErrorCode.INTERNAL_ERROR, 500))
182 return
183 self._send_json(result, HTTPStatus.OK)
184
185 def _send_error(self, error: ApiError) -> None:
186 self._send_json(error_payload(error), HTTPStatus(error.status))
187
188 def _send_json(self, payload: object, status: HTTPStatus) -> None:
189 body = json.dumps(payload, sort_keys=True).encode("utf-8")
190 self.send_response(status.value)
191 self.send_header("Content-Type", "application/json")
192 self.send_header("Content-Length", str(len(body)))
193 self.send_header("Cache-Control", "no-store")
194 self.end_headers()
195 self.wfile.write(body)
196
197 return ScoolingLabRequestHandler
198
199
200 def run_server(host: str, port: int, persistence_path: Path | None = None) -> None:
201 """Run the Scooling Lab API server until interrupted."""
202
203 service = TrainingApiService(TrainingJobStore(persistence_path=persistence_path))
204 server = ThreadingHTTPServer((host, port), make_handler(service))
205 server.serve_forever()
206
207
208 if __name__ == "__main__":
209 run_server("127.0.0.1", 8080)
File History 1 commit
sha256:fc4c9ad652d1fff3dc508cb6ea02ee710ee6dfc4cb3761291d9900b5e029ea8a feat(slice-7): T3 dataset review lifecycle, job queue, prov… Human minor 41 days ago