labtasker-server 2.0.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -0,0 +1,3 @@
1
+ """Labtasker v2 HTTP server package."""
2
+
3
+ __version__ = "2.0.0"
@@ -0,0 +1,3 @@
1
+ from labtasker_server.cli import app
2
+
3
+ app()
@@ -0,0 +1,416 @@
1
+ from __future__ import annotations
2
+
3
+ import asyncio
4
+ import hmac
5
+ import logging
6
+ from collections.abc import AsyncIterator, Callable
7
+ from contextlib import asynccontextmanager, suppress
8
+ from typing import Annotated, Any
9
+
10
+ from fastapi import Depends, FastAPI, Query, Request, Response
11
+ from fastapi.exceptions import RequestValidationError
12
+ from fastapi.responses import JSONResponse
13
+ from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
14
+ from sqlalchemy import text
15
+
16
+ from labtasker_server.config import ServerSettings
17
+ from labtasker_server.database import Database
18
+ from labtasker_server.errors import DomainError
19
+ from labtasker_server.middleware import RequestBodyLimitMiddleware
20
+ from labtasker_server.schemas import (
21
+ BulkUpdateRequest,
22
+ BulkUpdateResult,
23
+ ClaimRequest,
24
+ ClaimResponse,
25
+ CompleteRequest,
26
+ CountResponse,
27
+ ErrorEnvelope,
28
+ FailRequest,
29
+ HealthyResponse,
30
+ HeartbeatResponse,
31
+ Queue,
32
+ RunRequest,
33
+ Task,
34
+ TaskCreate,
35
+ TaskOrderField,
36
+ TaskPage,
37
+ TaskStatus,
38
+ TaskUpdate,
39
+ UnhealthyResponse,
40
+ )
41
+ from labtasker_server.services.queues import QueueService
42
+ from labtasker_server.services.tasks import TaskService, system_now_us
43
+ from labtasker_server.validation import MAX_TASK_DATA_BYTES
44
+
45
+ EXPIRY_SCAN_INTERVAL_SECONDS = 60
46
+ logger = logging.getLogger(__name__)
47
+ BEARER = HTTPBearer(auto_error=False)
48
+ API_ERROR_RESPONSES: dict[int | str, dict[str, Any]] = {
49
+ status: {"model": ErrorEnvelope} for status in (401, 404, 409, 413, 422, 503)
50
+ }
51
+
52
+
53
+ def create_app(
54
+ settings: ServerSettings,
55
+ *,
56
+ now_us: Callable[[], int] = system_now_us,
57
+ ) -> FastAPI:
58
+ database = Database(settings.database, ownership_fd=settings.database_fd)
59
+ try:
60
+ database.initialize()
61
+ except BaseException:
62
+ database.dispose()
63
+ raise
64
+ queue_service = QueueService(database)
65
+ task_service = TaskService(database, now_us=now_us)
66
+ task_service.expire_leases()
67
+
68
+ @asynccontextmanager
69
+ async def lifespan(_: FastAPI) -> AsyncIterator[None]:
70
+ scanner = asyncio.create_task(_expiry_scanner(task_service))
71
+ try:
72
+ yield
73
+ finally:
74
+ scanner.cancel()
75
+ with suppress(asyncio.CancelledError):
76
+ await scanner
77
+ database.dispose()
78
+
79
+ app = FastAPI(docs_url=None, redoc_url=None, lifespan=lifespan)
80
+ app.add_middleware(RequestBodyLimitMiddleware, max_bytes=MAX_TASK_DATA_BYTES)
81
+ app.state.database = database
82
+ app.state.settings = settings
83
+ app.state.task_service = task_service
84
+
85
+ @app.exception_handler(DomainError)
86
+ async def handle_domain_error(_: Request, exc: DomainError) -> JSONResponse:
87
+ headers = {"WWW-Authenticate": "Bearer"} if exc.status_code == 401 else None
88
+ return JSONResponse(
89
+ status_code=exc.status_code,
90
+ content={"error": {"code": exc.code, "message": exc.message, "details": exc.details}},
91
+ headers=headers,
92
+ )
93
+
94
+ @app.exception_handler(RequestValidationError)
95
+ async def handle_request_validation(
96
+ request: Request,
97
+ exc: RequestValidationError,
98
+ ) -> JSONResponse:
99
+ code, details = _validation_error(request, exc)
100
+ if details is not None:
101
+ return JSONResponse(
102
+ status_code=422,
103
+ content={
104
+ "error": {
105
+ "code": code,
106
+ "message": (
107
+ "Request validation failed."
108
+ if code == "invalid_request"
109
+ else _specific_validation_message(code)
110
+ ),
111
+ "details": details,
112
+ }
113
+ },
114
+ )
115
+ errors = []
116
+ for error in exc.errors():
117
+ location = list(error.get("loc", ()))
118
+ if not location:
119
+ location = ["body"]
120
+ errors.append(
121
+ {
122
+ "location": location,
123
+ "message": str(error.get("msg", "Invalid value.")),
124
+ }
125
+ )
126
+ return JSONResponse(
127
+ status_code=422,
128
+ content={
129
+ "error": {
130
+ "code": code,
131
+ "message": "Request validation failed.",
132
+ "details": {"errors": errors},
133
+ }
134
+ },
135
+ )
136
+
137
+ def require_auth(
138
+ credentials: Annotated[HTTPAuthorizationCredentials | None, Depends(BEARER)],
139
+ ) -> None:
140
+ token = settings.token
141
+ if token is None:
142
+ return
143
+ if credentials is None or credentials.scheme.lower() != "bearer":
144
+ raise _unauthorized()
145
+ if not hmac.compare_digest(credentials.credentials, token):
146
+ raise _unauthorized()
147
+
148
+ authenticated = [Depends(require_auth)]
149
+
150
+ @app.get(
151
+ "/health",
152
+ response_model=HealthyResponse,
153
+ responses={503: {"model": UnhealthyResponse}},
154
+ )
155
+ def health() -> JSONResponse:
156
+ try:
157
+ with database.read_session() as session:
158
+ session.execute(text("SELECT 1"))
159
+ except Exception:
160
+ return JSONResponse(
161
+ status_code=503,
162
+ content={"status": "error", "api_version": "2", "database": "error"},
163
+ )
164
+ return JSONResponse(
165
+ status_code=200,
166
+ content={"status": "ok", "api_version": "2", "database": "ok"},
167
+ )
168
+
169
+ @app.put(
170
+ "/api/v2/queues/{queue}",
171
+ response_model=Queue,
172
+ dependencies=authenticated,
173
+ responses={**API_ERROR_RESPONSES, 201: {"model": Queue}},
174
+ )
175
+ def create_queue(queue: str, response: Response) -> Queue:
176
+ result, created = queue_service.create(queue)
177
+ response.status_code = 201 if created else 200
178
+ return result
179
+
180
+ @app.get(
181
+ "/api/v2/queues",
182
+ response_model=list[Queue],
183
+ dependencies=authenticated,
184
+ responses=API_ERROR_RESPONSES,
185
+ )
186
+ def list_queues() -> list[Queue]:
187
+ return queue_service.list()
188
+
189
+ @app.delete(
190
+ "/api/v2/queues/{queue}",
191
+ status_code=204,
192
+ dependencies=authenticated,
193
+ responses=API_ERROR_RESPONSES,
194
+ )
195
+ def delete_queue(queue: str, cascade: bool = False) -> Response:
196
+ queue_service.delete(queue, cascade=cascade)
197
+ return Response(status_code=204)
198
+
199
+ @app.put(
200
+ "/api/v2/queues/{queue}/tasks/{task_id}",
201
+ response_model=Task,
202
+ dependencies=authenticated,
203
+ responses={**API_ERROR_RESPONSES, 201: {"model": Task}},
204
+ )
205
+ def create_task(queue: str, task_id: str, request: TaskCreate, response: Response) -> Task:
206
+ result, created = task_service.create(queue, task_id, request)
207
+ response.status_code = 201 if created else 200
208
+ return result
209
+
210
+ @app.get(
211
+ "/api/v2/queues/{queue}/tasks",
212
+ response_model=TaskPage,
213
+ dependencies=authenticated,
214
+ responses=API_ERROR_RESPONSES,
215
+ )
216
+ def list_tasks(
217
+ queue: str,
218
+ status: TaskStatus | None = None,
219
+ name: str | None = None,
220
+ filter_expression: Annotated[str | None, Query(alias="filter")] = None,
221
+ order_by: TaskOrderField = "created_at",
222
+ descending: bool = True,
223
+ limit: Annotated[int, Query(ge=1, le=1000)] = 100,
224
+ cursor: str | None = None,
225
+ ) -> TaskPage:
226
+ return task_service.list_tasks(
227
+ queue,
228
+ status=status,
229
+ name=name,
230
+ filter_expression=filter_expression,
231
+ order_by=order_by,
232
+ descending=descending,
233
+ limit=limit,
234
+ cursor=cursor,
235
+ )
236
+
237
+ @app.get(
238
+ "/api/v2/queues/{queue}/tasks/count",
239
+ response_model=CountResponse,
240
+ dependencies=authenticated,
241
+ responses=API_ERROR_RESPONSES,
242
+ )
243
+ def count_tasks(
244
+ queue: str,
245
+ status: TaskStatus | None = None,
246
+ name: str | None = None,
247
+ filter_expression: Annotated[str | None, Query(alias="filter")] = None,
248
+ ) -> CountResponse:
249
+ return CountResponse(
250
+ count=task_service.count_tasks(
251
+ queue,
252
+ status=status,
253
+ name=name,
254
+ filter_expression=filter_expression,
255
+ )
256
+ )
257
+
258
+ @app.get(
259
+ "/api/v2/queues/{queue}/tasks/{task_id}",
260
+ response_model=Task,
261
+ dependencies=authenticated,
262
+ responses=API_ERROR_RESPONSES,
263
+ )
264
+ def get_task(queue: str, task_id: str) -> Task:
265
+ return task_service.get(queue, task_id)
266
+
267
+ @app.patch(
268
+ "/api/v2/queues/{queue}/tasks/{task_id}",
269
+ response_model=Task,
270
+ dependencies=authenticated,
271
+ responses=API_ERROR_RESPONSES,
272
+ )
273
+ def update_task(queue: str, task_id: str, changes: TaskUpdate) -> Task:
274
+ return task_service.update_task(queue, task_id, changes)
275
+
276
+ @app.patch(
277
+ "/api/v2/queues/{queue}/tasks",
278
+ response_model=BulkUpdateResult,
279
+ dependencies=authenticated,
280
+ responses=API_ERROR_RESPONSES,
281
+ )
282
+ def update_tasks(queue: str, request: BulkUpdateRequest) -> BulkUpdateResult:
283
+ return task_service.update_tasks(
284
+ queue,
285
+ filter_expression=request.filter,
286
+ changes=request.changes,
287
+ )
288
+
289
+ @app.post(
290
+ "/api/v2/queues/{queue}/tasks/claim",
291
+ response_model=ClaimResponse,
292
+ dependencies=authenticated,
293
+ responses={**API_ERROR_RESPONSES, 204: {"description": "No eligible Task."}},
294
+ )
295
+ def claim_task(queue: str, request: ClaimRequest) -> ClaimResponse | Response:
296
+ claim = task_service.claim(queue, request.route, request.run_id)
297
+ return Response(status_code=204) if claim is None else claim
298
+
299
+ @app.post(
300
+ "/api/v2/queues/{queue}/tasks/{task_id}/heartbeat",
301
+ response_model=HeartbeatResponse,
302
+ dependencies=authenticated,
303
+ responses=API_ERROR_RESPONSES,
304
+ )
305
+ def heartbeat(queue: str, task_id: str, request: RunRequest) -> HeartbeatResponse:
306
+ return task_service.heartbeat(queue, task_id, request.run_id)
307
+
308
+ @app.post(
309
+ "/api/v2/queues/{queue}/tasks/{task_id}/complete",
310
+ status_code=204,
311
+ dependencies=authenticated,
312
+ responses=API_ERROR_RESPONSES,
313
+ )
314
+ def complete(queue: str, task_id: str, request: CompleteRequest) -> Response:
315
+ task_service.complete(queue, task_id, request.run_id, request.result)
316
+ return Response(status_code=204)
317
+
318
+ @app.post(
319
+ "/api/v2/queues/{queue}/tasks/{task_id}/fail",
320
+ status_code=204,
321
+ dependencies=authenticated,
322
+ responses=API_ERROR_RESPONSES,
323
+ )
324
+ def fail(queue: str, task_id: str, request: FailRequest) -> Response:
325
+ task_service.fail(queue, task_id, request.run_id, request.error)
326
+ return Response(status_code=204)
327
+
328
+ @app.post(
329
+ "/api/v2/queues/{queue}/tasks/{task_id}/unclaim",
330
+ status_code=204,
331
+ dependencies=authenticated,
332
+ responses=API_ERROR_RESPONSES,
333
+ )
334
+ def unclaim(queue: str, task_id: str, request: RunRequest) -> Response:
335
+ task_service.unclaim(queue, task_id, request.run_id)
336
+ return Response(status_code=204)
337
+
338
+ @app.post(
339
+ "/api/v2/queues/{queue}/tasks/{task_id}/cancel",
340
+ response_model=Task,
341
+ dependencies=authenticated,
342
+ responses=API_ERROR_RESPONSES,
343
+ )
344
+ def cancel_task(queue: str, task_id: str) -> Task:
345
+ return task_service.cancel(queue, task_id)
346
+
347
+ @app.post(
348
+ "/api/v2/queues/{queue}/tasks/{task_id}/requeue",
349
+ response_model=Task,
350
+ dependencies=authenticated,
351
+ responses=API_ERROR_RESPONSES,
352
+ )
353
+ def requeue_task(queue: str, task_id: str) -> Task:
354
+ return task_service.requeue(queue, task_id)
355
+
356
+ @app.delete(
357
+ "/api/v2/queues/{queue}/tasks/{task_id}",
358
+ status_code=204,
359
+ dependencies=authenticated,
360
+ responses=API_ERROR_RESPONSES,
361
+ )
362
+ def delete_task(queue: str, task_id: str) -> Response:
363
+ task_service.delete(queue, task_id)
364
+ return Response(status_code=204)
365
+
366
+ return app
367
+
368
+
369
+ async def _expiry_scanner(task_service: TaskService) -> None:
370
+ while True:
371
+ await asyncio.sleep(EXPIRY_SCAN_INTERVAL_SECONDS)
372
+ try:
373
+ await asyncio.to_thread(task_service.expire_leases)
374
+ except Exception:
375
+ logger.exception("Heartbeat expiry scan failed; it will retry in 60 seconds.")
376
+
377
+
378
+ def _unauthorized() -> DomainError:
379
+ return DomainError(401, "unauthorized", "Authentication is required.", {})
380
+
381
+
382
+ def _validation_code(request: Request) -> str:
383
+ if request.method == "PUT" and "/tasks/" in request.url.path:
384
+ return "invalid_task"
385
+ if request.method == "PATCH" and "/tasks" in request.url.path:
386
+ return "invalid_update"
387
+ return "invalid_request"
388
+
389
+
390
+ def _validation_error(
391
+ request: Request,
392
+ exc: RequestValidationError,
393
+ ) -> tuple[str, dict[str, object] | None]:
394
+ for error in exc.errors():
395
+ error_type = str(error.get("type", ""))
396
+ if error_type == "json_invalid":
397
+ return "invalid_request", {
398
+ "errors": [{"location": ["body"], "message": "Malformed JSON body."}]
399
+ }
400
+ if error_type in {"invalid_task_name", "json_too_deep"}:
401
+ context = error.get("ctx")
402
+ return error_type, dict(context) if isinstance(context, dict) else {}
403
+ if request.method == "PATCH" and request.url.path.endswith("/tasks"):
404
+ for error in exc.errors():
405
+ location = tuple(error.get("loc", ()))
406
+ if location[:2] == ("body", "filter"):
407
+ return "invalid_filter", None
408
+ return _validation_code(request), None
409
+
410
+
411
+ def _specific_validation_message(code: str) -> str:
412
+ if code == "json_too_deep":
413
+ return "JSON value is too deeply nested."
414
+ if code == "invalid_task_name":
415
+ return "Task name is invalid."
416
+ raise AssertionError(f"Unknown specific validation code: {code}")