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()