127 lines
4.1 KiB
Python
127 lines
4.1 KiB
Python
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()
|