zcbot/tests/test_software_job_creation.py

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()