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.
Files changed (32) hide show
  1. {labtasker_server-2.1.0 → labtasker_server-2.2.0}/.gitignore +4 -0
  2. {labtasker_server-2.1.0 → labtasker_server-2.2.0}/PKG-INFO +2 -2
  3. {labtasker_server-2.1.0 → labtasker_server-2.2.0}/pyproject.toml +3 -3
  4. {labtasker_server-2.1.0 → labtasker_server-2.2.0}/src/labtasker_server/__init__.py +1 -1
  5. {labtasker_server-2.1.0 → labtasker_server-2.2.0}/src/labtasker_server/app.py +166 -24
  6. {labtasker_server-2.1.0 → labtasker_server-2.2.0}/src/labtasker_server/cli.py +5 -7
  7. {labtasker_server-2.1.0 → labtasker_server-2.2.0}/src/labtasker_server/database.py +54 -40
  8. {labtasker_server-2.1.0 → labtasker_server-2.2.0}/src/labtasker_server/filtering.py +118 -30
  9. labtasker_server-2.2.0/src/labtasker_server/grouping.py +136 -0
  10. {labtasker_server-2.1.0 → labtasker_server-2.2.0}/src/labtasker_server/local.py +38 -26
  11. {labtasker_server-2.1.0 → labtasker_server-2.2.0}/src/labtasker_server/middleware.py +38 -0
  12. labtasker_server-2.2.0/src/labtasker_server/migrations/versions/0002_worker_observations.py +33 -0
  13. {labtasker_server-2.1.0 → labtasker_server-2.2.0}/src/labtasker_server/models.py +20 -0
  14. labtasker_server-2.2.0/src/labtasker_server/name_search.py +23 -0
  15. labtasker_server-2.2.0/src/labtasker_server/ownership.py +29 -0
  16. {labtasker_server-2.1.0 → labtasker_server-2.2.0}/src/labtasker_server/pagination.py +19 -10
  17. {labtasker_server-2.1.0 → labtasker_server-2.2.0}/src/labtasker_server/schemas.py +44 -0
  18. {labtasker_server-2.1.0 → labtasker_server-2.2.0}/src/labtasker_server/services/tasks.py +51 -3
  19. labtasker_server-2.2.0/src/labtasker_server/services/workers.py +166 -0
  20. {labtasker_server-2.1.0 → labtasker_server-2.2.0}/LICENSE +0 -0
  21. {labtasker_server-2.1.0 → labtasker_server-2.2.0}/src/labtasker_server/__main__.py +0 -0
  22. {labtasker_server-2.1.0 → labtasker_server-2.2.0}/src/labtasker_server/config.py +0 -0
  23. {labtasker_server-2.1.0 → labtasker_server-2.2.0}/src/labtasker_server/errors.py +0 -0
  24. {labtasker_server-2.1.0 → labtasker_server-2.2.0}/src/labtasker_server/logging.py +0 -0
  25. {labtasker_server-2.1.0 → labtasker_server-2.2.0}/src/labtasker_server/migrations/__init__.py +0 -0
  26. {labtasker_server-2.1.0 → labtasker_server-2.2.0}/src/labtasker_server/migrations/env.py +0 -0
  27. {labtasker_server-2.1.0 → labtasker_server-2.2.0}/src/labtasker_server/migrations/versions/0001_initial.py +0 -0
  28. {labtasker_server-2.1.0 → labtasker_server-2.2.0}/src/labtasker_server/migrations/versions/__init__.py +0 -0
  29. {labtasker_server-2.1.0 → labtasker_server-2.2.0}/src/labtasker_server/py.typed +0 -0
  30. {labtasker_server-2.1.0 → labtasker_server-2.2.0}/src/labtasker_server/services/__init__.py +0 -0
  31. {labtasker_server-2.1.0 → labtasker_server-2.2.0}/src/labtasker_server/services/queues.py +0 -0
  32. {labtasker_server-2.1.0 → labtasker_server-2.2.0}/src/labtasker_server/validation.py +0 -0
@@ -1,6 +1,7 @@
1
1
  .venv/
2
2
  .pytest_cache/
3
3
  .coverage
4
+ coverage/
4
5
  coverage.xml
5
6
  htmlcov/
6
7
  .mypy_cache/
@@ -12,3 +13,6 @@ build/
12
13
  dist/
13
14
  site/
14
15
  .labtasker/
16
+ .DS_Store
17
+ .idea/
18
+ tests/skill/runs/
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.5
2
2
  Name: labtasker-server
3
- Version: 2.1.0
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.10
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.27"]
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.1.0"
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.10,<3",
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,3 +1,3 @@
1
1
  """Labtasker v2 HTTP server package."""
2
2
 
3
- __version__ = "2.1.0"
3
+ __version__ = "2.2.0"
@@ -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.middleware import RequestBodyLimitMiddleware
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.validation import MAX_TASK_DATA_BYTES
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
- 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.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
- ) -> CountResponse:
249
- return CountResponse(
250
- count=task_service.count_tasks(
251
- queue,
252
- status=status,
253
- name=name,
254
- filter_expression=filter_expression,
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.to_thread(task_service.expire_leases)
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
- remove_generation_artifacts,
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
- remove_stale_artifacts(paths)
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
- remove_generation_artifacts(paths, metadata.generation)
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
- remove_generation_artifacts(paths, current.generation)
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
- remove_generation_socket(paths, generation)
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
- self.engine = _create_sqlite_engine(self.path)
41
- self._session_factory = sessionmaker(self.engine, expire_on_commit=False)
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
- str(Path(__file__).resolve().parent / "migrations"),
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
- _verify_sqlite_settings(self.engine)
58
-
59
- if is_fresh:
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
- self.engine.dispose()
91
- if self._ownership_fd is not None:
92
- os.close(self._ownership_fd)
93
- self._ownership_fd = None
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
- if fcntl is not None:
100
- try:
101
- fcntl.flock(fd, fcntl.LOCK_EX | fcntl.LOCK_NB)
102
- except BlockingIOError as error:
103
- os.close(fd)
104
- raise DatabaseOwnershipError(
105
- f"Another Server process already owns database {path}."
106
- ) from error
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
- f"sqlite+pysqlite:///{path}",
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(engine: Engine) -> None:
162
- with engine.connect() as connection:
163
- actual = {
164
- "journal_mode": connection.scalar(text("PRAGMA journal_mode")),
165
- "foreign_keys": connection.scalar(text("PRAGMA foreign_keys")),
166
- "busy_timeout": connection.scalar(text("PRAGMA busy_timeout")),
167
- "synchronous": connection.scalar(text("PRAGMA synchronous")),
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,