zcbot/tests/test_software_job_tools.py

128 lines
5.0 KiB
Python

from __future__ import annotations
import json
import unittest
from unittest.mock import patch
from uuid import uuid4
from tools.software_jobs import (
SoftwareCapabilityListTool,
SoftwareJobCancelTool,
SoftwareJobStatusTool,
SoftwareJobSubmitTool,
)
class SoftwareJobToolTests(unittest.TestCase):
def setUp(self):
self.user_id = uuid4()
self.task_id = uuid4()
def test_capability_list_reports_current_capacity(self):
nodes = [
{
"status": "online",
"capabilities": ["origin.plot@v2"],
"runtime": {"available_slots": 1},
},
{
"status": "offline",
"capabilities": ["origin.plot@v2"],
"runtime": {"available_slots": 1},
},
]
with patch("tools.software_jobs.list_nodes", return_value=nodes):
result = json.loads(
SoftwareCapabilityListTool(self.user_id, self.task_id).execute()
)
self.assertEqual(result["capabilities"][0]["available_nodes"], 1)
def test_submit_injects_current_user_and_task(self):
created = {"job_id": str(uuid4()), "status": "queued"}
artifact_id = uuid4()
tool = SoftwareJobSubmitTool(self.user_id, self.task_id)
with patch("tools.software_jobs.create_job", return_value=(created, True)) as create:
result = json.loads(tool.execute(
"origin.plot@v2",
inputs=[{"key": "sample", "artifact_id": str(artifact_id)}],
operation={"plot": {
"type": "scatter",
"series": [{"input": "sample", "x": "x", "y": "y"}],
}},
outputs=[
{"key": "project", "type": "project", "format": "opju"},
{"key": "figure_png", "type": "figure", "format": "png", "options": {"dpi": 300}},
],
))
self.assertTrue(result["created"])
self.assertEqual(create.call_args.args[:2], (self.user_id, self.task_id))
self.assertEqual(create.call_args.kwargs["capability"], "origin.plot@v2")
self.assertEqual(
create.call_args.kwargs["request"],
{
"schema_version": 2,
"inputs": [{"key": "sample", "artifact_id": str(artifact_id)}],
"operation": {"plot": {
"type": "scatter",
"series": [{"input": "sample", "x": "x", "y": "y"}],
}},
"outputs": [
{"key": "project", "type": "project", "format": "opju"},
{"key": "figure_png", "type": "figure", "format": "png", "options": {"dpi": 300}},
],
},
)
def test_submit_requires_artifact_uuid(self):
result = SoftwareJobSubmitTool(self.user_id, self.task_id).execute(
"origin.plot@v2",
inputs=[{"key": "sample", "artifact_id": "data/input.csv"}],
operation={"plot": {
"type": "scatter",
"series": [{"input": "sample", "x": "x", "y": "y"}],
}},
outputs=[{"key": "figure_png", "type": "figure", "format": "png"}],
)
self.assertIn("call register_artifact first", result)
def test_submit_schema_requires_artifact_input(self):
required = SoftwareJobSubmitTool.parameters["required"]
self.assertIn("inputs", required)
self.assertIn("operation", required)
self.assertIn("outputs", required)
self.assertNotIn("input_path", SoftwareJobSubmitTool.parameters["properties"])
self.assertNotIn("request", SoftwareJobSubmitTool.parameters["properties"])
def test_status_and_cancel_reject_cross_task_job(self):
foreign = {"job_id": str(uuid4()), "task_id": str(uuid4())}
with patch("tools.software_jobs.get_job", return_value=foreign):
status = SoftwareJobStatusTool(self.user_id, self.task_id).execute(
foreign["job_id"]
)
cancel = SoftwareJobCancelTool(self.user_id, self.task_id).execute(
foreign["job_id"]
)
self.assertIn("not found", status)
self.assertIn("not found", cancel)
def test_cancel_uses_user_scoped_service(self):
job_id = uuid4()
current = {"job_id": str(job_id), "task_id": str(self.task_id)}
cancelled = {**current, "status": "cancelled"}
with (
patch("tools.software_jobs.get_job", return_value=current),
patch(
"tools.software_jobs.request_job_cancel",
return_value=(cancelled, None),
) as request_cancel,
):
result = json.loads(
SoftwareJobCancelTool(self.user_id, self.task_id).execute(str(job_id))
)
self.assertEqual(result["status"], "cancelled")
request_cancel.assert_called_once_with(self.user_id, job_id)
if __name__ == "__main__":
unittest.main()