From 109d395356f8add93bbfd5ff4dc1bd25801cc8bb Mon Sep 17 00:00:00 2001 From: caoqianming Date: Wed, 12 Aug 2026 15:56:28 +0800 Subject: [PATCH] =?UTF-8?q?feat(compute):=20=E5=AE=9E=E7=8E=B0=20Windows?= =?UTF-8?q?=20Node=20=E6=B3=A8=E5=86=8C=E4=B8=8E=E5=BF=83=E8=B7=B3?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- DESIGN.md | 10 + RUN.md | 14 ++ core/compute_nodes.py | 201 ++++++++++++++++++ core/storage/models.py | 41 ++++ .../20260812_2000_0030_compute_nodes.py | 76 +++++++ tests/test_compute_nodes.py | 85 ++++++++ web/app.py | 6 +- web/routers/compute_nodes.py | 146 +++++++++++++ web/schemas.py | 19 ++ 9 files changed, 596 insertions(+), 2 deletions(-) create mode 100644 core/compute_nodes.py create mode 100644 db/migrations/versions/20260812_2000_0030_compute_nodes.py create mode 100644 tests/test_compute_nodes.py create mode 100644 web/routers/compute_nodes.py diff --git a/DESIGN.md b/DESIGN.md index c1295c1..38e5ac2 100644 --- a/DESIGN.md +++ b/DESIGN.md @@ -454,6 +454,16 @@ scheduled_jobs(§8.5) channel_bindings(§8.7,判别列+JSONB) **落地顺序/触发**:先 `ActionPolicy -> Attention Inbox(Web) -> 渠道响应 -> action audit`,四者作为一个完整外部写安全闭环;再按真实无人值守写需求增加 exact-target standing rule,按真实长计算中间步骤增加 self-wake;有明确第三方 MCP 接入对象后再做 adapter。任何阶段都不得先开放外部写、再用 prompt 要求模型“记得询问”补安全边界。实现时需回写 §3.1 loop、§7.5 sandbox/tool registry、§8.5 scheduler、§8.7 channel、§8.13 bg proc 与 §8.14 external systems 的最终契约,并以 migration 保持现有 API/数据兼容。 +### 8.16 Windows Node 内网 MVP(implementation,2026-08-12) + +第一阶段以 `docs/windows-node-mvp-intranet.md` 为实现契约:Windows Node 只作为受控执行节点,通过出站 HTTP/WS 主动连接 zcbot;首批能力固定为 `origin.plot@v1`。长期方案中的 mTLS、Service/DesktopRunner 双进程、完整租约与多节点调度暂不进入 MVP,但 URL path、Node ID、Bearer Header 和任务协议保留原位升级空间。 + +云端控制面使用独立的 `compute_node_enrollments` 与 `compute_nodes`,不复用用户外部系统连接。管理员创建的一次性注册码具有 128 bit 随机熵,数据库只保存 SHA-256 摘要;节点注册在行锁事务中校验有效期、预期名称和允许能力,成功后原子消费。每个节点获得独立高熵 Token,数据库只保存 bcrypt 强哈希,明文仅在注册响应出现一次。 + +Node 通过 `Authorization: Bearer` 与 `X-Node-Id` 建立 `/v1/compute/nodes/connect` WebSocket。进程内 Connection Manager 保证同一节点单活,新连接关闭旧连接;`hello`/`heartbeat` 更新版本、容量、软件健康与最后在线时间。管理员禁用节点时先持久化禁用态,再关闭现有连接;断线收尾不得覆盖禁用态。当前单活只覆盖单 Web 进程,生产启用多实例前必须增加 Redis/PG fencing 或将 Node API 固定路由到单一控制面实例。 + +首批只落注册、认证、心跳、状态与禁用基础链路。`compute_jobs`、任务 offer/accept、Origin Worker、输入输出传输、重连对账和 Token 轮换属于后续垂直闭环,不以任意命令或脚本接口临时代替。 + --- ## 附录:DeepSeek V4 关键事实(2026-04-24) diff --git a/RUN.md b/RUN.md index 89e300d..2c14505 100644 --- a/RUN.md +++ b/RUN.md @@ -1039,6 +1039,20 @@ sudo xfs_quota -x -c "limit -p bhard=10g zcbot_" /opt ## 关键路径与文件 +### Windows Node 内网 MVP(开发中) + +先执行 `alembic upgrade head` 创建 `compute_node_enrollments` 和 `compute_nodes`。不要在未确认目标数据库时运行迁移;本机 `.env` 的 `ZCBOT_DB_URL` 可能是生产隧道。 + +云端当前提供: + +- 管理员 `POST /v1/admin/compute-node-enrollments` 创建一次性注册码; +- Node `POST /v1/compute/nodes/enroll` 注册并一次性取得 `node_id`、`node_token`; +- Node 携带 `Authorization: Bearer ` 和 `X-Node-Id` 连接 `WS /v1/compute/nodes/connect`; +- 管理员 `GET /v1/admin/compute-nodes` 查看节点,`PATCH /v1/admin/compute-nodes/{node_id}` 启停节点。 + +Node API 只能绑定受控内网地址并由安全组限制来源 IP。当前 HTTP/WS 链路不加密;跨安全域、公网或不可信终端接入前,必须先升级 HTTPS/WSS。多 Web 实例部署时,Node API 暂时固定路由到单一实例,直至 Connection Manager 增加跨实例 fencing。 + + - **入口**:`main.py`(`web / db / probe / user`)→ `core/agent_builder.py::build_agent` - **核心**:`core/{agent_builder, loop, session, task, llm, memory, paths}.py` + `core/storage/{engine,models,utils}.py` + `db/migrations/` - **工具**:`tools/{fs, shell, run_python, skill_tool}.py` diff --git a/core/compute_nodes.py b/core/compute_nodes.py new file mode 100644 index 0000000..197cb46 --- /dev/null +++ b/core/compute_nodes.py @@ -0,0 +1,201 @@ +"""Windows Node MVP 的注册、认证与状态持久化。""" + +from __future__ import annotations + +import secrets +from datetime import datetime, timedelta, timezone +from hashlib import sha256 +from uuid import UUID, uuid4 + +import bcrypt +from sqlalchemy import select + +from core.storage.engine import session_scope +from core.storage.models import ComputeNode, ComputeNodeEnrollment + +SUPPORTED_CAPABILITIES = frozenset({"origin.plot@v1"}) +MAX_ENROLLMENT_FAILURES = 5 + + +class ComputeNodeError(Exception): + pass + + +def _hash_secret(value: str) -> str: + return bcrypt.hashpw(value.encode("utf-8"), bcrypt.gensalt()).decode("ascii") + + +def _verify_secret(value: str, digest: str) -> bool: + try: + return bcrypt.checkpw(value.encode("utf-8"), digest.encode("ascii")) + except (ValueError, TypeError, UnicodeError): + return False + + +def _enrollment_digest(value: str) -> str: + """注册码本身有 128 bit 熵;定长摘要允许数据库精确查找和行锁。""" + return sha256(value.encode("ascii", errors="ignore")).hexdigest() + + +def create_enrollment( + created_by: UUID, + *, + expected_name: str = "", + capabilities: list[str] | None = None, + ttl_seconds: int = 600, +) -> dict: + allowed = list(dict.fromkeys(capabilities or ["origin.plot@v1"])) + if not allowed or any(item not in SUPPORTED_CAPABILITIES for item in allowed): + raise ComputeNodeError("unsupported capability") + if not 60 <= ttl_seconds <= 3600: + raise ComputeNodeError("ttl_seconds must be between 60 and 3600") + code = "ZCN-" + secrets.token_hex(16).upper() + expires_at = datetime.now(timezone.utc) + timedelta(seconds=ttl_seconds) + row = ComputeNodeEnrollment( + enrollment_id=uuid4(), + code_hash=_enrollment_digest(code), + expected_name=(expected_name or "").strip() or None, + allowed_capabilities=allowed, + expires_at=expires_at, + created_by=created_by, + ) + with session_scope() as session: + session.add(row) + return { + "enrollment_code": code, + "expires_at": expires_at.isoformat(), + "capabilities": allowed, + } + + +def enroll_node( + *, + enrollment_code: str, + node_name: str, + install_id: UUID, + node_version: str, + os_version: str, + capabilities: list[str], +) -> dict: + name = node_name.strip() + requested = list(dict.fromkeys(capabilities)) + if not name or not requested: + raise ComputeNodeError("node_name and capabilities are required") + now = datetime.now(timezone.utc) + error: str | None = None + with session_scope() as session: + enrollment = session.execute( + select(ComputeNodeEnrollment) + .where( + ComputeNodeEnrollment.code_hash == _enrollment_digest(enrollment_code), + ComputeNodeEnrollment.consumed_at.is_(None), + ComputeNodeEnrollment.expires_at > now, + ComputeNodeEnrollment.failed_attempts < MAX_ENROLLMENT_FAILURES, + ) + .with_for_update() + ).scalar_one_or_none() + if enrollment is None: + raise ComputeNodeError("invalid or expired enrollment code") + enrollment.failed_attempts += 1 + if enrollment.expected_name and enrollment.expected_name != name: + error = "node name does not match enrollment" + elif any(item not in enrollment.allowed_capabilities for item in requested): + error = "capability is not allowed by enrollment" + else: + existing = session.execute( + select(ComputeNode.node_id).where(ComputeNode.install_id == install_id) + ).first() + if existing is not None: + error = "install is already enrolled" + else: + token = secrets.token_urlsafe(48) + node_id = uuid4() + session.add( + ComputeNode( + node_id=node_id, + name=name, + install_id=install_id, + token_hash=_hash_secret(token), + status="offline", + capabilities=requested, + node_version=node_version.strip(), + os_version=os_version.strip(), + runtime={}, + ) + ) + enrollment.consumed_at = now + if error is not None: + raise ComputeNodeError(error) + return { + "node_id": str(node_id), + "node_token": token, + "heartbeat_seconds": 15, + "max_concurrency": 1, + } + + +def authenticate_node(node_id: UUID, token: str) -> dict: + with session_scope() as session: + node = session.get(ComputeNode, node_id) + if ( + node is None + or node.status == "disabled" + or not _verify_secret(token, node.token_hash) + ): + raise ComputeNodeError("invalid node credentials") + return { + "node_id": node.node_id, + "install_id": node.install_id, + "status": node.status, + } + + +def update_node_runtime(node_id: UUID, *, status: str, runtime: dict) -> None: + with session_scope() as session: + node = session.get(ComputeNode, node_id) + if node is None or node.status == "disabled": + raise ComputeNodeError("node is disabled or missing") + node.status = status + node.runtime = runtime + node.last_seen_at = datetime.now(timezone.utc) + + +def mark_node_offline(node_id: UUID) -> None: + """仅把活动节点转离线;管理员禁用态不可被断线收尾覆盖。""" + with session_scope() as session: + node = session.get(ComputeNode, node_id) + if node is not None and node.status != "disabled": + node.status = "offline" + + +def set_node_disabled(node_id: UUID, disabled: bool) -> bool: + with session_scope() as session: + node = session.get(ComputeNode, node_id) + if node is None: + return False + node.status = "disabled" if disabled else "offline" + return True + + +def list_nodes() -> list[dict]: + with session_scope() as session: + rows = ( + session.execute(select(ComputeNode).order_by(ComputeNode.created_at)) + .scalars() + .all() + ) + return [ + { + "node_id": str(row.node_id), + "name": row.name, + "status": row.status, + "capabilities": row.capabilities, + "node_version": row.node_version, + "os_version": row.os_version, + "runtime": row.runtime, + "last_seen_at": row.last_seen_at.isoformat() + if row.last_seen_at + else None, + } + for row in rows + ] diff --git a/core/storage/models.py b/core/storage/models.py index a075f56..f01d2ee 100644 --- a/core/storage/models.py +++ b/core/storage/models.py @@ -420,6 +420,47 @@ class ChannelBinding(Base): ) +class ComputeNodeEnrollment(Base): + """Windows Node 一次性注册码;数据库只保存不可逆摘要。""" + + __tablename__ = "compute_node_enrollments" + enrollment_id: Mapped[UUID] = mapped_column(PG_UUID(as_uuid=True), primary_key=True, default=uuid4) + code_hash: Mapped[str] = mapped_column(Text, nullable=False, unique=True) + expected_name: Mapped[Optional[str]] = mapped_column(Text, nullable=True) + allowed_capabilities: Mapped[list[str]] = mapped_column(JSONB, nullable=False, default=list) + expires_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False) + failed_attempts: Mapped[int] = mapped_column(Integer, nullable=False, default=0, server_default="0") + consumed_at: Mapped[Optional[datetime]] = mapped_column(DateTime(timezone=True), nullable=True) + created_by: Mapped[Optional[UUID]] = mapped_column( + PG_UUID(as_uuid=True), ForeignKey("users.user_id", ondelete="SET NULL"), nullable=True + ) + created_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), server_default=func.now(), nullable=False + ) + + +class ComputeNode(Base): + """平台托管的 Windows 执行节点身份与最后一次运行态。""" + + __tablename__ = "compute_nodes" + node_id: Mapped[UUID] = mapped_column(PG_UUID(as_uuid=True), primary_key=True, default=uuid4) + name: Mapped[str] = mapped_column(Text, nullable=False) + install_id: Mapped[UUID] = mapped_column(PG_UUID(as_uuid=True), nullable=False, unique=True) + token_hash: Mapped[str] = mapped_column(Text, nullable=False) + status: Mapped[str] = mapped_column(Text, nullable=False, default="offline", server_default="offline") + capabilities: Mapped[list[Any]] = mapped_column(JSONB, nullable=False, default=list) + node_version: Mapped[str] = mapped_column(Text, nullable=False, server_default="") + os_version: Mapped[str] = mapped_column(Text, nullable=False, server_default="") + runtime: Mapped[dict[str, Any]] = mapped_column(JSONB, nullable=False, default=dict) + last_seen_at: Mapped[Optional[datetime]] = mapped_column(DateTime(timezone=True), nullable=True) + created_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), server_default=func.now(), nullable=False + ) + updated_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), server_default=func.now(), onupdate=func.now(), nullable=False + ) + + class ExternalSystemDefinition(Base): """管理员维护的可信外部系统目录;不含任何用户凭据。""" diff --git a/db/migrations/versions/20260812_2000_0030_compute_nodes.py b/db/migrations/versions/20260812_2000_0030_compute_nodes.py new file mode 100644 index 0000000..02193f6 --- /dev/null +++ b/db/migrations/versions/20260812_2000_0030_compute_nodes.py @@ -0,0 +1,76 @@ +"""Add Windows compute node enrollment and registry tables. + +Revision ID: 0030 +Revises: 0029 +Create Date: 2026-08-12 +""" + +from collections.abc import Sequence + +import sqlalchemy as sa +from alembic import op +from sqlalchemy.dialects import postgresql + +revision: str = "0030" +down_revision: str | None = "0029" +branch_labels: str | Sequence[str] | None = None +depends_on: str | Sequence[str] | None = None + + +def upgrade() -> None: + op.create_table( + "compute_node_enrollments", + sa.Column("enrollment_id", postgresql.UUID(as_uuid=True), primary_key=True), + sa.Column("code_hash", sa.Text(), nullable=False, unique=True), + sa.Column("expected_name", sa.Text(), nullable=True), + sa.Column("allowed_capabilities", postgresql.JSONB(), nullable=False), + sa.Column("expires_at", sa.DateTime(timezone=True), nullable=False), + sa.Column("failed_attempts", sa.Integer(), server_default="0", nullable=False), + sa.Column("consumed_at", sa.DateTime(timezone=True), nullable=True), + sa.Column( + "created_by", + postgresql.UUID(as_uuid=True), + sa.ForeignKey("users.user_id", ondelete="SET NULL"), + nullable=True, + ), + sa.Column( + "created_at", + sa.DateTime(timezone=True), + server_default=sa.func.now(), + nullable=False, + ), + ) + op.create_table( + "compute_nodes", + sa.Column("node_id", postgresql.UUID(as_uuid=True), primary_key=True), + sa.Column("name", sa.Text(), nullable=False), + sa.Column( + "install_id", postgresql.UUID(as_uuid=True), nullable=False, unique=True + ), + sa.Column("token_hash", sa.Text(), nullable=False), + sa.Column("status", sa.Text(), server_default="offline", nullable=False), + sa.Column("capabilities", postgresql.JSONB(), nullable=False), + sa.Column("node_version", sa.Text(), server_default="", nullable=False), + sa.Column("os_version", sa.Text(), server_default="", nullable=False), + sa.Column("runtime", postgresql.JSONB(), nullable=False), + sa.Column("last_seen_at", sa.DateTime(timezone=True), nullable=True), + sa.Column( + "created_at", + sa.DateTime(timezone=True), + server_default=sa.func.now(), + nullable=False, + ), + sa.Column( + "updated_at", + sa.DateTime(timezone=True), + server_default=sa.func.now(), + nullable=False, + ), + ) + op.create_index("ix_compute_nodes_status", "compute_nodes", ["status"]) + + +def downgrade() -> None: + op.drop_index("ix_compute_nodes_status", table_name="compute_nodes") + op.drop_table("compute_nodes") + op.drop_table("compute_node_enrollments") diff --git a/tests/test_compute_nodes.py b/tests/test_compute_nodes.py new file mode 100644 index 0000000..505b8a2 --- /dev/null +++ b/tests/test_compute_nodes.py @@ -0,0 +1,85 @@ +from __future__ import annotations + +import importlib +import unittest +from unittest.mock import AsyncMock, patch +from uuid import uuid4 + +from alembic.migration import MigrationContext +from alembic.operations import Operations +from sqlalchemy import create_mock_engine +from sqlalchemy.dialects import postgresql + +from core.compute_nodes import _enrollment_digest, _hash_secret, _verify_secret +from web.routers.compute_nodes import NodeConnectionManager, _bearer + + +class ComputeNodeSecurityTests(unittest.TestCase): + def test_secret_hash_is_salted_and_verifiable(self) -> None: + first = _hash_secret("node-secret") + second = _hash_secret("node-secret") + self.assertNotEqual(first, second) + self.assertNotIn("node-secret", first) + self.assertTrue(_verify_secret("node-secret", first)) + self.assertFalse(_verify_secret("wrong", first)) + + def test_bearer_parser_rejects_query_style_or_missing_token(self) -> None: + self.assertEqual(_bearer("Bearer abc"), "abc") + with self.assertRaisesRegex(Exception, "missing node bearer token"): + _bearer(None) + + def test_enrollment_digest_does_not_store_plaintext(self) -> None: + digest = _enrollment_digest("ZCN-ABC") + self.assertEqual(len(digest), 64) + self.assertNotIn("ZCN-ABC", digest) + + +class ComputeNodeConnectionTests(unittest.IsolatedAsyncioTestCase): + async def test_new_connection_replaces_old_without_removing_new(self) -> None: + manager = NodeConnectionManager() + node_id = uuid4() + old = AsyncMock() + new = AsyncMock() + + await manager.activate(node_id, old) + await manager.activate(node_id, new) + + old.close.assert_awaited_once_with( + code=4001, reason="replaced by a newer connection" + ) + self.assertFalse(await manager.remove(node_id, old)) + self.assertTrue(await manager.remove(node_id, new)) + + async def test_admin_close_removes_and_closes_connection(self) -> None: + manager = NodeConnectionManager() + node_id = uuid4() + websocket = AsyncMock() + await manager.activate(node_id, websocket) + await manager.close(node_id) + websocket.close.assert_awaited_once_with(code=4003, reason="node disabled") + self.assertFalse(await manager.remove(node_id, websocket)) + + +class ComputeNodeMigrationTests(unittest.TestCase): + def test_0030_upgrade_compiles_as_postgresql_ddl(self) -> None: + statements: list[str] = [] + + def capture(sql, *multiparams, **params): + statements.append(str(sql.compile(dialect=postgresql.dialect()))) + + engine = create_mock_engine("postgresql+psycopg://", capture) + operations = Operations(MigrationContext.configure(engine.connect())) + migration = importlib.import_module( + "db.migrations.versions.20260812_2000_0030_compute_nodes" + ) + with patch.object(migration, "op", operations): + migration.upgrade() + + rendered = "\n".join(statements) + self.assertIn("compute_node_enrollments", rendered) + self.assertIn("compute_nodes", rendered) + self.assertIn("ix_compute_nodes_status", rendered) + + +if __name__ == "__main__": + unittest.main() diff --git a/web/app.py b/web/app.py index e96d778..1f2c988 100644 --- a/web/app.py +++ b/web/app.py @@ -28,6 +28,7 @@ from fastapi.middleware.cors import CORSMiddleware from core import __version__ +from .admin import register_admin_routes from .auth import ( REFRESHED_TOKEN_HEADER, TOKEN_EXPIRES_HEADER, @@ -35,7 +36,6 @@ from .auth import ( make_require_admin, make_require_user, ) -from .admin import register_admin_routes from .background import ( cancel_and_wait, drain_inflight, @@ -49,8 +49,9 @@ from .background import ( from .broker import broker from .routers.asr import register_asr_routes from .routers.authroutes import register_auth_routes -from .routers.files import register_file_routes +from .routers.compute_nodes import register_compute_node_routes from .routers.external_systems import register_external_system_routes +from .routers.files import register_file_routes from .routers.kb import register_kb_routes from .routers.messages import register_message_routes from .routers.misc import register_misc_routes @@ -205,6 +206,7 @@ def create_app() -> FastAPI: register_asr_routes(app, require_user=require_user, auth_cfg=auth_cfg) register_task_routes(app, require_user=require_user) register_message_routes(app, require_user=require_user) + register_compute_node_routes(app, require_admin=require_admin) # ───────────── 管理后台(admin-only)───────────── register_admin_routes(app, require_admin) diff --git a/web/routers/compute_nodes.py b/web/routers/compute_nodes.py new file mode 100644 index 0000000..4a29b06 --- /dev/null +++ b/web/routers/compute_nodes.py @@ -0,0 +1,146 @@ +"""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", + } diff --git a/web/schemas.py b/web/schemas.py index 61da74b..28d53b6 100644 --- a/web/schemas.py +++ b/web/schemas.py @@ -107,3 +107,22 @@ class ExternalSystemCreateRequest(BaseModel): class ExternalSystemCredentialsRequest(BaseModel): credentials: dict[str, str] = Field(default_factory=dict) + + +class ComputeEnrollmentCreateRequest(BaseModel): + expected_name: str = "" + capabilities: list[str] = Field(default_factory=lambda: ["origin.plot@v1"]) + ttl_seconds: int = 600 + + +class ComputeNodeEnrollRequest(BaseModel): + enrollment_code: str + node_name: str + install_id: UUID + node_version: str = "" + os_version: str = "" + capabilities: list[str] + + +class ComputeNodeDisableRequest(BaseModel): + disabled: bool = True