from __future__ import annotations import unittest from contextlib import contextmanager, nullcontext from types import SimpleNamespace from unittest.mock import MagicMock, call, patch from uuid import uuid4 from sqlalchemy.exc import IntegrityError from core.software_jobs import SoftwareJobError, create_job from core.storage.models import SoftwareJob, SoftwareWorkspace class _Result: def __init__(self, *, first=None, scalar=None, scalars=None): self._first = first self._scalar = scalar self._scalars = scalars or [] def first(self): return self._first def scalar_one_or_none(self): return self._scalar def scalars(self): return SimpleNamespace(all=lambda: self._scalars) class SoftwareJobCreationTests(unittest.TestCase): def setUp(self) -> None: self.user_id = uuid4() self.task_id = uuid4() self.artifact_id = uuid4() self.request = { "schema_version": 2, "inputs": [{"key": "data", "artifact_id": str(self.artifact_id)}], "operation": {"plot": {"type": "line"}}, "outputs": [{"key": "figure_png", "type": "figure", "format": "png"}], } self.contract = SimpleNamespace( workspace=SimpleNamespace(), output_namespace="origin", input_policy={ "suffixes": [".xlsx"], "max_bytes": 1024, "max_total_bytes": 2048, }, normalize_request=lambda request: (request, "a" * 64), input_bindings=lambda request: request["inputs"], ) self.artifact = SimpleNamespace( artifact_id=self.artifact_id, user_id=self.user_id, status="active", current_path="demo.xlsx", size_bytes=128, content_sha256="b" * 64, ) def _session(self, *, fail_first_flush: bool = False): session = MagicMock() results = iter([ _Result(first=(self.task_id,)), _Result(scalar=None), _Result(scalars=[self.artifact]), _Result(scalar=None), _Result(scalar=None), ]) session.execute.side_effect = lambda _statement: next(results) session.begin_nested.return_value = nullcontext() if fail_first_flush: session.flush.side_effect = IntegrityError( "insert workspace", {}, Exception("foreign key violation") ) return session @staticmethod def _scope(session): @contextmanager def scope(): yield session return scope def _create(self, session): with ( patch("core.software_jobs.session_scope", self._scope(session)), patch("core.software_jobs.get_contract", return_value=self.contract), ): return create_job( self.user_id, self.task_id, idempotency_key="workspace-first-job", capability="origin.plot@v2", request=self.request, ) def test_first_workspace_job_flushes_workspace_before_job(self) -> None: session = self._session() job, created = self._create(session) self.assertTrue(created) self.assertEqual(job["workspace_id"], str(session.add.call_args_list[0].args[0].workspace_id)) workspace = session.add.call_args_list[0].args[0] row = session.add.call_args_list[1].args[0] self.assertIsInstance(workspace, SoftwareWorkspace) self.assertIsInstance(row, SoftwareJob) self.assertEqual(row.workspace_id, workspace.workspace_id) self.assertEqual(session.flush.call_args_list, [call([workspace]), call([row])]) def test_non_idempotency_integrity_error_is_not_masked_as_no_result(self) -> None: session = self._session(fail_first_flush=True) with self.assertRaisesRegex( SoftwareJobError, "database integrity validation failed" ): self._create(session) self.assertEqual(session.execute.call_count, 5) if __name__ == "__main__": unittest.main()