147 lines
5.3 KiB
Python
147 lines
5.3 KiB
Python
"""Windows Node MVP 的注册、管理与长连接端点。"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import asyncio
|
|
from uuid import UUID
|
|
|
|
from fastapi import Depends, HTTPException, WebSocket, WebSocketDisconnect, status
|
|
|
|
from core.compute_nodes import (
|
|
ComputeNodeError,
|
|
authenticate_node,
|
|
create_enrollment,
|
|
enroll_node,
|
|
list_nodes,
|
|
mark_node_offline,
|
|
set_node_disabled,
|
|
update_node_runtime,
|
|
)
|
|
from web.schemas import (
|
|
ComputeEnrollmentCreateRequest,
|
|
ComputeNodeDisableRequest,
|
|
ComputeNodeEnrollRequest,
|
|
)
|
|
|
|
|
|
class NodeConnectionManager:
|
|
def __init__(self) -> None:
|
|
self._connections: dict[UUID, WebSocket] = {}
|
|
self._lock = asyncio.Lock()
|
|
|
|
async def activate(self, node_id: UUID, websocket: WebSocket) -> None:
|
|
async with self._lock:
|
|
old = self._connections.get(node_id)
|
|
self._connections[node_id] = websocket
|
|
if old is not None and old is not websocket:
|
|
await old.close(code=4001, reason="replaced by a newer connection")
|
|
|
|
async def remove(self, node_id: UUID, websocket: WebSocket) -> bool:
|
|
async with self._lock:
|
|
if self._connections.get(node_id) is websocket:
|
|
self._connections.pop(node_id, None)
|
|
return True
|
|
return False
|
|
|
|
async def close(self, node_id: UUID) -> None:
|
|
async with self._lock:
|
|
websocket = self._connections.pop(node_id, None)
|
|
if websocket is not None:
|
|
await websocket.close(code=4003, reason="node disabled")
|
|
|
|
|
|
node_connections = NodeConnectionManager()
|
|
|
|
|
|
def _bearer(authorization: str | None) -> str:
|
|
scheme, _, token = (authorization or "").partition(" ")
|
|
if scheme.lower() != "bearer" or not token:
|
|
raise ComputeNodeError("missing node bearer token")
|
|
return token
|
|
|
|
|
|
def register_compute_node_routes(app, *, require_admin) -> None:
|
|
@app.post(
|
|
"/v1/compute/nodes/enroll",
|
|
tags=["compute-nodes"],
|
|
status_code=status.HTTP_201_CREATED,
|
|
)
|
|
def node_enroll(body: ComputeNodeEnrollRequest):
|
|
try:
|
|
return enroll_node(**body.model_dump())
|
|
except ComputeNodeError as exc:
|
|
raise HTTPException(400, str(exc)) from exc
|
|
|
|
@app.websocket("/v1/compute/nodes/connect")
|
|
async def node_connect(websocket: WebSocket):
|
|
try:
|
|
node_id = UUID(websocket.headers.get("x-node-id", ""))
|
|
token = _bearer(websocket.headers.get("authorization"))
|
|
identity = await asyncio.to_thread(authenticate_node, node_id, token)
|
|
except (ValueError, ComputeNodeError):
|
|
await websocket.close(code=1008, reason="invalid node credentials")
|
|
return
|
|
await websocket.accept()
|
|
await node_connections.activate(node_id, websocket)
|
|
try:
|
|
await websocket.send_json({"type": "connected", "heartbeat_seconds": 15})
|
|
while True:
|
|
message = await websocket.receive_json()
|
|
message_type = message.get("type")
|
|
payload = message.get("payload") or {}
|
|
if message_type not in {"hello", "heartbeat"} or not isinstance(
|
|
payload, dict
|
|
):
|
|
await websocket.send_json(
|
|
{"type": "error", "code": "unsupported_message"}
|
|
)
|
|
continue
|
|
if payload.get("install_id") and payload["install_id"] != str(
|
|
identity["install_id"]
|
|
):
|
|
await websocket.close(code=1008, reason="install identity mismatch")
|
|
return
|
|
await asyncio.to_thread(
|
|
update_node_runtime,
|
|
node_id,
|
|
status="online",
|
|
runtime=payload,
|
|
)
|
|
await websocket.send_json(
|
|
{"type": "ack", "message_id": message.get("message_id")}
|
|
)
|
|
except (ComputeNodeError, WebSocketDisconnect, RuntimeError, ValueError):
|
|
pass
|
|
finally:
|
|
if await node_connections.remove(node_id, websocket):
|
|
await asyncio.to_thread(mark_node_offline, node_id)
|
|
|
|
@app.post("/v1/admin/compute-node-enrollments", tags=["admin"])
|
|
def admin_create_compute_enrollment(
|
|
body: ComputeEnrollmentCreateRequest,
|
|
user_id: UUID = Depends(require_admin), # noqa: B008
|
|
):
|
|
try:
|
|
return create_enrollment(user_id, **body.model_dump())
|
|
except ComputeNodeError as exc:
|
|
raise HTTPException(400, str(exc)) from exc
|
|
|
|
@app.get("/v1/admin/compute-nodes", tags=["admin"])
|
|
def admin_compute_nodes(user_id: UUID = Depends(require_admin)): # noqa: B008
|
|
return {"results": list_nodes()}
|
|
|
|
@app.patch("/v1/admin/compute-nodes/{node_id}", tags=["admin"])
|
|
async def admin_disable_compute_node(
|
|
node_id: UUID,
|
|
body: ComputeNodeDisableRequest,
|
|
user_id: UUID = Depends(require_admin), # noqa: B008
|
|
):
|
|
if not await asyncio.to_thread(set_node_disabled, node_id, body.disabled):
|
|
raise HTTPException(404, "compute node not found")
|
|
if body.disabled:
|
|
await node_connections.close(node_id)
|
|
return {
|
|
"node_id": str(node_id),
|
|
"status": "disabled" if body.disabled else "offline",
|
|
}
|