labtasker-server 2.1.0__tar.gz → 2.2.0__tar.gz
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-2.1.0 → labtasker_server-2.2.0}/.gitignore +4 -0
- {labtasker_server-2.1.0 → labtasker_server-2.2.0}/PKG-INFO +2 -2
- {labtasker_server-2.1.0 → labtasker_server-2.2.0}/pyproject.toml +3 -3
- {labtasker_server-2.1.0 → labtasker_server-2.2.0}/src/labtasker_server/__init__.py +1 -1
- {labtasker_server-2.1.0 → labtasker_server-2.2.0}/src/labtasker_server/app.py +166 -24
- {labtasker_server-2.1.0 → labtasker_server-2.2.0}/src/labtasker_server/cli.py +5 -7
- {labtasker_server-2.1.0 → labtasker_server-2.2.0}/src/labtasker_server/database.py +54 -40
- {labtasker_server-2.1.0 → labtasker_server-2.2.0}/src/labtasker_server/filtering.py +118 -30
- labtasker_server-2.2.0/src/labtasker_server/grouping.py +136 -0
- {labtasker_server-2.1.0 → labtasker_server-2.2.0}/src/labtasker_server/local.py +38 -26
- {labtasker_server-2.1.0 → labtasker_server-2.2.0}/src/labtasker_server/middleware.py +38 -0
- labtasker_server-2.2.0/src/labtasker_server/migrations/versions/0002_worker_observations.py +33 -0
- {labtasker_server-2.1.0 → labtasker_server-2.2.0}/src/labtasker_server/models.py +20 -0
- labtasker_server-2.2.0/src/labtasker_server/name_search.py +23 -0
- labtasker_server-2.2.0/src/labtasker_server/ownership.py +29 -0
- {labtasker_server-2.1.0 → labtasker_server-2.2.0}/src/labtasker_server/pagination.py +19 -10
- {labtasker_server-2.1.0 → labtasker_server-2.2.0}/src/labtasker_server/schemas.py +44 -0
- {labtasker_server-2.1.0 → labtasker_server-2.2.0}/src/labtasker_server/services/tasks.py +51 -3
- labtasker_server-2.2.0/src/labtasker_server/services/workers.py +166 -0
- {labtasker_server-2.1.0 → labtasker_server-2.2.0}/LICENSE +0 -0
- {labtasker_server-2.1.0 → labtasker_server-2.2.0}/src/labtasker_server/__main__.py +0 -0
- {labtasker_server-2.1.0 → labtasker_server-2.2.0}/src/labtasker_server/config.py +0 -0
- {labtasker_server-2.1.0 → labtasker_server-2.2.0}/src/labtasker_server/errors.py +0 -0
- {labtasker_server-2.1.0 → labtasker_server-2.2.0}/src/labtasker_server/logging.py +0 -0
- {labtasker_server-2.1.0 → labtasker_server-2.2.0}/src/labtasker_server/migrations/__init__.py +0 -0
- {labtasker_server-2.1.0 → labtasker_server-2.2.0}/src/labtasker_server/migrations/env.py +0 -0
- {labtasker_server-2.1.0 → labtasker_server-2.2.0}/src/labtasker_server/migrations/versions/0001_initial.py +0 -0
- {labtasker_server-2.1.0 → labtasker_server-2.2.0}/src/labtasker_server/migrations/versions/__init__.py +0 -0
- {labtasker_server-2.1.0 → labtasker_server-2.2.0}/src/labtasker_server/py.typed +0 -0
- {labtasker_server-2.1.0 → labtasker_server-2.2.0}/src/labtasker_server/services/__init__.py +0 -0
- {labtasker_server-2.1.0 → labtasker_server-2.2.0}/src/labtasker_server/services/queues.py +0 -0
- {labtasker_server-2.1.0 → labtasker_server-2.2.0}/src/labtasker_server/validation.py +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.5
|
|
2
2
|
Name: labtasker-server
|
|
3
|
-
Version: 2.
|
|
3
|
+
Version: 2.2.0
|
|
4
4
|
Summary: SQLite/FastAPI server for parallel model inference and evaluation
|
|
5
5
|
Project-URL: Homepage, https://github.com/luocfprime/labtasker
|
|
6
6
|
Project-URL: Repository, https://github.com/luocfprime/labtasker.git
|
|
@@ -10,7 +10,7 @@ License-File: LICENSE
|
|
|
10
10
|
Requires-Python: >=3.11
|
|
11
11
|
Requires-Dist: alembic<2,>=1.14
|
|
12
12
|
Requires-Dist: fastapi<1,>=0.115
|
|
13
|
-
Requires-Dist: pydantic<3,>=2.
|
|
13
|
+
Requires-Dist: pydantic<3,>=2.13.5
|
|
14
14
|
Requires-Dist: sqlalchemy<3,>=2.0
|
|
15
15
|
Requires-Dist: typer<1,>=0.16
|
|
16
16
|
Requires-Dist: uvicorn<1,>=0.34
|
|
@@ -1,10 +1,10 @@
|
|
|
1
1
|
[build-system]
|
|
2
|
-
requires = ["hatchling>=1.
|
|
2
|
+
requires = ["hatchling>=1.32.0"]
|
|
3
3
|
build-backend = "hatchling.build"
|
|
4
4
|
|
|
5
5
|
[project]
|
|
6
6
|
name = "labtasker-server"
|
|
7
|
-
version = "2.
|
|
7
|
+
version = "2.2.0"
|
|
8
8
|
description = "SQLite/FastAPI server for parallel model inference and evaluation"
|
|
9
9
|
requires-python = ">=3.11"
|
|
10
10
|
license = "Apache-2.0"
|
|
@@ -13,7 +13,7 @@ authors = [{ name = "lcf", email = "luocfprime@gmail.com" }]
|
|
|
13
13
|
dependencies = [
|
|
14
14
|
"alembic>=1.14,<2",
|
|
15
15
|
"fastapi>=0.115,<1",
|
|
16
|
-
"pydantic>=2.
|
|
16
|
+
"pydantic>=2.13.5,<3",
|
|
17
17
|
"sqlalchemy>=2.0,<3",
|
|
18
18
|
"typer>=0.16,<1",
|
|
19
19
|
"uvicorn>=0.34,<1",
|
|
@@ -1,22 +1,30 @@
|
|
|
1
1
|
from __future__ import annotations
|
|
2
2
|
|
|
3
3
|
import asyncio
|
|
4
|
-
import hmac
|
|
5
4
|
import logging
|
|
6
5
|
from collections.abc import AsyncIterator, Callable
|
|
7
6
|
from contextlib import asynccontextmanager, suppress
|
|
8
7
|
from typing import Annotated, Any
|
|
9
8
|
|
|
10
9
|
from fastapi import Depends, FastAPI, Query, Request, Response
|
|
10
|
+
from fastapi.exception_handlers import http_exception_handler
|
|
11
11
|
from fastapi.exceptions import RequestValidationError
|
|
12
12
|
from fastapi.responses import JSONResponse
|
|
13
|
+
from fastapi.routing import APIRoute
|
|
13
14
|
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
|
|
14
15
|
from sqlalchemy import text
|
|
16
|
+
from starlette.exceptions import HTTPException
|
|
15
17
|
|
|
18
|
+
from labtasker_server import __version__
|
|
16
19
|
from labtasker_server.config import ServerSettings
|
|
17
20
|
from labtasker_server.database import Database
|
|
18
21
|
from labtasker_server.errors import DomainError
|
|
19
|
-
from labtasker_server.
|
|
22
|
+
from labtasker_server.grouping import request_error
|
|
23
|
+
from labtasker_server.middleware import (
|
|
24
|
+
RequestBodyLimitMiddleware,
|
|
25
|
+
ServerVersionMiddleware,
|
|
26
|
+
is_authenticated,
|
|
27
|
+
)
|
|
20
28
|
from labtasker_server.schemas import (
|
|
21
29
|
BulkUpdateRequest,
|
|
22
30
|
BulkUpdateResult,
|
|
@@ -26,6 +34,7 @@ from labtasker_server.schemas import (
|
|
|
26
34
|
CountResponse,
|
|
27
35
|
ErrorEnvelope,
|
|
28
36
|
FailRequest,
|
|
37
|
+
GroupCountPage,
|
|
29
38
|
HealthyResponse,
|
|
30
39
|
HeartbeatResponse,
|
|
31
40
|
Queue,
|
|
@@ -37,10 +46,13 @@ from labtasker_server.schemas import (
|
|
|
37
46
|
TaskStatus,
|
|
38
47
|
TaskUpdate,
|
|
39
48
|
UnhealthyResponse,
|
|
49
|
+
WorkerPage,
|
|
50
|
+
WorkerReport,
|
|
40
51
|
)
|
|
41
52
|
from labtasker_server.services.queues import QueueService
|
|
42
53
|
from labtasker_server.services.tasks import TaskService, system_now_us
|
|
43
|
-
from labtasker_server.
|
|
54
|
+
from labtasker_server.services.workers import WorkerService
|
|
55
|
+
from labtasker_server.validation import MAX_JSON_DEPTH, MAX_TASK_DATA_BYTES
|
|
44
56
|
|
|
45
57
|
EXPIRY_SCAN_INTERVAL_SECONDS = 60
|
|
46
58
|
logger = logging.getLogger(__name__)
|
|
@@ -58,16 +70,18 @@ def create_app(
|
|
|
58
70
|
database = Database(settings.database, ownership_fd=settings.database_fd)
|
|
59
71
|
try:
|
|
60
72
|
database.initialize()
|
|
73
|
+
queue_service = QueueService(database)
|
|
74
|
+
task_service = TaskService(database, now_us=now_us)
|
|
75
|
+
task_service.expire_leases()
|
|
76
|
+
worker_service = WorkerService(database, now_us=now_us)
|
|
77
|
+
worker_service.expire()
|
|
61
78
|
except BaseException:
|
|
62
79
|
database.dispose()
|
|
63
80
|
raise
|
|
64
|
-
queue_service = QueueService(database)
|
|
65
|
-
task_service = TaskService(database, now_us=now_us)
|
|
66
|
-
task_service.expire_leases()
|
|
67
81
|
|
|
68
82
|
@asynccontextmanager
|
|
69
83
|
async def lifespan(_: FastAPI) -> AsyncIterator[None]:
|
|
70
|
-
scanner = asyncio.create_task(_expiry_scanner(task_service))
|
|
84
|
+
scanner = asyncio.create_task(_expiry_scanner(task_service, worker_service))
|
|
71
85
|
try:
|
|
72
86
|
yield
|
|
73
87
|
finally:
|
|
@@ -78,9 +92,11 @@ def create_app(
|
|
|
78
92
|
|
|
79
93
|
app = FastAPI(docs_url=None, redoc_url=None, lifespan=lifespan)
|
|
80
94
|
app.add_middleware(RequestBodyLimitMiddleware, max_bytes=MAX_TASK_DATA_BYTES)
|
|
95
|
+
app.add_middleware(ServerVersionMiddleware, version=__version__, token=settings.token)
|
|
81
96
|
app.state.database = database
|
|
82
97
|
app.state.settings = settings
|
|
83
98
|
app.state.task_service = task_service
|
|
99
|
+
app.state.worker_service = worker_service
|
|
84
100
|
|
|
85
101
|
@app.exception_handler(DomainError)
|
|
86
102
|
async def handle_domain_error(_: Request, exc: DomainError) -> JSONResponse:
|
|
@@ -134,15 +150,37 @@ def create_app(
|
|
|
134
150
|
},
|
|
135
151
|
)
|
|
136
152
|
|
|
153
|
+
@app.exception_handler(HTTPException)
|
|
154
|
+
async def handle_http_error(request: Request, exc: HTTPException) -> Response:
|
|
155
|
+
# FastAPI wraps decoder limits and invalid byte encodings in HTTP 400
|
|
156
|
+
# before Pydantic sees the body. Keep those in the validation contract.
|
|
157
|
+
if (
|
|
158
|
+
exc.status_code == 400
|
|
159
|
+
and exc.detail == "There was an error parsing the body"
|
|
160
|
+
and isinstance(exc.__cause__, (ValueError, RecursionError))
|
|
161
|
+
):
|
|
162
|
+
if isinstance(exc.__cause__, RecursionError):
|
|
163
|
+
error = DomainError(
|
|
164
|
+
422,
|
|
165
|
+
"json_too_deep",
|
|
166
|
+
"JSON value is too deeply nested.",
|
|
167
|
+
{"max_depth": MAX_JSON_DEPTH},
|
|
168
|
+
)
|
|
169
|
+
else:
|
|
170
|
+
error = DomainError(
|
|
171
|
+
422,
|
|
172
|
+
"invalid_request",
|
|
173
|
+
"Request validation failed.",
|
|
174
|
+
{"errors": [{"location": ["body"], "message": "Malformed JSON body."}]},
|
|
175
|
+
)
|
|
176
|
+
return await handle_domain_error(request, error)
|
|
177
|
+
return await http_exception_handler(request, exc)
|
|
178
|
+
|
|
137
179
|
def require_auth(
|
|
180
|
+
request: Request,
|
|
138
181
|
credentials: Annotated[HTTPAuthorizationCredentials | None, Depends(BEARER)],
|
|
139
182
|
) -> None:
|
|
140
|
-
|
|
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.encode(), token.encode("ascii")):
|
|
183
|
+
if not is_authenticated(request.headers.get("authorization"), settings.token):
|
|
146
184
|
raise _unauthorized()
|
|
147
185
|
|
|
148
186
|
authenticated = [Depends(require_auth)]
|
|
@@ -217,6 +255,7 @@ def create_app(
|
|
|
217
255
|
queue: str,
|
|
218
256
|
status: TaskStatus | None = None,
|
|
219
257
|
name: str | None = None,
|
|
258
|
+
name_fuzzy: str | None = None,
|
|
220
259
|
filter_expression: Annotated[str | None, Query(alias="filter")] = None,
|
|
221
260
|
order_by: TaskOrderField = "created_at",
|
|
222
261
|
descending: bool = True,
|
|
@@ -227,6 +266,7 @@ def create_app(
|
|
|
227
266
|
queue,
|
|
228
267
|
status=status,
|
|
229
268
|
name=name,
|
|
269
|
+
name_fuzzy=name_fuzzy,
|
|
230
270
|
filter_expression=filter_expression,
|
|
231
271
|
order_by=order_by,
|
|
232
272
|
descending=descending,
|
|
@@ -236,24 +276,93 @@ def create_app(
|
|
|
236
276
|
|
|
237
277
|
@app.get(
|
|
238
278
|
"/api/v2/queues/{queue}/tasks/count",
|
|
239
|
-
response_model=CountResponse,
|
|
279
|
+
response_model=CountResponse | GroupCountPage,
|
|
240
280
|
dependencies=authenticated,
|
|
241
281
|
responses=API_ERROR_RESPONSES,
|
|
242
282
|
)
|
|
243
283
|
def count_tasks(
|
|
284
|
+
request: Request,
|
|
244
285
|
queue: str,
|
|
245
286
|
status: TaskStatus | None = None,
|
|
246
287
|
name: str | None = None,
|
|
288
|
+
name_fuzzy: str | None = None,
|
|
247
289
|
filter_expression: Annotated[str | None, Query(alias="filter")] = None,
|
|
248
|
-
|
|
249
|
-
|
|
250
|
-
|
|
251
|
-
|
|
252
|
-
|
|
253
|
-
|
|
254
|
-
|
|
255
|
-
|
|
290
|
+
group_by: str | None = None,
|
|
291
|
+
limit: Annotated[int | None, Query(ge=1, le=1000)] = None,
|
|
292
|
+
cursor: str | None = None,
|
|
293
|
+
) -> CountResponse | GroupCountPage:
|
|
294
|
+
_single_grouping(request)
|
|
295
|
+
result = task_service.count_tasks(
|
|
296
|
+
queue,
|
|
297
|
+
status=status,
|
|
298
|
+
name=name,
|
|
299
|
+
name_fuzzy=name_fuzzy,
|
|
300
|
+
filter_expression=filter_expression,
|
|
301
|
+
group_by=group_by,
|
|
302
|
+
limit=limit,
|
|
303
|
+
cursor=cursor,
|
|
304
|
+
)
|
|
305
|
+
return CountResponse(count=result) if isinstance(result, int) else result
|
|
306
|
+
|
|
307
|
+
@app.get(
|
|
308
|
+
"/api/v2/queues/{queue}/workers",
|
|
309
|
+
response_model=WorkerPage,
|
|
310
|
+
dependencies=authenticated,
|
|
311
|
+
responses=API_ERROR_RESPONSES,
|
|
312
|
+
)
|
|
313
|
+
def list_workers(
|
|
314
|
+
queue: str,
|
|
315
|
+
filter_expression: Annotated[str | None, Query(alias="filter")] = None,
|
|
316
|
+
limit: Annotated[int, Query(ge=1, le=1000)] = 100,
|
|
317
|
+
cursor: str | None = None,
|
|
318
|
+
) -> WorkerPage:
|
|
319
|
+
return worker_service.list(
|
|
320
|
+
queue, filter_expression=filter_expression, limit=limit, cursor=cursor
|
|
321
|
+
)
|
|
322
|
+
|
|
323
|
+
@app.get(
|
|
324
|
+
"/api/v2/queues/{queue}/workers/count",
|
|
325
|
+
response_model=CountResponse | GroupCountPage,
|
|
326
|
+
dependencies=authenticated,
|
|
327
|
+
responses=API_ERROR_RESPONSES,
|
|
328
|
+
)
|
|
329
|
+
def count_workers(
|
|
330
|
+
request: Request,
|
|
331
|
+
queue: str,
|
|
332
|
+
filter_expression: Annotated[str | None, Query(alias="filter")] = None,
|
|
333
|
+
group_by: str | None = None,
|
|
334
|
+
limit: Annotated[int | None, Query(ge=1, le=1000)] = None,
|
|
335
|
+
cursor: str | None = None,
|
|
336
|
+
) -> CountResponse | GroupCountPage:
|
|
337
|
+
_single_grouping(request)
|
|
338
|
+
result = worker_service.count(
|
|
339
|
+
queue,
|
|
340
|
+
filter_expression=filter_expression,
|
|
341
|
+
group_by=group_by,
|
|
342
|
+
limit=limit,
|
|
343
|
+
cursor=cursor,
|
|
256
344
|
)
|
|
345
|
+
return CountResponse(count=result) if isinstance(result, int) else result
|
|
346
|
+
|
|
347
|
+
@app.put(
|
|
348
|
+
"/api/v2/queues/{queue}/workers/{id}",
|
|
349
|
+
status_code=204,
|
|
350
|
+
dependencies=authenticated,
|
|
351
|
+
responses=API_ERROR_RESPONSES,
|
|
352
|
+
)
|
|
353
|
+
def report_worker(queue: str, id: str, report: WorkerReport) -> Response:
|
|
354
|
+
worker_service.report(queue, id, report)
|
|
355
|
+
return Response(status_code=204)
|
|
356
|
+
|
|
357
|
+
@app.delete(
|
|
358
|
+
"/api/v2/queues/{queue}/workers/{id}",
|
|
359
|
+
status_code=204,
|
|
360
|
+
dependencies=authenticated,
|
|
361
|
+
responses=API_ERROR_RESPONSES,
|
|
362
|
+
)
|
|
363
|
+
def withdraw_worker(queue: str, id: str) -> Response:
|
|
364
|
+
worker_service.withdraw(queue, id)
|
|
365
|
+
return Response(status_code=204)
|
|
257
366
|
|
|
258
367
|
@app.get(
|
|
259
368
|
"/api/v2/queues/{queue}/tasks/{task_id}",
|
|
@@ -363,18 +472,42 @@ def create_app(
|
|
|
363
472
|
task_service.delete(queue, task_id)
|
|
364
473
|
return Response(status_code=204)
|
|
365
474
|
|
|
475
|
+
for route in app.routes:
|
|
476
|
+
if isinstance(route, APIRoute) and route.path.startswith("/api/"):
|
|
477
|
+
for status in {route.status_code or 200, *route.responses}:
|
|
478
|
+
if str(status) == "401":
|
|
479
|
+
continue
|
|
480
|
+
route.responses.setdefault(status, {}).setdefault("headers", {})[
|
|
481
|
+
"Labtasker-Server-Version"
|
|
482
|
+
] = {
|
|
483
|
+
"description": "Server package version; present only with valid credentials "
|
|
484
|
+
"when authentication is enabled.",
|
|
485
|
+
"schema": {"type": "string"},
|
|
486
|
+
}
|
|
366
487
|
return app
|
|
367
488
|
|
|
368
489
|
|
|
369
|
-
async def _expiry_scanner(task_service: TaskService) -> None:
|
|
490
|
+
async def _expiry_scanner(task_service: TaskService, worker_service: WorkerService) -> None:
|
|
370
491
|
while True:
|
|
371
492
|
await asyncio.sleep(EXPIRY_SCAN_INTERVAL_SECONDS)
|
|
493
|
+
scan = asyncio.create_task(asyncio.to_thread(_expire_records, task_service, worker_service))
|
|
372
494
|
try:
|
|
373
|
-
await asyncio.
|
|
495
|
+
await asyncio.shield(scan)
|
|
496
|
+
except asyncio.CancelledError:
|
|
497
|
+
# Cancelling to_thread cannot stop its database command. Keep the
|
|
498
|
+
# ownership descriptor until that command has actually finished.
|
|
499
|
+
with suppress(Exception):
|
|
500
|
+
await scan
|
|
501
|
+
raise
|
|
374
502
|
except Exception:
|
|
375
503
|
logger.exception("Heartbeat expiry scan failed; it will retry in 60 seconds.")
|
|
376
504
|
|
|
377
505
|
|
|
506
|
+
def _expire_records(task_service: TaskService, worker_service: WorkerService) -> None:
|
|
507
|
+
task_service.expire_leases()
|
|
508
|
+
worker_service.expire()
|
|
509
|
+
|
|
510
|
+
|
|
378
511
|
def _unauthorized() -> DomainError:
|
|
379
512
|
return DomainError(401, "unauthorized", "Authentication is required.", {})
|
|
380
513
|
|
|
@@ -393,6 +526,10 @@ def _validation_error(
|
|
|
393
526
|
) -> tuple[str, dict[str, object] | None]:
|
|
394
527
|
for error in exc.errors():
|
|
395
528
|
error_type = str(error.get("type", ""))
|
|
529
|
+
if error_type == "recursion_loop":
|
|
530
|
+
# A JSON body cannot contain reference cycles. This is Pydantic's
|
|
531
|
+
# own nesting limit, reached before the domain depth validator.
|
|
532
|
+
return "json_too_deep", {"max_depth": MAX_JSON_DEPTH}
|
|
396
533
|
if error_type == "json_invalid":
|
|
397
534
|
return "invalid_request", {
|
|
398
535
|
"errors": [{"location": ["body"], "message": "Malformed JSON body."}]
|
|
@@ -414,3 +551,8 @@ def _specific_validation_message(code: str) -> str:
|
|
|
414
551
|
if code == "invalid_task_name":
|
|
415
552
|
return "Task name is invalid."
|
|
416
553
|
raise AssertionError(f"Unknown specific validation code: {code}")
|
|
554
|
+
|
|
555
|
+
|
|
556
|
+
def _single_grouping(request: Request) -> None:
|
|
557
|
+
if len(request.query_params.getlist("group_by")) > 1:
|
|
558
|
+
request_error("group_by", "Specify group_by only once.")
|
|
@@ -27,9 +27,7 @@ from labtasker_server.local import (
|
|
|
27
27
|
metadata_matches_database,
|
|
28
28
|
metadata_owner_is_verified,
|
|
29
29
|
read_metadata,
|
|
30
|
-
|
|
31
|
-
remove_generation_socket,
|
|
32
|
-
remove_stale_artifacts,
|
|
30
|
+
remove_stopped_artifacts,
|
|
33
31
|
require_local_capabilities,
|
|
34
32
|
socket_health,
|
|
35
33
|
startup_age,
|
|
@@ -209,7 +207,7 @@ def stop(
|
|
|
209
207
|
if database_is_free(paths):
|
|
210
208
|
if has_runtime_artifacts(paths):
|
|
211
209
|
try:
|
|
212
|
-
|
|
210
|
+
remove_stopped_artifacts(paths)
|
|
213
211
|
except RuntimeError as error:
|
|
214
212
|
typer.echo(f"[labtasker-server] Local Server stop error: {error}", err=True)
|
|
215
213
|
raise typer.Exit(1) from error
|
|
@@ -231,7 +229,7 @@ def stop(
|
|
|
231
229
|
os.kill(metadata.pid, signal.SIGTERM)
|
|
232
230
|
typer.echo(f"[labtasker-server] stopping local daemon pid={metadata.pid}", err=True)
|
|
233
231
|
if _wait_for_exit(paths.directory, timeout=30.0):
|
|
234
|
-
|
|
232
|
+
remove_stopped_artifacts(paths, generation=metadata.generation)
|
|
235
233
|
typer.echo(f"[labtasker-server] stopped local daemon pid={metadata.pid}", err=True)
|
|
236
234
|
return
|
|
237
235
|
if not force:
|
|
@@ -262,7 +260,7 @@ def stop(
|
|
|
262
260
|
err=True,
|
|
263
261
|
)
|
|
264
262
|
raise typer.Exit(1)
|
|
265
|
-
|
|
263
|
+
remove_stopped_artifacts(paths, generation=current.generation)
|
|
266
264
|
typer.echo(f"[labtasker-server] stopped local daemon pid={current.pid}", err=True)
|
|
267
265
|
|
|
268
266
|
|
|
@@ -331,7 +329,7 @@ def daemon(
|
|
|
331
329
|
listener.close()
|
|
332
330
|
# Preserve the attempt metadata so an unexpected exit remains throttled.
|
|
333
331
|
# An explicit successful stop removes the full generation itself.
|
|
334
|
-
|
|
332
|
+
remove_stopped_artifacts(paths, generation=generation, preserve_metadata=True)
|
|
335
333
|
|
|
336
334
|
|
|
337
335
|
def _local_status(directory: Path) -> dict[str, object]:
|
|
@@ -5,19 +5,16 @@ from collections.abc import Iterator
|
|
|
5
5
|
from contextlib import contextmanager
|
|
6
6
|
from pathlib import Path
|
|
7
7
|
|
|
8
|
-
try:
|
|
9
|
-
import fcntl
|
|
10
|
-
except ImportError: # pragma: no cover - explicit HTTP remains best effort off POSIX
|
|
11
|
-
fcntl = None # type: ignore[assignment]
|
|
12
|
-
|
|
13
8
|
from alembic import command
|
|
14
9
|
from alembic.config import Config
|
|
15
|
-
from sqlalchemy import Engine, create_engine, event, inspect, text
|
|
10
|
+
from sqlalchemy import URL, Connection, Engine, create_engine, event, insert, inspect, text
|
|
16
11
|
from sqlalchemy.exc import OperationalError
|
|
17
12
|
from sqlalchemy.orm import Session, sessionmaker
|
|
18
13
|
|
|
19
14
|
from labtasker_server.errors import DomainError
|
|
20
15
|
from labtasker_server.models import QueueRow
|
|
16
|
+
from labtasker_server.name_search import name_matches_fuzzy
|
|
17
|
+
from labtasker_server.ownership import lock_database
|
|
21
18
|
|
|
22
19
|
LOCAL_GITIGNORE = "*\n!.gitignore\n"
|
|
23
20
|
|
|
@@ -37,32 +34,43 @@ class Database:
|
|
|
37
34
|
if labtasker_dir is not None:
|
|
38
35
|
_ensure_local_gitignore(labtasker_dir)
|
|
39
36
|
self._ownership_fd: int | None = _acquire_database_ownership(self.path, ownership_fd)
|
|
40
|
-
|
|
41
|
-
|
|
37
|
+
try:
|
|
38
|
+
self.engine = _create_sqlite_engine(self.path)
|
|
39
|
+
self._session_factory = sessionmaker(self.engine, expire_on_commit=False)
|
|
40
|
+
except BaseException:
|
|
41
|
+
os.close(self._ownership_fd)
|
|
42
|
+
self._ownership_fd = None
|
|
43
|
+
raise
|
|
42
44
|
|
|
43
45
|
def initialize(self) -> None:
|
|
44
|
-
existing_tables = set(inspect(self.engine).get_table_names())
|
|
45
|
-
is_fresh = not existing_tables
|
|
46
|
-
if existing_tables and "alembic_version" not in existing_tables:
|
|
47
|
-
raise RuntimeError("Database has tables but is not a recognized Labtasker v2 schema.")
|
|
48
|
-
|
|
49
46
|
alembic_config = Config()
|
|
50
47
|
alembic_config.set_main_option(
|
|
51
48
|
"script_location",
|
|
52
|
-
|
|
49
|
+
# Alembic's ConfigParser interpolates percent signs in option values.
|
|
50
|
+
str(Path(__file__).resolve().parent / "migrations").replace("%", "%%"),
|
|
53
51
|
)
|
|
54
52
|
with self.engine.begin() as connection:
|
|
53
|
+
# sqlite3's legacy mode does not begin a transaction for DDL.
|
|
54
|
+
# Commit schema changes and initial Queue together, including when
|
|
55
|
+
# startup fails or is interrupted between migration operations.
|
|
56
|
+
connection.exec_driver_sql("BEGIN IMMEDIATE")
|
|
57
|
+
existing_tables = set(inspect(connection).get_table_names())
|
|
58
|
+
if existing_tables and "alembic_version" not in existing_tables:
|
|
59
|
+
raise RuntimeError(
|
|
60
|
+
"Database has tables but is not a recognized Labtasker v2 schema."
|
|
61
|
+
)
|
|
55
62
|
alembic_config.attributes["connection"] = connection
|
|
56
63
|
command.upgrade(alembic_config, "head")
|
|
57
|
-
|
|
58
|
-
|
|
59
|
-
|
|
60
|
-
with self.write_session() as session:
|
|
61
|
-
session.add(QueueRow(name="default"))
|
|
64
|
+
_verify_sqlite_settings(connection)
|
|
65
|
+
if not existing_tables:
|
|
66
|
+
connection.execute(insert(QueueRow).values(name="default"))
|
|
62
67
|
|
|
63
68
|
@contextmanager
|
|
64
69
|
def read_session(self) -> Iterator[Session]:
|
|
65
70
|
with self._session_factory() as session:
|
|
71
|
+
# sqlite3's legacy transaction mode does not begin on SELECT. Keep
|
|
72
|
+
# all reads (including relationship loaders) in one SQLite snapshot.
|
|
73
|
+
session.execute(text("BEGIN"))
|
|
66
74
|
yield session
|
|
67
75
|
|
|
68
76
|
@contextmanager
|
|
@@ -87,23 +95,27 @@ class Database:
|
|
|
87
95
|
raise
|
|
88
96
|
|
|
89
97
|
def dispose(self) -> None:
|
|
90
|
-
|
|
91
|
-
|
|
92
|
-
|
|
93
|
-
self._ownership_fd
|
|
98
|
+
try:
|
|
99
|
+
self.engine.dispose()
|
|
100
|
+
finally:
|
|
101
|
+
if self._ownership_fd is not None:
|
|
102
|
+
os.close(self._ownership_fd)
|
|
103
|
+
self._ownership_fd = None
|
|
94
104
|
|
|
95
105
|
|
|
96
106
|
def _acquire_database_ownership(path: Path, inherited_fd: int | None) -> int:
|
|
97
107
|
if inherited_fd is None:
|
|
98
108
|
fd = os.open(path, os.O_RDWR | os.O_CREAT, 0o600)
|
|
99
|
-
|
|
100
|
-
|
|
101
|
-
|
|
102
|
-
|
|
103
|
-
|
|
104
|
-
|
|
105
|
-
|
|
106
|
-
|
|
109
|
+
try:
|
|
110
|
+
lock_database(fd)
|
|
111
|
+
except BlockingIOError as error:
|
|
112
|
+
os.close(fd)
|
|
113
|
+
raise DatabaseOwnershipError(
|
|
114
|
+
f"Another Server process already owns database {path}."
|
|
115
|
+
) from error
|
|
116
|
+
except BaseException:
|
|
117
|
+
os.close(fd)
|
|
118
|
+
raise
|
|
107
119
|
else:
|
|
108
120
|
fd = os.dup(inherited_fd)
|
|
109
121
|
os.set_inheritable(fd, False)
|
|
@@ -127,12 +139,15 @@ def _acquire_database_ownership(path: Path, inherited_fd: int | None) -> int:
|
|
|
127
139
|
|
|
128
140
|
def _create_sqlite_engine(path: Path) -> Engine:
|
|
129
141
|
engine = create_engine(
|
|
130
|
-
|
|
142
|
+
URL.create("sqlite+pysqlite", database=str(path)),
|
|
131
143
|
connect_args={"check_same_thread": False, "timeout": 5.0},
|
|
132
144
|
)
|
|
133
145
|
|
|
134
146
|
@event.listens_for(engine, "connect")
|
|
135
147
|
def configure_connection(dbapi_connection: object, _: object) -> None:
|
|
148
|
+
dbapi_connection.create_function( # type: ignore[attr-defined]
|
|
149
|
+
"labtasker_name_fuzzy", 2, name_matches_fuzzy, deterministic=True
|
|
150
|
+
)
|
|
136
151
|
cursor = dbapi_connection.cursor() # type: ignore[attr-defined]
|
|
137
152
|
try:
|
|
138
153
|
cursor.execute("PRAGMA foreign_keys=ON")
|
|
@@ -158,14 +173,13 @@ def _is_sqlite_busy(error: OperationalError) -> bool:
|
|
|
158
173
|
return code in {5, 6} or "database is locked" in str(error.orig).lower()
|
|
159
174
|
|
|
160
175
|
|
|
161
|
-
def _verify_sqlite_settings(
|
|
162
|
-
|
|
163
|
-
|
|
164
|
-
|
|
165
|
-
|
|
166
|
-
|
|
167
|
-
|
|
168
|
-
}
|
|
176
|
+
def _verify_sqlite_settings(connection: Connection) -> None:
|
|
177
|
+
actual = {
|
|
178
|
+
"journal_mode": connection.scalar(text("PRAGMA journal_mode")),
|
|
179
|
+
"foreign_keys": connection.scalar(text("PRAGMA foreign_keys")),
|
|
180
|
+
"busy_timeout": connection.scalar(text("PRAGMA busy_timeout")),
|
|
181
|
+
"synchronous": connection.scalar(text("PRAGMA synchronous")),
|
|
182
|
+
}
|
|
169
183
|
expected = {
|
|
170
184
|
"journal_mode": "wal",
|
|
171
185
|
"foreign_keys": 1,
|