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.software_jobs import ( _canonical_request, _published_output_is_valid, abandon_offer, list_jobs, mark_node_jobs_disconnected, record_job_terminal, request_job_cancel, respond_to_offer, update_job_state, validate_output_manifest, ) from core.software_nodes import ( _enrollment_digest, _hash_secret, _verify_secret, delete_node, ) from web.routers.software_nodes import NodeConnectionManager, _bearer class SoftwareNodeSecurityTests(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" / "software_nodes.py" ).read_text(encoding="utf-8") rejection = source.split("except (ValueError, SoftwareNodeError):", 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 SoftwareNodeConnectionTests(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 SoftwareNodeMigrationTests(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) def test_0032_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.20260813_1600_0032_software_jobs" ) with patch.object(migration, "op", operations): migration.upgrade() rendered = "\n".join(statements) self.assertIn("ALTER TABLE compute_node_enrollments RENAME TO software_node_enrollments", rendered) self.assertIn("ALTER TABLE compute_nodes RENAME TO software_nodes", rendered) self.assertIn("software_jobs", rendered) self.assertIn("uq_software_jobs_user_idempotency", rendered) self.assertIn("ix_software_jobs_status_created", rendered) def test_0033_adds_software_job_artifact_source(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.20260814_0900_0033_artifact_software_job" ) with patch.object(migration, "op", operations): migration.upgrade() rendered = "\n".join(statements) self.assertIn("software_job_id", rendered) self.assertIn("ix_artifacts_software_job_id", rendered) self.assertIn("jsonb_array_elements", rendered) class SoftwareJobProtocolTests(unittest.TestCase): def test_published_output_distinguishes_artifacts_from_metadata(self) -> None: job_id = uuid4() self.assertTrue(_published_output_is_valid(job_id, { "source_artifact_id": "figure_png", "artifact_id": str(uuid4()), "path": f"origin/{job_id}/figure.png", })) self.assertTrue(_published_output_is_valid(job_id, { "source_artifact_id": "plot_spec", "artifact_id": None, "path": f"origin/{job_id}/.meta/plot-spec.json", })) self.assertFalse(_published_output_is_valid(job_id, { "source_artifact_id": ["plot_spec"], "artifact_id": None, "path": f"origin/{job_id}/.meta/plot-spec.json", })) @patch("core.software_jobs.session_scope") def test_queued_job_cancels_without_node_message(self, session_scope) -> None: session = session_scope.return_value.__enter__.return_value job = type("Job", (), {})() job.job_id = uuid4(); job.user_id = uuid4(); job.task_id = uuid4() job.capability = "origin.plot@v1"; job.request_digest = "a" * 64 job.node_id = None; job.lease_id = None; job.status = "queued"; job.stage = "" job.progress = 0; job.metrics = {}; job.error = {}; job.artifact_manifest = [] job.created_at = None; job.started_at = None; job.terminal_at = None session.execute.return_value.scalar_one_or_none.return_value = job result, node_message = request_job_cancel(job.user_id, job.job_id) self.assertEqual(result["status"], "cancelled") self.assertIsNone(node_message) self.assertEqual(job.error["code"], "USER_CANCELLED") @patch("core.software_jobs.session_scope") def test_running_job_persists_cancel_before_sending(self, session_scope) -> None: session = session_scope.return_value.__enter__.return_value job = type("Job", (), {})() job.job_id = uuid4(); job.user_id = uuid4(); job.task_id = uuid4() job.capability = "origin.plot@v1"; job.request_digest = "b" * 64 job.node_id = uuid4(); job.lease_id = uuid4(); job.status = "running" job.stage = "origin_running"; job.progress = 10; job.metrics = {}; job.error = {} job.artifact_manifest = []; job.created_at = None; job.started_at = None; job.terminal_at = None session.execute.return_value.scalar_one_or_none.return_value = job result, node_message = request_job_cancel(job.user_id, job.job_id) self.assertEqual(result["status"], "cancelling") self.assertEqual(node_message["node_id"], job.node_id) self.assertEqual(node_message["payload"]["lease_id"], str(job.lease_id)) @patch("core.software_jobs.session_scope") def test_job_list_is_enriched_for_job_center(self, session_scope) -> None: session = session_scope.return_value.__enter__.return_value job = type("Job", (), {})() job.job_id = uuid4(); job.task_id = uuid4(); job.capability = "origin.plot@v1" job.request_digest = "c" * 64; job.node_id = uuid4(); job.status = "running" job.stage = "origin_running"; job.progress = 20; job.metrics = {}; job.error = {} job.artifact_manifest = []; job.input_manifest = {"filename": "input.xlsx"} job.request = {"plot": {"title": "Test"}, "output": {"formats": ["png"]}} job.created_at = None; job.started_at = None; job.terminal_at = None session.execute.return_value.all.return_value = [(job, "材料仿真", "LAB-01")] results = list_jobs(uuid4(), limit=10) self.assertEqual(results[0]["task_name"], "材料仿真") self.assertEqual(results[0]["node_name"], "LAB-01") self.assertEqual(results[0]["request_summary"]["display_name"], "Origin 科研绘图") self.assertEqual(results[0]["output_dir"], f"origin/{job.job_id}") def test_origin_request_is_canonical_and_rejects_extra_fields(self) -> None: request = { "schema_version": 1, "input": {"input_id": str(uuid4()), "sheet": "Sheet1"}, "plot": {"type": "line", "x": "x", "y": ["y"]}, "output": {"formats": ["png", "opju"]}, } normalized, digest = _canonical_request(request) self.assertEqual(normalized, request) self.assertEqual(len(digest), 64) with self.assertRaisesRegex(Exception, "invalid origin plot request fields"): _canonical_request({**request, "script": "anything"}) with self.assertRaisesRegex(Exception, "unsupported origin plot fields"): _canonical_request({**request, "plot": {**request["plot"], "script": "anything"}}) with self.assertRaisesRegex(Exception, "input.input_id must be an artifact UUID"): _canonical_request({**request, "input": {"input_id": "C:\\data.csv"}}) def test_origin_request_rejects_unimplemented_plot_semantics(self) -> None: request = { "schema_version": 1, "input": {"input_id": str(uuid4())}, "plot": {"type": "scatter", "x": "time", "y": ["a", "b"]}, "output": {"formats": ["png"], "dpi": 300}, } for plot, message in ( ({**request["plot"], "template": "custom"}, "unsupported origin plot template"), ({**request["plot"], "x_axis": {"scale": "log10"}}, "invalid x_axis"), ({**request["plot"], "legend": {"enabled": False}}, "invalid plot.legend"), ({**request["plot"], "y": ["a", "a"]}, "plot.y must contain"), ): with self.subTest(message=message), self.assertRaisesRegex(Exception, message): _canonical_request({**request, "plot": plot}) with self.assertRaisesRegex(Exception, "video recording is not supported"): _canonical_request( {**request, "output": {"formats": ["png"], "record_video": True}} ) def test_output_manifest_matches_exact_requested_formats(self) -> None: request = {"output": {"formats": ["opju", "png"]}} manifest = [ {"artifact_id": "project", "filename": "project.opju", "media_type": "application/x-origin-project", "size_bytes": 10, "sha256": "a" * 64}, {"artifact_id": "figure_png", "filename": "figure.png", "media_type": "image/png", "size_bytes": 20, "sha256": "b" * 64}, {"artifact_id": "plot_spec", "filename": "plot-spec.json", "media_type": "application/json", "size_bytes": 30, "sha256": "c" * 64}, {"artifact_id": "provenance", "filename": "provenance.json", "media_type": "application/json", "size_bytes": 40, "sha256": "d" * 64}, ] self.assertEqual(validate_output_manifest(request, manifest), manifest) with self.assertRaisesRegex(Exception, "incomplete"): validate_output_manifest(request, manifest[:-1]) with self.assertRaisesRegex(Exception, "metadata"): validate_output_manifest( request, [{**manifest[0], "filename": "anything.opju"}, *manifest[1:]], ) @patch("core.software_jobs.session_scope") def test_stale_offer_cannot_be_accepted_by_another_node(self, session_scope) -> None: session = session_scope.return_value.__enter__.return_value job = type("Job", (), {})() job.node_id = uuid4() job.lease_id = uuid4() job.status = "offered" session.execute.return_value.scalar_one_or_none.return_value = job with self.assertRaisesRegex(Exception, "stale or does not belong"): respond_to_offer( uuid4(), accepted=True, payload={"job_id": str(uuid4()), "lease_id": str(job.lease_id)}, ) @patch("core.software_jobs.session_scope") def test_failed_delivery_only_abandons_matching_offer(self, session_scope) -> None: session = session_scope.return_value.__enter__.return_value node_id = uuid4() lease_id = uuid4() job = type("Job", (), {})() job.node_id = node_id job.lease_id = lease_id job.status = "offered" session.execute.return_value.scalar_one_or_none.return_value = job abandon_offer( node_id, {"job_id": str(uuid4()), "lease_id": str(lease_id)}, ) self.assertEqual(job.status, "queued") self.assertIsNone(job.node_id) def test_dispatcher_excludes_nodes_with_active_jobs(self) -> None: source = ( Path(__file__).resolve().parents[1] / "core" / "software_jobs.py" ).read_text(encoding="utf-8") self.assertIn('{"offered", "dispatched", "running"}', source) self.assertIn("item.node_id not in busy_node_ids", source) def test_input_download_rechecks_file_digest(self) -> None: source = ( Path(__file__).resolve().parents[1] / "web" / "routers" / "software_nodes.py" ).read_text(encoding="utf-8") self.assertIn("digest = sha256()", source) self.assertIn('digest.hexdigest() != item["sha256"]', source) @patch("core.software_jobs.session_scope") def test_job_state_restores_disconnected_job(self, session_scope) -> None: session = session_scope.return_value.__enter__.return_value node_id = uuid4() lease_id = uuid4() digest = "a" * 64 job = type("Job", (), {})() job.node_id = node_id job.lease_id = lease_id job.request_digest = digest job.status = "disconnected" job.started_at = None session.execute.return_value.scalar_one_or_none.return_value = job update_job_state(node_id, { "job_id": str(uuid4()), "lease_id": str(lease_id), "request_digest": digest, "stage": "waiting_input", "progress": 0, "metrics": {}, }) self.assertEqual(job.status, "dispatched") self.assertEqual(job.stage, "waiting_input") @patch("core.software_jobs.session_scope") def test_ready_to_run_is_not_reported_as_running(self, session_scope) -> None: session = session_scope.return_value.__enter__.return_value node_id = uuid4() lease_id = uuid4() digest = "c" * 64 job = type("Job", (), {})() job.node_id = node_id job.lease_id = lease_id job.request_digest = digest job.status = "dispatched" job.started_at = None session.execute.return_value.scalar_one_or_none.return_value = job update_job_state(node_id, { "job_id": str(uuid4()), "lease_id": str(lease_id), "request_digest": digest, "stage": "ready_to_run", "progress": 5, "metrics": {"input_bytes": 10}, }) self.assertEqual(job.status, "dispatched") @patch("core.software_jobs.session_scope") def test_terminal_replay_is_idempotent(self, session_scope) -> None: session = session_scope.return_value.__enter__.return_value node_id = uuid4() lease_id = uuid4() digest = "b" * 64 job = type("Job", (), {})() job.node_id = node_id job.lease_id = lease_id job.request_digest = digest job.status = "failed" session.execute.return_value.scalar_one_or_none.return_value = job record_job_terminal(node_id, { "job_id": str(uuid4()), "lease_id": str(lease_id), "request_digest": digest, "status": "failed", "error": {"code": "TEST"}, "artifact_manifest": [], }) self.assertEqual(job.status, "failed") @patch("core.software_jobs.session_scope") def test_disconnect_does_not_requeue_active_jobs(self, session_scope) -> None: session = session_scope.return_value.__enter__.return_value first = type("Job", (), {"status": "running"})() second = type("Job", (), {"status": "dispatched"})() session.execute.return_value.scalars.return_value = [first, second] mark_node_jobs_disconnected(uuid4()) self.assertEqual(first.status, "disconnected") self.assertEqual(second.status, "disconnected") class SoftwareNodeDeleteTests(unittest.TestCase): @patch("core.software_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.software_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()