"""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 SoftwareNode, SoftwareNodeEnrollment SUPPORTED_CAPABILITIES = frozenset({"origin.plot@v2"}) MAX_ENROLLMENT_FAILURES = 5 class SoftwareNodeError(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@v2"])) if not allowed or any(item not in SUPPORTED_CAPABILITIES for item in allowed): raise SoftwareNodeError("unsupported capability") if not 60 <= ttl_seconds <= 3600: raise SoftwareNodeError("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 = SoftwareNodeEnrollment( 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 SoftwareNodeError("node_name and capabilities are required") now = datetime.now(timezone.utc) error: str | None = None with session_scope() as session: enrollment = session.execute( select(SoftwareNodeEnrollment) .where( SoftwareNodeEnrollment.code_hash == _enrollment_digest(enrollment_code), SoftwareNodeEnrollment.consumed_at.is_(None), SoftwareNodeEnrollment.expires_at > now, SoftwareNodeEnrollment.failed_attempts < MAX_ENROLLMENT_FAILURES, ) .with_for_update() ).scalar_one_or_none() if enrollment is None: raise SoftwareNodeError("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(SoftwareNode.node_id).where(SoftwareNode.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( SoftwareNode( 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 SoftwareNodeError(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(SoftwareNode, node_id) if ( node is None or node.status == "disabled" or not _verify_secret(token, node.token_hash) ): raise SoftwareNodeError("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(SoftwareNode, node_id) if node is None or node.status == "disabled": raise SoftwareNodeError("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(SoftwareNode, 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(SoftwareNode, node_id) if node is None: return False node.status = "disabled" if disabled else "offline" return True def delete_node(node_id: UUID) -> bool: """撤销并物理删除节点身份;当前节点表没有任务历史外键。""" with session_scope() as session: node = session.get(SoftwareNode, node_id) if node is None: return False session.delete(node) return True def list_nodes() -> list[dict]: with session_scope() as session: rows = ( session.execute(select(SoftwareNode).order_by(SoftwareNode.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 ]