From fb3c502732d5b369d5d72cda039fb02be1f6dd63 Mon Sep 17 00:00:00 2001 From: basil-k-aji-dev <70605804+basil-k-aji-dev@users.noreply.github.com> Date: Mon, 14 Sep 2026 00:16:09 +0530 Subject: [PATCH] fix(client): reject nonpositive wait controls before submitting a task `ArtemisClient.run()` awaited `submit()` before `wait_for_task()`, and only the latter validated `timeout` and `poll_interval`. A caller passing a nonpositive value therefore received a `ValueError` with no task handle, while a real server may already have admitted and started the submitted work. `run_task()` delegates to `run()` and inherited the same behaviour. Move both checks into a `_resolve_wait_controls()` helper, called by `wait_for_task()` as before and by `run()` before it submits anything. Valid calls are unchanged: `run()` now resolves the poll interval once and hands the resolved value to `wait_for_task()`. Closes #83. --- .../src/artemis_client/client.py | 27 ++++++-- packages/artemis-client/tests/test_client.py | 63 +++++++++++++++++++ 2 files changed, 84 insertions(+), 6 deletions(-) diff --git a/packages/artemis-client/src/artemis_client/client.py b/packages/artemis-client/src/artemis_client/client.py index d1d3e2d4..d336b870 100644 --- a/packages/artemis-client/src/artemis_client/client.py +++ b/packages/artemis-client/src/artemis_client/client.py @@ -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, @@ -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: @@ -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, @@ -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: diff --git a/packages/artemis-client/tests/test_client.py b/packages/artemis-client/tests/test_client.py index a5480a7a..d39b1128 100644 --- a/packages/artemis-client/tests/test_client.py +++ b/packages/artemis-client/tests/test_client.py @@ -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",