"""Remote workers: how one attaches, and what is attached right now. A worker dials in rather than being dialled: the GPU box and the engine are usually on different networks, and only one of them can be reached. It presents a token minted here, says what it can do, and then answers calls on the socket it opened. """ from __future__ import annotations import asyncio import logging from datetime import timedelta from typing import Any from fastapi import APIRouter, Depends, HTTPException, Request, WebSocket from fastapi.responses import PlainTextResponse from fluksio_worker import worker_main from jwt.exceptions import InvalidTokenError from pydantic import BaseModel, Field from fluksio.api.deps import get_current_active_superuser, get_current_user from fluksio.core import security from fluksio.flow.remote import PROTOCOL, RemoteWorker, RemoteWorkerHub logger = logging.getLogger(__name__) router = APIRouter(prefix="/workers", tags=["workers"]) #: Long, because a worker is a machine somebody set up once and left running. TOKEN_DAYS = 365 class WorkerInfo(BaseModel): name: str labels: list[str] = Field(default_factory=list) max_parallel: int = 1 in_flight: int = 0 attached_at: float = 0.0 last_seen: float = 0.0 python: str = "" venv_digest: str = "" class TokenRequest(BaseModel): name: str class TokenIssued(BaseModel): name: str token: str expires_days: int = TOKEN_DAYS def _hub(app: Any) -> RemoteWorkerHub: hub: RemoteWorkerHub | None = getattr(app.state, "worker_hub", None) if hub is None: raise HTTPException(status_code=503, detail="Remote workers are not available") return hub @router.get( "", response_model=list[WorkerInfo], dependencies=[Depends(get_current_user)] ) def read_workers(request: Request) -> Any: """What is attached, and how busy it is.""" return [ WorkerInfo( name=worker.name, labels=sorted(worker.labels), max_parallel=worker.max_parallel, in_flight=worker.in_flight, attached_at=worker.attached_at, last_seen=worker.last_seen, python=str(worker.info.get("python") or ""), venv_digest=str(worker.info.get("venv_digest") or ""), ) for worker in _hub(request.app).workers() ] @router.post( "/tokens", response_model=TokenIssued, dependencies=[Depends(get_current_active_superuser)], ) def issue_token(body: TokenRequest) -> Any: """Mint the credential a worker presents when it dials in. Shown once. It is signed with the same keypair the agent tokens use, so rotating that key revokes every worker along with them. """ token = security.create_worker_token(body.name, timedelta(days=TOKEN_DAYS)) return TokenIssued(name=body.name, token=token) @router.get( "/runtime", response_class=PlainTextResponse, dependencies=[Depends(get_current_user)], ) def read_runtime() -> str: """The worker's own code, so a fresh host installs by fetching one file. It is the same module the engine's local workers run — deliberately standard library only, and with nothing of the engine importable in it. """ return worker_main.__file__ and open(worker_main.__file__).read() @router.websocket("/attach") async def attach(websocket: WebSocket, token: str = "") -> None: """A worker's connection, for as long as it holds. The token goes in the query string for the same reason the dashboard's does: a websocket handshake carries no headers of its own. """ try: claims = security.decode_worker_token(token) except InvalidTokenError: await websocket.close(code=1008) return await websocket.accept() try: hello = await asyncio.wait_for(websocket.receive_json(), timeout=30) except (TimeoutError, asyncio.TimeoutError, ValueError): await websocket.close(code=1002) return if hello.get("op") != "hello" or int(hello.get("protocol", 0)) != PROTOCOL: await websocket.send_json( {"op": "refused", "reason": f"this engine speaks protocol {PROTOCOL}"} ) await websocket.close(code=1002) return # The token names the worker; what it calls itself is a suggestion, so two # hosts cannot fight over one identity by claiming the same name. name = str(claims.get("sub") or hello.get("name") or "worker") hub = _hub(websocket.app) worker = RemoteWorker( name=name, labels=[str(label) for label in (hello.get("labels") or [])], send=websocket.send_json, loop=asyncio.get_running_loop(), max_parallel=max(1, int(hello.get("max_parallel") or 1)), info={ "python": hello.get("python"), "venv_digest": hello.get("venv_digest"), }, ) hub.attach(worker) await websocket.send_json({"op": "welcome", "protocol": PROTOCOL, "name": name}) try: while True: message = await websocket.receive_json() worker.deliver(message) except Exception: # Any way this ends is the same thing: the socket is gone, and whatever # was waiting on it has to be told rather than left hanging. logger.info("Worker '%s' disconnected", name) finally: hub.detach(name)