Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
27 changes: 21 additions & 6 deletions packages/artemis-client/src/artemis_client/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -268,6 +268,24 @@ async def get_task(self, task_id: str) -> TaskResult:
return TaskResult(task_id=task_id, status="launching")
return TaskResult.from_payload(live_task, task_id=task_id)

def _resolve_wait_controls(
self,
timeout: float,
poll_interval: float | None,
) -> tuple[float, float]:
"""Validate the wait controls and resolve the effective poll interval.

Kept separate so :meth:`run` can reject a nonpositive control *before*
it submits anything: a rejected call must not leave work running on
the host that the caller never received a handle for.
"""
if timeout <= 0:
raise ValueError("timeout must be greater than zero")
interval = self.poll_interval if poll_interval is None else float(poll_interval)
if interval <= 0:
raise ValueError("poll_interval must be greater than zero")
return float(timeout), interval

async def wait_for_task(
self,
task_id: str,
Expand All @@ -276,11 +294,7 @@ async def wait_for_task(
poll_interval: float | None = None,
) -> TaskResult:
"""Wait until a task reaches a terminal state."""
if timeout <= 0:
raise ValueError("timeout must be greater than zero")
interval = self.poll_interval if poll_interval is None else float(poll_interval)
if interval <= 0:
raise ValueError("poll_interval must be greater than zero")
timeout, interval = self._resolve_wait_controls(timeout, poll_interval)

started = time.monotonic()
while True:
Expand Down Expand Up @@ -311,6 +325,7 @@ async def run(
poll_interval: float | None = None,
) -> TaskResult:
"""Submit a task and wait for its terminal result (see :meth:`submit`)."""
timeout, interval = self._resolve_wait_controls(timeout, poll_interval)
handle = await self.submit(
goal,
profile=profile,
Expand All @@ -328,7 +343,7 @@ async def run(
return await self.wait_for_task(
handle.task_id,
timeout=timeout,
poll_interval=poll_interval,
poll_interval=interval,
)

async def run_task(self, task: Any, **overrides: Any) -> TaskResult:
Expand Down
63 changes: 63 additions & 0 deletions packages/artemis-client/tests/test_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -215,6 +215,69 @@ async def get_task(self, task_id: str): # type: ignore[override]
with self.assertRaises(TaskTimeoutError):
await client.wait_for_task("slow-task", timeout=0.003)

async def test_run_rejects_nonpositive_wait_controls_without_submitting(self) -> None:
for kwargs, message in (
({"timeout": 0}, "timeout must be greater than zero"),
({"timeout": -1.0}, "timeout must be greater than zero"),
({"poll_interval": 0}, "poll_interval must be greater than zero"),
({"poll_interval": -0.5}, "poll_interval must be greater than zero"),
):
with self.subTest(**kwargs):
transport = FakeTransport()
client = ArtemisClient(
"https://artemis.example.test",
poll_interval=0.001,
transport=transport,
)
with self.assertRaises(ValueError) as caught:
await client.run("Open Settings", **kwargs)
self.assertEqual(str(caught.exception), message)
# The task must never have been admitted: the caller holds no
# handle for work the host may already have started.
self.assertEqual(transport.calls, [])

async def test_run_task_rejects_nonpositive_wait_controls_without_submitting(self) -> None:
class Task:
goal = "Open Settings"

for kwargs in ({"timeout": 0}, {"poll_interval": 0}):
with self.subTest(**kwargs):
transport = FakeTransport()
client = ArtemisClient(
"https://artemis.example.test",
poll_interval=0.001,
transport=transport,
)
with self.assertRaises(ValueError):
await client.run_task(Task(), **kwargs)
self.assertEqual(transport.calls, [])

async def test_run_still_submits_and_completes_with_valid_wait_controls(self) -> None:
task_id = "00000000-0000-4000-8000-000000000456"
self.transport.add(
"POST",
"/api/run",
{"status": "started", "tasks": [{"session_id": task_id, "status": "pending"}]},
)
self.transport.add(
"GET",
f"/api/sessions/{task_id}",
{"session_id": task_id, "status": "completed", "summary": "done"},
)

result = await self.client.run(
"Open Settings",
task_id=task_id,
timeout=5,
poll_interval=0.001,
)

self.assertTrue(result.succeeded)
self.assertEqual(
[(method, path) for method, path, _ in self.transport.calls],
[("POST", "/api/run"), ("GET", f"/api/sessions/{task_id}")],
)

async def test_list_devices_accepts_legacy_shape(self) -> None:
self.transport.add(
"GET",
Expand Down