diff --git a/.env.example b/.env.example index c18fc2e..32c161a 100644 --- a/.env.example +++ b/.env.example @@ -11,6 +11,8 @@ PVE_API_VERSION=9.2.3 CONTRACT_SNAPSHOT=/app/contracts/pve-9.2.3.json CONTRACT_FALLBACK=error TICKET_SIGNING_KEY=development-only-signing-key-change-me +TASK_WORKER_CONCURRENCY=2 +TASK_LEASE_SECONDS=30 SIMULATION_SEED=42 SIMULATION_TIME_SCALE=10 SIMULATOR_ADMIN_ENABLED=false diff --git a/README.md b/README.md index 287cc99..858c60b 100644 --- a/README.md +++ b/README.md @@ -10,7 +10,9 @@ is claimed as compatible yet; the vertical slice is tracked in The bundled PVE 9.2.3 declared contract contains 444 paths and 675 methods. Implemented semantics currently include version, ticket login, node listing and -status, and cluster resources; all other declared methods return an explicit +status, cluster resources, basic QEMU list/config/status/start/stop, and task +list/status/log. Mutations require the ticket-bound CSRF header and execute +through PostgreSQL-leased workers; all other declared methods return an explicit unsupported error. ## Development diff --git a/app/api/registry.py b/app/api/registry.py index 5cb0390..5dc7cfc 100644 --- a/app/api/registry.py +++ b/app/api/registry.py @@ -4,14 +4,16 @@ from __future__ import annotations from collections.abc import Awaitable, Callable from dataclasses import dataclass, field -from typing import Any, Literal +from typing import Any, Literal, cast from urllib.parse import parse_qsl from fastapi import FastAPI, Request from fastapi.responses import JSONResponse -from app.api.errors import ContractValidationError +from app.api.errors import ApiError, ContractValidationError +from app.config import Settings from app.contracts.model import Method, Schema, Snapshot +from app.security.auth import verify_csrf, verify_ticket Handler = Callable[[Request, dict[str, Any]], Awaitable[Any]] FallbackMode = Literal["error", "schema-default", "fixture"] @@ -79,6 +81,7 @@ def _endpoint( fallback: FallbackMode, ) -> Callable[[Request], Awaitable[JSONResponse]]: async def dispatch(request: Request) -> JSONResponse: + _authenticate(request, semantic_path) handler = handlers.get(semantic_path, method.verb) inputs = await _parse_inputs(request, method) if handler is not None: @@ -92,13 +95,37 @@ def _endpoint( status_code=501, content={"data": None, "errors": "method semantics are not implemented"}, ) - if renderer == "extjs": - return JSONResponse({"data": data, "success": True}) - return JSONResponse({"data": data}) + content = {"data": data, "success": True} if renderer == "extjs" else {"data": data} + response = JSONResponse(content) + if semantic_path == "/access/ticket" and isinstance(data, dict): + ticket = data.get("ticket") + if isinstance(ticket, str): + response.set_cookie( + "PVEAuthCookie", ticket, httponly=True, samesite="strict", path="/" + ) + return response return dispatch +def _authenticate(request: Request, semantic_path: str) -> None: + if semantic_path in {"/version", "/access/ticket"}: + return + ticket = request.cookies.get("PVEAuthCookie") + if ticket is None: + raise ApiError(401, "authentication required") + settings = cast(Settings, request.app.state.settings) + key = settings.ticket_signing_key.get_secret_value().encode() + try: + verify_ticket(ticket, key) + except ValueError as error: + raise ApiError(401, "authentication failure") from error + if request.method not in {"GET", "HEAD", "OPTIONS"}: + token = request.headers.get("CSRFPreventionToken", "") + if not verify_csrf(ticket, token, key): + raise ApiError(403, "invalid CSRF prevention token") + + async def _parse_inputs(request: Request, method: Method) -> dict[str, Any]: supplied: dict[str, Any] = dict(request.query_params) supplied.update(request.path_params) diff --git a/app/config.py b/app/config.py index 4fd2237..705e4d5 100644 --- a/app/config.py +++ b/app/config.py @@ -35,6 +35,9 @@ class Settings(BaseSettings): contract_snapshot: Path | None = None contract_fallback: Literal["error", "schema-default", "fixture"] = "error" ticket_signing_key: SecretStr = SecretStr("development-only-signing-key-change-me") + task_worker_concurrency: int = Field(default=2, ge=1, le=32) + task_lease_seconds: float = Field(default=30.0, gt=1, le=300) + simulation_time_scale: float = Field(default=10.0, gt=0, le=10000) @lru_cache(maxsize=1) diff --git a/app/handlers/core.py b/app/handlers/core.py index 72bac04..77c470e 100644 --- a/app/handlers/core.py +++ b/app/handlers/core.py @@ -11,6 +11,7 @@ from app.api.errors import ApiError from app.api.registry import HandlerRegistry from app.config import Settings from app.db.pool import AsyncpgDatabase +from app.handlers.qemu import register_qemu_handlers from app.security.auth import csrf_token, issue_ticket, verify_secret @@ -92,4 +93,5 @@ def build_core_handlers(settings: Settings) -> HandlerRegistry: registry.register("/nodes", "GET", nodes) registry.register("/nodes/{node}/status", "GET", node_status) registry.register("/cluster/resources", "GET", resources) + register_qemu_handlers(registry) return registry diff --git a/app/handlers/qemu.py b/app/handlers/qemu.py new file mode 100644 index 0000000..2d79d5f --- /dev/null +++ b/app/handlers/qemu.py @@ -0,0 +1,133 @@ +"""Basic persistent QEMU and task semantic handlers.""" + +from __future__ import annotations + +import json +from collections.abc import Mapping +from typing import Any, cast + +from fastapi import Request + +from app.api.errors import ApiError +from app.api.registry import HandlerRegistry +from app.db.pool import AsyncpgDatabase +from app.tasks.repository import TaskRepository +from app.tasks.upid import Upid + + +def _database(request: Request) -> AsyncpgDatabase: + return cast(AsyncpgDatabase, request.app.state.database) + + +def _values(inputs: dict[str, Any]) -> dict[str, Any]: + return cast(dict[str, Any], inputs["values"]) + + +def _state(value: object) -> dict[str, Any]: + if isinstance(value, str): + return cast(dict[str, Any], json.loads(value)) + return dict(cast(Mapping[str, Any], value)) + + +def register_qemu_handlers(registry: HandlerRegistry) -> None: + async def qemu_list(request: Request, inputs: dict[str, Any]) -> list[dict[str, Any]]: + node = str(_values(inputs)["node"]) + rows = await _database(request).pool.fetch( + """SELECT r.external_id::integer AS vmid, r.state + FROM resources r JOIN nodes n ON n.id=r.node_id + WHERE n.name=$1 AND r.kind='qemu' ORDER BY r.external_id::integer""", + node, + ) + return [{"vmid": int(row["vmid"]), **_state(row["state"])} for row in rows] + + async def qemu_config(request: Request, inputs: dict[str, Any]) -> dict[str, Any]: + node, vmid = str(_values(inputs)["node"]), str(_values(inputs)["vmid"]) + row = await _database(request).pool.fetchrow( + """SELECT r.state FROM resources r JOIN nodes n ON n.id=r.node_id + WHERE n.name=$1 AND r.kind='qemu' AND r.external_id=$2""", + node, + vmid, + ) + if row is None: + raise ApiError(404, "virtual machine does not exist") + return {"vmid": int(vmid), **_state(row["state"])} + + async def qemu_status(request: Request, inputs: dict[str, Any]) -> dict[str, Any]: + return await qemu_config(request, inputs) + + async def mutate(operation: str, request: Request, inputs: dict[str, Any]) -> str: + values = _values(inputs) + node, vmid = str(values["node"]), str(values["vmid"]) + database = _database(request) + row = await database.pool.fetchrow( + """SELECT r.id, r.state FROM resources r JOIN nodes n ON n.id=r.node_id + WHERE n.name=$1 AND r.kind='qemu' AND r.external_id=$2""", + node, + vmid, + ) + if row is None: + raise ApiError(404, "virtual machine does not exist") + current = str(_state(row["state"]).get("status", "stopped")) + if (operation == "start" and current != "stopped") or ( + operation == "stop" and current != "running" + ): + raise ApiError(409, f"cannot {operation} VM while it is {current}") + timestamp = int(await database.pool.fetchval("SELECT extract(epoch from now())::bigint")) + pid = int(await database.pool.fetchval("SELECT pg_backend_pid()")) + upid = str(Upid(node, pid, pid, timestamp, f"qm{operation}", vmid, "root@pam")) + task = await TaskRepository(database.pool).create( + upid=upid, + task_type=f"qemu-{operation}", + payload={"node": node, "vmid": vmid, "resource_id": str(row["id"])}, + resource_key=f"qemu:{vmid}", + idempotency_key=request.headers.get("Idempotency-Key"), + ) + return task.upid + + async def start(request: Request, inputs: dict[str, Any]) -> str: + return await mutate("start", request, inputs) + + async def stop(request: Request, inputs: dict[str, Any]) -> str: + return await mutate("stop", request, inputs) + + async def task_list(request: Request, inputs: dict[str, Any]) -> list[dict[str, Any]]: + tasks = await TaskRepository(_database(request).pool).list_for_node( + str(_values(inputs)["node"]) + ) + return [ + {"upid": task.upid, "status": task.status, "type": task.task_type} for task in tasks + ] + + async def task_status(request: Request, inputs: dict[str, Any]) -> dict[str, Any]: + task = await TaskRepository(_database(request).pool).get_by_upid( + str(_values(inputs)["upid"]) + ) + if task is None: + raise ApiError(404, "task does not exist") + result: dict[str, Any] = { + "upid": task.upid, + "status": "stopped" if task.status in {"success", "error", "cancelled"} else "running", + "progress": task.progress, + } + if task.status in {"success", "error", "cancelled"}: + result["exitstatus"] = "OK" if task.status == "success" else task.status.upper() + return result + + async def task_log(request: Request, inputs: dict[str, Any]) -> list[dict[str, Any]]: + repository = TaskRepository(_database(request).pool) + task = await repository.get_by_upid(str(_values(inputs)["upid"])) + if task is None: + raise ApiError(404, "task does not exist") + return [ + {"n": index + 1, "t": message} + for index, message in enumerate(await repository.logs(task.id)) + ] + + registry.register("/nodes/{node}/qemu", "GET", qemu_list) + registry.register("/nodes/{node}/qemu/{vmid}/config", "GET", qemu_config) + registry.register("/nodes/{node}/qemu/{vmid}/status/current", "GET", qemu_status) + registry.register("/nodes/{node}/qemu/{vmid}/status/start", "POST", start) + registry.register("/nodes/{node}/qemu/{vmid}/status/stop", "POST", stop) + registry.register("/nodes/{node}/tasks", "GET", task_list) + registry.register("/nodes/{node}/tasks/{upid}/status", "GET", task_status) + registry.register("/nodes/{node}/tasks/{upid}/log", "GET", task_log) diff --git a/app/main.py b/app/main.py index 2d420d2..3f43ecb 100644 --- a/app/main.py +++ b/app/main.py @@ -2,6 +2,8 @@ from __future__ import annotations +from typing import cast + from fastapi import FastAPI from app.api.errors import ApiError, api_error_handler, unhandled_exception_handler @@ -10,27 +12,49 @@ from app.api.registry import HandlerRegistry, register_contract_routes from app.compatibility import build_report from app.config import Settings, get_settings from app.contracts.model import Snapshot +from app.db.pool import AsyncpgDatabase, Database from app.handlers.core import build_core_handlers from app.lifespan import DatabaseFactory, WorkerFactory, create_lifespan, default_database_factory from app.logging import configure_logging from app.observability.health import router as health_router +from app.simulation.clock import AcceleratedClock +from app.tasks.qemu import qemu_handler +from app.tasks.repository import TaskRepository +from app.tasks.worker import TaskWorker def create_app( settings: Settings | None = None, database_factory: DatabaseFactory = default_database_factory, handlers: HandlerRegistry | None = None, - worker_factories: tuple[WorkerFactory, ...] = (), + worker_factories: tuple[WorkerFactory, ...] | None = None, ) -> FastAPI: """Create an isolated application instance with explicit resource factories.""" resolved = settings or get_settings() configure_logging(resolved.log_level) + resolved_workers = worker_factories + if resolved_workers is None and resolved.contract_snapshot is not None and handlers is None: + + def task_worker(database: Database) -> TaskWorker: + adapter = cast(AsyncpgDatabase, database) + repository = TaskRepository(adapter.pool) + handler = qemu_handler(repository, AcceleratedClock(resolved.simulation_time_scale)) + return TaskWorker( + repository, + "simulator-worker", + {"qemu-start": handler, "qemu-stop": handler}, + concurrency=resolved.task_worker_concurrency, + lease_seconds=resolved.task_lease_seconds, + ) + + resolved_workers = (task_worker,) app = FastAPI( title=resolved.app_name, version="0.0.1", - lifespan=create_lifespan(resolved, database_factory, worker_factories), + lifespan=create_lifespan(resolved, database_factory, resolved_workers or ()), ) + app.state.settings = resolved app.add_middleware(RequestContextMiddleware, header_name=resolved.request_id_header) app.add_exception_handler(Exception, unhandled_exception_handler) app.add_exception_handler(ApiError, api_error_handler) @@ -44,7 +68,10 @@ def create_app( resolved_handlers, resolved.contract_fallback, ) - report = build_report(snapshot, implemented=resolved_handlers.keys()) + declared = frozenset( + (path.path, method.verb) for path in snapshot.paths for method in path.methods + ) + report = build_report(snapshot, implemented=resolved_handlers.keys() & declared) @app.get("/admin/compatibility", include_in_schema=False) async def compatibility_report() -> dict[str, object]: diff --git a/app/simulation/seed.py b/app/simulation/seed.py index d7b4a5a..aaedcab 100644 --- a/app/simulation/seed.py +++ b/app/simulation/seed.py @@ -61,6 +61,13 @@ def small_profile() -> SeedProfile: SeedResource( stable_id("qemu:100"), first.id, "qemu", "100", {"name": "demo", "status": "stopped"} ), + SeedResource( + stable_id("qemu:101"), + first.id, + "qemu", + "101", + {"name": "worker", "status": "stopped"}, + ), SeedResource( stable_id("storage:local"), first.id, "storage", "local", {"content": ["iso", "backup"]} ), diff --git a/app/tasks/qemu.py b/app/tasks/qemu.py new file mode 100644 index 0000000..c03ce56 --- /dev/null +++ b/app/tasks/qemu.py @@ -0,0 +1,43 @@ +"""Worker semantics for asynchronous QEMU transitions.""" + +from __future__ import annotations + +import json +import uuid + +from app.simulation.clock import Clock +from app.simulation.transitions import VmState, plan_transition +from app.tasks.repository import Task, TaskRepository +from app.tasks.worker import TaskHandler + + +def qemu_handler(repository: TaskRepository, clock: Clock) -> TaskHandler: + async def execute(task: Task) -> dict[str, str]: + operation = task.task_type.removeprefix("qemu-") + resource_id = uuid.UUID(str(task.payload["resource_id"])) + async with repository.pool.acquire() as connection: + row = await connection.fetchrow("SELECT state FROM resources WHERE id=$1", resource_id) + if row is None: + raise ValueError("resource disappeared") + raw = row["state"] + state = json.loads(raw) if isinstance(raw, str) else dict(raw) + transition = plan_transition(VmState(str(state["status"])), operation) + state["status"] = transition.intermediate + await connection.execute( + "UPDATE resources SET state=$2::jsonb WHERE id=$1", + resource_id, + json.dumps(state), + ) + await repository.append_log(task.id, f"VM {operation} started") + await clock.sleep(1.0) + async with repository.pool.acquire() as connection: + state["status"] = transition.after + await connection.execute( + "UPDATE resources SET state=$2::jsonb WHERE id=$1", + resource_id, + json.dumps(state), + ) + await repository.append_log(task.id, f"VM {operation} completed") + return {"status": str(transition.after)} + + return execute diff --git a/app/tasks/repository.py b/app/tasks/repository.py index 9a6beab..3a4c0e6 100644 --- a/app/tasks/repository.py +++ b/app/tasks/repository.py @@ -178,6 +178,18 @@ class TaskRepository: 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 get_by_upid(self, upid: str) -> Task | None: + row = await self.pool.fetchrow("SELECT * FROM tasks WHERE upid=$1", upid) + return _task(row) if row is not None else None + + async def list_for_node(self, node: str) -> tuple[Task, ...]: + rows = await self.pool.fetch( + """SELECT * FROM tasks WHERE payload->>'node'=$1 + ORDER BY created_at DESC LIMIT 1000""", + node, + ) + return tuple(_task(row) for row in rows) + 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 diff --git a/tests/unit/test_compatible_io.py b/tests/unit/test_compatible_io.py index d024637..e75d3b7 100644 --- a/tests/unit/test_compatible_io.py +++ b/tests/unit/test_compatible_io.py @@ -11,6 +11,7 @@ from app.api.registry import HandlerRegistry from app.config import Settings from app.contracts.model import Method, Parameter, PathContract, Schema, Snapshot from app.main import create_app +from app.security.auth import csrf_token, issue_ticket from tests.unit.test_health import FakeDatabase @@ -44,9 +45,19 @@ async def client_for(tmp_path: Path) -> AsyncClient: handlers.register("/nodes/{node}/test", "POST", handler) app = create_app( - Settings(contract_snapshot=path), lambda _settings: FakeDatabase(True), handlers + Settings(contract_snapshot=path), + lambda _settings: FakeDatabase(True), + handlers, + worker_factories=(), + ) + key = Settings().ticket_signing_key.get_secret_value().encode() + ticket = issue_ticket("root@pam", key) + return AsyncClient( + transport=ASGITransport(app=app), + base_url="http://test", + cookies={"PVEAuthCookie": ticket}, + headers={"CSRFPreventionToken": csrf_token(ticket, key)}, ) - return AsyncClient(transport=ASGITransport(app=app), base_url="http://test") async def test_json_input_and_null_envelope(tmp_path: Path) -> None: diff --git a/tests/unit/test_core_handlers.py b/tests/unit/test_core_handlers.py index 9bc5af5..bdd14a2 100644 --- a/tests/unit/test_core_handlers.py +++ b/tests/unit/test_core_handlers.py @@ -1,14 +1,17 @@ """First vertical read/login handler tests.""" +import uuid from datetime import UTC, datetime from pathlib import Path +import pytest from httpx import ASGITransport, AsyncClient from app.config import Settings from app.contracts.model import Method, Parameter, PathContract, Schema, Snapshot from app.main import create_app from app.security.auth import hash_secret +from app.tasks.repository import Task class FakePool: @@ -20,12 +23,21 @@ class FakePool: } if "FROM nodes" in sql and args[0] == "pve1": return {"name": "pve1", "status": "online"} + if "FROM resources r" in sql and args == ("pve1", "100"): + if "SELECT r.id" in sql: + return { + "id": uuid.UUID("00000000-0000-0000-0000-000000000100"), + "state": '{"name":"demo","status":"stopped"}', + } + return {"state": '{"name":"demo","status":"stopped"}'} return None async def fetch(self, sql: str, *args: object) -> list[dict[str, object]]: del args if "FROM nodes" in sql: return [{"node": "pve1", "status": "online"}] + if "r.kind='qemu'" in sql: + return [{"vmid": 100, "state": '{"name":"demo","status":"stopped"}'}] return [ { "type": "qemu", @@ -35,6 +47,9 @@ class FakePool: } ] + async def fetchval(self, sql: str) -> int: + return 100 if "pg_backend_pid" in sql else 1_700_000_000 + class FakeDatabase: pool = FakePool() @@ -82,6 +97,36 @@ def write_snapshot(path: Path) -> None: methods=(method("GET", "status", (Parameter(name="node", definition=string),)),), ), PathContract(path="/cluster/resources", methods=(method("GET", "resources"),)), + PathContract( + path="/nodes/{node}/qemu", + methods=(method("GET", "qemu", (Parameter(name="node", definition=string),)),), + ), + PathContract( + path="/nodes/{node}/qemu/{vmid}/config", + methods=( + method( + "GET", + "config", + ( + Parameter(name="node", definition=string), + Parameter(name="vmid", definition=Schema(type="integer")), + ), + ), + ), + ), + PathContract( + path="/nodes/{node}/qemu/{vmid}/status/start", + methods=( + method( + "POST", + "start", + ( + Parameter(name="node", definition=string), + Parameter(name="vmid", definition=Schema(type="integer")), + ), + ), + ), + ), ) snapshot = Snapshot( source_version="test", @@ -89,16 +134,39 @@ def write_snapshot(path: Path) -> None: raw_sha256="0" * 64, paths=paths, path_count=len(paths), - method_count=5, + method_count=sum(len(item.methods) for item in paths), ) path.write_bytes(snapshot.canonical_bytes()) -async def test_core_login_and_read_endpoints(tmp_path: Path) -> None: +async def test_core_login_and_read_endpoints( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + class FakeTaskRepository: + def __init__(self, pool: object) -> None: + del pool + + async def create(self, **kwargs: object) -> Task: + return Task( + uuid.uuid4(), + str(kwargs["upid"]), + str(kwargs["task_type"]), + "queued", + {}, + 0, + False, + 0, + ) + + monkeypatch.setattr("app.handlers.qemu.TaskRepository", FakeTaskRepository) snapshot_path = tmp_path / "snapshot.json" write_snapshot(snapshot_path) database = FakeDatabase() - app = create_app(Settings(contract_snapshot=snapshot_path), lambda _settings: database) + app = create_app( + Settings(contract_snapshot=snapshot_path), + lambda _settings: database, + worker_factories=(), + ) async with app.router.lifespan_context(app): async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: login = await client.post( @@ -106,10 +174,17 @@ async def test_core_login_and_read_endpoints(tmp_path: Path) -> None: content="username=root%40pam&password=secret", headers={"content-type": "application/x-www-form-urlencoded"}, ) + csrf = login.json()["data"]["CSRFPreventionToken"] version = await client.get("/api2/json/version") nodes = await client.get("/api2/json/nodes") status = await client.get("/api2/json/nodes/pve1/status") resources = await client.get("/api2/json/cluster/resources") + qemu = await client.get("/api2/json/nodes/pve1/qemu") + config = await client.get("/api2/json/nodes/pve1/qemu/100/config") + start = await client.post( + "/api2/json/nodes/pve1/qemu/100/status/start", + headers={"CSRFPreventionToken": csrf}, + ) assert login.status_code == 200 assert login.json()["data"]["username"] == "root@pam" @@ -118,3 +193,6 @@ async def test_core_login_and_read_endpoints(tmp_path: Path) -> None: assert nodes.json()["data"][0]["node"] == "pve1" assert status.json()["data"]["status"] == "online" assert resources.json()["data"][0]["type"] == "qemu" + assert qemu.json()["data"][0]["vmid"] == 100 + assert config.json()["data"]["name"] == "demo" + assert start.json()["data"].startswith("UPID:pve1:") diff --git a/tests/unit/test_dynamic_routes.py b/tests/unit/test_dynamic_routes.py index 314e330..2ff86cf 100644 --- a/tests/unit/test_dynamic_routes.py +++ b/tests/unit/test_dynamic_routes.py @@ -48,6 +48,7 @@ async def request_app( settings, lambda _settings: database, handlers if handlers is not None else HandlerRegistry(), + worker_factories=(), ) async with app.router.lifespan_context(app): async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: diff --git a/tests/unit/test_qemu_task.py b/tests/unit/test_qemu_task.py new file mode 100644 index 0000000..ce929d7 --- /dev/null +++ b/tests/unit/test_qemu_task.py @@ -0,0 +1,84 @@ +"""QEMU worker transition semantics.""" + +import uuid +from datetime import UTC, datetime +from typing import cast + +from app.simulation.clock import Clock +from app.tasks.qemu import qemu_handler +from app.tasks.repository import Task, TaskRepository + + +class ImmediateClock: + async def now(self) -> datetime: + return datetime(2026, 1, 1, tzinfo=UTC) + + async def sleep(self, seconds: float) -> None: + assert seconds == 1.0 + + +class Connection: + def __init__(self) -> None: + self.states: list[str] = [] + + async def fetchrow(self, sql: str, resource_id: uuid.UUID) -> dict[str, object]: + del sql, resource_id + return {"state": '{"status":"stopped"}'} + + async def execute(self, sql: str, resource_id: uuid.UUID, state: str) -> str: + del sql, resource_id + self.states.append(state) + return "UPDATE 1" + + +class Acquire: + def __init__(self, connection: Connection) -> None: + self.connection = connection + + async def __aenter__(self) -> Connection: + return self.connection + + async def __aexit__(self, *args: object) -> None: + return None + + +class Pool: + def __init__(self, connection: Connection) -> None: + self.connection = connection + + def acquire(self) -> Acquire: + return Acquire(self.connection) + + +class Repository: + def __init__(self) -> None: + self.connection = Connection() + self.pool = Pool(self.connection) + self.logs: list[str] = [] + + async def append_log(self, task_id: uuid.UUID, message: str) -> None: + del task_id + self.logs.append(message) + + +async def test_qemu_worker_applies_intermediate_and_final_states() -> None: + repository = Repository() + task = Task( + uuid.uuid4(), + "UPID:test", + "qemu-start", + "running", + {"resource_id": str(uuid.uuid4())}, + 0, + False, + 1, + ) + + result = await qemu_handler(cast(TaskRepository, repository), cast(Clock, ImmediateClock()))( + task + ) + + assert result == {"status": "running"} + assert '"starting"' in repository.connection.states[0] + assert '"running"' in repository.connection.states[1] + assert repository.logs == ["VM start started", "VM start completed"] diff --git a/tests/unit/test_seed.py b/tests/unit/test_seed.py index 055018f..7432317 100644 --- a/tests/unit/test_seed.py +++ b/tests/unit/test_seed.py @@ -21,6 +21,12 @@ def test_small_profile_has_stable_logical_state() -> None: "node": "pve1", "state": {"name": "demo", "status": "stopped"}, }, + { + "kind": "qemu", + "external_id": "101", + "node": "pve1", + "state": {"name": "worker", "status": "stopped"}, + }, { "kind": "storage", "external_id": "local",