feat(compute): 实现 Windows Node 注册与心跳
This commit is contained in:
parent
1ac9aa1ba0
commit
109d395356
10
DESIGN.md
10
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/数据兼容。
|
**落地顺序/触发**:先 `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)
|
## 附录:DeepSeek V4 关键事实(2026-04-24)
|
||||||
|
|
|
||||||
14
RUN.md
14
RUN.md
|
|
@ -1039,6 +1039,20 @@ sudo xfs_quota -x -c "limit -p bhard=10g zcbot_<user_uuid>" /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 <node_token>` 和 `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`
|
- **入口**:`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/`
|
- **核心**:`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`
|
- **工具**:`tools/{fs, shell, run_python, skill_tool}.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
|
||||||
|
]
|
||||||
|
|
@ -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):
|
class ExternalSystemDefinition(Base):
|
||||||
"""管理员维护的可信外部系统目录;不含任何用户凭据。"""
|
"""管理员维护的可信外部系统目录;不含任何用户凭据。"""
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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")
|
||||||
|
|
@ -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()
|
||||||
|
|
@ -28,6 +28,7 @@ from fastapi.middleware.cors import CORSMiddleware
|
||||||
|
|
||||||
from core import __version__
|
from core import __version__
|
||||||
|
|
||||||
|
from .admin import register_admin_routes
|
||||||
from .auth import (
|
from .auth import (
|
||||||
REFRESHED_TOKEN_HEADER,
|
REFRESHED_TOKEN_HEADER,
|
||||||
TOKEN_EXPIRES_HEADER,
|
TOKEN_EXPIRES_HEADER,
|
||||||
|
|
@ -35,7 +36,6 @@ from .auth import (
|
||||||
make_require_admin,
|
make_require_admin,
|
||||||
make_require_user,
|
make_require_user,
|
||||||
)
|
)
|
||||||
from .admin import register_admin_routes
|
|
||||||
from .background import (
|
from .background import (
|
||||||
cancel_and_wait,
|
cancel_and_wait,
|
||||||
drain_inflight,
|
drain_inflight,
|
||||||
|
|
@ -49,8 +49,9 @@ from .background import (
|
||||||
from .broker import broker
|
from .broker import broker
|
||||||
from .routers.asr import register_asr_routes
|
from .routers.asr import register_asr_routes
|
||||||
from .routers.authroutes import register_auth_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.external_systems import register_external_system_routes
|
||||||
|
from .routers.files import register_file_routes
|
||||||
from .routers.kb import register_kb_routes
|
from .routers.kb import register_kb_routes
|
||||||
from .routers.messages import register_message_routes
|
from .routers.messages import register_message_routes
|
||||||
from .routers.misc import register_misc_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_asr_routes(app, require_user=require_user, auth_cfg=auth_cfg)
|
||||||
register_task_routes(app, require_user=require_user)
|
register_task_routes(app, require_user=require_user)
|
||||||
register_message_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)─────────────
|
# ───────────── 管理后台(admin-only)─────────────
|
||||||
register_admin_routes(app, require_admin)
|
register_admin_routes(app, require_admin)
|
||||||
|
|
|
||||||
|
|
@ -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",
|
||||||
|
}
|
||||||
|
|
@ -107,3 +107,22 @@ class ExternalSystemCreateRequest(BaseModel):
|
||||||
|
|
||||||
class ExternalSystemCredentialsRequest(BaseModel):
|
class ExternalSystemCredentialsRequest(BaseModel):
|
||||||
credentials: dict[str, str] = Field(default_factory=dict)
|
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
|
||||||
|
|
|
||||||
Loading…
Reference in New Issue