zcbot/tests/test_software_job_tools.py

262 lines
10 KiB
Python

from __future__ import annotations
import json
import unittest
from contextlib import contextmanager
from types import SimpleNamespace
from unittest.mock import patch
from uuid import uuid4
from core.software_jobs import revise_job
from tools.software_jobs import (
SoftwareCapabilityListTool,
SoftwareJobCancelTool,
SoftwareJobReviseTool,
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(result["completion_delivery"], "automatic")
self.assertEqual(
result["next_action"], "end_turn_after_reporting_queued_job_id"
)
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["completion_action"], "report")
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"])
self.assertEqual(
SoftwareJobSubmitTool.parameters["properties"]["completion_action"]["enum"],
["report", "analyze"],
)
self.assertIn("terminal action for the current run", SoftwareJobSubmitTool.description)
self.assertIn("automatically", SoftwareJobSubmitTool.description)
def test_submit_can_request_automatic_analysis(self):
artifact_id = uuid4()
with patch(
"tools.software_jobs.create_job",
return_value=({"job_id": str(uuid4())}, True),
) as create:
SoftwareJobSubmitTool(self.user_id, self.task_id).execute(
"origin.plot@v2",
inputs=[{"key": "sample", "artifact_id": str(artifact_id)}],
operation={"plot": {"type": "line", "series": []}},
outputs=[{"key": "figure_png", "type": "figure", "format": "png"}],
completion_action="analyze",
)
self.assertEqual(create.call_args.kwargs["completion_action"], "analyze")
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_status_includes_editable_request_for_revision(self):
job_id = uuid4()
current = {"job_id": str(job_id), "task_id": str(self.task_id)}
editable = {
"schema_version": 2,
"inputs": [{"key": "sample", "artifact_id": str(uuid4())}],
"operation": {"plot": {"type": "line"}},
"outputs": [{"key": "figure_png", "format": "png"}],
}
with (
patch("tools.software_jobs.get_job", return_value=current),
patch("tools.software_jobs.get_job_request", return_value=editable),
):
result = json.loads(
SoftwareJobStatusTool(self.user_id, self.task_id).execute(str(job_id))
)
self.assertEqual(result["editable_request"], editable)
def test_revise_reuses_source_job_through_user_scoped_service(self):
source_job_id = uuid4()
revised = {"job_id": str(uuid4()), "status": "queued"}
operation = {"plot": {
"type": "recipe",
"recipe_version": 1,
"layout": {"rows": 1, "columns": 1},
"panels": [],
}}
outputs = [{"key": "figure_png", "type": "figure", "format": "png"}]
with patch(
"tools.software_jobs.revise_job", return_value=(revised, True)
) as revise:
result = json.loads(SoftwareJobReviseTool(
self.user_id, self.task_id
).execute(
str(source_job_id),
operation=operation,
outputs=outputs,
idempotency_key="revision-1",
))
self.assertTrue(result["created"])
self.assertEqual(result["completion_delivery"], "automatic")
self.assertEqual(
result["next_action"], "end_turn_after_reporting_queued_job_id"
)
revise.assert_called_once_with(
self.user_id,
self.task_id,
source_job_id,
idempotency_key="revision-1",
operation=operation,
outputs=outputs,
)
def test_revise_service_copies_inputs_and_revalidates_as_new_job(self):
source_job_id = uuid4()
source = SimpleNamespace(
capability="origin.plot@v2",
request={
"schema_version": 2,
"inputs": [{"key": "sample", "artifact_id": str(uuid4())}],
"operation": {"plot": {"type": "line"}},
"outputs": [{"key": "figure_png", "format": "png"}],
},
)
class Result:
@staticmethod
def scalar_one_or_none():
return source
class Session:
@staticmethod
def execute(_statement):
return Result()
@contextmanager
def fake_session_scope():
yield Session()
operation = {"plot": {"type": "scatter"}}
outputs = [{"key": "figure_png", "format": "png"}]
created = {"job_id": str(uuid4())}
with (
patch("core.software_jobs.session_scope", fake_session_scope),
patch("core.software_jobs.create_job", return_value=(created, True)) as create,
):
result = revise_job(
self.user_id,
self.task_id,
source_job_id,
idempotency_key="revision-service-1",
operation=operation,
outputs=outputs,
)
self.assertEqual(result, (created, True))
self.assertEqual(create.call_args.kwargs["capability"], "origin.plot@v2")
self.assertEqual(create.call_args.kwargs["completion_action"], "report")
self.assertEqual(create.call_args.kwargs["request"], {
"schema_version": 2,
"inputs": source.request["inputs"],
"operation": operation,
"outputs": outputs,
})
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()