feat: add durable leased task engine
This commit is contained in:
@@ -0,0 +1,32 @@
|
|||||||
|
ALTER TABLE tasks
|
||||||
|
ADD COLUMN task_type text NOT NULL DEFAULT 'generic',
|
||||||
|
ADD COLUMN progress integer NOT NULL DEFAULT 0 CHECK (progress BETWEEN 0 AND 100),
|
||||||
|
ADD COLUMN result jsonb,
|
||||||
|
ADD COLUMN error text,
|
||||||
|
ADD COLUMN worker_id text,
|
||||||
|
ADD COLUMN lease_expires_at timestamptz,
|
||||||
|
ADD COLUMN cancel_requested boolean NOT NULL DEFAULT false,
|
||||||
|
ADD COLUMN idempotency_key text UNIQUE,
|
||||||
|
ADD COLUMN attempt integer NOT NULL DEFAULT 0,
|
||||||
|
ADD CONSTRAINT tasks_status_check CHECK (status IN ('queued', 'running', 'success', 'error', 'cancelled'));
|
||||||
|
CREATE INDEX tasks_claim_idx ON tasks(status, lease_expires_at, created_at);
|
||||||
|
CREATE TABLE resource_locks (
|
||||||
|
resource_key text PRIMARY KEY,
|
||||||
|
task_id uuid NOT NULL UNIQUE REFERENCES tasks(id) ON DELETE CASCADE,
|
||||||
|
acquired_at timestamptz NOT NULL DEFAULT now()
|
||||||
|
);
|
||||||
|
CREATE TABLE task_logs (
|
||||||
|
task_id uuid NOT NULL REFERENCES tasks(id) ON DELETE CASCADE,
|
||||||
|
sequence bigint GENERATED ALWAYS AS IDENTITY,
|
||||||
|
created_at timestamptz NOT NULL DEFAULT now(),
|
||||||
|
message text NOT NULL,
|
||||||
|
PRIMARY KEY (task_id, sequence)
|
||||||
|
);
|
||||||
|
CREATE TABLE task_events (
|
||||||
|
task_id uuid NOT NULL REFERENCES tasks(id) ON DELETE CASCADE,
|
||||||
|
sequence bigint GENERATED ALWAYS AS IDENTITY,
|
||||||
|
created_at timestamptz NOT NULL DEFAULT now(),
|
||||||
|
kind text NOT NULL,
|
||||||
|
data jsonb NOT NULL DEFAULT '{}'::jsonb,
|
||||||
|
PRIMARY KEY (task_id, sequence)
|
||||||
|
);
|
||||||
+22
-1
@@ -2,8 +2,10 @@
|
|||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
from collections.abc import AsyncIterator, Callable
|
from collections.abc import AsyncIterator, Callable
|
||||||
from contextlib import AbstractAsyncContextManager, asynccontextmanager
|
from contextlib import AbstractAsyncContextManager, asynccontextmanager
|
||||||
|
from typing import Protocol
|
||||||
|
|
||||||
from fastapi import FastAPI
|
from fastapi import FastAPI
|
||||||
|
|
||||||
@@ -14,7 +16,20 @@ DatabaseFactory = Callable[[Settings], Database]
|
|||||||
Lifespan = Callable[[FastAPI], AbstractAsyncContextManager[None]]
|
Lifespan = Callable[[FastAPI], AbstractAsyncContextManager[None]]
|
||||||
|
|
||||||
|
|
||||||
def create_lifespan(settings: Settings, database_factory: DatabaseFactory) -> Lifespan:
|
class LifespanWorker(Protocol):
|
||||||
|
async def run(self) -> None: ...
|
||||||
|
|
||||||
|
def stop(self) -> None: ...
|
||||||
|
|
||||||
|
|
||||||
|
WorkerFactory = Callable[[Database], LifespanWorker]
|
||||||
|
|
||||||
|
|
||||||
|
def create_lifespan(
|
||||||
|
settings: Settings,
|
||||||
|
database_factory: DatabaseFactory,
|
||||||
|
worker_factories: tuple[WorkerFactory, ...] = (),
|
||||||
|
) -> Lifespan:
|
||||||
"""Build a lifespan context so tests can inject a database implementation."""
|
"""Build a lifespan context so tests can inject a database implementation."""
|
||||||
|
|
||||||
@asynccontextmanager
|
@asynccontextmanager
|
||||||
@@ -22,9 +37,15 @@ def create_lifespan(settings: Settings, database_factory: DatabaseFactory) -> Li
|
|||||||
database = database_factory(settings)
|
database = database_factory(settings)
|
||||||
await database.connect()
|
await database.connect()
|
||||||
app.state.database = database
|
app.state.database = database
|
||||||
|
workers = tuple(factory(database) for factory in worker_factories)
|
||||||
|
worker_tasks = tuple(asyncio.create_task(worker.run()) for worker in workers)
|
||||||
try:
|
try:
|
||||||
yield
|
yield
|
||||||
finally:
|
finally:
|
||||||
|
for worker in workers:
|
||||||
|
worker.stop()
|
||||||
|
if worker_tasks:
|
||||||
|
await asyncio.gather(*worker_tasks)
|
||||||
await database.close()
|
await database.close()
|
||||||
|
|
||||||
return lifespan
|
return lifespan
|
||||||
|
|||||||
+3
-2
@@ -10,7 +10,7 @@ from app.api.registry import HandlerRegistry, register_contract_routes
|
|||||||
from app.compatibility import build_report
|
from app.compatibility import build_report
|
||||||
from app.config import Settings, get_settings
|
from app.config import Settings, get_settings
|
||||||
from app.contracts.model import Snapshot
|
from app.contracts.model import Snapshot
|
||||||
from app.lifespan import DatabaseFactory, create_lifespan, default_database_factory
|
from app.lifespan import DatabaseFactory, WorkerFactory, create_lifespan, default_database_factory
|
||||||
from app.logging import configure_logging
|
from app.logging import configure_logging
|
||||||
from app.observability.health import router as health_router
|
from app.observability.health import router as health_router
|
||||||
|
|
||||||
@@ -19,6 +19,7 @@ def create_app(
|
|||||||
settings: Settings | None = None,
|
settings: Settings | None = None,
|
||||||
database_factory: DatabaseFactory = default_database_factory,
|
database_factory: DatabaseFactory = default_database_factory,
|
||||||
handlers: HandlerRegistry | None = None,
|
handlers: HandlerRegistry | None = None,
|
||||||
|
worker_factories: tuple[WorkerFactory, ...] = (),
|
||||||
) -> FastAPI:
|
) -> FastAPI:
|
||||||
"""Create an isolated application instance with explicit resource factories."""
|
"""Create an isolated application instance with explicit resource factories."""
|
||||||
|
|
||||||
@@ -27,7 +28,7 @@ def create_app(
|
|||||||
app = FastAPI(
|
app = FastAPI(
|
||||||
title=resolved.app_name,
|
title=resolved.app_name,
|
||||||
version="0.0.1",
|
version="0.0.1",
|
||||||
lifespan=create_lifespan(resolved, database_factory),
|
lifespan=create_lifespan(resolved, database_factory, worker_factories),
|
||||||
)
|
)
|
||||||
app.add_middleware(RequestContextMiddleware, header_name=resolved.request_id_header)
|
app.add_middleware(RequestContextMiddleware, header_name=resolved.request_id_header)
|
||||||
app.add_exception_handler(Exception, unhandled_exception_handler)
|
app.add_exception_handler(Exception, unhandled_exception_handler)
|
||||||
|
|||||||
@@ -0,0 +1 @@
|
|||||||
|
"""Durable asynchronous task engine."""
|
||||||
@@ -0,0 +1,185 @@
|
|||||||
|
"""PostgreSQL repository for durable leased tasks."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import json
|
||||||
|
import uuid
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
import asyncpg # type: ignore[import-untyped] # noqa: F401
|
||||||
|
from asyncpg import Pool, Record
|
||||||
|
|
||||||
|
from app.db.primitives import ConflictError, require_affected, transaction
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class Task:
|
||||||
|
id: uuid.UUID
|
||||||
|
upid: str
|
||||||
|
task_type: str
|
||||||
|
status: str
|
||||||
|
payload: dict[str, Any]
|
||||||
|
progress: int
|
||||||
|
cancel_requested: bool
|
||||||
|
attempt: int
|
||||||
|
|
||||||
|
|
||||||
|
def _task(row: Record) -> Task:
|
||||||
|
return Task(
|
||||||
|
id=row["id"],
|
||||||
|
upid=str(row["upid"]),
|
||||||
|
task_type=str(row["task_type"]),
|
||||||
|
status=str(row["status"]),
|
||||||
|
payload=json.loads(row["payload"])
|
||||||
|
if isinstance(row["payload"], str)
|
||||||
|
else dict(row["payload"]),
|
||||||
|
progress=int(row["progress"]),
|
||||||
|
cancel_requested=bool(row["cancel_requested"]),
|
||||||
|
attempt=int(row["attempt"]),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class TaskRepository:
|
||||||
|
pool: Pool
|
||||||
|
|
||||||
|
async def create(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
upid: str,
|
||||||
|
task_type: str,
|
||||||
|
payload: dict[str, Any],
|
||||||
|
resource_key: str | None = None,
|
||||||
|
idempotency_key: str | None = None,
|
||||||
|
) -> Task:
|
||||||
|
task_id = uuid.uuid4()
|
||||||
|
async with transaction(self.pool) as connection:
|
||||||
|
if idempotency_key is not None:
|
||||||
|
existing = await connection.fetchrow(
|
||||||
|
"SELECT * FROM tasks WHERE idempotency_key=$1", idempotency_key
|
||||||
|
)
|
||||||
|
if existing is not None:
|
||||||
|
return _task(existing)
|
||||||
|
row = await connection.fetchrow(
|
||||||
|
"""INSERT INTO tasks(id, upid, task_type, status, payload, idempotency_key)
|
||||||
|
VALUES($1,$2,$3,'queued',$4::jsonb,$5) RETURNING *""",
|
||||||
|
task_id,
|
||||||
|
upid,
|
||||||
|
task_type,
|
||||||
|
json.dumps(payload),
|
||||||
|
idempotency_key,
|
||||||
|
)
|
||||||
|
if resource_key is not None:
|
||||||
|
try:
|
||||||
|
await connection.execute(
|
||||||
|
"INSERT INTO resource_locks(resource_key, task_id) VALUES($1,$2)",
|
||||||
|
resource_key,
|
||||||
|
task_id,
|
||||||
|
)
|
||||||
|
except Exception as error:
|
||||||
|
raise ConflictError(f"resource is locked: {resource_key}") from error
|
||||||
|
await connection.execute(
|
||||||
|
"INSERT INTO task_events(task_id, kind) VALUES($1,'created')", task_id
|
||||||
|
)
|
||||||
|
if row is None:
|
||||||
|
raise RuntimeError("task insert returned no row")
|
||||||
|
return _task(row)
|
||||||
|
|
||||||
|
async def claim(self, worker_id: str, lease_seconds: float) -> Task | None:
|
||||||
|
async with transaction(self.pool) as connection:
|
||||||
|
row = await connection.fetchrow(
|
||||||
|
"""WITH candidate AS (
|
||||||
|
SELECT id FROM tasks
|
||||||
|
WHERE status='queued' OR (status='running' AND lease_expires_at < now())
|
||||||
|
ORDER BY created_at FOR UPDATE SKIP LOCKED LIMIT 1
|
||||||
|
) UPDATE tasks SET status='running', worker_id=$1,
|
||||||
|
lease_expires_at=now() + $2 * interval '1 second', attempt=attempt+1,
|
||||||
|
updated_at=now()
|
||||||
|
WHERE id=(SELECT id FROM candidate) RETURNING *""",
|
||||||
|
worker_id,
|
||||||
|
lease_seconds,
|
||||||
|
)
|
||||||
|
if row is None:
|
||||||
|
return None
|
||||||
|
await connection.execute(
|
||||||
|
"INSERT INTO task_events(task_id, kind, data) VALUES($1,'claimed',$2::jsonb)",
|
||||||
|
row["id"],
|
||||||
|
json.dumps({"worker": worker_id}),
|
||||||
|
)
|
||||||
|
return _task(row)
|
||||||
|
|
||||||
|
async def heartbeat(self, task_id: uuid.UUID, worker_id: str, lease_seconds: float) -> None:
|
||||||
|
status = await self.pool.execute(
|
||||||
|
"""UPDATE tasks SET lease_expires_at=now()+$3*interval '1 second', updated_at=now()
|
||||||
|
WHERE id=$1 AND worker_id=$2 AND status='running'""",
|
||||||
|
task_id,
|
||||||
|
worker_id,
|
||||||
|
lease_seconds,
|
||||||
|
)
|
||||||
|
require_affected(status)
|
||||||
|
|
||||||
|
async def progress(self, task_id: uuid.UUID, worker_id: str, value: int) -> None:
|
||||||
|
status = await self.pool.execute(
|
||||||
|
"""UPDATE tasks SET progress=$3, updated_at=now()
|
||||||
|
WHERE id=$1 AND worker_id=$2 AND status='running'""",
|
||||||
|
task_id,
|
||||||
|
worker_id,
|
||||||
|
value,
|
||||||
|
)
|
||||||
|
require_affected(status)
|
||||||
|
|
||||||
|
async def append_log(self, task_id: uuid.UUID, message: str) -> None:
|
||||||
|
await self.pool.execute(
|
||||||
|
"INSERT INTO task_logs(task_id, message) VALUES($1,$2)", task_id, message
|
||||||
|
)
|
||||||
|
|
||||||
|
async def request_cancel(self, task_id: uuid.UUID) -> None:
|
||||||
|
status = await self.pool.execute(
|
||||||
|
"""UPDATE tasks SET cancel_requested=true, updated_at=now()
|
||||||
|
WHERE id=$1 AND status IN ('queued','running')""",
|
||||||
|
task_id,
|
||||||
|
)
|
||||||
|
require_affected(status)
|
||||||
|
|
||||||
|
async def finish(
|
||||||
|
self,
|
||||||
|
task_id: uuid.UUID,
|
||||||
|
worker_id: str,
|
||||||
|
*,
|
||||||
|
status: str,
|
||||||
|
result: dict[str, Any] | None = None,
|
||||||
|
error: str | None = None,
|
||||||
|
) -> None:
|
||||||
|
if status not in {"success", "error", "cancelled"}:
|
||||||
|
raise ValueError("invalid terminal task status")
|
||||||
|
async with transaction(self.pool) as connection:
|
||||||
|
command = await connection.execute(
|
||||||
|
"""UPDATE tasks SET status=$3, result=$4::jsonb, error=$5,
|
||||||
|
progress=CASE WHEN $3='success' THEN 100 ELSE progress END,
|
||||||
|
lease_expires_at=NULL, updated_at=now()
|
||||||
|
WHERE id=$1 AND worker_id=$2 AND status='running'""",
|
||||||
|
task_id,
|
||||||
|
worker_id,
|
||||||
|
status,
|
||||||
|
json.dumps(result) if result is not None else None,
|
||||||
|
error,
|
||||||
|
)
|
||||||
|
require_affected(command)
|
||||||
|
await connection.execute("DELETE FROM resource_locks WHERE task_id=$1", task_id)
|
||||||
|
await connection.execute(
|
||||||
|
"INSERT INTO task_events(task_id, kind, data) VALUES($1,$2,$3::jsonb)",
|
||||||
|
task_id,
|
||||||
|
status,
|
||||||
|
json.dumps({"error": error} if error else {}),
|
||||||
|
)
|
||||||
|
|
||||||
|
async def get(self, task_id: uuid.UUID) -> Task | None:
|
||||||
|
row = await self.pool.fetchrow("SELECT * FROM tasks WHERE id=$1", task_id)
|
||||||
|
return _task(row) if row is not None else None
|
||||||
|
|
||||||
|
async def logs(self, task_id: uuid.UUID) -> tuple[str, ...]:
|
||||||
|
rows = await self.pool.fetch(
|
||||||
|
"SELECT message FROM task_logs WHERE task_id=$1 ORDER BY sequence", task_id
|
||||||
|
)
|
||||||
|
return tuple(str(row["message"]) for row in rows)
|
||||||
@@ -0,0 +1,59 @@
|
|||||||
|
"""Proxmox-compatible unique process/task identifiers."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import re
|
||||||
|
from dataclasses import dataclass
|
||||||
|
|
||||||
|
UPID_RE = re.compile(
|
||||||
|
r"^UPID:(?P<node>[A-Za-z0-9][A-Za-z0-9_-]*):"
|
||||||
|
r"(?P<pid>[0-9A-Fa-f]{8}):(?P<pstart>[0-9A-Fa-f]{8}):"
|
||||||
|
r"(?P<start>[0-9A-Fa-f]{8}):(?P<type>[A-Za-z0-9_-]+):"
|
||||||
|
r"(?P<task_id>[^:]*):(?P<user>[^:]+):$"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class Upid:
|
||||||
|
node: str
|
||||||
|
pid: int
|
||||||
|
process_start: int
|
||||||
|
start_time: int
|
||||||
|
task_type: str
|
||||||
|
task_id: str
|
||||||
|
user: str
|
||||||
|
|
||||||
|
def __post_init__(self) -> None:
|
||||||
|
for name, value in (
|
||||||
|
("pid", self.pid),
|
||||||
|
("process_start", self.process_start),
|
||||||
|
("start_time", self.start_time),
|
||||||
|
):
|
||||||
|
if not 0 <= value <= 0xFFFFFFFF:
|
||||||
|
raise ValueError(f"{name} is outside the 32-bit UPID range")
|
||||||
|
if not self.node or ":" in self.node or not self.task_type or ":" in self.task_type:
|
||||||
|
raise ValueError("invalid UPID node or task type")
|
||||||
|
if ":" in self.task_id or not self.user or ":" in self.user:
|
||||||
|
raise ValueError("invalid UPID task id or user")
|
||||||
|
|
||||||
|
def __str__(self) -> str:
|
||||||
|
return (
|
||||||
|
f"UPID:{self.node}:{self.pid:08X}:{self.process_start:08X}:"
|
||||||
|
f"{self.start_time:08X}:{self.task_type}:{self.task_id}:{self.user}:"
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def parse(cls, value: str) -> Upid:
|
||||||
|
match = UPID_RE.fullmatch(value)
|
||||||
|
if match is None:
|
||||||
|
raise ValueError("invalid UPID")
|
||||||
|
values = match.groupdict()
|
||||||
|
return cls(
|
||||||
|
node=values["node"],
|
||||||
|
pid=int(values["pid"], 16),
|
||||||
|
process_start=int(values["pstart"], 16),
|
||||||
|
start_time=int(values["start"], 16),
|
||||||
|
task_type=values["type"],
|
||||||
|
task_id=values["task_id"],
|
||||||
|
user=values["user"],
|
||||||
|
)
|
||||||
@@ -0,0 +1,90 @@
|
|||||||
|
"""Bounded durable task worker with cooperative cancellation."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
from collections.abc import Awaitable, Callable
|
||||||
|
from dataclasses import dataclass, field
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from app.tasks.repository import Task, TaskRepository
|
||||||
|
|
||||||
|
TaskHandler = Callable[[Task], Awaitable[dict[str, Any] | None]]
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(slots=True)
|
||||||
|
class TaskWorker:
|
||||||
|
repository: TaskRepository
|
||||||
|
worker_id: str
|
||||||
|
handlers: dict[str, TaskHandler]
|
||||||
|
concurrency: int = 2
|
||||||
|
lease_seconds: float = 30.0
|
||||||
|
poll_seconds: float = 0.1
|
||||||
|
_running: set[asyncio.Task[None]] = field(default_factory=set, init=False)
|
||||||
|
_stopping: asyncio.Event = field(default_factory=asyncio.Event, init=False)
|
||||||
|
|
||||||
|
async def run(self) -> None:
|
||||||
|
self._stopping.clear()
|
||||||
|
try:
|
||||||
|
while not self._stopping.is_set():
|
||||||
|
self._reap()
|
||||||
|
if len(self._running) >= self.concurrency:
|
||||||
|
await asyncio.sleep(self.poll_seconds)
|
||||||
|
continue
|
||||||
|
task = await self.repository.claim(self.worker_id, self.lease_seconds)
|
||||||
|
if task is None:
|
||||||
|
await asyncio.sleep(self.poll_seconds)
|
||||||
|
continue
|
||||||
|
execution = asyncio.create_task(self._execute(task))
|
||||||
|
self._running.add(execution)
|
||||||
|
finally:
|
||||||
|
if self._running:
|
||||||
|
await asyncio.gather(*self._running, return_exceptions=True)
|
||||||
|
self._running.clear()
|
||||||
|
|
||||||
|
def stop(self) -> None:
|
||||||
|
self._stopping.set()
|
||||||
|
|
||||||
|
def _reap(self) -> None:
|
||||||
|
self._running = {task for task in self._running if not task.done()}
|
||||||
|
|
||||||
|
async def _execute(self, task: Task) -> None:
|
||||||
|
handler = self.handlers.get(task.task_type)
|
||||||
|
if handler is None:
|
||||||
|
await self.repository.finish(
|
||||||
|
task.id, self.worker_id, status="error", error="unsupported task type"
|
||||||
|
)
|
||||||
|
return
|
||||||
|
try:
|
||||||
|
current = await self.repository.get(task.id)
|
||||||
|
if current is not None and current.cancel_requested:
|
||||||
|
await self.repository.finish(task.id, self.worker_id, status="cancelled")
|
||||||
|
return
|
||||||
|
execution: asyncio.Future[dict[str, Any] | None] = asyncio.ensure_future(handler(task))
|
||||||
|
heartbeat = asyncio.create_task(self._heartbeat(task))
|
||||||
|
try:
|
||||||
|
while not execution.done():
|
||||||
|
await asyncio.sleep(self.poll_seconds)
|
||||||
|
current = await self.repository.get(task.id)
|
||||||
|
if current is not None and current.cancel_requested:
|
||||||
|
execution.cancel()
|
||||||
|
await asyncio.gather(execution, return_exceptions=True)
|
||||||
|
await self.repository.finish(task.id, self.worker_id, status="cancelled")
|
||||||
|
return
|
||||||
|
result = await execution
|
||||||
|
finally:
|
||||||
|
heartbeat.cancel()
|
||||||
|
await asyncio.gather(heartbeat, return_exceptions=True)
|
||||||
|
await self.repository.finish(task.id, self.worker_id, status="success", result=result)
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
raise
|
||||||
|
except Exception as error: # task failures are persisted, not leaked
|
||||||
|
await self.repository.finish(
|
||||||
|
task.id, self.worker_id, status="error", error=type(error).__name__
|
||||||
|
)
|
||||||
|
|
||||||
|
async def _heartbeat(self, task: Task) -> None:
|
||||||
|
interval = max(self.lease_seconds / 3, 0.01)
|
||||||
|
while True:
|
||||||
|
await asyncio.sleep(interval)
|
||||||
|
await self.repository.heartbeat(task.id, self.worker_id, self.lease_seconds)
|
||||||
@@ -90,6 +90,13 @@ domain models and repositories that do not depend on FastAPI or source-specific
|
|||||||
contract structures. PostgreSQL is the system of record for resources, security
|
contract structures. PostgreSQL is the system of record for resources, security
|
||||||
state, locks, scenarios, and tasks.
|
state, locks, scenarios, and tasks.
|
||||||
|
|
||||||
|
Durable tasks are acknowledged only after the task row, event, idempotency key,
|
||||||
|
and optional resource lock commit together. Workers claim with `SKIP LOCKED`,
|
||||||
|
renew real-time leases, persist progress and append-only logs/events, and allow
|
||||||
|
expired work to be reclaimed after process failure. Lifespan owns a bounded set
|
||||||
|
of asyncio workers and waits for orderly shutdown; PostgreSQL remains the queue
|
||||||
|
and source of truth across replicas.
|
||||||
|
|
||||||
Authentication secrets use salted scrypt hashes. Session tickets are signed and
|
Authentication secrets use salted scrypt hashes. Session tickets are signed and
|
||||||
expiring; mutation requests use ticket-bound CSRF tokens. API-token privileges
|
expiring; mutation requests use ticket-bound CSRF tokens. API-token privileges
|
||||||
are intersected with their owning principal's effective propagated ACLs, so a
|
are intersected with their owning principal's effective propagated ACLs, so a
|
||||||
|
|||||||
@@ -0,0 +1,81 @@
|
|||||||
|
"""Durable task concurrency and recovery tests."""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import os
|
||||||
|
import uuid
|
||||||
|
|
||||||
|
import asyncpg # type: ignore[import-untyped]
|
||||||
|
import pytest
|
||||||
|
from asyncpg import Pool
|
||||||
|
|
||||||
|
from app.db.migrations import migrate
|
||||||
|
from app.tasks.repository import TaskRepository
|
||||||
|
|
||||||
|
pytestmark = [
|
||||||
|
pytest.mark.integration,
|
||||||
|
pytest.mark.skipif(not os.getenv("TEST_DATABASE_URL"), reason="TEST_DATABASE_URL is required"),
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
async def repository() -> tuple[Pool, TaskRepository]:
|
||||||
|
pool = await asyncpg.create_pool(os.environ["TEST_DATABASE_URL"], min_size=1, max_size=4)
|
||||||
|
async with pool.acquire() as connection:
|
||||||
|
await migrate(connection)
|
||||||
|
return pool, TaskRepository(pool)
|
||||||
|
|
||||||
|
|
||||||
|
async def test_two_worker_exclusion_idempotency_and_logs() -> None:
|
||||||
|
pool, tasks = await repository()
|
||||||
|
key = uuid.uuid4().hex
|
||||||
|
try:
|
||||||
|
created = await tasks.create(
|
||||||
|
upid=f"UPID:pve1:00000001:00000001:00000001:test:{key}:root@pam:",
|
||||||
|
task_type="test",
|
||||||
|
payload={"value": 1},
|
||||||
|
resource_key=f"vm:{key}",
|
||||||
|
idempotency_key=key,
|
||||||
|
)
|
||||||
|
repeated = await tasks.create(
|
||||||
|
upid=f"ignored-{key}", task_type="test", payload={}, idempotency_key=key
|
||||||
|
)
|
||||||
|
assert repeated.id == created.id
|
||||||
|
|
||||||
|
first, second = await asyncio.gather(
|
||||||
|
tasks.claim("worker-a", 30), tasks.claim("worker-b", 30)
|
||||||
|
)
|
||||||
|
claimed = first or second
|
||||||
|
assert claimed is not None
|
||||||
|
assert (first is None) != (second is None)
|
||||||
|
worker = "worker-a" if first is not None else "worker-b"
|
||||||
|
await tasks.append_log(claimed.id, "started")
|
||||||
|
await tasks.progress(claimed.id, worker, 50)
|
||||||
|
await tasks.finish(claimed.id, worker, status="success", result={"ok": True})
|
||||||
|
assert await tasks.logs(claimed.id) == ("started",)
|
||||||
|
finished = await tasks.get(claimed.id)
|
||||||
|
assert finished is not None
|
||||||
|
assert finished.status == "success"
|
||||||
|
finally:
|
||||||
|
await pool.close()
|
||||||
|
|
||||||
|
|
||||||
|
async def test_expired_lease_is_reclaimed_after_restart() -> None:
|
||||||
|
pool, tasks = await repository()
|
||||||
|
key = uuid.uuid4().hex
|
||||||
|
try:
|
||||||
|
created = await tasks.create(
|
||||||
|
upid=f"UPID:pve1:00000001:00000001:00000001:test:{key}:root@pam:",
|
||||||
|
task_type="test",
|
||||||
|
payload={},
|
||||||
|
)
|
||||||
|
assert await tasks.claim("dead-worker", 0) is not None
|
||||||
|
recovered = await tasks.claim("new-worker", 30)
|
||||||
|
assert recovered is not None
|
||||||
|
assert recovered.id == created.id
|
||||||
|
assert recovered.attempt == 2
|
||||||
|
await tasks.request_cancel(recovered.id)
|
||||||
|
cancelled = await tasks.get(recovered.id)
|
||||||
|
assert cancelled is not None
|
||||||
|
assert cancelled.cancel_requested
|
||||||
|
await tasks.finish(recovered.id, "new-worker", status="cancelled")
|
||||||
|
finally:
|
||||||
|
await pool.close()
|
||||||
@@ -1,5 +1,6 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
from typing import Self
|
from typing import Self
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
@@ -56,3 +57,27 @@ async def test_health_endpoints(database_ready: bool, status_code: int) -> None:
|
|||||||
assert ready.headers["X-Request-ID"] == "test-request"
|
assert ready.headers["X-Request-ID"] == "test-request"
|
||||||
assert database.connected
|
assert database.connected
|
||||||
assert database.closed
|
assert database.closed
|
||||||
|
|
||||||
|
|
||||||
|
async def test_lifespan_starts_and_stops_injected_workers() -> None:
|
||||||
|
database = FakeDatabase(True)
|
||||||
|
started = asyncio.Event()
|
||||||
|
stopping = asyncio.Event()
|
||||||
|
|
||||||
|
class Worker:
|
||||||
|
async def run(self) -> None:
|
||||||
|
started.set()
|
||||||
|
await stopping.wait()
|
||||||
|
|
||||||
|
def stop(self) -> None:
|
||||||
|
stopping.set()
|
||||||
|
|
||||||
|
application = create_app(
|
||||||
|
Settings(),
|
||||||
|
lambda _settings: database,
|
||||||
|
worker_factories=(lambda _database: Worker(),),
|
||||||
|
)
|
||||||
|
async with application.router.lifespan_context(application):
|
||||||
|
await started.wait()
|
||||||
|
|
||||||
|
assert stopping.is_set()
|
||||||
|
|||||||
@@ -0,0 +1,74 @@
|
|||||||
|
"""Bounded task worker outcome tests."""
|
||||||
|
|
||||||
|
import uuid
|
||||||
|
from typing import cast
|
||||||
|
|
||||||
|
from app.tasks.repository import Task, TaskRepository
|
||||||
|
from app.tasks.worker import TaskWorker
|
||||||
|
|
||||||
|
|
||||||
|
class FakeRepository:
|
||||||
|
def __init__(self, task: Task) -> None:
|
||||||
|
self.task = task
|
||||||
|
self.finishes: list[tuple[str, str | None]] = []
|
||||||
|
|
||||||
|
async def get(self, _task_id: uuid.UUID) -> Task:
|
||||||
|
return self.task
|
||||||
|
|
||||||
|
async def finish(
|
||||||
|
self,
|
||||||
|
_task_id: uuid.UUID,
|
||||||
|
_worker_id: str,
|
||||||
|
*,
|
||||||
|
status: str,
|
||||||
|
result: dict[str, object] | None = None,
|
||||||
|
error: str | None = None,
|
||||||
|
) -> None:
|
||||||
|
del result
|
||||||
|
self.finishes.append((status, error))
|
||||||
|
|
||||||
|
|
||||||
|
def make_task(*, task_type: str = "test", cancelled: bool = False) -> Task:
|
||||||
|
return Task(uuid.uuid4(), "UPID:test", task_type, "running", {}, 0, cancelled, 1)
|
||||||
|
|
||||||
|
|
||||||
|
async def test_worker_persists_success_error_and_unsupported() -> None:
|
||||||
|
task = make_task()
|
||||||
|
repository = FakeRepository(task)
|
||||||
|
|
||||||
|
async def success(_task: Task) -> dict[str, object]:
|
||||||
|
return {"ok": True}
|
||||||
|
|
||||||
|
worker = TaskWorker(cast(TaskRepository, repository), "worker", {"test": success})
|
||||||
|
await worker._execute(task)
|
||||||
|
assert repository.finishes == [("success", None)]
|
||||||
|
|
||||||
|
unsupported = make_task(task_type="missing")
|
||||||
|
repository.task = unsupported
|
||||||
|
await worker._execute(unsupported)
|
||||||
|
assert repository.finishes[-1] == ("error", "unsupported task type")
|
||||||
|
|
||||||
|
async def failure(_task: Task) -> None:
|
||||||
|
raise RuntimeError("private detail")
|
||||||
|
|
||||||
|
failed = make_task()
|
||||||
|
repository.task = failed
|
||||||
|
worker.handlers["test"] = failure
|
||||||
|
await worker._execute(failed)
|
||||||
|
assert repository.finishes[-1] == ("error", "RuntimeError")
|
||||||
|
|
||||||
|
|
||||||
|
async def test_worker_honors_persisted_cancellation() -> None:
|
||||||
|
task = make_task(cancelled=True)
|
||||||
|
repository = FakeRepository(task)
|
||||||
|
called = False
|
||||||
|
|
||||||
|
async def handler(_task: Task) -> None:
|
||||||
|
nonlocal called
|
||||||
|
called = True
|
||||||
|
|
||||||
|
worker = TaskWorker(cast(TaskRepository, repository), "worker", {"test": handler})
|
||||||
|
await worker._execute(task)
|
||||||
|
|
||||||
|
assert not called
|
||||||
|
assert repository.finishes == [("cancelled", None)]
|
||||||
@@ -0,0 +1,71 @@
|
|||||||
|
"""UPID examples and round-trip properties."""
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
from hypothesis import given
|
||||||
|
from hypothesis import strategies as st
|
||||||
|
|
||||||
|
from app.tasks.upid import Upid
|
||||||
|
|
||||||
|
SAFE = st.from_regex(r"[a-z0-9][a-z0-9_-]{0,19}", fullmatch=True)
|
||||||
|
|
||||||
|
|
||||||
|
@given(
|
||||||
|
node=SAFE,
|
||||||
|
pid=st.integers(min_value=0, max_value=0xFFFFFFFF),
|
||||||
|
process_start=st.integers(min_value=0, max_value=0xFFFFFFFF),
|
||||||
|
start_time=st.integers(min_value=0, max_value=0xFFFFFFFF),
|
||||||
|
task_type=SAFE,
|
||||||
|
task_id=st.text(alphabet="abcdefghijklmnopqrstuvwxyz0123456789_-", max_size=20),
|
||||||
|
user=SAFE,
|
||||||
|
)
|
||||||
|
def test_upid_round_trip(
|
||||||
|
node: str,
|
||||||
|
pid: int,
|
||||||
|
process_start: int,
|
||||||
|
start_time: int,
|
||||||
|
task_type: str,
|
||||||
|
task_id: str,
|
||||||
|
user: str,
|
||||||
|
) -> None:
|
||||||
|
upid = Upid(node, pid, process_start, start_time, task_type, task_id, user)
|
||||||
|
|
||||||
|
assert Upid.parse(str(upid)) == upid
|
||||||
|
|
||||||
|
|
||||||
|
def test_known_upid_shape() -> None:
|
||||||
|
value = "UPID:pve1:0000002A:00000010:65A1B2C3:qmstart:100:root@pam:"
|
||||||
|
|
||||||
|
parsed = Upid.parse(value)
|
||||||
|
|
||||||
|
assert parsed.pid == 42
|
||||||
|
assert parsed.task_id == "100"
|
||||||
|
assert str(parsed) == value
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("value", ["", "UPID:broken", "UPID:pve:GGGGGGGG:00000000:00000000:x::u:"])
|
||||||
|
def test_invalid_upids_are_rejected(value: str) -> None:
|
||||||
|
with pytest.raises(ValueError):
|
||||||
|
Upid.parse(value)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"kwargs",
|
||||||
|
[
|
||||||
|
{"pid": -1},
|
||||||
|
{"node": "bad:node"},
|
||||||
|
{"task_id": "bad:id"},
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_invalid_upid_components_are_rejected(kwargs: dict[str, object]) -> None:
|
||||||
|
values: dict[str, object] = {
|
||||||
|
"node": "pve1",
|
||||||
|
"pid": 1,
|
||||||
|
"process_start": 1,
|
||||||
|
"start_time": 1,
|
||||||
|
"task_type": "test",
|
||||||
|
"task_id": "100",
|
||||||
|
"user": "root@pam",
|
||||||
|
}
|
||||||
|
values.update(kwargs)
|
||||||
|
with pytest.raises(ValueError):
|
||||||
|
Upid(**values) # type: ignore[arg-type]
|
||||||
Reference in New Issue
Block a user