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.
- labtasker_server/__init__.py +3 -0
- labtasker_server/__main__.py +3 -0
- labtasker_server/app.py +416 -0
- labtasker_server/cli.py +371 -0
- labtasker_server/config.py +42 -0
- labtasker_server/database.py +176 -0
- labtasker_server/errors.py +27 -0
- labtasker_server/filtering.py +528 -0
- labtasker_server/local.py +453 -0
- labtasker_server/logging.py +47 -0
- labtasker_server/middleware.py +73 -0
- labtasker_server/migrations/__init__.py +1 -0
- labtasker_server/migrations/env.py +20 -0
- labtasker_server/migrations/versions/0001_initial.py +135 -0
- labtasker_server/migrations/versions/__init__.py +1 -0
- labtasker_server/models.py +139 -0
- labtasker_server/pagination.py +126 -0
- labtasker_server/py.typed +1 -0
- labtasker_server/schemas.py +265 -0
- labtasker_server/services/__init__.py +1 -0
- labtasker_server/services/queues.py +66 -0
- labtasker_server/services/tasks.py +859 -0
- labtasker_server/validation.py +150 -0
- labtasker_server-2.0.0.dist-info/METADATA +16 -0
- labtasker_server-2.0.0.dist-info/RECORD +28 -0
- labtasker_server-2.0.0.dist-info/WHEEL +4 -0
- labtasker_server-2.0.0.dist-info/entry_points.txt +2 -0
- labtasker_server-2.0.0.dist-info/licenses/LICENSE +201 -0
labtasker_server/app.py
ADDED
|
@@ -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}")
|