123 lines
4.5 KiB
Python
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()
|