zcbot/tests/test_compute_nodes.py

330 lines
14 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 core.compute_jobs import (
_canonical_request,
abandon_offer,
mark_node_jobs_disconnected,
record_job_terminal,
respond_to_offer,
update_job_state,
validate_output_manifest,
)
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)
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_compute_jobs"
)
with patch.object(migration, "op", operations):
migration.upgrade()
rendered = "\n".join(statements)
self.assertIn("compute_jobs", rendered)
self.assertIn("uq_compute_jobs_user_idempotency", rendered)
self.assertIn("ix_compute_jobs_status_created", rendered)
class ComputeJobProtocolTests(unittest.TestCase):
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.compute_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.compute_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" / "compute_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" / "compute_nodes.py"
).read_text(encoding="utf-8")
self.assertIn("digest = sha256()", source)
self.assertIn('digest.hexdigest() != item["sha256"]', source)
@patch("core.compute_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.compute_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.compute_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.compute_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 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()