zcbot/core/software_nodes.py

212 lines
7.1 KiB
Python

"""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.software_contracts import DEFAULT_CAPABILITIES, SUPPORTED_CAPABILITIES
from core.storage.engine import session_scope
from core.storage.models import SoftwareNode, SoftwareNodeEnrollment
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 DEFAULT_CAPABILITIES))
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
]