from __future__ import annotations import importlib import unittest 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 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) 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) if __name__ == "__main__": unittest.main()