feat: add basic QEMU and task vertical slice
This commit is contained in:
+32
-5
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
+30
-3
@@ -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]:
|
||||
|
||||
@@ -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"]}
|
||||
),
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user