212 lines
7.1 KiB
Python
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
|
|
]
|