zcbot/tests/test_compute_nodes.py

123 lines
4.5 KiB
Python

from __future__ import annotations
import importlib
import unittest
from pathlib import Path
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,
delete_node,
)
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)
def test_websocket_auth_rejection_uses_explicit_application_close_code(self) -> None:
source = (
Path(__file__).resolve().parents[1] / "web" / "routers" / "compute_nodes.py"
).read_text(encoding="utf-8")
rejection = source.split("except (ValueError, ComputeNodeError):", 1)[1].split(
"await node_connections.activate", 1
)[0]
self.assertLess(
rejection.index("await websocket.accept()"), rejection.index("await websocket.close")
)
self.assertIn('code=4003, reason="invalid node credentials"', rejection)
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)
class ComputeNodeDeleteTests(unittest.TestCase):
@patch("core.compute_nodes.session_scope")
def test_delete_node_removes_existing_identity(self, session_scope) -> None:
session = session_scope.return_value.__enter__.return_value
node = object()
session.get.return_value = node
self.assertTrue(delete_node(uuid4()))
session.delete.assert_called_once_with(node)
@patch("core.compute_nodes.session_scope")
def test_delete_node_reports_missing_identity(self, session_scope) -> None:
session = session_scope.return_value.__enter__.return_value
session.get.return_value = None
self.assertFalse(delete_node(uuid4()))
session.delete.assert_not_called()
if __name__ == "__main__":
unittest.main()