From 22b222eb98da04407f84e7d13dd318f319e75c3f Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Wed, 23 Sep 2026 21:42:54 +0700 Subject: [PATCH 01/71] feat: add user-specified retries to the API executor The API executor marks transient failures (5xx, 408, 429, connection errors) as retryable but never retried them. Add a spec.api.retries field (default 0, preserving current behavior) that controls how many times a transient failure is re-issued before the task fails, with a fixed 1s backoff between attempts. Non-retryable 4xx statuses and cancelled tasks are never retried. Co-Authored-By: Claude Code Signed-off-by: Zhengyuan Su --- docs/WORKFLOWS.md | 3 + src/worker/executors/api_executor.py | 76 +++++++++++++++++++-- tests/worker/test_api_executor.py | 98 +++++++++++++++++++++++++++- 3 files changed, 169 insertions(+), 8 deletions(-) diff --git a/docs/WORKFLOWS.md b/docs/WORKFLOWS.md index 404013e4a..404081018 100644 --- a/docs/WORKFLOWS.md +++ b/docs/WORKFLOWS.md @@ -77,6 +77,8 @@ contract. Credential handling: a caller-supplied `Authorization` header is always used as-is and never overwritten. With no header, `NEBULA_API_TOKEN` is injected only when the call is on the Nebula url (no custom `spec.api.url`) — the Nebula token is never sent to a custom endpoint. A Nebula-path call with no token available fails closed. +`spec.api.retries` (default `0`) sets how many times a transient failure is retried before the task fails. A transient failure is a connection error or an HTTP status of 5xx, 408, or 429; other 4xx statuses are never retried. Each retry waits a fixed 1s backoff. A cancelled task stops retrying immediately. + ```yaml spec: taskType: api @@ -89,6 +91,7 @@ spec: messages: - role: user content: Hello + retries: 3 response: parse_json: true ``` diff --git a/src/worker/executors/api_executor.py b/src/worker/executors/api_executor.py index 6c7086161..521e6984a 100644 --- a/src/worker/executors/api_executor.py +++ b/src/worker/executors/api_executor.py @@ -1,6 +1,7 @@ import logging import os import threading +import time from pathlib import Path from typing import Any, ClassVar @@ -11,13 +12,21 @@ from shared.tasks.task_type import TaskType from shared.utils.redact import is_credential_key -from .base_executor import ExecutionError, Executor, ExecutorTask +from .base_executor import ExecutionError, Executor, ExecutorTask, TaskCancelledError logger = logging.getLogger(__name__) # Cache key: (base_url, timeout_seconds, verify_tls, follow_redirects) _ClientKey = tuple[str, float, bool, bool] +# Fixed delay between retry attempts. +_RETRY_BACKOFF_SEC = 1.0 + + +def _is_retryable_status(status_code: int) -> bool: + """Whether an HTTP status is transient and worth retrying.""" + return status_code >= 500 or status_code in (408, 429) + class APIExecutor(Executor): """Performs a single HTTP request defined by task YAML. @@ -36,6 +45,14 @@ class APIExecutor(Executor): _clients: ClassVar[dict[_ClientKey, httpx.Client]] = {} _clients_lock: ClassVar[threading.Lock] = threading.Lock() + def __init__(self, *args: Any, **kwargs: Any) -> None: + super().__init__(*args, **kwargs) + self._cancel_event = threading.Event() + + def cancel(self, task_id: str) -> None: + """Signal the executor to abort the current request and any retries.""" + self._cancel_event.set() + @classmethod def _base_url(cls, url: str) -> str: """Extract scheme + host + port from a URL for pool keying.""" @@ -78,6 +95,47 @@ def _get_client( ) return client + def _request_with_retries( + self, + client: httpx.Client, + method: str, + url: str, + headers: dict[str, Any], + params: dict[str, Any] | None, + request_kwargs: dict[str, Any], + retries: int, + ) -> httpx.Response: + """Issue the request, retrying transient failures up to ``retries`` times. + + A retryable failure is a connection error or a transient HTTP status + (5xx, 408, 429). Non-retryable failures and a cancelled task stop the + loop immediately. The final attempt's failure propagates to the caller. + """ + attempt = 0 + while True: + if self._cancel_event.is_set(): + raise TaskCancelledError("API request cancelled") + try: + resp = client.request( + method, + url, + headers=headers, + params=params, + **request_kwargs, + ) + except httpx.RequestError: + if attempt < retries: + attempt += 1 + time.sleep(_RETRY_BACKOFF_SEC) + continue + raise + if resp.is_error and _is_retryable_status(resp.status_code): + if attempt < retries: + attempt += 1 + time.sleep(_RETRY_BACKOFF_SEC) + continue + return resp + @classmethod def close_all_clients(cls) -> None: """Close and discard all cached HTTP clients.""" @@ -164,15 +222,21 @@ def run(self, task: ExecutorTask, out_dir: Path) -> APIResult: raise_for_status = bool(response_cfg.get("raise_for_status", True)) max_body_bytes = int(response_cfg.get("max_body_bytes", 200000)) + retries = api_cfg.get("retries", 0) + if not isinstance(retries, int) or isinstance(retries, bool) or retries < 0: + raise ExecutionError("spec.api.retries must be a non-negative integer") + try: base = self._base_url(str(url)) client = self._get_client(base, timeout, verify_tls, follow_redirects) - resp = client.request( + resp = self._request_with_retries( + client, method, str(url), - headers=headers, - params=params, - **request_kwargs, + headers, + params, + request_kwargs, + retries, ) except httpx.RequestError as exc: raise ExecutionError(f"API request failed: {exc}", retryable=True) from exc @@ -225,7 +289,7 @@ def run(self, task: ExecutorTask, out_dir: Path) -> APIResult: message = f"API request returned status {resp.status_code}" if body_text: message = f"{message}: {body_text[:200]}" - retryable = resp.status_code >= 500 or resp.status_code in (408, 429) + retryable = _is_retryable_status(resp.status_code) raise ExecutionError(message, retryable=retryable) return result diff --git a/tests/worker/test_api_executor.py b/tests/worker/test_api_executor.py index 2f36a0235..214a32ec9 100644 --- a/tests/worker/test_api_executor.py +++ b/tests/worker/test_api_executor.py @@ -1,5 +1,6 @@ """Tests for the API executor url override and Nebula credential handling.""" +import threading from pathlib import Path from unittest.mock import patch @@ -8,7 +9,7 @@ from shared.tasks.worker_message import WorkerTaskMessage from worker.executors.api_executor import APIExecutor -from worker.executors.base_executor import ExecutionError +from worker.executors.base_executor import ExecutionError, TaskCancelledError def _task_message(**spec_updates: object) -> WorkerTaskMessage: @@ -54,8 +55,9 @@ def _handler(self, request: httpx.Request) -> httpx.Response: def _run( - executor: APIExecutor, task: WorkerTaskMessage, transport: _RecordingTransport + executor: APIExecutor, task: WorkerTaskMessage, transport: httpx.MockTransport ) -> None: + executor._cancel_event = threading.Event() with patch.object( APIExecutor, "_get_client", return_value=httpx.Client(transport=transport) ): @@ -162,3 +164,95 @@ def test_custom_url_with_only_innocent_header_stays_unauthenticated( assert transport.request is not None assert transport.request.headers["Content-Type"] == "application/json" assert "Authorization" not in transport.request.headers + + +class _SequenceTransport(httpx.MockTransport): + """MockTransport that serves a fixed sequence of responses.""" + + def __init__(self, responses: list[httpx.Response]) -> None: + self.responses = list(responses) + self.calls = 0 + super().__init__(self._handler) + + def _handler(self, request: httpx.Request) -> httpx.Response: + self.calls += 1 + return self.responses.pop(0) + + +def _ok_response() -> httpx.Response: + return httpx.Response( + 200, + json={ + "choices": [{"message": {"content": "hello"}}], + "usage": {"total_tokens": 3}, + }, + ) + + +def _error_response(status_code: int) -> httpx.Response: + return httpx.Response(status_code, json={"error": "boom"}) + + +class TestRetries: + def _task(self, **spec_updates: object) -> WorkerTaskMessage: + return _task_message( + url="https://custom.example.com/v1/chat/completions", + response={"parse_json": False}, + **spec_updates, + ) + + def test_retry_succeeds_after_transient_failures(self) -> None: + """A 504 followed by a 200 succeeds when retries are configured.""" + task = self._task(retries=2) + transport = _SequenceTransport( + [_error_response(504), _error_response(504), _ok_response()] + ) + _run(APIExecutor.__new__(APIExecutor), task, transport) + assert transport.calls == 3 + + def test_retries_exhausted_still_fails(self) -> None: + """Persistent 5xx failures exhaust retries and raise loudly.""" + task = self._task(retries=2) + transport = _SequenceTransport( + [_error_response(504), _error_response(504), _error_response(504)] + ) + with pytest.raises(ExecutionError, match="status 504"): + _run(APIExecutor.__new__(APIExecutor), task, transport) + assert transport.calls == 3 + + def test_no_retry_by_default(self) -> None: + """Without a retries field, a transient failure fails immediately.""" + task = self._task() + transport = _SequenceTransport([_error_response(504)]) + with pytest.raises(ExecutionError, match="status 504"): + _run(APIExecutor.__new__(APIExecutor), task, transport) + assert transport.calls == 1 + + def test_non_retryable_status_not_retried(self) -> None: + """A 4xx (other than 408/429) is never retried.""" + task = self._task(retries=3) + transport = _SequenceTransport([_error_response(400)]) + with pytest.raises(ExecutionError, match="status 400"): + _run(APIExecutor.__new__(APIExecutor), task, transport) + assert transport.calls == 1 + + def test_cancelled_task_stops_retrying(self) -> None: + """A cancelled task does not keep retrying.""" + executor = APIExecutor.__new__(APIExecutor) + executor._cancel_event = threading.Event() + executor._cancel_event.set() + task = self._task(retries=3) + transport = _SequenceTransport([_error_response(504)]) + with patch.object( + APIExecutor, "_get_client", return_value=httpx.Client(transport=transport) + ): + with pytest.raises(TaskCancelledError): + executor.run(task, Path("/tmp/out")) + assert transport.calls == 0 + + def test_invalid_retries_rejected(self) -> None: + """A negative or non-integer retries value is rejected.""" + for bad in (-1, "2", 1.5, True): + task = self._task(retries=bad) + with pytest.raises(ExecutionError, match="spec.api.retries"): + _run(APIExecutor.__new__(APIExecutor), task, _RecordingTransport()) From ad4752f4649e2a9642b9377a3b46590f56be005e Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Sun, 27 Sep 2026 11:16:29 +0700 Subject: [PATCH 02/71] fix(worker): reset stale API executor cancellations and honor cancel during backoff A cancellation left over from a previous task no longer leaks into the next task on a reused warm executor: cancel() records the task it targets, and run() clears a cancellation addressed to a different task while one aimed at the task now starting still stands, all under a single lock so a racing cancel is never lost. The retry backoff now waits on the cancel event instead of sleeping, so a cancelled task raises its cancellation as soon as it is signalled rather than sitting out the full backoff. Signed-off-by: Zhengyuan Su Co-Authored-By: Claude Code --- src/worker/executors/api_executor.py | 22 +++++++-- tests/worker/test_api_executor.py | 73 +++++++++++++++++++++++++++- 2 files changed, 90 insertions(+), 5 deletions(-) diff --git a/src/worker/executors/api_executor.py b/src/worker/executors/api_executor.py index 521e6984a..7e3daa29c 100644 --- a/src/worker/executors/api_executor.py +++ b/src/worker/executors/api_executor.py @@ -1,7 +1,6 @@ import logging import os import threading -import time from pathlib import Path from typing import Any, ClassVar @@ -48,10 +47,14 @@ class APIExecutor(Executor): def __init__(self, *args: Any, **kwargs: Any) -> None: super().__init__(*args, **kwargs) self._cancel_event = threading.Event() + self._cancel_lock = threading.Lock() + self._cancelled_task_id: str | None = None def cancel(self, task_id: str) -> None: """Signal the executor to abort the current request and any retries.""" - self._cancel_event.set() + with self._cancel_lock: + self._cancelled_task_id = task_id + self._cancel_event.set() @classmethod def _base_url(cls, url: str) -> str: @@ -126,16 +129,21 @@ def _request_with_retries( except httpx.RequestError: if attempt < retries: attempt += 1 - time.sleep(_RETRY_BACKOFF_SEC) + self._wait_for_backoff() continue raise if resp.is_error and _is_retryable_status(resp.status_code): if attempt < retries: attempt += 1 - time.sleep(_RETRY_BACKOFF_SEC) + self._wait_for_backoff() continue return resp + def _wait_for_backoff(self) -> None: + """Wait out the retry backoff, aborting early if the task is cancelled.""" + if self._cancel_event.wait(_RETRY_BACKOFF_SEC): + raise TaskCancelledError("API request cancelled") + @classmethod def close_all_clients(cls) -> None: """Close and discard all cached HTTP clients.""" @@ -153,6 +161,12 @@ def cleanup_after_run(self) -> None: self.close_all_clients() def run(self, task: ExecutorTask, out_dir: Path) -> APIResult: + with self._cancel_lock: + # A cancellation left over from a previous task must not leak into + # this one; one addressed to this task still stands. + if self._cancelled_task_id != task.task_id: + self._cancel_event.clear() + self._cancelled_task_id = None spec = self.require_spec(task, ApiSpecStrict) api_cfg = spec.api or {} if not isinstance(api_cfg, dict): diff --git a/tests/worker/test_api_executor.py b/tests/worker/test_api_executor.py index 214a32ec9..5ab96d558 100644 --- a/tests/worker/test_api_executor.py +++ b/tests/worker/test_api_executor.py @@ -1,6 +1,7 @@ """Tests for the API executor url override and Nebula credential handling.""" import threading +import time from pathlib import Path from unittest.mock import patch @@ -58,6 +59,8 @@ def _run( executor: APIExecutor, task: WorkerTaskMessage, transport: httpx.MockTransport ) -> None: executor._cancel_event = threading.Event() + executor._cancel_lock = threading.Lock() + executor._cancelled_task_id = None with patch.object( APIExecutor, "_get_client", return_value=httpx.Client(transport=transport) ): @@ -240,8 +243,10 @@ def test_cancelled_task_stops_retrying(self) -> None: """A cancelled task does not keep retrying.""" executor = APIExecutor.__new__(APIExecutor) executor._cancel_event = threading.Event() - executor._cancel_event.set() + executor._cancel_lock = threading.Lock() + executor._cancelled_task_id = None task = self._task(retries=3) + executor.cancel(task.task_id) transport = _SequenceTransport([_error_response(504)]) with patch.object( APIExecutor, "_get_client", return_value=httpx.Client(transport=transport) @@ -256,3 +261,69 @@ def test_invalid_retries_rejected(self) -> None: task = self._task(retries=bad) with pytest.raises(ExecutionError, match="spec.api.retries"): _run(APIExecutor.__new__(APIExecutor), task, _RecordingTransport()) + + def test_cancel_previous_task_does_not_cancel_next(self) -> None: + """A cancellation left over from a prior task does not cancel the next.""" + executor = APIExecutor.__new__(APIExecutor) + executor._cancel_event = threading.Event() + executor._cancel_lock = threading.Lock() + executor._cancelled_task_id = None + + task_a = self._task(retries=0) + task_a.task_id = "task-a" + executor.cancel(task_a.task_id) + + task_b = self._task(retries=0) + task_b.task_id = "task-b" + transport = _RecordingTransport() + with patch.object( + APIExecutor, "_get_client", return_value=httpx.Client(transport=transport) + ): + executor.run(task_b, Path("/tmp/out")) + assert transport.request is not None + + def test_cancel_before_start_still_cancels(self) -> None: + """A cancellation addressed to a task before it starts still cancels it.""" + executor = APIExecutor.__new__(APIExecutor) + executor._cancel_event = threading.Event() + executor._cancel_lock = threading.Lock() + executor._cancelled_task_id = None + + task = self._task(retries=3) + task.task_id = "task-b" + executor.cancel(task.task_id) + transport = _SequenceTransport([_error_response(504)]) + with patch.object( + APIExecutor, "_get_client", return_value=httpx.Client(transport=transport) + ): + with pytest.raises(TaskCancelledError): + executor.run(task, Path("/tmp/out")) + assert transport.calls == 0 + + def test_cancel_during_backoff_stops_retrying(self) -> None: + """A cancellation during the retry backoff aborts well before it ends.""" + executor = APIExecutor.__new__(APIExecutor) + executor._cancel_event = threading.Event() + executor._cancel_lock = threading.Lock() + executor._cancelled_task_id = None + + task = self._task(retries=3) + task.task_id = "task-b" + transport = _SequenceTransport([_error_response(503)]) + + def _cancel_after_delay() -> None: + time.sleep(0.05) + executor.cancel(task.task_id) + + canceller = threading.Thread(target=_cancel_after_delay) + canceller.start() + start = time.monotonic() + with patch.object( + APIExecutor, "_get_client", return_value=httpx.Client(transport=transport) + ): + with pytest.raises(TaskCancelledError): + executor.run(task, Path("/tmp/out")) + elapsed = time.monotonic() - start + canceller.join() + assert elapsed < 0.5 + assert transport.calls == 1 From 275ca455d64a8814d409fc280ace213b0e3cd2c3 Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Sun, 27 Sep 2026 11:19:00 +0700 Subject: [PATCH 03/71] fix(worker): shorten stale cancellation comment to a single line Signed-off-by: Zhengyuan Su Co-Authored-By: Claude Code --- src/worker/executors/api_executor.py | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) diff --git a/src/worker/executors/api_executor.py b/src/worker/executors/api_executor.py index 7e3daa29c..13f4adb9f 100644 --- a/src/worker/executors/api_executor.py +++ b/src/worker/executors/api_executor.py @@ -162,8 +162,7 @@ def cleanup_after_run(self) -> None: def run(self, task: ExecutorTask, out_dir: Path) -> APIResult: with self._cancel_lock: - # A cancellation left over from a previous task must not leak into - # this one; one addressed to this task still stands. + # A prior task's cancellation must not leak into this one. if self._cancelled_task_id != task.task_id: self._cancel_event.clear() self._cancelled_task_id = None From 67b795bb2b5477e879608f3e600e9f4b0b3b36c8 Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Sun, 27 Sep 2026 11:39:40 +0700 Subject: [PATCH 04/71] fix(worker): reject cancellations addressed to a different active task The interrupt monitor checks the runner's current task id and then calls executor.cancel(task_id) without a lock spanning both, so a late cancellation for a finished task can reach a warm executor that has since started another task. The executor now records the active task id under its cancel lock, sets the cancel event only when no run is in flight or the cancellation matches the active task, and clears the active id when the run returns or raises. Signed-off-by: Zhengyuan Su Co-Authored-By: Claude Code --- src/worker/executors/api_executor.py | 19 ++++- tests/worker/test_api_executor.py | 100 +++++++++++++++++++++++---- 2 files changed, 105 insertions(+), 14 deletions(-) diff --git a/src/worker/executors/api_executor.py b/src/worker/executors/api_executor.py index 13f4adb9f..b16bbe7fc 100644 --- a/src/worker/executors/api_executor.py +++ b/src/worker/executors/api_executor.py @@ -48,13 +48,20 @@ def __init__(self, *args: Any, **kwargs: Any) -> None: super().__init__(*args, **kwargs) self._cancel_event = threading.Event() self._cancel_lock = threading.Lock() + self._active_task_id: str | None = None self._cancelled_task_id: str | None = None def cancel(self, task_id: str) -> None: """Signal the executor to abort the current request and any retries.""" with self._cancel_lock: - self._cancelled_task_id = task_id - self._cancel_event.set() + if self._active_task_id is None: + # No run in flight: record the id so a cancellation addressed to + # a task that has not started yet still lands when it starts. + self._cancelled_task_id = task_id + self._cancel_event.set() + elif self._active_task_id == task_id: + self._cancel_event.set() + # A cancellation for a different task than the active one is ignored. @classmethod def _base_url(cls, url: str) -> str: @@ -162,10 +169,18 @@ def cleanup_after_run(self) -> None: def run(self, task: ExecutorTask, out_dir: Path) -> APIResult: with self._cancel_lock: + self._active_task_id = task.task_id # A prior task's cancellation must not leak into this one. if self._cancelled_task_id != task.task_id: self._cancel_event.clear() self._cancelled_task_id = None + try: + return self._run(task, out_dir) + finally: + with self._cancel_lock: + self._active_task_id = None + + def _run(self, task: ExecutorTask, out_dir: Path) -> APIResult: spec = self.require_spec(task, ApiSpecStrict) api_cfg = spec.api or {} if not isinstance(api_cfg, dict): diff --git a/tests/worker/test_api_executor.py b/tests/worker/test_api_executor.py index 5ab96d558..765c0f2f1 100644 --- a/tests/worker/test_api_executor.py +++ b/tests/worker/test_api_executor.py @@ -55,11 +55,30 @@ def _handler(self, request: httpx.Request) -> httpx.Response: ) +class _BlockingTransport(httpx.MockTransport): + """MockTransport that blocks on the first request, then serves a sequence.""" + + def __init__(self, responses: list[httpx.Response]) -> None: + self.started = threading.Event() + self.release = threading.Event() + self.responses = list(responses) + self.calls = 0 + super().__init__(self._handler) + + def _handler(self, request: httpx.Request) -> httpx.Response: + self.calls += 1 + if self.calls == 1: + self.started.set() + self.release.wait() + return self.responses.pop(0) + + def _run( executor: APIExecutor, task: WorkerTaskMessage, transport: httpx.MockTransport ) -> None: executor._cancel_event = threading.Event() executor._cancel_lock = threading.Lock() + executor._active_task_id = None executor._cancelled_task_id = None with patch.object( APIExecutor, "_get_client", return_value=httpx.Client(transport=transport) @@ -244,6 +263,7 @@ def test_cancelled_task_stops_retrying(self) -> None: executor = APIExecutor.__new__(APIExecutor) executor._cancel_event = threading.Event() executor._cancel_lock = threading.Lock() + executor._active_task_id = None executor._cancelled_task_id = None task = self._task(retries=3) executor.cancel(task.task_id) @@ -267,6 +287,7 @@ def test_cancel_previous_task_does_not_cancel_next(self) -> None: executor = APIExecutor.__new__(APIExecutor) executor._cancel_event = threading.Event() executor._cancel_lock = threading.Lock() + executor._active_task_id = None executor._cancelled_task_id = None task_a = self._task(retries=0) @@ -282,29 +303,84 @@ def test_cancel_previous_task_does_not_cancel_next(self) -> None: executor.run(task_b, Path("/tmp/out")) assert transport.request is not None - def test_cancel_before_start_still_cancels(self) -> None: - """A cancellation addressed to a task before it starts still cancels it.""" + def test_delayed_cancel_of_previous_task_does_not_cancel_next(self) -> None: + """A late cancellation for a prior task does not cancel a running task.""" executor = APIExecutor.__new__(APIExecutor) executor._cancel_event = threading.Event() executor._cancel_lock = threading.Lock() + executor._active_task_id = None executor._cancelled_task_id = None - task = self._task(retries=3) - task.task_id = "task-b" - executor.cancel(task.task_id) - transport = _SequenceTransport([_error_response(504)]) - with patch.object( - APIExecutor, "_get_client", return_value=httpx.Client(transport=transport) - ): - with pytest.raises(TaskCancelledError): - executor.run(task, Path("/tmp/out")) - assert transport.calls == 0 + task_b = self._task(retries=1) + task_b.task_id = "task-b" + # First request blocks; once released it returns a retryable 503 so the + # loop re-checks the cancel event, then a 200 succeeds. + transport = _BlockingTransport([_error_response(503), _ok_response()]) + + errors: list[BaseException] = [] + + def _run_b() -> None: + try: + with patch.object( + APIExecutor, + "_get_client", + return_value=httpx.Client(transport=transport), + ): + executor.run(task_b, Path("/tmp/out")) + except BaseException as exc: + errors.append(exc) + + thread = threading.Thread(target=_run_b) + thread.start() + assert transport.started.wait(2.0) + executor.cancel("task-a") + transport.release.set() + thread.join(2.0) + assert not thread.is_alive() + assert errors == [] + assert transport.calls == 2 + + def test_cancel_of_active_task_still_cancels(self) -> None: + """A cancellation addressed to the running task still cancels it.""" + executor = APIExecutor.__new__(APIExecutor) + executor._cancel_event = threading.Event() + executor._cancel_lock = threading.Lock() + executor._active_task_id = None + executor._cancelled_task_id = None + + task_b = self._task(retries=3) + task_b.task_id = "task-b" + transport = _BlockingTransport([_error_response(503)]) + + errors: list[BaseException] = [] + + def _run_b() -> None: + try: + with patch.object( + APIExecutor, + "_get_client", + return_value=httpx.Client(transport=transport), + ): + executor.run(task_b, Path("/tmp/out")) + except BaseException as exc: + errors.append(exc) + + thread = threading.Thread(target=_run_b) + thread.start() + assert transport.started.wait(2.0) + executor.cancel("task-b") + transport.release.set() + thread.join(2.0) + assert not thread.is_alive() + assert len(errors) == 1 + assert isinstance(errors[0], TaskCancelledError) def test_cancel_during_backoff_stops_retrying(self) -> None: """A cancellation during the retry backoff aborts well before it ends.""" executor = APIExecutor.__new__(APIExecutor) executor._cancel_event = threading.Event() executor._cancel_lock = threading.Lock() + executor._active_task_id = None executor._cancelled_task_id = None task = self._task(retries=3) From 38a2a07bc2757ce551fdba2d6667a5d99d107a48 Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Sun, 27 Sep 2026 11:41:02 +0700 Subject: [PATCH 05/71] fix(worker): shorten stale cancellation comment to a single line Signed-off-by: Zhengyuan Su Co-Authored-By: Claude Code --- src/worker/executors/api_executor.py | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) diff --git a/src/worker/executors/api_executor.py b/src/worker/executors/api_executor.py index b16bbe7fc..ff953ce50 100644 --- a/src/worker/executors/api_executor.py +++ b/src/worker/executors/api_executor.py @@ -55,8 +55,7 @@ def cancel(self, task_id: str) -> None: """Signal the executor to abort the current request and any retries.""" with self._cancel_lock: if self._active_task_id is None: - # No run in flight: record the id so a cancellation addressed to - # a task that has not started yet still lands when it starts. + # No run in flight: record the id so a pre-start cancel lands. self._cancelled_task_id = task_id self._cancel_event.set() elif self._active_task_id == task_id: From 096a58a40067f32b56e777e974c8aa9d33118cb3 Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Sun, 27 Sep 2026 11:47:20 +0700 Subject: [PATCH 06/71] fix(worker): keep pending cancellations by task id in the API executor A single recorded cancellation id let a late cancel for a prior task overwrite a recorded cancel for the task about to start, so that task ran despite its own cancellation. Track pending cancelled ids in a set guarded by the cancel lock; run() consumes only its own id and clears the rest as stale. Signed-off-by: Zhengyuan Su Co-Authored-By: Claude Code --- src/worker/executors/api_executor.py | 10 +++--- tests/worker/test_api_executor.py | 52 +++++++++++++++++++++------- 2 files changed, 44 insertions(+), 18 deletions(-) diff --git a/src/worker/executors/api_executor.py b/src/worker/executors/api_executor.py index ff953ce50..1eddbbdaa 100644 --- a/src/worker/executors/api_executor.py +++ b/src/worker/executors/api_executor.py @@ -49,14 +49,14 @@ def __init__(self, *args: Any, **kwargs: Any) -> None: self._cancel_event = threading.Event() self._cancel_lock = threading.Lock() self._active_task_id: str | None = None - self._cancelled_task_id: str | None = None + self._pending_cancelled_ids: set[str] = set() def cancel(self, task_id: str) -> None: """Signal the executor to abort the current request and any retries.""" with self._cancel_lock: if self._active_task_id is None: # No run in flight: record the id so a pre-start cancel lands. - self._cancelled_task_id = task_id + self._pending_cancelled_ids.add(task_id) self._cancel_event.set() elif self._active_task_id == task_id: self._cancel_event.set() @@ -169,10 +169,10 @@ def cleanup_after_run(self) -> None: def run(self, task: ExecutorTask, out_dir: Path) -> APIResult: with self._cancel_lock: self._active_task_id = task.task_id - # A prior task's cancellation must not leak into this one. - if self._cancelled_task_id != task.task_id: + # Event is set iff this task's id was pending; other ids are stale. + if task.task_id not in self._pending_cancelled_ids: self._cancel_event.clear() - self._cancelled_task_id = None + self._pending_cancelled_ids.clear() try: return self._run(task, out_dir) finally: diff --git a/tests/worker/test_api_executor.py b/tests/worker/test_api_executor.py index 765c0f2f1..b07caa60b 100644 --- a/tests/worker/test_api_executor.py +++ b/tests/worker/test_api_executor.py @@ -69,7 +69,8 @@ def _handler(self, request: httpx.Request) -> httpx.Response: self.calls += 1 if self.calls == 1: self.started.set() - self.release.wait() + if not self.release.wait(5.0): + raise AssertionError("blocking transport was not released") return self.responses.pop(0) @@ -79,7 +80,7 @@ def _run( executor._cancel_event = threading.Event() executor._cancel_lock = threading.Lock() executor._active_task_id = None - executor._cancelled_task_id = None + executor._pending_cancelled_ids = set() with patch.object( APIExecutor, "_get_client", return_value=httpx.Client(transport=transport) ): @@ -264,7 +265,7 @@ def test_cancelled_task_stops_retrying(self) -> None: executor._cancel_event = threading.Event() executor._cancel_lock = threading.Lock() executor._active_task_id = None - executor._cancelled_task_id = None + executor._pending_cancelled_ids = set() task = self._task(retries=3) executor.cancel(task.task_id) transport = _SequenceTransport([_error_response(504)]) @@ -288,7 +289,7 @@ def test_cancel_previous_task_does_not_cancel_next(self) -> None: executor._cancel_event = threading.Event() executor._cancel_lock = threading.Lock() executor._active_task_id = None - executor._cancelled_task_id = None + executor._pending_cancelled_ids = set() task_a = self._task(retries=0) task_a.task_id = "task-a" @@ -303,13 +304,34 @@ def test_cancel_previous_task_does_not_cancel_next(self) -> None: executor.run(task_b, Path("/tmp/out")) assert transport.request is not None + def test_late_cancel_of_previous_task_does_not_overwrite_next(self) -> None: + """A late cancel for A cannot overwrite a recorded cancel for B.""" + executor = APIExecutor.__new__(APIExecutor) + executor._cancel_event = threading.Event() + executor._cancel_lock = threading.Lock() + executor._active_task_id = None + executor._pending_cancelled_ids = set() + + task_b = self._task(retries=0) + task_b.task_id = "task-b" + executor.cancel(task_b.task_id) + executor.cancel("task-a") + + transport = _RecordingTransport() + with patch.object( + APIExecutor, "_get_client", return_value=httpx.Client(transport=transport) + ): + with pytest.raises(TaskCancelledError): + executor.run(task_b, Path("/tmp/out")) + assert transport.request is None + def test_delayed_cancel_of_previous_task_does_not_cancel_next(self) -> None: """A late cancellation for a prior task does not cancel a running task.""" executor = APIExecutor.__new__(APIExecutor) executor._cancel_event = threading.Event() executor._cancel_lock = threading.Lock() executor._active_task_id = None - executor._cancelled_task_id = None + executor._pending_cancelled_ids = set() task_b = self._task(retries=1) task_b.task_id = "task-b" @@ -332,9 +354,11 @@ def _run_b() -> None: thread = threading.Thread(target=_run_b) thread.start() - assert transport.started.wait(2.0) - executor.cancel("task-a") - transport.release.set() + try: + assert transport.started.wait(2.0) + executor.cancel("task-a") + finally: + transport.release.set() thread.join(2.0) assert not thread.is_alive() assert errors == [] @@ -346,7 +370,7 @@ def test_cancel_of_active_task_still_cancels(self) -> None: executor._cancel_event = threading.Event() executor._cancel_lock = threading.Lock() executor._active_task_id = None - executor._cancelled_task_id = None + executor._pending_cancelled_ids = set() task_b = self._task(retries=3) task_b.task_id = "task-b" @@ -367,9 +391,11 @@ def _run_b() -> None: thread = threading.Thread(target=_run_b) thread.start() - assert transport.started.wait(2.0) - executor.cancel("task-b") - transport.release.set() + try: + assert transport.started.wait(2.0) + executor.cancel("task-b") + finally: + transport.release.set() thread.join(2.0) assert not thread.is_alive() assert len(errors) == 1 @@ -381,7 +407,7 @@ def test_cancel_during_backoff_stops_retrying(self) -> None: executor._cancel_event = threading.Event() executor._cancel_lock = threading.Lock() executor._active_task_id = None - executor._cancelled_task_id = None + executor._pending_cancelled_ids = set() task = self._task(retries=3) task.task_id = "task-b" From 0fb3d26c4063156b4165068c89cd52328a0d18aa Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Fri, 18 Sep 2026 10:37:14 +0700 Subject: [PATCH 07/71] feat: batch multiple HTTP requests in the API executor When spec.data is present, the API executor issues one request per row, substituting each row's prompt for the {{prompt}} placeholder in the request body, and returns the responses row-aligned in APIResult.items. The single-request path is unchanged. Reuses the DataMixin parsing infra shared with the vLLM executor. Co-Authored-By: Claude Code Signed-off-by: Zhengyuan Su --- docs/EXECUTORS.md | 1 + docs/WORKFLOWS.md | 28 +++ examples/templates/api_batch.yaml | 43 ++++ src/shared/schemas/result/__init__.py | 2 + src/shared/schemas/result/catalog.py | 8 +- src/shared/schemas/result/payloads.py | 18 ++ src/shared/tasks/specs/misc.py | 2 + src/worker/executors/api_executor.py | 315 +++++++++++++++++------- tests/worker/test_api_executor_batch.py | 178 +++++++++++++ 9 files changed, 505 insertions(+), 90 deletions(-) create mode 100644 examples/templates/api_batch.yaml create mode 100644 tests/worker/test_api_executor_batch.py diff --git a/docs/EXECUTORS.md b/docs/EXECUTORS.md index a5a822d81..be91c39e9 100644 --- a/docs/EXECUTORS.md +++ b/docs/EXECUTORS.md @@ -17,6 +17,7 @@ The worker resolves `spec.taskType` against an executor registry in | `data_profiling` | `DataProfilingExecutor` | DataFrame profiling | | `data_retrieval` | `DataRetrievalExecutor` | DataFrame loading from sources (`type: sql`, `type: s3`, `type: lumid` with `mode: sql\|s3\|agent` via lumid-data-app; `type: lumid` (mode `sql`/`s3`/`agent`) requires `lumid_data_token`, the bearer forwarded to lumid-data-app) | | `ssh` | `SSHExecutor` | Interactive SSH session or non-interactive container job | +| `api` | `APIExecutor` | HTTP request(s); batches one request per `spec.data` row | | `serve` | `VLLMServeExecutor` | Persistent vLLM API server for a single model | Helper utilities live in `src/worker/executors/utils/` (`artifacts`, diff --git a/docs/WORKFLOWS.md b/docs/WORKFLOWS.md index 404081018..a53e24c6b 100644 --- a/docs/WORKFLOWS.md +++ b/docs/WORKFLOWS.md @@ -96,6 +96,34 @@ spec: parse_json: true ``` +### Batching + +When `spec.data` is present, the task batches: one request is issued per row, +and the results are returned row-aligned in `APIResult.items`. Each row's +prompt is substituted for the `{{prompt}}` placeholder in the request body. +Server-side stage references are `${...}`; `{{prompt}}` is a worker-side +per-row slot, so it is not touched by server-side resolution. A failure in any +row fails the whole task rather than shifting the remaining rows. + +```yaml +spec: + taskType: api + data: + type: list + items: + - Explain vector databases + - Explain attention + api: + method: POST + body: + model: gpt-4o + messages: + - role: user + content: "{{prompt}}" + response: + parse_json: true +``` + ## data_retrieval: type lumid `type: lumid` routes the retrieval through lumid-data-app (HTTP). Three diff --git a/examples/templates/api_batch.yaml b/examples/templates/api_batch.yaml new file mode 100644 index 000000000..5f95a5588 --- /dev/null +++ b/examples/templates/api_batch.yaml @@ -0,0 +1,43 @@ +# api_batch.yaml +# +# Batched API executor demo: one task issues one HTTP request per row in +# spec.data, returning the responses row-aligned in APIResult.items. +# +# Each row's prompt is substituted for the {{prompt}} placeholder in the +# request body. Server-side stage references are ${...}; {{prompt}} is a +# worker-side per-row slot, so it is not touched by server-side resolution. + +apiVersion: flowmesh/v1 +kind: APITask +metadata: + name: api-batch + +spec: + taskType: api + + data: + type: list + items: + - Explain vector databases in simple terms. + - Explain attention in simple terms. + - Explain backpropagation in simple terms. + + api: + method: POST + headers: + Content-Type: application/json + body: + model: gpt-4o + messages: + - role: user + content: "{{prompt}}" + response: + parse_json: true + return_body: true + raise_for_status: true + + output: + destination: + type: http + artifacts: + - results.json diff --git a/src/shared/schemas/result/__init__.py b/src/shared/schemas/result/__init__.py index 2ebe15b76..8839005c6 100644 --- a/src/shared/schemas/result/__init__.py +++ b/src/shared/schemas/result/__init__.py @@ -39,6 +39,7 @@ AgentItem, AgentMetadata, AgentUsage, + APIItem, CostEstimates, DataRetrievalItem, EchoItem, @@ -87,6 +88,7 @@ _model.model_rebuild() __all__ = [ + "APIItem", "APIResult", "AgentBatchSummary", "AgentItem", diff --git a/src/shared/schemas/result/catalog.py b/src/shared/schemas/result/catalog.py index bcdc24ed9..0c950c787 100644 --- a/src/shared/schemas/result/catalog.py +++ b/src/shared/schemas/result/catalog.py @@ -20,6 +20,7 @@ AgentItem, AgentMetadata, AgentUsage, + APIItem, CostEstimates, DataRetrievalItem, EchoItem, @@ -245,7 +246,11 @@ class EchoResult(StrictExecutorResult): class APIResult(StrictExecutorResult): """HTTP request output. ``response_json``/``usage``/``headers`` are the - upstream API's own payloads and stay open mappings.""" + upstream API's own payloads and stay open mappings. + + ``items`` carries one entry per row when the task batches multiple requests + (``spec.data`` present); a single-request task leaves it empty and populates + the scalar fields instead.""" task_type: Literal[TaskType.API] = TaskType.API executor: str @@ -257,6 +262,7 @@ class APIResult(StrictExecutorResult): response_json: Any = Field(default=None, alias="json") usage: dict[str, Any] | None = None text: str | None = None + items: list[APIItem] = Field(default_factory=list) class SSHResult(StrictExecutorResult): diff --git a/src/shared/schemas/result/payloads.py b/src/shared/schemas/result/payloads.py index 7307b2a17..9d2e1e951 100644 --- a/src/shared/schemas/result/payloads.py +++ b/src/shared/schemas/result/payloads.py @@ -207,3 +207,21 @@ class EchoItem(StrictModel): """One echoed value.""" output: JsonValue = None + + +class APIItem(StrictModel): + """One row's HTTP response in a batched API task. + + Mirrors the per-response fields of :class:`APIResult`; ``response_json`` is + the upstream API's own payload and stays an open mapping. + """ + + index: int + url: str + status_code: int + truncated: bool = False + headers: dict[str, str] | None = None + response_json: Any = Field(default=None, alias="json") + usage: dict[str, Any] | None = None + text: str | None = None + prompt: str | None = None diff --git a/src/shared/tasks/specs/misc.py b/src/shared/tasks/specs/misc.py index 92a9253b8..1f8b1963f 100644 --- a/src/shared/tasks/specs/misc.py +++ b/src/shared/tasks/specs/misc.py @@ -14,6 +14,7 @@ class ApiSpecStrict(TaskSpecStrictBase): taskType: Literal[TaskType.API] api: dict[str, Any] | None = None + data: dict[str, Any] | None = None def redact_credentials(self) -> Self: spec = super().redact_credentials() @@ -30,6 +31,7 @@ def has_redacted_credentials(self) -> bool: class ApiSpecTemplate(TaskSpecTemplateBase): taskType: Literal[TaskType.API] api: dict[str, Any] | None = None + data: dict[str, Any] | None = None def redact_credentials(self) -> Self: spec = super().redact_credentials() diff --git a/src/worker/executors/api_executor.py b/src/worker/executors/api_executor.py index 1eddbbdaa..c563aa65f 100644 --- a/src/worker/executors/api_executor.py +++ b/src/worker/executors/api_executor.py @@ -1,3 +1,4 @@ +import json import logging import os import threading @@ -6,18 +7,23 @@ import httpx -from shared.schemas.result import APIResult +from shared.schemas.result import APIItem, APIResult from shared.tasks.specs import ApiSpecStrict from shared.tasks.task_type import TaskType from shared.utils.redact import is_credential_key from .base_executor import ExecutionError, Executor, ExecutorTask, TaskCancelledError +from .mixins.data import DataMixin logger = logging.getLogger(__name__) # Cache key: (base_url, timeout_seconds, verify_tls, follow_redirects) _ClientKey = tuple[str, float, bool, bool] +# Worker-side per-row slot in a batched request body. Server-side stage +# references are ${...}; this is a worker-side token, hence {{...}}. +_PROMPT_PLACEHOLDER = "{{prompt}}" + # Fixed delay between retry attempts. _RETRY_BACKOFF_SEC = 1.0 @@ -27,8 +33,14 @@ def _is_retryable_status(status_code: int) -> bool: return status_code >= 500 or status_code in (408, 429) -class APIExecutor(Executor): - """Performs a single HTTP request defined by task YAML. +class APIExecutor(DataMixin, Executor): + """Performs HTTP requests defined by task YAML. + + A single-request task issues one request from ``spec.api`` and returns one + ``APIResult`` with the scalar fields populated. When ``spec.data`` is + present the task batches: each row's prompt is substituted for the + ``{{prompt}}`` placeholder in the request body and one request is issued + per row, returned row-aligned in ``APIResult.items``. Defaults to the Nebula endpoint via ``NEBULA_API_BASE_URL`` and authenticates with ``NEBULA_API_TOKEN``. ``spec.api.url`` overrides the endpoint and @@ -166,6 +178,123 @@ def cleanup_after_run(self) -> None: """Close the connection pool when the runner deactivates this executor.""" self.close_all_clients() + @staticmethod + def _prompt_to_str(prompt: Any) -> str: + """Render a row's prompt as a string for body substitution.""" + if isinstance(prompt, str): + return prompt + return json.dumps(prompt) + + @classmethod + def _substitute_prompt(cls, value: Any, prompt: str) -> Any: + """Replace ``{{prompt}}`` in the request body with a row's prompt.""" + if isinstance(value, str): + if value == _PROMPT_PLACEHOLDER: + return prompt + return value.replace(_PROMPT_PLACEHOLDER, prompt) + if isinstance(value, dict): + return {k: cls._substitute_prompt(v, prompt) for k, v in value.items()} + if isinstance(value, list): + return [cls._substitute_prompt(v, prompt) for v in value] + return value + + def _build_request_kwargs( + self, api_cfg: dict[str, Any], prompt: str | None + ) -> dict[str, Any]: + """Build httpx request kwargs from ``spec.api``, substituting the row + prompt when batching.""" + json_payload = api_cfg.get("json") + body = api_cfg.get("body") + data_payload = api_cfg.get("data") + + if json_payload is not None and body is not None: + raise ExecutionError( + "spec.api.json and spec.api.body are mutually exclusive" + ) + + request_kwargs: dict[str, Any] = {} + if json_payload is not None: + request_kwargs["json"] = ( + self._substitute_prompt(json_payload, prompt) + if prompt is not None + else json_payload + ) + elif body is not None: + if isinstance(body, (dict, list)): + request_kwargs["json"] = ( + self._substitute_prompt(body, prompt) + if prompt is not None + else body + ) + else: + request_kwargs["content"] = ( + self._substitute_prompt(body, prompt) + if prompt is not None + else body + ) + elif data_payload is not None: + request_kwargs["data"] = ( + self._substitute_prompt(data_payload, prompt) + if prompt is not None + else data_payload + ) + return request_kwargs + + def _parse_response( + self, + resp: httpx.Response, + *, + response_cfg: dict[str, Any], + max_body_bytes: int, + ) -> tuple[APIItem, str | None]: + """Turn one HTTP response into an APIItem, applying response config. + + Returns the item and the raw body text (used for error messages). + """ + body_bytes = resp.content + truncated = False + if max_body_bytes is not None and len(body_bytes) > max_body_bytes: + body_bytes = body_bytes[:max_body_bytes] + truncated = True + + item = APIItem( + index=0, + url=str(resp.url), + status_code=resp.status_code, + truncated=truncated, + ) + + if response_cfg.get("include_headers", False): + item.headers = dict(resp.headers) + + body_text: str | None = None + if response_cfg.get("return_body", True): + encoding = resp.encoding or "utf-8" + body_text = body_bytes.decode(encoding, errors="replace") + + if response_cfg.get("parse_json", True): + item.response_json = resp.json() + if not isinstance(item.response_json, dict): + raise ExecutionError("Response is not a valid JSON mapping") + usage = item.response_json.get("usage") + if not isinstance(usage, dict): + raise ExecutionError( + "spec.api.response.parse_json is true but response JSON " + f"does not contain usage info: {item.response_json}" + ) + item.usage = usage + try: + item.text = item.response_json["choices"][0]["message"]["content"] + except Exception as exc: + raise ExecutionError( + "spec.api.response.parse_json is true but response JSON " + f"does not contain message.content: {item.response_json}" + ) from exc + elif response_cfg.get("return_body", True): + item.text = body_text + + return item, body_text + def run(self, task: ExecutorTask, out_dir: Path) -> APIResult: with self._cancel_lock: self._active_task_id = task.task_id @@ -218,105 +347,113 @@ def _run(self, task: ExecutorTask, out_dir: Path) -> APIResult: verify_tls = api_cfg.get("verify_tls", True) follow_redirects = api_cfg.get("follow_redirects", True) - body = api_cfg.get("body") - json_payload = api_cfg.get("json") - data_payload = api_cfg.get("data") - - if json_payload is not None and body is not None: - raise ExecutionError( - "spec.api.json and spec.api.body are mutually exclusive" - ) - - request_kwargs: dict[str, Any] = {} - if json_payload is not None: - request_kwargs["json"] = json_payload - elif body is not None: - if isinstance(body, (dict, list)): - request_kwargs["json"] = body - else: - request_kwargs["content"] = body - elif data_payload is not None: - request_kwargs["data"] = data_payload - response_cfg = api_cfg.get("response") or {} if response_cfg and not isinstance(response_cfg, dict): raise ExecutionError("spec.api.response must be a mapping") - include_headers = bool(response_cfg.get("include_headers", False)) - # return_body is a JSON backdoor: keep raw text when JSON isn't usable. - return_body = bool(response_cfg.get("return_body", True)) - parse_json = bool(response_cfg.get("parse_json", True)) - raise_for_status = bool(response_cfg.get("raise_for_status", True)) max_body_bytes = int(response_cfg.get("max_body_bytes", 200000)) + raise_for_status = bool(response_cfg.get("raise_for_status", True)) retries = api_cfg.get("retries", 0) if not isinstance(retries, int) or isinstance(retries, bool) or retries < 0: raise ExecutionError("spec.api.retries must be a non-negative integer") - try: - base = self._base_url(str(url)) - client = self._get_client(base, timeout, verify_tls, follow_redirects) - resp = self._request_with_retries( - client, - method, - str(url), - headers, - params, - request_kwargs, - retries, - ) - except httpx.RequestError as exc: - raise ExecutionError(f"API request failed: {exc}", retryable=True) from exc - - body_bytes = resp.content - truncated = False - if max_body_bytes is not None and len(body_bytes) > max_body_bytes: - body_bytes = body_bytes[:max_body_bytes] - truncated = True - - result = APIResult( - ok=resp.is_success, - executor=self.name, - method=method, - url=str(resp.url), - status_code=resp.status_code, - truncated=truncated, - ) - - if include_headers: - result.headers = dict(resp.headers) - - body_text: str | None = None - if return_body: - encoding = resp.encoding or "utf-8" - body_text = body_bytes.decode(encoding, errors="replace") + base = self._base_url(str(url)) + client = self._get_client(base, timeout, verify_tls, follow_redirects) - if parse_json: - result.response_json = resp.json() - if not isinstance(result.response_json, dict): - raise ExecutionError("Response is not a valid JSON mapping") - usage = result.response_json.get("usage") - if not isinstance(usage, dict): - raise ExecutionError( - "spec.api.response.parse_json is true but response JSON " - f"does not contain usage info: {result.response_json}" - ) - result.usage = usage + data_cfg = spec.data + if data_cfg is None: + # Single-request path: build kwargs once, issue one request. + request_kwargs = self._build_request_kwargs(api_cfg, None) try: - result.text = result.response_json["choices"][0]["message"]["content"] - except Exception as exc: + resp = self._request_with_retries( + client, + method, + str(url), + headers, + params, + request_kwargs, + retries, + ) + except httpx.RequestError as exc: raise ExecutionError( - "spec.api.response.parse_json is true but response JSON " - f"does not contain message.content: {result.response_json}" + f"API request failed: {exc}", retryable=True ) from exc - elif return_body: - result.text = body_text - if raise_for_status and resp.is_error: - message = f"API request returned status {resp.status_code}" - if body_text: - message = f"{message}: {body_text[:200]}" - retryable = _is_retryable_status(resp.status_code) - raise ExecutionError(message, retryable=retryable) + item, body_text = self._parse_response( + resp, + response_cfg=response_cfg, + max_body_bytes=max_body_bytes, + ) + result = APIResult( + ok=resp.is_success, + executor=self.name, + method=method, + url=str(resp.url), + status_code=resp.status_code, + truncated=item.truncated, + headers=item.headers, + text=item.text, + ) + result.response_json = item.response_json + result.usage = item.usage + + if raise_for_status and resp.is_error: + message = f"API request returned status {resp.status_code}" + if body_text: + message = f"{message}: {body_text[:200]}" + retryable = _is_retryable_status(resp.status_code) + raise ExecutionError(message, retryable=retryable) + + return result + + # Batch path: one request per row, row-aligned in items. + entry = self._collect_prompts_for_spec(spec, task_id=task.task_id) + prompts = entry.prompts + if not prompts: + raise ExecutionError("spec.data produced no rows to batch") + + items: list[APIItem] = [] + for idx, prompt in enumerate(prompts): + prompt_str = self._prompt_to_str(prompt) + request_kwargs = self._build_request_kwargs(api_cfg, prompt_str) + try: + resp = self._request_with_retries( + client, + method, + str(url), + headers, + params, + request_kwargs, + retries, + ) + except httpx.RequestError as exc: + raise ExecutionError( + f"API request failed (row {idx}): {exc}", retryable=True + ) from exc - return result + item, body_text = self._parse_response( + resp, + response_cfg=response_cfg, + max_body_bytes=max_body_bytes, + ) + item.index = idx + item.prompt = prompt_str + items.append(item) + + if raise_for_status and resp.is_error: + message = f"API request returned status {resp.status_code} (row {idx})" + if body_text: + message = f"{message}: {body_text[:200]}" + retryable = _is_retryable_status(resp.status_code) + raise ExecutionError(message, retryable=retryable) + + return APIResult( + ok=True, + executor=self.name, + method=method, + url=str(url), + status_code=items[0].status_code, + truncated=any(item.truncated for item in items), + items=items, + ) diff --git a/tests/worker/test_api_executor_batch.py b/tests/worker/test_api_executor_batch.py new file mode 100644 index 000000000..7a438e02d --- /dev/null +++ b/tests/worker/test_api_executor_batch.py @@ -0,0 +1,178 @@ +"""Tests for the API executor's batch mode (one task, N row-aligned requests).""" + +from pathlib import Path +from unittest.mock import patch + +import httpx +import pytest + +from shared.tasks.worker_message import WorkerTaskMessage +from worker.executors.api_executor import APIExecutor +from worker.executors.base_executor import ExecutionError + + +def _task_message(**spec_updates: object) -> WorkerTaskMessage: + payload = { + "task_id": "task-api-batch", + "workflow_id": "wf-1", + "owner_id": "owner", + "assigned_worker": "worker-1", + "dispatched_at": "2026-03-22T00:00:00Z", + "task": { + "apiVersion": "mloc/v1", + "kind": "Task", + "metadata": {"name": "wf:api"}, + "spec": { + "taskType": "api", + "api": { + "method": "POST", + "body": {"messages": [{"role": "user", "content": "{{prompt}}"}]}, + **spec_updates, + }, + }, + }, + } + return WorkerTaskMessage.model_validate(payload) + + +class _RecordingTransport(httpx.MockTransport): + """MockTransport that records every request it served, in order.""" + + def __init__(self) -> None: + self.requests: list[httpx.Request] = [] + super().__init__(self._handler) + + def _handler(self, request: httpx.Request) -> httpx.Response: + self.requests.append(request) + return httpx.Response( + 200, + json={ + "choices": [{"message": {"content": "hello"}}], + "usage": {"total_tokens": 3}, + }, + ) + + +def _run( + executor: APIExecutor, task: WorkerTaskMessage, transport: httpx.MockTransport +): + with patch.object( + APIExecutor, "_get_client", return_value=httpx.Client(transport=transport) + ): + return executor.run(task, Path("/tmp/out")) + + +def _batch_task(items: list[object]) -> WorkerTaskMessage: + payload = { + "task_id": "task-api-batch", + "workflow_id": "wf-1", + "owner_id": "owner", + "assigned_worker": "worker-1", + "dispatched_at": "2026-03-22T00:00:00Z", + "task": { + "apiVersion": "mloc/v1", + "kind": "Task", + "metadata": {"name": "wf:api"}, + "spec": { + "taskType": "api", + "api": { + "method": "POST", + "body": {"messages": [{"role": "user", "content": "{{prompt}}"}]}, + }, + "data": {"type": "list", "items": items}, + }, + }, + } + return WorkerTaskMessage.model_validate(payload) + + +class TestBatch: + @pytest.fixture(autouse=True) + def _nebula_env(self, monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("NEBULA_API_BASE_URL", "https://nebula.example.com") + monkeypatch.setenv("NEBULA_API_TOKEN", "nebula-token") + + def test_issues_one_request_per_row_in_order(self) -> None: + task = _batch_task(["first", "second", "third"]) + transport = _RecordingTransport() + result = _run(APIExecutor.__new__(APIExecutor), task, transport) + assert len(transport.requests) == 3 + # Row order is preserved: request i carries row i's prompt. + for idx, prompt in enumerate(["first", "second", "third"]): + body = transport.requests[idx].read() + assert prompt.encode() in body + assert [item.index for item in result.items] == [0, 1, 2] + assert [item.text for item in result.items] == ["hello"] * 3 + + def test_single_row_batches_to_one_item(self) -> None: + task = _batch_task(["only"]) + transport = _RecordingTransport() + result = _run(APIExecutor.__new__(APIExecutor), task, transport) + assert len(transport.requests) == 1 + assert len(result.items) == 1 + assert result.items[0].index == 0 + assert result.items[0].prompt == "only" + + def test_row_failure_fails_whole_task_without_shifting(self) -> None: + class _FailSecond(httpx.MockTransport): + def __init__(self) -> None: + self.requests: list[httpx.Request] = [] + super().__init__(self._handler) + + def _handler(self, request: httpx.Request) -> httpx.Response: + self.requests.append(request) + if len(self.requests) == 2: + return httpx.Response( + 500, + json={ + "choices": [{"message": {"content": "boom"}}], + "usage": {"total_tokens": 1}, + }, + ) + return httpx.Response( + 200, + json={ + "choices": [{"message": {"content": "ok"}}], + "usage": {"total_tokens": 1}, + }, + ) + + task = _batch_task(["a", "b", "c"]) + transport = _FailSecond() + with pytest.raises(ExecutionError, match="row 1"): + _run(APIExecutor.__new__(APIExecutor), task, transport) + # The failing row aborts the task; no partial result is returned. + assert len(transport.requests) == 2 + + def test_placeholder_not_required_for_scalar_body(self) -> None: + """A batch task whose body has no placeholder still issues N requests.""" + payload = { + "task_id": "task-api-batch", + "workflow_id": "wf-1", + "owner_id": "owner", + "assigned_worker": "worker-1", + "dispatched_at": "2026-03-22T00:00:00Z", + "task": { + "apiVersion": "mloc/v1", + "kind": "Task", + "metadata": {"name": "wf:api"}, + "spec": { + "taskType": "api", + "api": { + "method": "POST", + "body": {"messages": [{"role": "user", "content": "static"}]}, + }, + "data": {"type": "list", "items": ["a", "b"]}, + }, + }, + } + task = WorkerTaskMessage.model_validate(payload) + transport = _RecordingTransport() + result = _run(APIExecutor.__new__(APIExecutor), task, transport) + assert len(transport.requests) == 2 + assert len(result.items) == 2 + + def test_no_rows_raises(self) -> None: + task = _batch_task([]) + with pytest.raises(ExecutionError, match="no rows"): + _run(APIExecutor.__new__(APIExecutor), task, _RecordingTransport()) From d2fe09c599f277b2e3c951722ee248bb6d370579 Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Fri, 18 Sep 2026 12:23:27 +0700 Subject: [PATCH 08/71] fix(sdk): mirror APIResult.items and APIItem in the SDK models The batch change added `items: list[APIItem]` to the server-side APIResult without the SDK counterpart, so tests/sdk/test_schema_compat.py failed with "APIResult missing server fields: ['items']". The SDK mirrors every server result model and a compat test enforces that they stay in step. This adds APIItem to the SDK payloads beside the existing InferenceItem and Omni* items, carrying the identical `alias="json"` so the wire-alias check passes too, adds `items` to the SDK APIResult, and registers APIItem in _RESULT_MODEL_NAMES so the drift guard covers it from now on rather than only the field that happened to break. Co-Authored-By: Claude Opus 5 (1M context) Signed-off-by: Zhengyuan Su --- sdk/src/flowmesh/models/__init__.py | 2 ++ sdk/src/flowmesh/models/result/__init__.py | 2 ++ sdk/src/flowmesh/models/result/catalog.py | 2 ++ sdk/src/flowmesh/models/result/payloads.py | 12 ++++++++++++ tests/sdk/test_schema_compat.py | 1 + 5 files changed, 19 insertions(+) diff --git a/sdk/src/flowmesh/models/__init__.py b/sdk/src/flowmesh/models/__init__.py index ae22cefc4..69506c497 100644 --- a/sdk/src/flowmesh/models/__init__.py +++ b/sdk/src/flowmesh/models/__init__.py @@ -28,6 +28,7 @@ AgentResult, AgentUsage, AnyExecutorResult, + APIItem, APIResult, BaseExecutorResult, CostEstimates, @@ -104,6 +105,7 @@ ) __all__ = [ + "APIItem", "APIResult", "ActiveWaitBreakdown", "AgentBatchSummary", diff --git a/sdk/src/flowmesh/models/result/__init__.py b/sdk/src/flowmesh/models/result/__init__.py index 8d94b23d7..ecd772519 100644 --- a/sdk/src/flowmesh/models/result/__init__.py +++ b/sdk/src/flowmesh/models/result/__init__.py @@ -39,6 +39,7 @@ AgentItem, AgentMetadata, AgentUsage, + APIItem, CostEstimates, DataRetrievalItem, EchoItem, @@ -88,6 +89,7 @@ _model.model_rebuild() __all__ = [ + "APIItem", "APIResult", "AgentBatchSummary", "AgentItem", diff --git a/sdk/src/flowmesh/models/result/catalog.py b/sdk/src/flowmesh/models/result/catalog.py index 12948deac..dcdafa54b 100644 --- a/sdk/src/flowmesh/models/result/catalog.py +++ b/sdk/src/flowmesh/models/result/catalog.py @@ -18,6 +18,7 @@ AgentItem, AgentMetadata, AgentUsage, + APIItem, CostEstimates, DataRetrievalItem, EchoItem, @@ -212,6 +213,7 @@ class APIResult(StrictExecutorResult): response_json: Any = Field(default=None, alias="json") usage: dict[str, Any] | None = None text: str | None = None + items: list[APIItem] = Field(default_factory=list) class SSHResult(StrictExecutorResult): diff --git a/sdk/src/flowmesh/models/result/payloads.py b/sdk/src/flowmesh/models/result/payloads.py index c512668e4..49d0da8ac 100644 --- a/sdk/src/flowmesh/models/result/payloads.py +++ b/sdk/src/flowmesh/models/result/payloads.py @@ -158,3 +158,15 @@ class RagQuery(StrictModel): class EchoItem(StrictModel): output: JsonValue = None + + +class APIItem(StrictModel): + index: int + url: str + status_code: int + truncated: bool = False + headers: dict[str, str] | None = None + response_json: Any = Field(default=None, alias="json") + usage: dict[str, Any] | None = None + text: str | None = None + prompt: str | None = None diff --git a/tests/sdk/test_schema_compat.py b/tests/sdk/test_schema_compat.py index 744d12f38..b30e557f0 100644 --- a/tests/sdk/test_schema_compat.py +++ b/tests/sdk/test_schema_compat.py @@ -152,6 +152,7 @@ "RagHit", "RagQuery", "EchoItem", + "APIItem", ] RESULT_MODEL_PAIRS = [ From b5583eadee2e064ae90c4fcebae4950bb3ef6f85 Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Fri, 18 Sep 2026 13:01:07 +0700 Subject: [PATCH 09/71] refactor: single parallel path for the API executor Require spec.data (as the vLLM executor does) and issue one request per row in parallel, returning row-aligned APIResult.items. A single request is a one-row spec.data. The request skeleton is built once and each worker only substitutes its prompt into the prepared body. Concurrency is bounded by spec.api.concurrency (default 8) and the client connection pool is sized to it. Migrate api_two_stage.yaml and the n8n parser to carry spec.data. Co-Authored-By: Claude Code Signed-off-by: Zhengyuan Su --- docs/EXECUTORS.md | 2 +- docs/WORKFLOWS.md | 9 ++- examples/templates/api_two_stage.yaml | 20 +++-- src/server/task/n8n_parser.py | 15 ++-- src/worker/executors/api_executor.py | 103 ++++++++++-------------- tests/server/task/test_n8n_parser.py | 8 +- tests/worker/test_api_executor.py | 3 +- tests/worker/test_api_executor_batch.py | 93 ++++++++++++++++++++- 8 files changed, 175 insertions(+), 78 deletions(-) diff --git a/docs/EXECUTORS.md b/docs/EXECUTORS.md index be91c39e9..4e04fa2c6 100644 --- a/docs/EXECUTORS.md +++ b/docs/EXECUTORS.md @@ -17,7 +17,7 @@ The worker resolves `spec.taskType` against an executor registry in | `data_profiling` | `DataProfilingExecutor` | DataFrame profiling | | `data_retrieval` | `DataRetrievalExecutor` | DataFrame loading from sources (`type: sql`, `type: s3`, `type: lumid` with `mode: sql\|s3\|agent` via lumid-data-app; `type: lumid` (mode `sql`/`s3`/`agent`) requires `lumid_data_token`, the bearer forwarded to lumid-data-app) | | `ssh` | `SSHExecutor` | Interactive SSH session or non-interactive container job | -| `api` | `APIExecutor` | HTTP request(s); batches one request per `spec.data` row | +| `api` | `APIExecutor` | One parallel HTTP request per `spec.data` row (spec.data required) | | `serve` | `VLLMServeExecutor` | Persistent vLLM API server for a single model | Helper utilities live in `src/worker/executors/utils/` (`artifacts`, diff --git a/docs/WORKFLOWS.md b/docs/WORKFLOWS.md index a53e24c6b..6fc6c9164 100644 --- a/docs/WORKFLOWS.md +++ b/docs/WORKFLOWS.md @@ -71,7 +71,13 @@ contract. ## API task -`taskType: api` performs a single HTTP request. By default it routes to the Nebula endpoint and authenticates with the worker's `NEBULA_API_TOKEN`. +`taskType: api` issues one HTTP request per row of `spec.data`, in parallel, +and returns the responses row-aligned in `APIResult.items`. A single request +is a one-row `spec.data`. `spec.data` is required, exactly as for the vLLM +executor; it supports the same data types (`list`, `dataset`, `graph_template`, +`dataframe`). + +By default it routes to the Nebula endpoint and authenticates with the worker's `NEBULA_API_TOKEN`. `spec.api.url` overrides the endpoint; when absent, the executor uses `NEBULA_API_BASE_URL` (appending `/v1/chat/completions`). `spec.api.headers` may supply an `Authorization` header directly. @@ -104,6 +110,7 @@ prompt is substituted for the `{{prompt}}` placeholder in the request body. Server-side stage references are `${...}`; `{{prompt}}` is a worker-side per-row slot, so it is not touched by server-side resolution. A failure in any row fails the whole task rather than shifting the remaining rows. +`spec.api.concurrency` bounds the number of in-flight requests (default 8). ```yaml spec: diff --git a/examples/templates/api_two_stage.yaml b/examples/templates/api_two_stage.yaml index a8a83f812..7eb30e913 100644 --- a/examples/templates/api_two_stage.yaml +++ b/examples/templates/api_two_stage.yaml @@ -4,9 +4,11 @@ # Stage 1 calls a chat completion endpoint and returns raw text. # Stage 2 sends Stage 1's returned text as the next prompt. # -# Stage 1 names a custom endpoint via spec.api.url and supplies its own -# Authorization header. Stage 2 omits url and header, so it uses the Nebula -# defaults (NEBULA_API_BASE_URL + NEBULA_API_TOKEN). +# spec.data is required for an api task and yields one row per request; a +# single request is a one-row list. Stage 1 names a custom endpoint via +# spec.api.url and supplies its own Authorization header. Stage 2 omits url +# and header, so it uses the Nebula defaults (NEBULA_API_BASE_URL + +# NEBULA_API_TOKEN). apiVersion: flowmesh/v1 kind: APITask @@ -19,6 +21,10 @@ spec: stages: - name: stage-1 spec: + data: + type: list + items: + - Please explain vector databases in simple terms. api: url: https://api.example.com/v1/chat/completions method: POST @@ -29,7 +35,7 @@ spec: model: gpt-4o messages: - role: user - content: Please explain vector databases in simple terms. + content: "{{prompt}}" response: parse_json: false # return_body is a JSON backdoor: keep raw text when JSON isn't usable. @@ -38,6 +44,10 @@ spec: - name: stage-2 spec: + data: + type: list + items: + - "The previous stage's response is as follows. Please provide a simpler explanation: \n${stage-1.text}" api: method: POST headers: @@ -46,7 +56,7 @@ spec: model: gpt-4o messages: - role: user - content: "The previous stage's response is as follows. Please provide a simpler explanation: \n${stage-1.text}" + content: "{{prompt}}" response: parse_json: true # return_body is a JSON backdoor: keep raw text when JSON isn't usable. diff --git a/src/server/task/n8n_parser.py b/src/server/task/n8n_parser.py index 35485a0b9..fac95a4d0 100644 --- a/src/server/task/n8n_parser.py +++ b/src/server/task/n8n_parser.py @@ -144,7 +144,7 @@ def translate_n8n_workflow(payload: dict[str, Any]) -> dict[str, Any]: name = node["name"] spec = { "taskType": "api", - "api": _build_api_node_spec( + **_build_api_node_spec( node, openai_model_nodes, incoming, @@ -239,12 +239,12 @@ def _build_api_node_spec( headers = { "Content-Type": "application/json", } - spec = { + api_spec = { "method": "POST", "headers": headers, "body": { "model": model_id, - "messages": [{"role": "user", "content": prompt_text}], + "messages": [{"role": "user", "content": "{{prompt}}"}], }, "response": { "parse_json": True, @@ -253,11 +253,14 @@ def _build_api_node_spec( }, } if api_key := credential_data.get("api_key"): - spec["key"] = api_key + api_spec["key"] = api_key headers["Authorization"] = f"Bearer {api_key}" if api_url := credential_data.get("url"): - spec["url"] = api_url - return spec + api_spec["url"] = api_url + return { + "data": {"type": "list", "items": [prompt_text]}, + "api": api_spec, + } def _resolve_api_model_id( diff --git a/src/worker/executors/api_executor.py b/src/worker/executors/api_executor.py index c563aa65f..915574631 100644 --- a/src/worker/executors/api_executor.py +++ b/src/worker/executors/api_executor.py @@ -2,6 +2,7 @@ import logging import os import threading +from concurrent.futures import ThreadPoolExecutor, as_completed from pathlib import Path from typing import Any, ClassVar @@ -27,6 +28,10 @@ # Fixed delay between retry attempts. _RETRY_BACKOFF_SEC = 1.0 +# Default cap on parallel row requests and the client connection pool. The +# pool is sized to this so parallel requests never queue on connections. +_MAX_CONCURRENCY = 8 + def _is_retryable_status(status_code: int) -> bool: """Whether an HTTP status is transient and worth retrying.""" @@ -36,11 +41,10 @@ def _is_retryable_status(status_code: int) -> bool: class APIExecutor(DataMixin, Executor): """Performs HTTP requests defined by task YAML. - A single-request task issues one request from ``spec.api`` and returns one - ``APIResult`` with the scalar fields populated. When ``spec.data`` is - present the task batches: each row's prompt is substituted for the - ``{{prompt}}`` placeholder in the request body and one request is issued - per row, returned row-aligned in ``APIResult.items``. + ``spec.data`` is required and yields one row per request: each row's prompt + is substituted for the ``{{prompt}}`` placeholder in the request body and + one request is issued per row, returned row-aligned in ``APIResult.items``. + A single request is a one-row ``spec.data``. Defaults to the Nebula endpoint via ``NEBULA_API_BASE_URL`` and authenticates with ``NEBULA_API_TOKEN``. ``spec.api.url`` overrides the endpoint and @@ -101,11 +105,17 @@ def _get_client( client = cls._clients.get(key) if client is not None and not client.is_closed: return client - # Create a new client for this combination + # Create a new client for this combination. The connection pool is + # sized to the max concurrency so parallel row requests never queue + # on connections (which would make the parallelism imaginary). client = httpx.Client( timeout=timeout, verify=verify_tls, follow_redirects=follow_redirects, + limits=httpx.Limits( + max_connections=_MAX_CONCURRENCY, + max_keepalive_connections=_MAX_CONCURRENCY, + ), ) cls._clients[key] = client logger.debug( @@ -361,62 +371,24 @@ def _run(self, task: ExecutorTask, out_dir: Path) -> APIResult: base = self._base_url(str(url)) client = self._get_client(base, timeout, verify_tls, follow_redirects) - data_cfg = spec.data - if data_cfg is None: - # Single-request path: build kwargs once, issue one request. - request_kwargs = self._build_request_kwargs(api_cfg, None) - try: - resp = self._request_with_retries( - client, - method, - str(url), - headers, - params, - request_kwargs, - retries, - ) - except httpx.RequestError as exc: - raise ExecutionError( - f"API request failed: {exc}", retryable=True - ) from exc - - item, body_text = self._parse_response( - resp, - response_cfg=response_cfg, - max_body_bytes=max_body_bytes, - ) - result = APIResult( - ok=resp.is_success, - executor=self.name, - method=method, - url=str(resp.url), - status_code=resp.status_code, - truncated=item.truncated, - headers=item.headers, - text=item.text, - ) - result.response_json = item.response_json - result.usage = item.usage - - if raise_for_status and resp.is_error: - message = f"API request returned status {resp.status_code}" - if body_text: - message = f"{message}: {body_text[:200]}" - retryable = _is_retryable_status(resp.status_code) - raise ExecutionError(message, retryable=retryable) - - return result - - # Batch path: one request per row, row-aligned in items. + # spec.data is required and yields one row per request. entry = self._collect_prompts_for_spec(spec, task_id=task.task_id) prompts = entry.prompts if not prompts: - raise ExecutionError("spec.data produced no rows to batch") + raise ExecutionError("spec.data produced no rows") + + # Build the request skeleton once; each worker only substitutes its + # prompt into the prepared body and issues the request. + request_kwargs = self._build_request_kwargs(api_cfg, None) - items: list[APIItem] = [] - for idx, prompt in enumerate(prompts): + concurrency = int(api_cfg.get("concurrency", _MAX_CONCURRENCY)) + if concurrency < 1: + raise ExecutionError("spec.api.concurrency must be >= 1") + concurrency = min(concurrency, _MAX_CONCURRENCY) + + def _issue(idx: int, prompt: Any) -> APIItem: prompt_str = self._prompt_to_str(prompt) - request_kwargs = self._build_request_kwargs(api_cfg, prompt_str) + kwargs = self._substitute_prompt(request_kwargs, prompt_str) try: resp = self._request_with_retries( client, @@ -424,7 +396,7 @@ def _run(self, task: ExecutorTask, out_dir: Path) -> APIResult: str(url), headers, params, - request_kwargs, + kwargs, retries, ) except httpx.RequestError as exc: @@ -439,7 +411,6 @@ def _run(self, task: ExecutorTask, out_dir: Path) -> APIResult: ) item.index = idx item.prompt = prompt_str - items.append(item) if raise_for_status and resp.is_error: message = f"API request returned status {resp.status_code} (row {idx})" @@ -448,6 +419,20 @@ def _run(self, task: ExecutorTask, out_dir: Path) -> APIResult: retryable = _is_retryable_status(resp.status_code) raise ExecutionError(message, retryable=retryable) + return item + + results: dict[int, APIItem] = {} + with ThreadPoolExecutor(max_workers=concurrency) as pool: + futures = { + pool.submit(_issue, idx, prompt): idx + for idx, prompt in enumerate(prompts) + } + for future in as_completed(futures): + idx = futures[future] + results[idx] = future.result() + + items = [results[idx] for idx in range(len(prompts))] + return APIResult( ok=True, executor=self.name, diff --git a/tests/server/task/test_n8n_parser.py b/tests/server/task/test_n8n_parser.py index cd8bf7bf5..e8d3abe04 100644 --- a/tests/server/task/test_n8n_parser.py +++ b/tests/server/task/test_n8n_parser.py @@ -36,9 +36,11 @@ def test_simple_openai_node(self) -> None: assert api["method"] == "POST" assert api["body"]["model"] == "gpt-4" - # Prompt content preserved in messages - messages = api["body"]["messages"] - assert any("Hello, world!" in m.get("content", "") for m in messages) + # Prompt content preserved as a one-row spec.data + assert spec["data"]["type"] == "list" + assert spec["data"]["items"] == ["Hello, world!"] + # Body carries the worker-side per-row placeholder + assert api["body"]["messages"][0]["content"] == "{{prompt}}" def test_no_task_nodes_raises_value_error(self) -> None: """Workflow with no recognized task nodes should raise ValueError.""" diff --git a/tests/worker/test_api_executor.py b/tests/worker/test_api_executor.py index b07caa60b..cd9c545a2 100644 --- a/tests/worker/test_api_executor.py +++ b/tests/worker/test_api_executor.py @@ -26,9 +26,10 @@ def _task_message(**spec_updates: object) -> WorkerTaskMessage: "metadata": {"name": "wf:api"}, "spec": { "taskType": "api", + "data": {"type": "list", "items": ["hi"]}, "api": { "method": "POST", - "body": {"messages": [{"role": "user", "content": "hi"}]}, + "body": {"messages": [{"role": "user", "content": "{{prompt}}"}]}, **spec_updates, }, }, diff --git a/tests/worker/test_api_executor_batch.py b/tests/worker/test_api_executor_batch.py index 7a438e02d..9e07ca147 100644 --- a/tests/worker/test_api_executor_batch.py +++ b/tests/worker/test_api_executor_batch.py @@ -1,5 +1,6 @@ """Tests for the API executor's batch mode (one task, N row-aligned requests).""" +import time from pathlib import Path from unittest.mock import patch @@ -141,8 +142,9 @@ def _handler(self, request: httpx.Request) -> httpx.Response: transport = _FailSecond() with pytest.raises(ExecutionError, match="row 1"): _run(APIExecutor.__new__(APIExecutor), task, transport) - # The failing row aborts the task; no partial result is returned. - assert len(transport.requests) == 2 + # All rows are issued in parallel; the failing row aborts the task and + # no partial result is returned. + assert len(transport.requests) == 3 def test_placeholder_not_required_for_scalar_body(self) -> None: """A batch task whose body has no placeholder still issues N requests.""" @@ -176,3 +178,90 @@ def test_no_rows_raises(self) -> None: task = _batch_task([]) with pytest.raises(ExecutionError, match="no rows"): _run(APIExecutor.__new__(APIExecutor), task, _RecordingTransport()) + + def test_missing_data_raises(self) -> None: + """spec.data is required; an api task without it fails closed.""" + payload = { + "task_id": "task-api-batch", + "workflow_id": "wf-1", + "owner_id": "owner", + "assigned_worker": "worker-1", + "dispatched_at": "2026-03-22T00:00:00Z", + "task": { + "apiVersion": "mloc/v1", + "kind": "Task", + "metadata": {"name": "wf:api"}, + "spec": { + "taskType": "api", + "api": { + "method": "POST", + "body": {"messages": [{"role": "user", "content": "hi"}]}, + }, + }, + }, + } + task = WorkerTaskMessage.model_validate(payload) + with pytest.raises(ExecutionError, match="spec.data is required"): + _run(APIExecutor.__new__(APIExecutor), task, _RecordingTransport()) + + def test_requests_issue_in_parallel(self) -> None: + """N rows must take ~one row's latency, not N x, on a network-bound path. + + A transport that sleeps per request proves the requests overlap: with + serial issue this would take N x the sleep, with parallel issue it takes + roughly one sleep (plus scheduling overhead). + """ + + class _SlowTransport(httpx.MockTransport): + def __init__(self) -> None: + self.requests: list[httpx.Request] = [] + super().__init__(self._handler) + + def _handler(self, request: httpx.Request) -> httpx.Response: + self.requests.append(request) + time.sleep(0.2) + return httpx.Response( + 200, + json={ + "choices": [{"message": {"content": "hello"}}], + "usage": {"total_tokens": 3}, + }, + ) + + n_rows = 4 + task = _batch_task([f"row-{i}" for i in range(n_rows)]) + transport = _SlowTransport() + start = time.monotonic() + result = _run(APIExecutor.__new__(APIExecutor), task, transport) + elapsed = time.monotonic() - start + + assert len(transport.requests) == n_rows + # Serial issue would take ~0.8s; parallel takes ~0.2s. Allow generous + # headroom for thread scheduling while still failing a serial loop. + assert elapsed < 0.2 * n_rows * 0.6 + # Row order is preserved regardless of completion order. + assert [item.index for item in result.items] == list(range(n_rows)) + + def test_request_skeleton_constructed_once(self) -> None: + """The request template is built once, not once per row.""" + task = _batch_task(["a", "b", "c"]) + transport = _RecordingTransport() + real_build = APIExecutor._build_request_kwargs + + with ( + patch.object( + APIExecutor, + "_get_client", + return_value=httpx.Client(transport=transport), + ), + patch.object( + APIExecutor, "_build_request_kwargs", autospec=True + ) as mock_build, + ): + mock_build.side_effect = lambda *a, **k: real_build(*a, **k) + APIExecutor.__new__(APIExecutor).run(task, Path("/tmp/out")) + + # One call with prompt=None (the skeleton); per-row substitution happens + # in _substitute_prompt, not by rebuilding the request. + assert mock_build.call_count == 1 + assert mock_build.call_args.args[2] is None From 630d55abb31e432f6aae421bb489de5325fa7f79 Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Fri, 18 Sep 2026 13:33:46 +0700 Subject: [PATCH 10/71] docs: fold batching into the two-stage api template Remove the separate api_batch.yaml template and make api_two_stage.yaml demonstrate batching: multiple rows fan out through both stages, each row's prompt fills the {{prompt}} worker-side slot, and stage two consumes stage one. A single request is a one-row spec.data, so a dedicated batch template no longer represents a distinct feature. Co-Authored-By: Claude Code Signed-off-by: Zhengyuan Su --- examples/templates/api_batch.yaml | 43 --------------------------- examples/templates/api_two_stage.yaml | 21 ++++++++----- 2 files changed, 13 insertions(+), 51 deletions(-) delete mode 100644 examples/templates/api_batch.yaml diff --git a/examples/templates/api_batch.yaml b/examples/templates/api_batch.yaml deleted file mode 100644 index 5f95a5588..000000000 --- a/examples/templates/api_batch.yaml +++ /dev/null @@ -1,43 +0,0 @@ -# api_batch.yaml -# -# Batched API executor demo: one task issues one HTTP request per row in -# spec.data, returning the responses row-aligned in APIResult.items. -# -# Each row's prompt is substituted for the {{prompt}} placeholder in the -# request body. Server-side stage references are ${...}; {{prompt}} is a -# worker-side per-row slot, so it is not touched by server-side resolution. - -apiVersion: flowmesh/v1 -kind: APITask -metadata: - name: api-batch - -spec: - taskType: api - - data: - type: list - items: - - Explain vector databases in simple terms. - - Explain attention in simple terms. - - Explain backpropagation in simple terms. - - api: - method: POST - headers: - Content-Type: application/json - body: - model: gpt-4o - messages: - - role: user - content: "{{prompt}}" - response: - parse_json: true - return_body: true - raise_for_status: true - - output: - destination: - type: http - artifacts: - - results.json diff --git a/examples/templates/api_two_stage.yaml b/examples/templates/api_two_stage.yaml index 7eb30e913..d9d64b590 100644 --- a/examples/templates/api_two_stage.yaml +++ b/examples/templates/api_two_stage.yaml @@ -1,14 +1,17 @@ # api_two_stage.yaml # -# Two-stage API executor demo. -# Stage 1 calls a chat completion endpoint and returns raw text. -# Stage 2 sends Stage 1's returned text as the next prompt. +# Two-stage batched API executor demo. # # spec.data is required for an api task and yields one row per request; a -# single request is a one-row list. Stage 1 names a custom endpoint via -# spec.api.url and supplies its own Authorization header. Stage 2 omits url -# and header, so it uses the Nebula defaults (NEBULA_API_BASE_URL + -# NEBULA_API_TOKEN). +# single request is a one-row list. Each row's prompt is substituted for the +# {{prompt}} placeholder in the request body and one request is issued per +# row, returned row-aligned in APIResult.items. Server-side stage references +# are ${...}; {{prompt}} is a worker-side per-row slot, so it is not touched +# by server-side resolution. +# +# Stage 1 fans out over three rows and calls a chat completion endpoint. +# Stage 2 consumes Stage 1's returned text (${stage-1.text}) as the next +# prompt, so the two-stage character is kept while exercising the batch path. apiVersion: flowmesh/v1 kind: APITask @@ -24,7 +27,9 @@ spec: data: type: list items: - - Please explain vector databases in simple terms. + - Explain vector databases in simple terms. + - Explain attention in simple terms. + - Explain backpropagation in simple terms. api: url: https://api.example.com/v1/chat/completions method: POST From aff4e63b8f76f032b026af360ec461f4e36ed250 Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Fri, 18 Sep 2026 17:50:51 +0700 Subject: [PATCH 11/71] fix: let APIItem round-trip its own result output The committed APIItem could not round-trip its own output. The executor constructs the item by field name (response_json=...), which the json alias plus extra="forbid" rejects. The worker serialises results with model_dump_json() and no by_alias, so it emits the field name response_json; the server then re-validates against the json alias and rejects it. Every API-executor result would 422 on ingest. populate_by_name=True fixes construction and ingest while keeping the json alias accepted on input, so it is backward compatible in both directions. Co-Authored-By: Claude Code Signed-off-by: Zhengyuan Su --- .gitignore | 4 ++- sdk/src/flowmesh/models/result/payloads.py | 2 ++ src/shared/schemas/result/payloads.py | 5 +++- tests/shared/test_executor_result.py | 32 ++++++++++++++++++++++ 4 files changed, 41 insertions(+), 2 deletions(-) diff --git a/.gitignore b/.gitignore index 8a3f0c34b..5d7ce3094 100644 --- a/.gitignore +++ b/.gitignore @@ -180,4 +180,6 @@ secrets/ dump.rdb plugins/ -plugin-data/ \ No newline at end of file +plugin-data/ +# Isolated e2e stack env (contains a live credential) +.env.apibatch diff --git a/sdk/src/flowmesh/models/result/payloads.py b/sdk/src/flowmesh/models/result/payloads.py index 49d0da8ac..9f3eee8b7 100644 --- a/sdk/src/flowmesh/models/result/payloads.py +++ b/sdk/src/flowmesh/models/result/payloads.py @@ -161,6 +161,8 @@ class EchoItem(StrictModel): class APIItem(StrictModel): + model_config = ConfigDict(extra="forbid", populate_by_name=True) + index: int url: str status_code: int diff --git a/src/shared/schemas/result/payloads.py b/src/shared/schemas/result/payloads.py index 9d2e1e951..7769f8926 100644 --- a/src/shared/schemas/result/payloads.py +++ b/src/shared/schemas/result/payloads.py @@ -213,9 +213,12 @@ class APIItem(StrictModel): """One row's HTTP response in a batched API task. Mirrors the per-response fields of :class:`APIResult`; ``response_json`` is - the upstream API's own payload and stays an open mapping. + the upstream API's own payload and stays an open mapping. ``populate_by_name`` + lets code construct by field name while the wire key stays ``json``. """ + model_config = ConfigDict(extra="forbid", populate_by_name=True) + index: int url: str status_code: int diff --git a/tests/shared/test_executor_result.py b/tests/shared/test_executor_result.py index dc1aea478..a52d38fe9 100644 --- a/tests/shared/test_executor_result.py +++ b/tests/shared/test_executor_result.py @@ -8,6 +8,7 @@ from shared.schemas.artifact import ArtifactContext, ArtifactRef from shared.schemas.result import ( + APIItem, APIResult, BaseExecutorResult, DataRetrievalItem, @@ -190,3 +191,34 @@ def test_upstream_results_preserve_subclass_payload_over_the_wire() -> None: assert reloaded.upstreamResults is not None injected = reloaded.upstreamResults["echo-a"] assert injected.model_dump()["items"][0]["output"] == "literal_from_a" + + +def test_api_item_round_trip_construct_serialize_validate() -> None: + """An APIItem constructed by field name must round-trip through the worker's + serialization and the server's ingest validation. + + The executor builds ``APIItem(response_json=...)`` by field name; the worker + writes it with ``model_dump_json()`` (no ``by_alias``); the server re-validates + it on ingest. Without ``populate_by_name`` the field name is rejected as extra + (the field's validation name is the ``json`` alias), so this round trip breaks. + """ + # mypy cannot see populate_by_name; the field's declared name is the json alias. + item = APIItem( # type: ignore[call-arg] + index=0, + url="http://example.com/v1/chat/completions", + status_code=200, + response_json={"choices": [{"message": {"content": "hello"}}]}, + text="hello", + ) + # Serialize the way the worker does (envelope.model_dump_json, no by_alias). + wire = item.model_dump_json() + # Re-validate the way the server does on ingest. + reloaded = APIItem.model_validate_json(wire) + assert reloaded.index == 0 + assert reloaded.response_json["choices"][0]["message"]["content"] == "hello" + assert reloaded.text == "hello" + # The wire alias is still accepted on input (backward compatible). + by_alias = APIItem.model_validate( + {"index": 1, "url": "u", "status_code": 200, "json": {"a": 1}} + ) + assert by_alias.response_json == {"a": 1} From da6a2df18e4938ee0ee1937a647f6fe2b328f72f Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Fri, 18 Sep 2026 21:11:56 +0700 Subject: [PATCH 12/71] fix: size the API HTTP pool from the capped concurrency The connection pool was hard-coded at 8 while the ThreadPoolExecutor used the effective spec.api.concurrency, and the client cache key omitted concurrency so a pool built for one value was reused for another. Size both limits from the capped configured value and include concurrency in the cache key. Co-Authored-By: Claude Code Signed-off-by: Zhengyuan Su --- src/worker/executors/api_executor.py | 46 +++--- tests/worker/test_api_executor_batch.py | 183 ++++++++++++++++++------ 2 files changed, 166 insertions(+), 63 deletions(-) diff --git a/src/worker/executors/api_executor.py b/src/worker/executors/api_executor.py index 915574631..b11298a7e 100644 --- a/src/worker/executors/api_executor.py +++ b/src/worker/executors/api_executor.py @@ -18,8 +18,8 @@ logger = logging.getLogger(__name__) -# Cache key: (base_url, timeout_seconds, verify_tls, follow_redirects) -_ClientKey = tuple[str, float, bool, bool] +# Cache key: (base_url, timeout_seconds, verify_tls, follow_redirects, concurrency) +_ClientKey = tuple[str, float, bool, bool, int] # Worker-side per-row slot in a batched request body. Server-side stage # references are ${...}; this is a worker-side token, hence {{...}}. @@ -49,8 +49,9 @@ class APIExecutor(DataMixin, Executor): Defaults to the Nebula endpoint via ``NEBULA_API_BASE_URL`` and authenticates with ``NEBULA_API_TOKEN``. ``spec.api.url`` overrides the endpoint and ``spec.api.headers`` may supply a credential header (``Authorization``, - ``X-API-Key``, etc.) directly. A custom ``spec.api.url`` requires its own - credential: the Nebula token is never sent to an endpoint the caller chose. + ``X-API-Key``, etc.) directly. A custom ``spec.api.url`` may be + unauthenticated; the Nebula token is never sent to an endpoint the caller + chose. """ name = "api" @@ -95,34 +96,40 @@ def _get_client( timeout: httpx.Timeout, verify_tls: bool, follow_redirects: bool, + concurrency: int, ) -> httpx.Client: """Return a cached client or create a new one for the given parameters.""" timeout_sec = timeout.connect # all four fields are set to same value if timeout_sec is None: timeout_sec = 0.0 - key: _ClientKey = (base_url, float(timeout_sec), verify_tls, follow_redirects) + key: _ClientKey = ( + base_url, + float(timeout_sec), + verify_tls, + follow_redirects, + concurrency, + ) with cls._clients_lock: client = cls._clients.get(key) if client is not None and not client.is_closed: return client - # Create a new client for this combination. The connection pool is - # sized to the max concurrency so parallel row requests never queue - # on connections (which would make the parallelism imaginary). client = httpx.Client( timeout=timeout, verify=verify_tls, follow_redirects=follow_redirects, limits=httpx.Limits( - max_connections=_MAX_CONCURRENCY, - max_keepalive_connections=_MAX_CONCURRENCY, + max_connections=concurrency, + max_keepalive_connections=concurrency, ), ) cls._clients[key] = client logger.debug( - "Created new HTTP client for %s (verify=%s, timeout=%s)", + "Created new HTTP client for %s (verify=%s, timeout=%s, " + "concurrency=%s)", base_url, verify_tls, timeout_sec, + concurrency, ) return client @@ -368,24 +375,23 @@ def _run(self, task: ExecutorTask, out_dir: Path) -> APIResult: if not isinstance(retries, int) or isinstance(retries, bool) or retries < 0: raise ExecutionError("spec.api.retries must be a non-negative integer") + concurrency = int(api_cfg.get("concurrency", _MAX_CONCURRENCY)) + if concurrency < 1: + raise ExecutionError("spec.api.concurrency must be >= 1") + concurrency = min(concurrency, _MAX_CONCURRENCY) + base = self._base_url(str(url)) - client = self._get_client(base, timeout, verify_tls, follow_redirects) + client = self._get_client( + base, timeout, verify_tls, follow_redirects, concurrency + ) - # spec.data is required and yields one row per request. entry = self._collect_prompts_for_spec(spec, task_id=task.task_id) prompts = entry.prompts if not prompts: raise ExecutionError("spec.data produced no rows") - # Build the request skeleton once; each worker only substitutes its - # prompt into the prepared body and issues the request. request_kwargs = self._build_request_kwargs(api_cfg, None) - concurrency = int(api_cfg.get("concurrency", _MAX_CONCURRENCY)) - if concurrency < 1: - raise ExecutionError("spec.api.concurrency must be >= 1") - concurrency = min(concurrency, _MAX_CONCURRENCY) - def _issue(idx: int, prompt: Any) -> APIItem: prompt_str = self._prompt_to_str(prompt) kwargs = self._substitute_prompt(request_kwargs, prompt_str) diff --git a/tests/worker/test_api_executor_batch.py b/tests/worker/test_api_executor_batch.py index 9e07ca147..36c6dba74 100644 --- a/tests/worker/test_api_executor_batch.py +++ b/tests/worker/test_api_executor_batch.py @@ -1,7 +1,9 @@ """Tests for the API executor's batch mode (one task, N row-aligned requests).""" +import json import time from pathlib import Path +from typing import Any from unittest.mock import patch import httpx @@ -12,7 +14,7 @@ from worker.executors.base_executor import ExecutionError -def _task_message(**spec_updates: object) -> WorkerTaskMessage: +def _task_message(**spec_updates: Any) -> WorkerTaskMessage: payload = { "task_id": "task-api-batch", "workflow_id": "wf-1", @@ -37,7 +39,7 @@ def _task_message(**spec_updates: object) -> WorkerTaskMessage: class _RecordingTransport(httpx.MockTransport): - """MockTransport that records every request it served, in order.""" + """MockTransport that echoes each row's prompt back as its response text.""" def __init__(self) -> None: self.requests: list[httpx.Request] = [] @@ -45,25 +47,30 @@ def __init__(self) -> None: def _handler(self, request: httpx.Request) -> httpx.Response: self.requests.append(request) + body = request.read() + prompt = json.loads(body)["messages"][0]["content"] return httpx.Response( 200, json={ - "choices": [{"message": {"content": "hello"}}], + "choices": [{"message": {"content": f"echo:{prompt}"}}], "usage": {"total_tokens": 3}, }, ) def _run( - executor: APIExecutor, task: WorkerTaskMessage, transport: httpx.MockTransport + executor: APIExecutor, + task: WorkerTaskMessage, + transport: httpx.MockTransport, + out_dir: Path, ): with patch.object( APIExecutor, "_get_client", return_value=httpx.Client(transport=transport) ): - return executor.run(task, Path("/tmp/out")) + return executor.run(task, out_dir) -def _batch_task(items: list[object]) -> WorkerTaskMessage: +def _batch_task(items: list[Any], **api_updates: Any) -> WorkerTaskMessage: payload = { "task_id": "task-api-batch", "workflow_id": "wf-1", @@ -79,6 +86,7 @@ def _batch_task(items: list[object]) -> WorkerTaskMessage: "api": { "method": "POST", "body": {"messages": [{"role": "user", "content": "{{prompt}}"}]}, + **api_updates, }, "data": {"type": "list", "items": items}, }, @@ -93,36 +101,80 @@ def _nebula_env(self, monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setenv("NEBULA_API_BASE_URL", "https://nebula.example.com") monkeypatch.setenv("NEBULA_API_TOKEN", "nebula-token") - def test_issues_one_request_per_row_in_order(self) -> None: + def test_issues_one_request_per_row_in_order(self, tmp_path: Path) -> None: task = _batch_task(["first", "second", "third"]) transport = _RecordingTransport() - result = _run(APIExecutor.__new__(APIExecutor), task, transport) + result = _run(APIExecutor.__new__(APIExecutor), task, transport, tmp_path) assert len(transport.requests) == 3 - # Row order is preserved: request i carries row i's prompt. for idx, prompt in enumerate(["first", "second", "third"]): body = transport.requests[idx].read() assert prompt.encode() in body - assert [item.index for item in result.items] == [0, 1, 2] - assert [item.text for item in result.items] == ["hello"] * 3 + item = result.items[idx] + assert item.index == idx + assert item.prompt == prompt + assert item.text == f"echo:{prompt}" + assert item.response_json["choices"][0]["message"]["content"] == ( + f"echo:{prompt}" + ) + + def test_rows_stay_aligned_when_requests_complete_out_of_order( + self, tmp_path: Path + ) -> None: + """Output row i corresponds to input row i even when requests finish + in reverse order.""" + + class _ReverseTransport(httpx.MockTransport): + def __init__(self) -> None: + self.requests: list[httpx.Request] = [] + super().__init__(self._handler) - def test_single_row_batches_to_one_item(self) -> None: + def _handler(self, request: httpx.Request) -> httpx.Response: + self.requests.append(request) + prompt = json.loads(request.read())["messages"][0]["content"] + delay = {"a": 0.3, "b": 0.2, "c": 0.1}[prompt] + time.sleep(delay) + return httpx.Response( + 200, + json={ + "choices": [{"message": {"content": f"echo:{prompt}"}}], + "usage": {"total_tokens": 3}, + }, + ) + + task = _batch_task(["a", "b", "c"]) + transport = _ReverseTransport() + result = _run(APIExecutor.__new__(APIExecutor), task, transport, tmp_path) + for idx, prompt in enumerate(["a", "b", "c"]): + item = result.items[idx] + assert item.index == idx + assert item.prompt == prompt + assert item.text == f"echo:{prompt}" + assert item.response_json["choices"][0]["message"]["content"] == ( + f"echo:{prompt}" + ) + + def test_single_row_batches_to_one_item(self, tmp_path: Path) -> None: task = _batch_task(["only"]) transport = _RecordingTransport() - result = _run(APIExecutor.__new__(APIExecutor), task, transport) + result = _run(APIExecutor.__new__(APIExecutor), task, transport, tmp_path) assert len(transport.requests) == 1 assert len(result.items) == 1 assert result.items[0].index == 0 assert result.items[0].prompt == "only" - def test_row_failure_fails_whole_task_without_shifting(self) -> None: - class _FailSecond(httpx.MockTransport): - def __init__(self) -> None: + def test_row_failure_fails_whole_task_without_shifting( + self, tmp_path: Path + ) -> None: + class _FailRow(httpx.MockTransport): + def __init__(self, failing_prompt: str) -> None: + self.failing_prompt = failing_prompt self.requests: list[httpx.Request] = [] super().__init__(self._handler) def _handler(self, request: httpx.Request) -> httpx.Response: self.requests.append(request) - if len(self.requests) == 2: + prompt = json.loads(request.read())["messages"][0]["content"] + if prompt == self.failing_prompt: return httpx.Response( 500, json={ @@ -139,14 +191,12 @@ def _handler(self, request: httpx.Request) -> httpx.Response: ) task = _batch_task(["a", "b", "c"]) - transport = _FailSecond() + transport = _FailRow("b") with pytest.raises(ExecutionError, match="row 1"): - _run(APIExecutor.__new__(APIExecutor), task, transport) - # All rows are issued in parallel; the failing row aborts the task and - # no partial result is returned. + _run(APIExecutor.__new__(APIExecutor), task, transport, tmp_path) assert len(transport.requests) == 3 - def test_placeholder_not_required_for_scalar_body(self) -> None: + def test_placeholder_not_required_for_scalar_body(self, tmp_path: Path) -> None: """A batch task whose body has no placeholder still issues N requests.""" payload = { "task_id": "task-api-batch", @@ -170,16 +220,21 @@ def test_placeholder_not_required_for_scalar_body(self) -> None: } task = WorkerTaskMessage.model_validate(payload) transport = _RecordingTransport() - result = _run(APIExecutor.__new__(APIExecutor), task, transport) + result = _run(APIExecutor.__new__(APIExecutor), task, transport, tmp_path) assert len(transport.requests) == 2 assert len(result.items) == 2 - def test_no_rows_raises(self) -> None: + def test_no_rows_raises(self, tmp_path: Path) -> None: task = _batch_task([]) with pytest.raises(ExecutionError, match="no rows"): - _run(APIExecutor.__new__(APIExecutor), task, _RecordingTransport()) - - def test_missing_data_raises(self) -> None: + _run( + APIExecutor.__new__(APIExecutor), + task, + _RecordingTransport(), + tmp_path, + ) + + def test_missing_data_raises(self, tmp_path: Path) -> None: """spec.data is required; an api task without it fails closed.""" payload = { "task_id": "task-api-batch", @@ -202,15 +257,15 @@ def test_missing_data_raises(self) -> None: } task = WorkerTaskMessage.model_validate(payload) with pytest.raises(ExecutionError, match="spec.data is required"): - _run(APIExecutor.__new__(APIExecutor), task, _RecordingTransport()) + _run( + APIExecutor.__new__(APIExecutor), + task, + _RecordingTransport(), + tmp_path, + ) - def test_requests_issue_in_parallel(self) -> None: - """N rows must take ~one row's latency, not N x, on a network-bound path. - - A transport that sleeps per request proves the requests overlap: with - serial issue this would take N x the sleep, with parallel issue it takes - roughly one sleep (plus scheduling overhead). - """ + def test_requests_issue_in_parallel(self, tmp_path: Path) -> None: + """N rows take ~one row's latency, not N x, on a network-bound path.""" class _SlowTransport(httpx.MockTransport): def __init__(self) -> None: @@ -232,17 +287,14 @@ def _handler(self, request: httpx.Request) -> httpx.Response: task = _batch_task([f"row-{i}" for i in range(n_rows)]) transport = _SlowTransport() start = time.monotonic() - result = _run(APIExecutor.__new__(APIExecutor), task, transport) + result = _run(APIExecutor.__new__(APIExecutor), task, transport, tmp_path) elapsed = time.monotonic() - start assert len(transport.requests) == n_rows - # Serial issue would take ~0.8s; parallel takes ~0.2s. Allow generous - # headroom for thread scheduling while still failing a serial loop. assert elapsed < 0.2 * n_rows * 0.6 - # Row order is preserved regardless of completion order. assert [item.index for item in result.items] == list(range(n_rows)) - def test_request_skeleton_constructed_once(self) -> None: + def test_request_skeleton_constructed_once(self, tmp_path: Path) -> None: """The request template is built once, not once per row.""" task = _batch_task(["a", "b", "c"]) transport = _RecordingTransport() @@ -259,9 +311,54 @@ def test_request_skeleton_constructed_once(self) -> None: ) as mock_build, ): mock_build.side_effect = lambda *a, **k: real_build(*a, **k) - APIExecutor.__new__(APIExecutor).run(task, Path("/tmp/out")) + APIExecutor.__new__(APIExecutor).run(task, tmp_path) - # One call with prompt=None (the skeleton); per-row substitution happens - # in _substitute_prompt, not by rebuilding the request. assert mock_build.call_count == 1 assert mock_build.call_args.args[2] is None + + @pytest.mark.parametrize("concurrency", [1, 4, 8]) + def test_client_pool_sized_to_concurrency(self, concurrency: int) -> None: + """The connection pool matches the effective concurrency.""" + APIExecutor.close_all_clients() + try: + client = APIExecutor._get_client( + "https://example.com", + httpx.Timeout(60), + True, + True, + concurrency, + ) + pool = client._transport._pool # type: ignore[attr-defined] + assert pool._max_connections == concurrency + assert pool._max_keepalive_connections == concurrency + finally: + APIExecutor.close_all_clients() + + def test_concurrency_capped_at_max(self, tmp_path: Path) -> None: + """A configured concurrency above the cap is clamped to the cap.""" + task = _batch_task(["a", "b", "c"], concurrency=100) + captured: dict[str, Any] = {} + + def _fake_get_client(*args: Any, **kwargs: Any) -> httpx.Client: + captured["concurrency"] = kwargs.get("concurrency", args[4]) + return httpx.Client(transport=_RecordingTransport()) + + with patch.object(APIExecutor, "_get_client", side_effect=_fake_get_client): + APIExecutor.__new__(APIExecutor).run(task, tmp_path) + + assert captured["concurrency"] == 8 + + def test_client_cache_key_includes_concurrency(self) -> None: + """Pools built for different concurrency values are not shared.""" + APIExecutor.close_all_clients() + try: + c1 = APIExecutor._get_client( + "https://example.com", httpx.Timeout(60), True, True, 1 + ) + c4 = APIExecutor._get_client( + "https://example.com", httpx.Timeout(60), True, True, 4 + ) + assert c1 is not c4 + assert len(APIExecutor._clients) == 2 + finally: + APIExecutor.close_all_clients() From d3e31b86f79910cfd9cfe434b505c3b8231dc836 Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Fri, 18 Sep 2026 21:12:13 +0700 Subject: [PATCH 13/71] fix: resolve API stage dependencies from the first row's text APIResult is batch-only: text lives per row, so a dependent stage's placeholder must address items.0.text rather than a scalar text field that is never populated. Update the n8n parser, the two-stage example, and add an end-to-end dependent-stage resolution test. Co-Authored-By: Claude Code Signed-off-by: Zhengyuan Su --- examples/templates/api_two_stage.yaml | 12 ++-- src/server/task/n8n_parser.py | 4 +- tests/server/task/test_n8n_parser.py | 38 +++++++++-- tests/server/task/test_ssh_result_mounting.py | 64 +++++++++++++++++++ 4 files changed, 108 insertions(+), 10 deletions(-) diff --git a/examples/templates/api_two_stage.yaml b/examples/templates/api_two_stage.yaml index d9d64b590..136aaad50 100644 --- a/examples/templates/api_two_stage.yaml +++ b/examples/templates/api_two_stage.yaml @@ -10,8 +10,12 @@ # by server-side resolution. # # Stage 1 fans out over three rows and calls a chat completion endpoint. -# Stage 2 consumes Stage 1's returned text (${stage-1.text}) as the next -# prompt, so the two-stage character is kept while exercising the batch path. +# Stage 2 consumes Stage 1's returned text as the next prompt, so the +# two-stage character is kept while exercising the batch path. A dependent +# stage reads a single upstream value: APIResult is batch-only, so the +# reference addresses one row explicitly (${stage-1.items.0.text} is the +# first row's text), and stage-2's one-row spec.data supplies that value as +# its prompt. apiVersion: flowmesh/v1 kind: APITask @@ -43,7 +47,6 @@ spec: content: "{{prompt}}" response: parse_json: false - # return_body is a JSON backdoor: keep raw text when JSON isn't usable. return_body: true raise_for_status: true @@ -52,7 +55,7 @@ spec: data: type: list items: - - "The previous stage's response is as follows. Please provide a simpler explanation: \n${stage-1.text}" + - "The previous stage's response is as follows. Please provide a simpler explanation: \n${stage-1.items.0.text}" api: method: POST headers: @@ -64,7 +67,6 @@ spec: content: "{{prompt}}" response: parse_json: true - # return_body is a JSON backdoor: keep raw text when JSON isn't usable. return_body: true raise_for_status: true diff --git a/src/server/task/n8n_parser.py b/src/server/task/n8n_parser.py index fac95a4d0..2edd2904b 100644 --- a/src/server/task/n8n_parser.py +++ b/src/server/task/n8n_parser.py @@ -414,7 +414,9 @@ def _inject_dependency_prompt(prompt_text: str, placeholder: str) -> str: def _dependency_placeholder(dep_name: str, dep_task_type: str) -> str: if dep_task_type == "api": - return f"${{{dep_name}.text}}" + # APIResult is batch-only: text lives per row, so a dependent stage + # reads the first row's text. + return f"${{{dep_name}.items.0.text}}" if dep_task_type == "inference": return f"${{{dep_name}.items.0.output}}" raise ValueError( diff --git a/tests/server/task/test_n8n_parser.py b/tests/server/task/test_n8n_parser.py index e8d3abe04..dc58f7ba2 100644 --- a/tests/server/task/test_n8n_parser.py +++ b/tests/server/task/test_n8n_parser.py @@ -23,12 +23,10 @@ def test_simple_openai_node(self) -> None: ] result = translate_n8n_workflow({"nodes": nodes, "connections": {}}) - # Top-level shape assert result["kind"] == "APITask" assert result["apiVersion"] == "flowmesh/v1" assert "spec" in result - # Task type and API spec spec = result["spec"] assert spec["taskType"] == "api" assert "api" in spec @@ -36,10 +34,8 @@ def test_simple_openai_node(self) -> None: assert api["method"] == "POST" assert api["body"]["model"] == "gpt-4" - # Prompt content preserved as a one-row spec.data assert spec["data"]["type"] == "list" assert spec["data"]["items"] == ["Hello, world!"] - # Body carries the worker-side per-row placeholder assert api["body"]["messages"][0]["content"] == "{{prompt}}" def test_no_task_nodes_raises_value_error(self) -> None: @@ -52,6 +48,40 @@ def test_invalid_json_via_parse_workflow(self) -> None: with pytest.raises(ValueError, match="Invalid JSON"): parse_workflow("not json at all {{{", format="n8n") + def test_api_dependency_resolves_first_row_text(self) -> None: + """A dependent API node reads the upstream API stage's first-row text.""" + nodes = [ + { + "name": "Upstream", + "type": "@n8n/n8n-nodes-langchain.openAi", + "parameters": { + "modelId": {"value": "gpt-4"}, + "responses": {"values": [{"content": "First answer"}]}, + }, + }, + { + "name": "Downstream", + "type": "@n8n/n8n-nodes-langchain.openAi", + "parameters": { + "modelId": {"value": "gpt-4"}, + "responses": {"values": [{"content": "Simplify this"}]}, + }, + }, + ] + connections = { + "Upstream": {"ai_languageModel": [[{"node": "Downstream"}]]}, + } + result = translate_n8n_workflow({"nodes": nodes, "connections": connections}) + + graph = result["spec"]["graph"]["nodes"] + assert [n["name"] for n in graph] == ["Upstream", "Downstream"] + downstream = next(n for n in graph if n["name"] == "Downstream") + assert downstream["dependsOn"] == ["Upstream"] + assert downstream["spec"]["data"]["items"] == [ + "The previous stage's response is as follows. Simplify this\n" + "${Upstream.items.0.text}" + ] + class TestDecodeSecretPart: def test_hex_decode(self) -> None: diff --git a/tests/server/task/test_ssh_result_mounting.py b/tests/server/task/test_ssh_result_mounting.py index 1953be090..cd1217d76 100644 --- a/tests/server/task/test_ssh_result_mounting.py +++ b/tests/server/task/test_ssh_result_mounting.py @@ -14,6 +14,7 @@ from server.task.models import TaskRecord, TaskStatus from server.task.parser import parse_workflow from server.task.runtime import TaskRuntime +from shared.schemas.result import ResultEnvelope from shared.tasks import TaskEnvelopeTemplate, TaskType from shared.tasks.specs import SSHSpecStrict @@ -386,3 +387,66 @@ def test_stage_reference_uses_payload_root_for_local_and_http_results( http_value == "http://flowmesh.example/api/v1/results/task-http/files/final_lora.tar.gz" ) + + +def test_api_dependent_stage_resolves_first_row_text(tmp_path: Path) -> None: + """A dependent stage's ${stage.items.0.text} resolves to the first row's + text of a batch-only APIResult.""" + from shared.schemas.result import APIItem, APIResult + + stage_dir = tmp_path / "task-api" + stage_dir.mkdir() + result = APIResult( + ok=True, + executor="api", + method="POST", + url="https://api.example.com/v1/chat/completions", + status_code=200, + items=[ + APIItem( + index=0, + url="https://api.example.com/v1/chat/completions", + status_code=200, + text="first row text", + prompt="first", + ), + APIItem( + index=1, + url="https://api.example.com/v1/chat/completions", + status_code=200, + text="second row text", + prompt="second", + ), + ], + ) + (stage_dir / "results.json").write_text( + json.dumps( + { + "task_id": "task-api", + "result": json.loads( + ResultEnvelope(task_id="task-api", result=result).model_dump_json() + )["result"], + } + ), + encoding="utf-8", + ) + + upstream = TaskRecord( + task_id="task-api", + workflow_id="wf-1", + owner_id="owner", + source="raw", + task=_task_template(TaskType.API), + status=TaskStatus.DONE, + task_type="api", + local_name="stage", + ) + dispatcher = Dispatcher( + runtime=cast(TaskRuntime, _DummyRuntime({})), + worker_registry=cast(WorkerRegistry, object()), + results_dir=tmp_path, + logger=logging.getLogger("test-api-dependent-stage"), + ) + + value = dispatcher._resolve_reference("stage.items.0.text", {"stage": upstream}) + assert value == "first row text" From 4358a0be81ae8d7fab5a353a7b4b7e42735fa5bf Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Fri, 18 Sep 2026 21:12:22 +0700 Subject: [PATCH 14/71] docs: correct the APIResult docstring and guard SDK by-name construction APIResult is batch-only, so the docstring no longer describes a scalar single-request path. Add an SDK-side by-name construct/serialize/revalidate test so the populate_by_name setting is guarded on both sides of the wire. Co-Authored-By: Claude Code Signed-off-by: Zhengyuan Su --- src/shared/schemas/result/catalog.py | 5 ++--- tests/sdk/test_models.py | 23 +++++++++++++++++++++++ tests/shared/test_executor_result.py | 10 ++-------- 3 files changed, 27 insertions(+), 11 deletions(-) diff --git a/src/shared/schemas/result/catalog.py b/src/shared/schemas/result/catalog.py index 0c950c787..d7c0aeb19 100644 --- a/src/shared/schemas/result/catalog.py +++ b/src/shared/schemas/result/catalog.py @@ -248,9 +248,8 @@ class APIResult(StrictExecutorResult): """HTTP request output. ``response_json``/``usage``/``headers`` are the upstream API's own payloads and stay open mappings. - ``items`` carries one entry per row when the task batches multiple requests - (``spec.data`` present); a single-request task leaves it empty and populates - the scalar fields instead.""" + ``items`` carries one entry per row: every API task batches over + ``spec.data``, so even a single request yields one item.""" task_type: Literal[TaskType.API] = TaskType.API executor: str diff --git a/tests/sdk/test_models.py b/tests/sdk/test_models.py index df2aa9889..a9a720ce8 100644 --- a/tests/sdk/test_models.py +++ b/tests/sdk/test_models.py @@ -5,6 +5,7 @@ import pytest from flowmesh.models import ( ActiveWaitBreakdown, + APIItem, AssetSummary, CriticalPathSummary, E2EBreakdown, @@ -490,3 +491,25 @@ def test_profile_summary_rejects_extra_fields(self) -> None: payload["unexpected"] = 1 with pytest.raises(Exception): ProfileSummary.model_validate(payload) + + +class TestAPIItem: + def test_construct_by_name_serialize_revalidate(self) -> None: + """An SDK APIItem built by field name round-trips through the worker's + serialization and the server's ingest validation.""" + item = APIItem( # type: ignore[call-arg] + index=0, + url="http://example.com/v1/chat/completions", + status_code=200, + response_json={"choices": [{"message": {"content": "hello"}}]}, + text="hello", + ) + wire = item.model_dump_json() + reloaded = APIItem.model_validate_json(wire) + assert reloaded.index == 0 + assert reloaded.response_json["choices"][0]["message"]["content"] == "hello" + assert reloaded.text == "hello" + by_alias = APIItem.model_validate( + {"index": 1, "url": "u", "status_code": 200, "json": {"a": 1}} + ) + assert by_alias.response_json == {"a": 1} diff --git a/tests/shared/test_executor_result.py b/tests/shared/test_executor_result.py index a52d38fe9..494b78435 100644 --- a/tests/shared/test_executor_result.py +++ b/tests/shared/test_executor_result.py @@ -194,14 +194,8 @@ def test_upstream_results_preserve_subclass_payload_over_the_wire() -> None: def test_api_item_round_trip_construct_serialize_validate() -> None: - """An APIItem constructed by field name must round-trip through the worker's - serialization and the server's ingest validation. - - The executor builds ``APIItem(response_json=...)`` by field name; the worker - writes it with ``model_dump_json()`` (no ``by_alias``); the server re-validates - it on ingest. Without ``populate_by_name`` the field name is rejected as extra - (the field's validation name is the ``json`` alias), so this round trip breaks. - """ + """An APIItem built by field name round-trips through the worker's + serialization and the server's ingest validation.""" # mypy cannot see populate_by_name; the field's declared name is the json alias. item = APIItem( # type: ignore[call-arg] index=0, From dc9da0569149cb345cbbdfc8ec652e2d7d33b089 Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Fri, 18 Sep 2026 21:17:40 +0700 Subject: [PATCH 15/71] refactor: fold API dependency and cache-key notes into docstrings Replace the added inline comments with condensed docstring notes so the non-obvious decisions (batch-only APIResult row addressing, the client cache key and pool sizing) stay documented without inline comment blocks. Co-Authored-By: Claude Code Signed-off-by: Zhengyuan Su --- src/server/task/n8n_parser.py | 6 ++++-- src/worker/executors/api_executor.py | 7 +++++-- 2 files changed, 9 insertions(+), 4 deletions(-) diff --git a/src/server/task/n8n_parser.py b/src/server/task/n8n_parser.py index 2edd2904b..aa44750fd 100644 --- a/src/server/task/n8n_parser.py +++ b/src/server/task/n8n_parser.py @@ -413,9 +413,11 @@ def _inject_dependency_prompt(prompt_text: str, placeholder: str) -> str: def _dependency_placeholder(dep_name: str, dep_task_type: str) -> str: + """Return the stage reference a dependent node injects for an upstream task. + + APIResult is batch-only, so an api dependency reads the first row's text + (``items.0.text``) rather than a scalar ``text`` field.""" if dep_task_type == "api": - # APIResult is batch-only: text lives per row, so a dependent stage - # reads the first row's text. return f"${{{dep_name}.items.0.text}}" if dep_task_type == "inference": return f"${{{dep_name}.items.0.output}}" diff --git a/src/worker/executors/api_executor.py b/src/worker/executors/api_executor.py index b11298a7e..7503bc69b 100644 --- a/src/worker/executors/api_executor.py +++ b/src/worker/executors/api_executor.py @@ -18,7 +18,6 @@ logger = logging.getLogger(__name__) -# Cache key: (base_url, timeout_seconds, verify_tls, follow_redirects, concurrency) _ClientKey = tuple[str, float, bool, bool, int] # Worker-side per-row slot in a batched request body. Server-side stage @@ -98,7 +97,11 @@ def _get_client( follow_redirects: bool, concurrency: int, ) -> httpx.Client: - """Return a cached client or create a new one for the given parameters.""" + """Return a cached client or create a new one for the given parameters. + + The cache key is ``(base_url, timeout_sec, verify_tls, follow_redirects, + concurrency)``; the pool is sized to ``concurrency`` so parallel row + requests never queue on connections.""" timeout_sec = timeout.connect # all four fields are set to same value if timeout_sec is None: timeout_sec = 0.0 From da2aef14af9ed3df1105de0f8c7724126435a19a Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Fri, 18 Sep 2026 22:02:28 +0700 Subject: [PATCH 16/71] fix: make the GPU test import-safe and clean up test style Move the module-scope nvmlInit() call in the GPU cleanup test behind a skip guard so importing the module is side-effect free on GPU-less hosts, including CI. Replace the inline import and inline type: ignore comments with top-level imports and a cast, correct the pool-sizing comment to say the pool is sized to the effective concurrency (capped), and assert the issued requests as an unordered collection instead of assuming thread-pool start order. Co-Authored-By: Claude Code Signed-off-by: Zhengyuan Su --- src/worker/executors/api_executor.py | 3 ++- tests/sdk/test_models.py | 14 ++++++++------ tests/server/task/test_ssh_result_mounting.py | 4 +--- tests/worker/test_api_executor_batch.py | 11 +++++++---- tests/worker/test_mp_executor_cleanup_gpu.py | 7 ++++++- 5 files changed, 24 insertions(+), 15 deletions(-) diff --git a/src/worker/executors/api_executor.py b/src/worker/executors/api_executor.py index 7503bc69b..496fa52d3 100644 --- a/src/worker/executors/api_executor.py +++ b/src/worker/executors/api_executor.py @@ -28,7 +28,8 @@ _RETRY_BACKOFF_SEC = 1.0 # Default cap on parallel row requests and the client connection pool. The -# pool is sized to this so parallel requests never queue on connections. +# pool is sized to the effective concurrency, capped at this value so parallel +# requests never queue on connections. _MAX_CONCURRENCY = 8 diff --git a/tests/sdk/test_models.py b/tests/sdk/test_models.py index a9a720ce8..c863442bd 100644 --- a/tests/sdk/test_models.py +++ b/tests/sdk/test_models.py @@ -497,12 +497,14 @@ class TestAPIItem: def test_construct_by_name_serialize_revalidate(self) -> None: """An SDK APIItem built by field name round-trips through the worker's serialization and the server's ingest validation.""" - item = APIItem( # type: ignore[call-arg] - index=0, - url="http://example.com/v1/chat/completions", - status_code=200, - response_json={"choices": [{"message": {"content": "hello"}}]}, - text="hello", + item = APIItem.model_validate( + { + "index": 0, + "url": "http://example.com/v1/chat/completions", + "status_code": 200, + "response_json": {"choices": [{"message": {"content": "hello"}}]}, + "text": "hello", + } ) wire = item.model_dump_json() reloaded = APIItem.model_validate_json(wire) diff --git a/tests/server/task/test_ssh_result_mounting.py b/tests/server/task/test_ssh_result_mounting.py index cd1217d76..d0aeaef4c 100644 --- a/tests/server/task/test_ssh_result_mounting.py +++ b/tests/server/task/test_ssh_result_mounting.py @@ -14,7 +14,7 @@ from server.task.models import TaskRecord, TaskStatus from server.task.parser import parse_workflow from server.task.runtime import TaskRuntime -from shared.schemas.result import ResultEnvelope +from shared.schemas.result import APIItem, APIResult, ResultEnvelope from shared.tasks import TaskEnvelopeTemplate, TaskType from shared.tasks.specs import SSHSpecStrict @@ -392,8 +392,6 @@ def test_stage_reference_uses_payload_root_for_local_and_http_results( def test_api_dependent_stage_resolves_first_row_text(tmp_path: Path) -> None: """A dependent stage's ${stage.items.0.text} resolves to the first row's text of a batch-only APIResult.""" - from shared.schemas.result import APIItem, APIResult - stage_dir = tmp_path / "task-api" stage_dir.mkdir() result = APIResult( diff --git a/tests/worker/test_api_executor_batch.py b/tests/worker/test_api_executor_batch.py index 36c6dba74..1dba221b5 100644 --- a/tests/worker/test_api_executor_batch.py +++ b/tests/worker/test_api_executor_batch.py @@ -3,7 +3,7 @@ import json import time from pathlib import Path -from typing import Any +from typing import Any, cast from unittest.mock import patch import httpx @@ -106,9 +106,12 @@ def test_issues_one_request_per_row_in_order(self, tmp_path: Path) -> None: transport = _RecordingTransport() result = _run(APIExecutor.__new__(APIExecutor), task, transport, tmp_path) assert len(transport.requests) == 3 + issued = { + json.loads(req.read())["messages"][0]["content"] + for req in transport.requests + } + assert issued == {"first", "second", "third"} for idx, prompt in enumerate(["first", "second", "third"]): - body = transport.requests[idx].read() - assert prompt.encode() in body item = result.items[idx] assert item.index == idx assert item.prompt == prompt @@ -328,7 +331,7 @@ def test_client_pool_sized_to_concurrency(self, concurrency: int) -> None: True, concurrency, ) - pool = client._transport._pool # type: ignore[attr-defined] + pool = cast(Any, client._transport)._pool assert pool._max_connections == concurrency assert pool._max_keepalive_connections == concurrency finally: diff --git a/tests/worker/test_mp_executor_cleanup_gpu.py b/tests/worker/test_mp_executor_cleanup_gpu.py index 338e5ac56..62f879109 100644 --- a/tests/worker/test_mp_executor_cleanup_gpu.py +++ b/tests/worker/test_mp_executor_cleanup_gpu.py @@ -13,7 +13,11 @@ from worker.executors.mp_executor import MPExecutor from worker.executors.vllm_executor import VLLMExecutor -pynvml.nvmlInit() +try: + pynvml.nvmlInit() + _NVML_AVAILABLE = True +except Exception: + _NVML_AVAILABLE = False def _descendants_of(pid: int) -> set[int]: @@ -25,6 +29,7 @@ def _descendants_of(pid: int) -> set[int]: @pytest.mark.gpu +@pytest.mark.skipif(not _NVML_AVAILABLE, reason="NVML unavailable (no GPU)") def test_mp_executor_cleans_up_vllm(caplog, tmp_path: Path) -> None: """Start MPExecutor with the real executors, run a minimal task to trigger engine startup, and ensure cleanup removes the worker process From 0b4f12ce1451fe8afc676d7989a308fe1917fe1e Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Fri, 18 Sep 2026 22:15:19 +0700 Subject: [PATCH 17/71] docs: fold the concurrency cap note into the API executor docstring Co-Authored-By: Claude Code Signed-off-by: Zhengyuan Su --- src/worker/executors/api_executor.py | 7 ++++--- 1 file changed, 4 insertions(+), 3 deletions(-) diff --git a/src/worker/executors/api_executor.py b/src/worker/executors/api_executor.py index 496fa52d3..73114c3b4 100644 --- a/src/worker/executors/api_executor.py +++ b/src/worker/executors/api_executor.py @@ -27,9 +27,6 @@ # Fixed delay between retry attempts. _RETRY_BACKOFF_SEC = 1.0 -# Default cap on parallel row requests and the client connection pool. The -# pool is sized to the effective concurrency, capped at this value so parallel -# requests never queue on connections. _MAX_CONCURRENCY = 8 @@ -52,6 +49,10 @@ class APIExecutor(DataMixin, Executor): ``X-API-Key``, etc.) directly. A custom ``spec.api.url`` may be unauthenticated; the Nebula token is never sent to an endpoint the caller chose. + + Parallel row requests and the HTTP connection pool are capped at + ``_MAX_CONCURRENCY``; the pool is sized to the effective concurrency so + parallel requests never queue on connections. """ name = "api" From 7dff652eb9a92a992a5cbc1db0ad69027ef2b28d Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Sat, 19 Sep 2026 01:40:22 +0700 Subject: [PATCH 18/71] test: cover row-specific fields, effective concurrency, and NVML deferral Give the recording transport distinct status, usage, and headers per row (with raise_for_status false and include_headers true) and assert them per row so a mix-up in those fields is caught. Add a run-level test at an uncapped concurrency (1 and 4) that isolates the effective value forwarded to the client, and a test that observes concurrency: 1 serializing requests. Defer nvmlInit() to a module fixture so importing the GPU test is side-effect free, and add an integrated translated-n8n-to-stage-resolution test covering the items.0.text production change end to end. Co-Authored-By: Claude Code Signed-off-by: Zhengyuan Su --- tests/server/task/test_ssh_result_mounting.py | 114 +++++++++++++++++- tests/worker/test_api_executor_batch.py | 84 ++++++++++++- tests/worker/test_mp_executor_cleanup_gpu.py | 20 +-- 3 files changed, 203 insertions(+), 15 deletions(-) diff --git a/tests/server/task/test_ssh_result_mounting.py b/tests/server/task/test_ssh_result_mounting.py index d0aeaef4c..1bfe80bb2 100644 --- a/tests/server/task/test_ssh_result_mounting.py +++ b/tests/server/task/test_ssh_result_mounting.py @@ -16,7 +16,7 @@ from server.task.runtime import TaskRuntime from shared.schemas.result import APIItem, APIResult, ResultEnvelope from shared.tasks import TaskEnvelopeTemplate, TaskType -from shared.tasks.specs import SSHSpecStrict +from shared.tasks.specs import ApiSpecTemplate, SSHSpecStrict class _DummyRuntime: @@ -448,3 +448,115 @@ def test_api_dependent_stage_resolves_first_row_text(tmp_path: Path) -> None: value = dispatcher._resolve_reference("stage.items.0.text", {"stage": upstream}) assert value == "first row text" + + +def test_translated_n8n_dependent_api_stage_resolves(tmp_path: Path) -> None: + """A translated n8n workflow's dependent API stage resolves end to end. + + The injected ${Upstream.items.0.text} placeholder reads the upstream + stage's first-row text through the dispatcher, covering the n8n + production change through to stage resolution. + """ + payload = { + "nodes": [ + { + "name": "Upstream", + "type": "@n8n/n8n-nodes-langchain.openAi", + "parameters": { + "modelId": {"value": "gpt-4"}, + "responses": {"values": [{"content": "First answer"}]}, + }, + }, + { + "name": "Downstream", + "type": "@n8n/n8n-nodes-langchain.openAi", + "parameters": { + "modelId": {"value": "gpt-4"}, + "responses": {"values": [{"content": "Simplify this"}]}, + }, + }, + ], + "connections": { + "Upstream": {"ai_languageModel": [[{"node": "Downstream"}]]}, + }, + } + parsed = parse_workflow(json.dumps(payload), "n8n") + by_name = {t.graph_node_name: t for t in parsed.tasks} + upstream = by_name["Upstream"] + downstream = by_name["Downstream"] + assert downstream.depends_on == [upstream.task_id] + + stage_dir = tmp_path / upstream.task_id + stage_dir.mkdir() + result = APIResult( + ok=True, + executor="api", + method="POST", + url="https://api.example.com/v1/chat/completions", + status_code=200, + items=[ + APIItem( + index=0, + url="https://api.example.com/v1/chat/completions", + status_code=200, + text="first row text", + prompt="first", + ), + ], + ) + (stage_dir / "results.json").write_text( + json.dumps( + { + "task_id": upstream.task_id, + "result": json.loads( + ResultEnvelope( + task_id=upstream.task_id, result=result + ).model_dump_json() + )["result"], + } + ), + encoding="utf-8", + ) + + upstream_record = TaskRecord( + task_id=upstream.task_id, + workflow_id="wf-1", + owner_id="owner", + source="raw", + task=upstream.task, + status=TaskStatus.DONE, + task_type="api", + graph_node_name="Upstream", + ) + downstream_record = TaskRecord( + task_id=downstream.task_id, + workflow_id="wf-1", + owner_id="owner", + source="raw", + task=downstream.task, + status=TaskStatus.PENDING, + task_type="api", + graph_node_name="Downstream", + ) + dispatcher = Dispatcher( + runtime=cast( + TaskRuntime, + _DummyRuntime( + { + upstream.task_id: upstream_record, + downstream.task_id: downstream_record, + }, + depends_on={downstream.task_id: [upstream.task_id]}, + ), + ), + worker_registry=cast(WorkerRegistry, object()), + results_dir=tmp_path, + logger=logging.getLogger("test-n8n-dependent-stage"), + ) + + context = dispatcher._build_stage_context(downstream_record) + spec = cast(ApiSpecTemplate, downstream.task.spec) + resolved = dispatcher._resolve_placeholders(spec.data, context) + assert resolved["items"][0] == ( + "The previous stage's response is as follows. Simplify this\n" "first row text" + ) diff --git a/tests/worker/test_api_executor_batch.py b/tests/worker/test_api_executor_batch.py index 1dba221b5..f1e7e7b69 100644 --- a/tests/worker/test_api_executor_batch.py +++ b/tests/worker/test_api_executor_batch.py @@ -1,6 +1,7 @@ """Tests for the API executor's batch mode (one task, N row-aligned requests).""" import json +import threading import time from pathlib import Path from typing import Any, cast @@ -39,7 +40,8 @@ def _task_message(**spec_updates: Any) -> WorkerTaskMessage: class _RecordingTransport(httpx.MockTransport): - """MockTransport that echoes each row's prompt back as its response text.""" + """MockTransport that echoes each row's prompt back with row-specific + status, usage, and headers.""" def __init__(self) -> None: self.requests: list[httpx.Request] = [] @@ -50,10 +52,11 @@ def _handler(self, request: httpx.Request) -> httpx.Response: body = request.read() prompt = json.loads(body)["messages"][0]["content"] return httpx.Response( - 200, + 200 + len(prompt) % 3, + headers={"X-Row": prompt}, json={ "choices": [{"message": {"content": f"echo:{prompt}"}}], - "usage": {"total_tokens": 3}, + "usage": {"total_tokens": len(prompt)}, }, ) @@ -86,6 +89,10 @@ def _batch_task(items: list[Any], **api_updates: Any) -> WorkerTaskMessage: "api": { "method": "POST", "body": {"messages": [{"role": "user", "content": "{{prompt}}"}]}, + "response": { + "raise_for_status": False, + "include_headers": True, + }, **api_updates, }, "data": {"type": "list", "items": items}, @@ -119,6 +126,9 @@ def test_issues_one_request_per_row_in_order(self, tmp_path: Path) -> None: assert item.response_json["choices"][0]["message"]["content"] == ( f"echo:{prompt}" ) + assert item.status_code == 200 + len(prompt) % 3 + assert item.usage == {"total_tokens": len(prompt)} + assert item.headers["x-row"] == prompt def test_rows_stay_aligned_when_requests_complete_out_of_order( self, tmp_path: Path @@ -137,10 +147,11 @@ def _handler(self, request: httpx.Request) -> httpx.Response: delay = {"a": 0.3, "b": 0.2, "c": 0.1}[prompt] time.sleep(delay) return httpx.Response( - 200, + 200 + len(prompt) % 3, + headers={"X-Row": prompt}, json={ "choices": [{"message": {"content": f"echo:{prompt}"}}], - "usage": {"total_tokens": 3}, + "usage": {"total_tokens": len(prompt)}, }, ) @@ -155,6 +166,9 @@ def _handler(self, request: httpx.Request) -> httpx.Response: assert item.response_json["choices"][0]["message"]["content"] == ( f"echo:{prompt}" ) + assert item.status_code == 200 + len(prompt) % 3 + assert item.usage == {"total_tokens": len(prompt)} + assert item.headers["x-row"] == prompt def test_single_row_batches_to_one_item(self, tmp_path: Path) -> None: task = _batch_task(["only"]) @@ -193,7 +207,7 @@ def _handler(self, request: httpx.Request) -> httpx.Response: }, ) - task = _batch_task(["a", "b", "c"]) + task = _batch_task(["a", "b", "c"], response={"raise_for_status": True}) transport = _FailRow("b") with pytest.raises(ExecutionError, match="row 1"): _run(APIExecutor.__new__(APIExecutor), task, transport, tmp_path) @@ -297,6 +311,36 @@ def _handler(self, request: httpx.Request) -> httpx.Response: assert elapsed < 0.2 * n_rows * 0.6 assert [item.index for item in result.items] == list(range(n_rows)) + def test_concurrency_one_serializes_requests(self, tmp_path: Path) -> None: + """concurrency: 1 limits the worker pool so requests never overlap.""" + + class _OverlapTransport(httpx.MockTransport): + def __init__(self) -> None: + self.max_in_flight = 0 + self._in_flight = 0 + self._lock = threading.Lock() + super().__init__(self._handler) + + def _handler(self, request: httpx.Request) -> httpx.Response: + with self._lock: + self._in_flight += 1 + self.max_in_flight = max(self.max_in_flight, self._in_flight) + time.sleep(0.05) + with self._lock: + self._in_flight -= 1 + return httpx.Response( + 200, + json={ + "choices": [{"message": {"content": "hello"}}], + "usage": {"total_tokens": 3}, + }, + ) + + task = _batch_task(["a", "b", "c", "d"], concurrency=1) + transport = _OverlapTransport() + _run(APIExecutor.__new__(APIExecutor), task, transport, tmp_path) + assert transport.max_in_flight == 1 + def test_request_skeleton_constructed_once(self, tmp_path: Path) -> None: """The request template is built once, not once per row.""" task = _batch_task(["a", "b", "c"]) @@ -351,6 +395,34 @@ def _fake_get_client(*args: Any, **kwargs: Any) -> httpx.Client: assert captured["concurrency"] == 8 + @pytest.mark.parametrize("concurrency", [1, 4]) + def test_run_passes_effective_concurrency_to_client( + self, tmp_path: Path, concurrency: int + ) -> None: + """run() forwards the uncapped configured concurrency to the client.""" + task = _batch_task(["a", "b", "c"], concurrency=concurrency) + captured: dict[str, Any] = {} + + def _fake_get_client(*args: Any, **kwargs: Any) -> httpx.Client: + captured["concurrency"] = kwargs.get("concurrency", args[4]) + return httpx.Client(transport=_RecordingTransport()) + + with patch.object(APIExecutor, "_get_client", side_effect=_fake_get_client): + APIExecutor.__new__(APIExecutor).run(task, tmp_path) + + assert captured["concurrency"] == concurrency + + @pytest.mark.parametrize("concurrency", [0, -1]) + def test_concurrency_below_one_rejected( + self, tmp_path: Path, concurrency: int + ) -> None: + """A configured concurrency below 1 is rejected.""" + task = _batch_task(["a", "b", "c"], concurrency=concurrency) + with pytest.raises(ExecutionError, match="spec.api.concurrency must be >= 1"): + _run( + APIExecutor.__new__(APIExecutor), task, _RecordingTransport(), tmp_path + ) + def test_client_cache_key_includes_concurrency(self) -> None: """Pools built for different concurrency values are not shared.""" APIExecutor.close_all_clients() diff --git a/tests/worker/test_mp_executor_cleanup_gpu.py b/tests/worker/test_mp_executor_cleanup_gpu.py index 62f879109..10874002c 100644 --- a/tests/worker/test_mp_executor_cleanup_gpu.py +++ b/tests/worker/test_mp_executor_cleanup_gpu.py @@ -2,10 +2,11 @@ import tempfile import time import uuid +from collections.abc import Iterator from pathlib import Path import psutil -import pynvml # type: ignore +import pynvml # type: ignore[import-not-found] import pytest from shared.tasks.worker_message import WorkerTaskMessage @@ -13,11 +14,15 @@ from worker.executors.mp_executor import MPExecutor from worker.executors.vllm_executor import VLLMExecutor -try: - pynvml.nvmlInit() - _NVML_AVAILABLE = True -except Exception: - _NVML_AVAILABLE = False + +@pytest.fixture(scope="module") +def _nvml() -> Iterator[None]: + try: + pynvml.nvmlInit() + except Exception: + pytest.skip("NVML unavailable (no GPU)") + yield + pynvml.nvmlShutdown() def _descendants_of(pid: int) -> set[int]: @@ -29,8 +34,7 @@ def _descendants_of(pid: int) -> set[int]: @pytest.mark.gpu -@pytest.mark.skipif(not _NVML_AVAILABLE, reason="NVML unavailable (no GPU)") -def test_mp_executor_cleans_up_vllm(caplog, tmp_path: Path) -> None: +def test_mp_executor_cleans_up_vllm(caplog, tmp_path: Path, _nvml: None) -> None: """Start MPExecutor with the real executors, run a minimal task to trigger engine startup, and ensure cleanup removes the worker process and any descendants it spawned. From 9524d1c417ae3186494b62345b2a15a8607b34a7 Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Sat, 19 Sep 2026 03:09:58 +0700 Subject: [PATCH 19/71] fix: classify API HTTP errors before strict payload parsing A 5xx error body (e.g. a 503 {"error": ...}) has no OpenAI usage or choices, so parsing it before the retryable classification raised a non-retryable ExecutionError and a retryable 503 was permanently failed. Classify HTTP errors first, and accept absent usage/text in the success path so a no-usage 2xx or a 5xx under raise_for_status: false still produces a row-aligned item. Co-Authored-By: Claude Code Signed-off-by: Zhengyuan Su --- src/worker/executors/api_executor.py | 32 ++++++-------- tests/worker/test_api_executor_batch.py | 58 +++++++++++++++++++++++++ 2 files changed, 71 insertions(+), 19 deletions(-) diff --git a/src/worker/executors/api_executor.py b/src/worker/executors/api_executor.py index 73114c3b4..b24e51980 100644 --- a/src/worker/executors/api_executor.py +++ b/src/worker/executors/api_executor.py @@ -299,19 +299,12 @@ def _parse_response( if not isinstance(item.response_json, dict): raise ExecutionError("Response is not a valid JSON mapping") usage = item.response_json.get("usage") - if not isinstance(usage, dict): - raise ExecutionError( - "spec.api.response.parse_json is true but response JSON " - f"does not contain usage info: {item.response_json}" - ) - item.usage = usage + if isinstance(usage, dict): + item.usage = usage try: item.text = item.response_json["choices"][0]["message"]["content"] - except Exception as exc: - raise ExecutionError( - "spec.api.response.parse_json is true but response JSON " - f"does not contain message.content: {item.response_json}" - ) from exc + except Exception: + item.text = None elif response_cfg.get("return_body", True): item.text = body_text @@ -415,7 +408,15 @@ def _issue(idx: int, prompt: Any) -> APIItem: f"API request failed (row {idx}): {exc}", retryable=True ) from exc - item, body_text = self._parse_response( + if raise_for_status and resp.is_error: + message = f"API request returned status {resp.status_code} (row {idx})" + body_text = resp.text[:200] + if body_text: + message = f"{message}: {body_text}" + retryable = resp.status_code >= 500 or resp.status_code in (408, 429) + raise ExecutionError(message, retryable=retryable) + + item, _ = self._parse_response( resp, response_cfg=response_cfg, max_body_bytes=max_body_bytes, @@ -423,13 +424,6 @@ def _issue(idx: int, prompt: Any) -> APIItem: item.index = idx item.prompt = prompt_str - if raise_for_status and resp.is_error: - message = f"API request returned status {resp.status_code} (row {idx})" - if body_text: - message = f"{message}: {body_text[:200]}" - retryable = _is_retryable_status(resp.status_code) - raise ExecutionError(message, retryable=retryable) - return item results: dict[int, APIItem] = {} diff --git a/tests/worker/test_api_executor_batch.py b/tests/worker/test_api_executor_batch.py index f1e7e7b69..026064aea 100644 --- a/tests/worker/test_api_executor_batch.py +++ b/tests/worker/test_api_executor_batch.py @@ -437,3 +437,61 @@ def test_client_cache_key_includes_concurrency(self) -> None: assert len(APIExecutor._clients) == 2 finally: APIExecutor.close_all_clients() + + def test_no_usage_2xx_produces_aligned_item(self, tmp_path: Path) -> None: + """A 2xx response without usage still yields a row-aligned item.""" + + class _NoUsage(httpx.MockTransport): + def __init__(self) -> None: + super().__init__(self._handler) + + def _handler(self, request: httpx.Request) -> httpx.Response: + return httpx.Response( + 200, + json={"choices": [{"message": {"content": "ok"}}]}, + ) + + task = _batch_task(["a", "b"]) + result = _run(APIExecutor.__new__(APIExecutor), task, _NoUsage(), tmp_path) + assert len(result.items) == 2 + assert result.items[0].text == "ok" + assert result.items[0].usage is None + + def test_5xx_with_raise_for_status_false_produces_aligned_item( + self, tmp_path: Path + ) -> None: + """A 5xx with raise_for_status false still yields a row-aligned item.""" + + class _ErrorBody(httpx.MockTransport): + def __init__(self) -> None: + super().__init__(self._handler) + + def _handler(self, request: httpx.Request) -> httpx.Response: + return httpx.Response( + 503, + json={"error": {"message": "overloaded"}}, + ) + + task = _batch_task(["a", "b"], response={"raise_for_status": False}) + result = _run(APIExecutor.__new__(APIExecutor), task, _ErrorBody(), tmp_path) + assert len(result.items) == 2 + assert result.items[0].status_code == 503 + assert result.items[0].response_json == {"error": {"message": "overloaded"}} + + def test_retryable_503_raises_retryable(self, tmp_path: Path) -> None: + """A 503 is classified retryable even without a success payload.""" + + class _ErrorBody(httpx.MockTransport): + def __init__(self) -> None: + super().__init__(self._handler) + + def _handler(self, request: httpx.Request) -> httpx.Response: + return httpx.Response( + 503, + json={"error": {"message": "overloaded"}}, + ) + + task = _batch_task(["a", "b"], response={"raise_for_status": True}) + with pytest.raises(ExecutionError, match="row 0") as excinfo: + _run(APIExecutor.__new__(APIExecutor), task, _ErrorBody(), tmp_path) + assert excinfo.value.retryable is True From 03410047a7037c25d0a9ccd6b28ddce65ef82a21 Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Sat, 19 Sep 2026 03:45:50 +0700 Subject: [PATCH 20/71] feat: add task-scoped cancellation to the API executor The executor inherited the no-op Executor.cancel(), so an interrupt during a batch let queued rows continue and the batch could report success. Add a per-instance cancel event, set it from cancel(), and check it before issuing each request and before submitting queued futures, raising TaskCancelledError so the runner aborts the batch. Co-Authored-By: Claude Code Signed-off-by: Zhengyuan Su --- src/worker/executors/api_executor.py | 18 +++++++--- tests/worker/test_api_executor_batch.py | 45 +++++++++++++++++++++++-- 2 files changed, 56 insertions(+), 7 deletions(-) diff --git a/src/worker/executors/api_executor.py b/src/worker/executors/api_executor.py index b24e51980..e04865847 100644 --- a/src/worker/executors/api_executor.py +++ b/src/worker/executors/api_executor.py @@ -13,7 +13,12 @@ from shared.tasks.task_type import TaskType from shared.utils.redact import is_credential_key -from .base_executor import ExecutionError, Executor, ExecutorTask, TaskCancelledError +from .base_executor import ( + ExecutionError, + Executor, + ExecutorTask, + TaskCancelledError, +) from .mixins.data import DataMixin logger = logging.getLogger(__name__) @@ -391,6 +396,8 @@ def _run(self, task: ExecutorTask, out_dir: Path) -> APIResult: request_kwargs = self._build_request_kwargs(api_cfg, None) def _issue(idx: int, prompt: Any) -> APIItem: + if self._cancel_event.is_set(): + raise TaskCancelledError("API task cancelled") prompt_str = self._prompt_to_str(prompt) kwargs = self._substitute_prompt(request_kwargs, prompt_str) try: @@ -428,10 +435,11 @@ def _issue(idx: int, prompt: Any) -> APIItem: results: dict[int, APIItem] = {} with ThreadPoolExecutor(max_workers=concurrency) as pool: - futures = { - pool.submit(_issue, idx, prompt): idx - for idx, prompt in enumerate(prompts) - } + futures = {} + for idx, prompt in enumerate(prompts): + if self._cancel_event.is_set(): + raise TaskCancelledError("API task cancelled") + futures[pool.submit(_issue, idx, prompt)] = idx for future in as_completed(futures): idx = futures[future] results[idx] = future.result() diff --git a/tests/worker/test_api_executor_batch.py b/tests/worker/test_api_executor_batch.py index 026064aea..841cbd871 100644 --- a/tests/worker/test_api_executor_batch.py +++ b/tests/worker/test_api_executor_batch.py @@ -12,7 +12,7 @@ from shared.tasks.worker_message import WorkerTaskMessage from worker.executors.api_executor import APIExecutor -from worker.executors.base_executor import ExecutionError +from worker.executors.base_executor import ExecutionError, TaskCancelledError def _task_message(**spec_updates: Any) -> WorkerTaskMessage: @@ -492,6 +492,47 @@ def _handler(self, request: httpx.Request) -> httpx.Response: ) task = _batch_task(["a", "b"], response={"raise_for_status": True}) - with pytest.raises(ExecutionError, match="row 0") as excinfo: + with pytest.raises(ExecutionError, match="status 503") as excinfo: _run(APIExecutor.__new__(APIExecutor), task, _ErrorBody(), tmp_path) assert excinfo.value.retryable is True + + def test_cancel_during_batch_raises_task_cancelled(self, tmp_path: Path) -> None: + """Cancelling a running batch raises TaskCancelledError.""" + + class _BlockingTransport(httpx.MockTransport): + def __init__(self) -> None: + self.started = threading.Event() + self.release = threading.Event() + super().__init__(self._handler) + + def _handler(self, request: httpx.Request) -> httpx.Response: + self.started.set() + self.release.wait(timeout=5) + return httpx.Response( + 200, + json={ + "choices": [{"message": {"content": "ok"}}], + "usage": {"total_tokens": 3}, + }, + ) + + executor = APIExecutor.__new__(APIExecutor) + task = _batch_task(["a", "b", "c", "d"]) + transport = _BlockingTransport() + errors: list[BaseException] = [] + + def _run_in_thread() -> None: + try: + _run(executor, task, transport, tmp_path) + except BaseException as exc: # noqa: BLE001 - captured for assertion + errors.append(exc) + + thread = threading.Thread(target=_run_in_thread) + thread.start() + assert transport.started.wait(timeout=5) + executor.cancel("task-api-batch") + transport.release.set() + thread.join(timeout=10) + + assert len(errors) == 1 + assert isinstance(errors[0], TaskCancelledError) From 3f9e5c7d3ff1f8816da9369e81733f93bfe94718 Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Sat, 19 Sep 2026 03:56:02 +0700 Subject: [PATCH 21/71] test: correct the SDK APIItem docstring and narrow the NVML skip The SDK by-name round-trip test only validates and revalidates the SDK model, so its docstring no longer claims a cross-package worker/server boundary. Narrow the NVML fixture's skip to pynvml.NVMLError so a real initialization bug is not swallowed as a missing GPU. Co-Authored-By: Claude Code Signed-off-by: Zhengyuan Su --- tests/sdk/test_models.py | 4 ++-- tests/worker/test_mp_executor_cleanup_gpu.py | 2 +- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/tests/sdk/test_models.py b/tests/sdk/test_models.py index c863442bd..505a5a33d 100644 --- a/tests/sdk/test_models.py +++ b/tests/sdk/test_models.py @@ -495,8 +495,8 @@ def test_profile_summary_rejects_extra_fields(self) -> None: class TestAPIItem: def test_construct_by_name_serialize_revalidate(self) -> None: - """An SDK APIItem built by field name round-trips through the worker's - serialization and the server's ingest validation.""" + """An SDK APIItem built by field name round-trips through the SDK's own + serialization and revalidation.""" item = APIItem.model_validate( { "index": 0, diff --git a/tests/worker/test_mp_executor_cleanup_gpu.py b/tests/worker/test_mp_executor_cleanup_gpu.py index 10874002c..64c85c7e6 100644 --- a/tests/worker/test_mp_executor_cleanup_gpu.py +++ b/tests/worker/test_mp_executor_cleanup_gpu.py @@ -19,7 +19,7 @@ def _nvml() -> Iterator[None]: try: pynvml.nvmlInit() - except Exception: + except pynvml.NVMLError: pytest.skip("NVML unavailable (no GPU)") yield pynvml.nvmlShutdown() From 40bf0a202f8222a4feec97237b24655c37fc5e23 Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Sat, 19 Sep 2026 05:20:44 +0700 Subject: [PATCH 22/71] pin pre-start and in-flight cancellation guards separately Signed-off-by: Zhengyuan Su --- tests/worker/test_api_executor_batch.py | 72 +++++++++++++++++++++++-- 1 file changed, 68 insertions(+), 4 deletions(-) diff --git a/tests/worker/test_api_executor_batch.py b/tests/worker/test_api_executor_batch.py index 841cbd871..340e10456 100644 --- a/tests/worker/test_api_executor_batch.py +++ b/tests/worker/test_api_executor_batch.py @@ -1,5 +1,6 @@ """Tests for the API executor's batch mode (one task, N row-aligned requests).""" +import concurrent.futures import json import threading import time @@ -496,16 +497,18 @@ def _handler(self, request: httpx.Request) -> httpx.Response: _run(APIExecutor.__new__(APIExecutor), task, _ErrorBody(), tmp_path) assert excinfo.value.retryable is True - def test_cancel_during_batch_raises_task_cancelled(self, tmp_path: Path) -> None: - """Cancelling a running batch raises TaskCancelledError.""" + def test_cancel_prevents_queued_rows_from_issuing(self, tmp_path: Path) -> None: + """After cancel, a row that has not started never issues its request.""" class _BlockingTransport(httpx.MockTransport): def __init__(self) -> None: + self.requests: list[httpx.Request] = [] self.started = threading.Event() self.release = threading.Event() super().__init__(self._handler) def _handler(self, request: httpx.Request) -> httpx.Response: + self.requests.append(request) self.started.set() self.release.wait(timeout=5) return httpx.Response( @@ -517,22 +520,83 @@ def _handler(self, request: httpx.Request) -> httpx.Response: ) executor = APIExecutor.__new__(APIExecutor) - task = _batch_task(["a", "b", "c", "d"]) + task = _batch_task(["a", "b", "c", "d"], concurrency=1) transport = _BlockingTransport() errors: list[BaseException] = [] + submitted: list[Any] = [] + all_submitted = threading.Event() + real_submit = concurrent.futures.ThreadPoolExecutor.submit + + def _recording_submit(self: Any, fn: Any, *args: Any, **kwargs: Any) -> Any: + submitted.append(fn) + if len(submitted) == 4: + all_submitted.set() + return real_submit(self, fn, *args, **kwargs) def _run_in_thread() -> None: try: - _run(executor, task, transport, tmp_path) + with patch.object( + concurrent.futures.ThreadPoolExecutor, + "submit", + _recording_submit, + ): + _run(executor, task, transport, tmp_path) except BaseException as exc: # noqa: BLE001 - captured for assertion errors.append(exc) thread = threading.Thread(target=_run_in_thread) thread.start() assert transport.started.wait(timeout=5) + assert all_submitted.wait(timeout=5) executor.cancel("task-api-batch") transport.release.set() thread.join(timeout=10) + assert len(transport.requests) == 1 + assert len(errors) == 1 + assert isinstance(errors[0], TaskCancelledError) + + def test_cancel_before_submission_prevents_any_future(self, tmp_path: Path) -> None: + """Cancelling before futures are submitted surfaces TaskCancelledError + without submitting any future.""" + executor = APIExecutor.__new__(APIExecutor) + task = _batch_task(["a", "b", "c", "d"]) + transport = _RecordingTransport() + errors: list[BaseException] = [] + submitted: list[Any] = [] + real_submit = concurrent.futures.ThreadPoolExecutor.submit + real_base_url = APIExecutor._base_url + + def _recording_submit(self: Any, fn: Any, *args: Any, **kwargs: Any) -> Any: + submitted.append(fn) + return real_submit(self, fn, *args, **kwargs) + + def _run_in_thread() -> None: + def _cancel_then_base_url(url: str) -> str: + executor.cancel("task-api-batch") + return real_base_url(url) + + try: + with ( + patch.object( + concurrent.futures.ThreadPoolExecutor, + "submit", + _recording_submit, + ), + patch.object( + APIExecutor, + "_base_url", + side_effect=_cancel_then_base_url, + ), + ): + _run(executor, task, transport, tmp_path) + except BaseException as exc: # noqa: BLE001 - captured for assertion + errors.append(exc) + + thread = threading.Thread(target=_run_in_thread) + thread.start() + thread.join(timeout=10) + assert len(errors) == 1 assert isinstance(errors[0], TaskCancelledError) + assert submitted == [] From 9691501651d13f1e18c8472b210da8d2fde5b309 Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Sat, 19 Sep 2026 07:18:24 +0700 Subject: [PATCH 23/71] fix: recheck cancellation during in-flight requests and before run Signed-off-by: Zhengyuan Su --- src/worker/executors/api_executor.py | 5 ++ tests/worker/test_api_executor_batch.py | 83 +++++++++++++++++++++++++ 2 files changed, 88 insertions(+) diff --git a/src/worker/executors/api_executor.py b/src/worker/executors/api_executor.py index e04865847..c1fae82e7 100644 --- a/src/worker/executors/api_executor.py +++ b/src/worker/executors/api_executor.py @@ -431,6 +431,9 @@ def _issue(idx: int, prompt: Any) -> APIItem: item.index = idx item.prompt = prompt_str + if self._cancel_event.is_set(): + raise TaskCancelledError("API task cancelled") + return item results: dict[int, APIItem] = {} @@ -442,6 +445,8 @@ def _issue(idx: int, prompt: Any) -> APIItem: futures[pool.submit(_issue, idx, prompt)] = idx for future in as_completed(futures): idx = futures[future] + if self._cancel_event.is_set(): + raise TaskCancelledError("API task cancelled") results[idx] = future.result() items = [results[idx] for idx in range(len(prompts))] diff --git a/tests/worker/test_api_executor_batch.py b/tests/worker/test_api_executor_batch.py index 340e10456..204c8b002 100644 --- a/tests/worker/test_api_executor_batch.py +++ b/tests/worker/test_api_executor_batch.py @@ -12,6 +12,7 @@ import pytest from shared.tasks.worker_message import WorkerTaskMessage +from worker.executors import api_executor as api_executor_module from worker.executors.api_executor import APIExecutor from worker.executors.base_executor import ExecutionError, TaskCancelledError @@ -600,3 +601,85 @@ def _cancel_then_base_url(url: str) -> str: assert len(errors) == 1 assert isinstance(errors[0], TaskCancelledError) assert submitted == [] + + def test_cancel_during_in_flight_requests_not_done(self, tmp_path: Path) -> None: + """A cancel arriving while requests are in flight fails the task rather + than returning DONE, even when every row already passed the pre-request + guard.""" + + class _RecordingTransport(httpx.MockTransport): + def __init__(self) -> None: + self.requests: list[httpx.Request] = [] + super().__init__(self._handler) + + def _handler(self, request: httpx.Request) -> httpx.Response: + self.requests.append(request) + return httpx.Response( + 200, + json={ + "choices": [{"message": {"content": "ok"}}], + "usage": {"total_tokens": 3}, + }, + ) + + executor = APIExecutor.__new__(APIExecutor) + task = _batch_task(["a", "b"], concurrency=2) + transport = _RecordingTransport() + errors: list[BaseException] = [] + futures: list[Any] = [] + collect_release = threading.Event() + real_submit = concurrent.futures.ThreadPoolExecutor.submit + real_as_completed = concurrent.futures.as_completed + + def _recording_submit(self: Any, fn: Any, *args: Any, **kwargs: Any) -> Any: + future = real_submit(self, fn, *args, **kwargs) + futures.append(future) + return future + + def _blocking_as_completed(fs: Any, timeout: float | None = None) -> Any: + collect_release.wait(timeout=5) + return real_as_completed(fs, timeout=timeout) + + def _run_in_thread() -> None: + try: + with ( + patch.object( + concurrent.futures.ThreadPoolExecutor, + "submit", + _recording_submit, + ), + patch.object( + api_executor_module, "as_completed", _blocking_as_completed + ), + ): + _run(executor, task, transport, tmp_path) + except BaseException as exc: # noqa: BLE001 - captured for assertion + errors.append(exc) + + thread = threading.Thread(target=_run_in_thread) + thread.start() + assert transport.requests or True + while len(futures) < 2: + time.sleep(0.01) + for future in futures: + assert future.done() + executor.cancel("task-api-batch") + collect_release.set() + thread.join(timeout=10) + + assert len(errors) == 1 + assert isinstance(errors[0], TaskCancelledError) + + def test_cancel_before_run_cancels(self, tmp_path: Path) -> None: + """A cancel that lands before run() starts still cancels the run.""" + executor = APIExecutor.__new__(APIExecutor) + executor._cancel_event = threading.Event() + executor._cancel_task_id = None + task = _batch_task(["a", "b"]) + transport = _RecordingTransport() + + executor.cancel("task-api-batch") + + with pytest.raises(TaskCancelledError): + _run(executor, task, transport, tmp_path) + assert transport.requests == [] From 76e218c611fa41ceacef209ea88c779d531c1c83 Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Sat, 19 Sep 2026 07:49:32 +0700 Subject: [PATCH 24/71] fix: make API executor cancel check-and-reset atomic under a lock A cancel landing between run()'s pre-run check and its event reset was dropped. Hold a lock across the check-and-clear and in cancel() so no interval exists where a cancel is accepted and then discarded. Document the capped concurrency and cancellation behavior in WORKFLOWS.md. Signed-off-by: Zhengyuan Su --- docs/WORKFLOWS.md | 5 +- tests/worker/test_api_executor.py | 24 +++++--- tests/worker/test_api_executor_batch.py | 79 +++++++++++++++++-------- 3 files changed, 76 insertions(+), 32 deletions(-) diff --git a/docs/WORKFLOWS.md b/docs/WORKFLOWS.md index 6fc6c9164..475c0cb36 100644 --- a/docs/WORKFLOWS.md +++ b/docs/WORKFLOWS.md @@ -110,7 +110,10 @@ prompt is substituted for the `{{prompt}}` placeholder in the request body. Server-side stage references are `${...}`; `{{prompt}}` is a worker-side per-row slot, so it is not touched by server-side resolution. A failure in any row fails the whole task rather than shifting the remaining rows. -`spec.api.concurrency` bounds the number of in-flight requests (default 8). +`spec.api.concurrency` bounds the number of in-flight requests and is capped +at 8 (the default); values above 8 are clamped down. Cancelling the task +aborts both in-flight and not-yet-started rows rather than letting them +complete. ```yaml spec: diff --git a/tests/worker/test_api_executor.py b/tests/worker/test_api_executor.py index cd9c545a2..1a34256be 100644 --- a/tests/worker/test_api_executor.py +++ b/tests/worker/test_api_executor.py @@ -88,6 +88,16 @@ def _run( executor.run(task, Path("/tmp/out")) +def _executor() -> APIExecutor: + """Construct an APIExecutor without a WorkerConfig, mirroring __init__'s + cancellation state so run() works under __new__.""" + executor = APIExecutor.__new__(APIExecutor) + executor._cancel_event = threading.Event() + executor._cancel_task_id = None + executor._cancel_lock = threading.Lock() + return executor + + class TestNebulaPath: def test_no_url_no_header_uses_nebula_url_and_token( self, monkeypatch: pytest.MonkeyPatch @@ -96,7 +106,7 @@ def test_no_url_no_header_uses_nebula_url_and_token( monkeypatch.setenv("NEBULA_API_TOKEN", "nebula-token") task = _task_message() transport = _RecordingTransport() - _run(APIExecutor.__new__(APIExecutor), task, transport) + _run(_executor(), task, transport) assert transport.request is not None assert transport.request.url == "https://nebula.example.com/v1/chat/completions" assert transport.request.headers["Authorization"] == "Bearer nebula-token" @@ -108,7 +118,7 @@ def test_no_url_with_header_preserves_header_and_skips_token( monkeypatch.setenv("NEBULA_API_TOKEN", "nebula-token") task = _task_message(headers={"Authorization": "Bearer custom"}) transport = _RecordingTransport() - _run(APIExecutor.__new__(APIExecutor), task, transport) + _run(_executor(), task, transport) assert transport.request is not None assert transport.request.url == "https://nebula.example.com/v1/chat/completions" assert transport.request.headers["Authorization"] == "Bearer custom" @@ -119,7 +129,7 @@ def test_neither_url_nor_base_url_raises( monkeypatch.delenv("NEBULA_API_BASE_URL", raising=False) task = _task_message() with pytest.raises(ExecutionError, match="spec.api.url or NEBULA_API_BASE_URL"): - _run(APIExecutor.__new__(APIExecutor), task, _RecordingTransport()) + _run(_executor(), task, _RecordingTransport()) class TestCustomUrl: @@ -136,7 +146,7 @@ def test_custom_url_without_credential_sends_no_nebula_token( monkeypatch.setenv("NEBULA_API_TOKEN", "nebula-token") task = _task_message(url="https://custom.example.com/v1/chat/completions") transport = _RecordingTransport() - _run(APIExecutor.__new__(APIExecutor), task, transport) + _run(_executor(), task, transport) assert transport.request is not None assert transport.request.url == "https://custom.example.com/v1/chat/completions" assert "Authorization" not in transport.request.headers @@ -150,7 +160,7 @@ def test_custom_url_with_header_preserves_header( headers={"Authorization": "Bearer custom"}, ) transport = _RecordingTransport() - _run(APIExecutor.__new__(APIExecutor), task, transport) + _run(_executor(), task, transport) assert transport.request is not None assert transport.request.url == "https://custom.example.com/v1/chat/completions" assert transport.request.headers["Authorization"] == "Bearer custom" @@ -165,7 +175,7 @@ def test_custom_url_with_x_api_key_accepted( headers={"X-API-Key": "custom-key"}, ) transport = _RecordingTransport() - _run(APIExecutor.__new__(APIExecutor), task, transport) + _run(_executor(), task, transport) assert transport.request is not None assert transport.request.url == "https://custom.example.com/v1/chat/completions" assert transport.request.headers["X-API-Key"] == "custom-key" @@ -184,7 +194,7 @@ def test_custom_url_with_only_innocent_header_stays_unauthenticated( headers={"Content-Type": "application/json"}, ) transport = _RecordingTransport() - _run(APIExecutor.__new__(APIExecutor), task, transport) + _run(_executor(), task, transport) assert transport.request is not None assert transport.request.headers["Content-Type"] == "application/json" assert "Authorization" not in transport.request.headers diff --git a/tests/worker/test_api_executor_batch.py b/tests/worker/test_api_executor_batch.py index 204c8b002..2bb21a222 100644 --- a/tests/worker/test_api_executor_batch.py +++ b/tests/worker/test_api_executor_batch.py @@ -75,6 +75,16 @@ def _run( return executor.run(task, out_dir) +def _executor() -> APIExecutor: + """Construct an APIExecutor without a WorkerConfig, mirroring __init__'s + cancellation state so run()/cancel() work under __new__.""" + executor = APIExecutor.__new__(APIExecutor) + executor._cancel_event = threading.Event() + executor._cancel_task_id = None + executor._cancel_lock = threading.Lock() + return executor + + def _batch_task(items: list[Any], **api_updates: Any) -> WorkerTaskMessage: payload = { "task_id": "task-api-batch", @@ -113,7 +123,7 @@ def _nebula_env(self, monkeypatch: pytest.MonkeyPatch) -> None: def test_issues_one_request_per_row_in_order(self, tmp_path: Path) -> None: task = _batch_task(["first", "second", "third"]) transport = _RecordingTransport() - result = _run(APIExecutor.__new__(APIExecutor), task, transport, tmp_path) + result = _run(_executor(), task, transport, tmp_path) assert len(transport.requests) == 3 issued = { json.loads(req.read())["messages"][0]["content"] @@ -159,7 +169,7 @@ def _handler(self, request: httpx.Request) -> httpx.Response: task = _batch_task(["a", "b", "c"]) transport = _ReverseTransport() - result = _run(APIExecutor.__new__(APIExecutor), task, transport, tmp_path) + result = _run(_executor(), task, transport, tmp_path) for idx, prompt in enumerate(["a", "b", "c"]): item = result.items[idx] assert item.index == idx @@ -175,7 +185,7 @@ def _handler(self, request: httpx.Request) -> httpx.Response: def test_single_row_batches_to_one_item(self, tmp_path: Path) -> None: task = _batch_task(["only"]) transport = _RecordingTransport() - result = _run(APIExecutor.__new__(APIExecutor), task, transport, tmp_path) + result = _run(_executor(), task, transport, tmp_path) assert len(transport.requests) == 1 assert len(result.items) == 1 assert result.items[0].index == 0 @@ -212,7 +222,7 @@ def _handler(self, request: httpx.Request) -> httpx.Response: task = _batch_task(["a", "b", "c"], response={"raise_for_status": True}) transport = _FailRow("b") with pytest.raises(ExecutionError, match="row 1"): - _run(APIExecutor.__new__(APIExecutor), task, transport, tmp_path) + _run(_executor(), task, transport, tmp_path) assert len(transport.requests) == 3 def test_placeholder_not_required_for_scalar_body(self, tmp_path: Path) -> None: @@ -239,7 +249,7 @@ def test_placeholder_not_required_for_scalar_body(self, tmp_path: Path) -> None: } task = WorkerTaskMessage.model_validate(payload) transport = _RecordingTransport() - result = _run(APIExecutor.__new__(APIExecutor), task, transport, tmp_path) + result = _run(_executor(), task, transport, tmp_path) assert len(transport.requests) == 2 assert len(result.items) == 2 @@ -247,7 +257,7 @@ def test_no_rows_raises(self, tmp_path: Path) -> None: task = _batch_task([]) with pytest.raises(ExecutionError, match="no rows"): _run( - APIExecutor.__new__(APIExecutor), + _executor(), task, _RecordingTransport(), tmp_path, @@ -277,7 +287,7 @@ def test_missing_data_raises(self, tmp_path: Path) -> None: task = WorkerTaskMessage.model_validate(payload) with pytest.raises(ExecutionError, match="spec.data is required"): _run( - APIExecutor.__new__(APIExecutor), + _executor(), task, _RecordingTransport(), tmp_path, @@ -306,7 +316,7 @@ def _handler(self, request: httpx.Request) -> httpx.Response: task = _batch_task([f"row-{i}" for i in range(n_rows)]) transport = _SlowTransport() start = time.monotonic() - result = _run(APIExecutor.__new__(APIExecutor), task, transport, tmp_path) + result = _run(_executor(), task, transport, tmp_path) elapsed = time.monotonic() - start assert len(transport.requests) == n_rows @@ -340,7 +350,7 @@ def _handler(self, request: httpx.Request) -> httpx.Response: task = _batch_task(["a", "b", "c", "d"], concurrency=1) transport = _OverlapTransport() - _run(APIExecutor.__new__(APIExecutor), task, transport, tmp_path) + _run(_executor(), task, transport, tmp_path) assert transport.max_in_flight == 1 def test_request_skeleton_constructed_once(self, tmp_path: Path) -> None: @@ -360,7 +370,7 @@ def test_request_skeleton_constructed_once(self, tmp_path: Path) -> None: ) as mock_build, ): mock_build.side_effect = lambda *a, **k: real_build(*a, **k) - APIExecutor.__new__(APIExecutor).run(task, tmp_path) + _executor().run(task, tmp_path) assert mock_build.call_count == 1 assert mock_build.call_args.args[2] is None @@ -393,7 +403,7 @@ def _fake_get_client(*args: Any, **kwargs: Any) -> httpx.Client: return httpx.Client(transport=_RecordingTransport()) with patch.object(APIExecutor, "_get_client", side_effect=_fake_get_client): - APIExecutor.__new__(APIExecutor).run(task, tmp_path) + _executor().run(task, tmp_path) assert captured["concurrency"] == 8 @@ -410,7 +420,7 @@ def _fake_get_client(*args: Any, **kwargs: Any) -> httpx.Client: return httpx.Client(transport=_RecordingTransport()) with patch.object(APIExecutor, "_get_client", side_effect=_fake_get_client): - APIExecutor.__new__(APIExecutor).run(task, tmp_path) + _executor().run(task, tmp_path) assert captured["concurrency"] == concurrency @@ -421,9 +431,7 @@ def test_concurrency_below_one_rejected( """A configured concurrency below 1 is rejected.""" task = _batch_task(["a", "b", "c"], concurrency=concurrency) with pytest.raises(ExecutionError, match="spec.api.concurrency must be >= 1"): - _run( - APIExecutor.__new__(APIExecutor), task, _RecordingTransport(), tmp_path - ) + _run(_executor(), task, _RecordingTransport(), tmp_path) def test_client_cache_key_includes_concurrency(self) -> None: """Pools built for different concurrency values are not shared.""" @@ -454,7 +462,7 @@ def _handler(self, request: httpx.Request) -> httpx.Response: ) task = _batch_task(["a", "b"]) - result = _run(APIExecutor.__new__(APIExecutor), task, _NoUsage(), tmp_path) + result = _run(_executor(), task, _NoUsage(), tmp_path) assert len(result.items) == 2 assert result.items[0].text == "ok" assert result.items[0].usage is None @@ -475,7 +483,7 @@ def _handler(self, request: httpx.Request) -> httpx.Response: ) task = _batch_task(["a", "b"], response={"raise_for_status": False}) - result = _run(APIExecutor.__new__(APIExecutor), task, _ErrorBody(), tmp_path) + result = _run(_executor(), task, _ErrorBody(), tmp_path) assert len(result.items) == 2 assert result.items[0].status_code == 503 assert result.items[0].response_json == {"error": {"message": "overloaded"}} @@ -495,7 +503,7 @@ def _handler(self, request: httpx.Request) -> httpx.Response: task = _batch_task(["a", "b"], response={"raise_for_status": True}) with pytest.raises(ExecutionError, match="status 503") as excinfo: - _run(APIExecutor.__new__(APIExecutor), task, _ErrorBody(), tmp_path) + _run(_executor(), task, _ErrorBody(), tmp_path) assert excinfo.value.retryable is True def test_cancel_prevents_queued_rows_from_issuing(self, tmp_path: Path) -> None: @@ -520,7 +528,7 @@ def _handler(self, request: httpx.Request) -> httpx.Response: }, ) - executor = APIExecutor.__new__(APIExecutor) + executor = _executor() task = _batch_task(["a", "b", "c", "d"], concurrency=1) transport = _BlockingTransport() errors: list[BaseException] = [] @@ -560,7 +568,7 @@ def _run_in_thread() -> None: def test_cancel_before_submission_prevents_any_future(self, tmp_path: Path) -> None: """Cancelling before futures are submitted surfaces TaskCancelledError without submitting any future.""" - executor = APIExecutor.__new__(APIExecutor) + executor = _executor() task = _batch_task(["a", "b", "c", "d"]) transport = _RecordingTransport() errors: list[BaseException] = [] @@ -622,7 +630,7 @@ def _handler(self, request: httpx.Request) -> httpx.Response: }, ) - executor = APIExecutor.__new__(APIExecutor) + executor = _executor() task = _batch_task(["a", "b"], concurrency=2) transport = _RecordingTransport() errors: list[BaseException] = [] @@ -672,9 +680,7 @@ def _run_in_thread() -> None: def test_cancel_before_run_cancels(self, tmp_path: Path) -> None: """A cancel that lands before run() starts still cancels the run.""" - executor = APIExecutor.__new__(APIExecutor) - executor._cancel_event = threading.Event() - executor._cancel_task_id = None + executor = _executor() task = _batch_task(["a", "b"]) transport = _RecordingTransport() @@ -683,3 +689,28 @@ def test_cancel_before_run_cancels(self, tmp_path: Path) -> None: with pytest.raises(TaskCancelledError): _run(executor, task, transport, tmp_path) assert transport.requests == [] + + def test_cancel_during_run_setup_not_lost(self) -> None: + """A cancel arriving while run() holds the lock across check-and-clear + is not dropped: cancel() blocks on the same lock and sets the event + once run() releases it.""" + executor = _executor() + task = _batch_task(["a", "b"]) + + executor._cancel_lock.acquire() + cancelled: list[bool] = [] + + def _cancel_in_thread() -> None: + executor.cancel(task.task_id) + cancelled.append(executor._cancel_event.is_set()) + + thread = threading.Thread(target=_cancel_in_thread) + thread.start() + thread.join(timeout=0.2) + assert not cancelled + executor._cancel_lock.release() + thread.join(timeout=5) + + assert cancelled == [True] + assert executor._cancel_event.is_set() + assert executor._cancel_task_id == task.task_id From d4ced9f0b874aa2f37b2e884a159e223875d0b4d Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Sat, 19 Sep 2026 08:39:02 +0700 Subject: [PATCH 25/71] fix: correct cancellation docs and rename collection-guard test The docs claimed in-flight requests are aborted on cancel; they are not - a request inside the synchronous HTTP call is not interrupted. Say that cancellation prevents queued rows and marks the task cancelled once in-flight requests return. Rename the collection-guard test to describe what it actually asserts. Signed-off-by: Zhengyuan Su --- docs/WORKFLOWS.md | 5 ++- tests/worker/test_api_executor_batch.py | 58 +++++++++++++++---------- 2 files changed, 39 insertions(+), 24 deletions(-) diff --git a/docs/WORKFLOWS.md b/docs/WORKFLOWS.md index 475c0cb36..5d1debc61 100644 --- a/docs/WORKFLOWS.md +++ b/docs/WORKFLOWS.md @@ -112,8 +112,9 @@ per-row slot, so it is not touched by server-side resolution. A failure in any row fails the whole task rather than shifting the remaining rows. `spec.api.concurrency` bounds the number of in-flight requests and is capped at 8 (the default); values above 8 are clamped down. Cancelling the task -aborts both in-flight and not-yet-started rows rather than letting them -complete. +prevents not-yet-started rows from issuing and marks the task cancelled once +in-flight requests return; a request already inside the HTTP call is not +interrupted. ```yaml spec: diff --git a/tests/worker/test_api_executor_batch.py b/tests/worker/test_api_executor_batch.py index 2bb21a222..81bae82a7 100644 --- a/tests/worker/test_api_executor_batch.py +++ b/tests/worker/test_api_executor_batch.py @@ -610,10 +610,11 @@ def _cancel_then_base_url(url: str) -> str: assert isinstance(errors[0], TaskCancelledError) assert submitted == [] - def test_cancel_during_in_flight_requests_not_done(self, tmp_path: Path) -> None: - """A cancel arriving while requests are in flight fails the task rather - than returning DONE, even when every row already passed the pre-request - guard.""" + def test_cancel_after_requests_complete_before_collection_not_done( + self, tmp_path: Path + ) -> None: + """A cancel arriving after every request has completed but before results + are collected fails the task rather than returning DONE.""" class _RecordingTransport(httpx.MockTransport): def __init__(self) -> None: @@ -690,27 +691,40 @@ def test_cancel_before_run_cancels(self, tmp_path: Path) -> None: _run(executor, task, transport, tmp_path) assert transport.requests == [] - def test_cancel_during_run_setup_not_lost(self) -> None: - """A cancel arriving while run() holds the lock across check-and-clear - is not dropped: cancel() blocks on the same lock and sets the event - once run() releases it.""" + def test_cancel_during_run_setup_not_lost(self, tmp_path: Path) -> None: + """A cancel arriving while run() is mid check-and-clear is not dropped. + + run() checks the event, then clears it under the lock; cancel() sets it + under the same lock. Pausing run() inside clear() and firing cancel() + while it is paused must still cancel the run (via the in-flight guards), + not let it complete as if never cancelled.""" executor = _executor() task = _batch_task(["a", "b"]) + transport = _RecordingTransport() + errors: list[BaseException] = [] + in_clear = threading.Event() + release_clear = threading.Event() + real_clear = executor._cancel_event.clear - executor._cancel_lock.acquire() - cancelled: list[bool] = [] + def _blocking_clear() -> None: + in_clear.set() + release_clear.wait(timeout=5) + real_clear() - def _cancel_in_thread() -> None: - executor.cancel(task.task_id) - cancelled.append(executor._cancel_event.is_set()) + executor._cancel_event.clear = _blocking_clear # type: ignore[method-assign] - thread = threading.Thread(target=_cancel_in_thread) + def _run_in_thread() -> None: + try: + _run(executor, task, transport, tmp_path) + except BaseException as exc: # noqa: BLE001 - captured for assertion + errors.append(exc) + + thread = threading.Thread(target=_run_in_thread) thread.start() - thread.join(timeout=0.2) - assert not cancelled - executor._cancel_lock.release() - thread.join(timeout=5) - - assert cancelled == [True] - assert executor._cancel_event.is_set() - assert executor._cancel_task_id == task.task_id + assert in_clear.wait(timeout=5) + executor.cancel(task.task_id) + release_clear.set() + thread.join(timeout=10) + + assert len(errors) == 1 + assert isinstance(errors[0], TaskCancelledError) From b6c01a6a75d34621e198bc93c357c7f9f3ee418e Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Sat, 19 Sep 2026 11:53:11 +0700 Subject: [PATCH 26/71] test: merge API executor batch tests into test_api_executor.py Signed-off-by: Zhengyuan Su Co-Authored-By: Claude Code --- tests/worker/test_api_executor.py | 684 +++++++++++++++++++++- tests/worker/test_api_executor_batch.py | 730 ------------------------ 2 files changed, 679 insertions(+), 735 deletions(-) delete mode 100644 tests/worker/test_api_executor_batch.py diff --git a/tests/worker/test_api_executor.py b/tests/worker/test_api_executor.py index 1a34256be..274904afb 100644 --- a/tests/worker/test_api_executor.py +++ b/tests/worker/test_api_executor.py @@ -1,14 +1,19 @@ -"""Tests for the API executor url override and Nebula credential handling.""" +"""Tests for the API executor's url override, Nebula credential handling, and +batch mode (one task, N row-aligned requests).""" +import concurrent.futures +import json import threading import time from pathlib import Path +from typing import Any, cast from unittest.mock import patch import httpx import pytest from shared.tasks.worker_message import WorkerTaskMessage +from worker.executors import api_executor as api_executor_module from worker.executors.api_executor import APIExecutor from worker.executors.base_executor import ExecutionError, TaskCancelledError @@ -75,9 +80,34 @@ def _handler(self, request: httpx.Request) -> httpx.Response: return self.responses.pop(0) +class _EchoTransport(httpx.MockTransport): + """MockTransport that echoes each row's prompt back with row-specific + status, usage, and headers.""" + + def __init__(self) -> None: + self.requests: list[httpx.Request] = [] + super().__init__(self._handler) + + def _handler(self, request: httpx.Request) -> httpx.Response: + self.requests.append(request) + body = request.read() + prompt = json.loads(body)["messages"][0]["content"] + return httpx.Response( + 200 + len(prompt) % 3, + headers={"X-Row": prompt}, + json={ + "choices": [{"message": {"content": f"echo:{prompt}"}}], + "usage": {"total_tokens": len(prompt)}, + }, + ) + + def _run( - executor: APIExecutor, task: WorkerTaskMessage, transport: httpx.MockTransport -) -> None: + executor: APIExecutor, + task: WorkerTaskMessage, + transport: httpx.MockTransport, + out_dir: Path = Path("/tmp/out"), +): executor._cancel_event = threading.Event() executor._cancel_lock = threading.Lock() executor._active_task_id = None @@ -85,12 +115,12 @@ def _run( with patch.object( APIExecutor, "_get_client", return_value=httpx.Client(transport=transport) ): - executor.run(task, Path("/tmp/out")) + return executor.run(task, out_dir) def _executor() -> APIExecutor: """Construct an APIExecutor without a WorkerConfig, mirroring __init__'s - cancellation state so run() works under __new__.""" + cancellation state so run()/cancel() work under __new__.""" executor = APIExecutor.__new__(APIExecutor) executor._cancel_event = threading.Event() executor._cancel_task_id = None @@ -98,6 +128,35 @@ def _executor() -> APIExecutor: return executor +def _batch_task(items: list[Any], **api_updates: Any) -> WorkerTaskMessage: + payload = { + "task_id": "task-api-batch", + "workflow_id": "wf-1", + "owner_id": "owner", + "assigned_worker": "worker-1", + "dispatched_at": "2026-03-22T00:00:00Z", + "task": { + "apiVersion": "mloc/v1", + "kind": "Task", + "metadata": {"name": "wf:api"}, + "spec": { + "taskType": "api", + "api": { + "method": "POST", + "body": {"messages": [{"role": "user", "content": "{{prompt}}"}]}, + "response": { + "raise_for_status": False, + "include_headers": True, + }, + **api_updates, + }, + "data": {"type": "list", "items": items}, + }, + }, + } + return WorkerTaskMessage.model_validate(payload) + + class TestNebulaPath: def test_no_url_no_header_uses_nebula_url_and_token( self, monkeypatch: pytest.MonkeyPatch @@ -440,3 +499,618 @@ def _cancel_after_delay() -> None: canceller.join() assert elapsed < 0.5 assert transport.calls == 1 + +class TestBatch: + @pytest.fixture(autouse=True) + def _nebula_env(self, monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("NEBULA_API_BASE_URL", "https://nebula.example.com") + monkeypatch.setenv("NEBULA_API_TOKEN", "nebula-token") + + def test_issues_one_request_per_row_in_order(self, tmp_path: Path) -> None: + task = _batch_task(["first", "second", "third"]) + transport = _EchoTransport() + result = _run(_executor(), task, transport, tmp_path) + assert len(transport.requests) == 3 + issued = { + json.loads(req.read())["messages"][0]["content"] + for req in transport.requests + } + assert issued == {"first", "second", "third"} + for idx, prompt in enumerate(["first", "second", "third"]): + item = result.items[idx] + assert item.index == idx + assert item.prompt == prompt + assert item.text == f"echo:{prompt}" + assert item.response_json["choices"][0]["message"]["content"] == ( + f"echo:{prompt}" + ) + assert item.status_code == 200 + len(prompt) % 3 + assert item.usage == {"total_tokens": len(prompt)} + assert item.headers["x-row"] == prompt + + def test_rows_stay_aligned_when_requests_complete_out_of_order( + self, tmp_path: Path + ) -> None: + """Output row i corresponds to input row i even when requests finish + in reverse order.""" + + class _ReverseTransport(httpx.MockTransport): + def __init__(self) -> None: + self.requests: list[httpx.Request] = [] + super().__init__(self._handler) + + def _handler(self, request: httpx.Request) -> httpx.Response: + self.requests.append(request) + prompt = json.loads(request.read())["messages"][0]["content"] + delay = {"a": 0.3, "b": 0.2, "c": 0.1}[prompt] + time.sleep(delay) + return httpx.Response( + 200 + len(prompt) % 3, + headers={"X-Row": prompt}, + json={ + "choices": [{"message": {"content": f"echo:{prompt}"}}], + "usage": {"total_tokens": len(prompt)}, + }, + ) + + task = _batch_task(["a", "b", "c"]) + transport = _ReverseTransport() + result = _run(_executor(), task, transport, tmp_path) + for idx, prompt in enumerate(["a", "b", "c"]): + item = result.items[idx] + assert item.index == idx + assert item.prompt == prompt + assert item.text == f"echo:{prompt}" + assert item.response_json["choices"][0]["message"]["content"] == ( + f"echo:{prompt}" + ) + assert item.status_code == 200 + len(prompt) % 3 + assert item.usage == {"total_tokens": len(prompt)} + assert item.headers["x-row"] == prompt + + def test_single_row_batches_to_one_item(self, tmp_path: Path) -> None: + task = _batch_task(["only"]) + transport = _EchoTransport() + result = _run(_executor(), task, transport, tmp_path) + assert len(transport.requests) == 1 + assert len(result.items) == 1 + assert result.items[0].index == 0 + assert result.items[0].prompt == "only" + + def test_row_failure_fails_whole_task_without_shifting( + self, tmp_path: Path + ) -> None: + class _FailRow(httpx.MockTransport): + def __init__(self, failing_prompt: str) -> None: + self.failing_prompt = failing_prompt + self.requests: list[httpx.Request] = [] + super().__init__(self._handler) + + def _handler(self, request: httpx.Request) -> httpx.Response: + self.requests.append(request) + prompt = json.loads(request.read())["messages"][0]["content"] + if prompt == self.failing_prompt: + return httpx.Response( + 500, + json={ + "choices": [{"message": {"content": "boom"}}], + "usage": {"total_tokens": 1}, + }, + ) + return httpx.Response( + 200, + json={ + "choices": [{"message": {"content": "ok"}}], + "usage": {"total_tokens": 1}, + }, + ) + + task = _batch_task(["a", "b", "c"], response={"raise_for_status": True}) + transport = _FailRow("b") + with pytest.raises(ExecutionError, match="row 1"): + _run(_executor(), task, transport, tmp_path) + assert len(transport.requests) == 3 + + def test_placeholder_not_required_for_scalar_body(self, tmp_path: Path) -> None: + """A batch task whose body has no placeholder still issues N requests.""" + payload = { + "task_id": "task-api-batch", + "workflow_id": "wf-1", + "owner_id": "owner", + "assigned_worker": "worker-1", + "dispatched_at": "2026-03-22T00:00:00Z", + "task": { + "apiVersion": "mloc/v1", + "kind": "Task", + "metadata": {"name": "wf:api"}, + "spec": { + "taskType": "api", + "api": { + "method": "POST", + "body": {"messages": [{"role": "user", "content": "static"}]}, + }, + "data": {"type": "list", "items": ["a", "b"]}, + }, + }, + } + task = WorkerTaskMessage.model_validate(payload) + transport = _EchoTransport() + result = _run(_executor(), task, transport, tmp_path) + assert len(transport.requests) == 2 + assert len(result.items) == 2 + + def test_no_rows_raises(self, tmp_path: Path) -> None: + task = _batch_task([]) + with pytest.raises(ExecutionError, match="no rows"): + _run( + _executor(), + task, + _EchoTransport(), + tmp_path, + ) + + def test_missing_data_raises(self, tmp_path: Path) -> None: + """spec.data is required; an api task without it fails closed.""" + payload = { + "task_id": "task-api-batch", + "workflow_id": "wf-1", + "owner_id": "owner", + "assigned_worker": "worker-1", + "dispatched_at": "2026-03-22T00:00:00Z", + "task": { + "apiVersion": "mloc/v1", + "kind": "Task", + "metadata": {"name": "wf:api"}, + "spec": { + "taskType": "api", + "api": { + "method": "POST", + "body": {"messages": [{"role": "user", "content": "hi"}]}, + }, + }, + }, + } + task = WorkerTaskMessage.model_validate(payload) + with pytest.raises(ExecutionError, match="spec.data is required"): + _run( + _executor(), + task, + _EchoTransport(), + tmp_path, + ) + + def test_requests_issue_in_parallel(self, tmp_path: Path) -> None: + """N rows take ~one row's latency, not N x, on a network-bound path.""" + + class _SlowTransport(httpx.MockTransport): + def __init__(self) -> None: + self.requests: list[httpx.Request] = [] + super().__init__(self._handler) + + def _handler(self, request: httpx.Request) -> httpx.Response: + self.requests.append(request) + time.sleep(0.2) + return httpx.Response( + 200, + json={ + "choices": [{"message": {"content": "hello"}}], + "usage": {"total_tokens": 3}, + }, + ) + + n_rows = 4 + task = _batch_task([f"row-{i}" for i in range(n_rows)]) + transport = _SlowTransport() + start = time.monotonic() + result = _run(_executor(), task, transport, tmp_path) + elapsed = time.monotonic() - start + + assert len(transport.requests) == n_rows + assert elapsed < 0.2 * n_rows * 0.6 + assert [item.index for item in result.items] == list(range(n_rows)) + + def test_concurrency_one_serializes_requests(self, tmp_path: Path) -> None: + """concurrency: 1 limits the worker pool so requests never overlap.""" + + class _OverlapTransport(httpx.MockTransport): + def __init__(self) -> None: + self.max_in_flight = 0 + self._in_flight = 0 + self._lock = threading.Lock() + super().__init__(self._handler) + + def _handler(self, request: httpx.Request) -> httpx.Response: + with self._lock: + self._in_flight += 1 + self.max_in_flight = max(self.max_in_flight, self._in_flight) + time.sleep(0.05) + with self._lock: + self._in_flight -= 1 + return httpx.Response( + 200, + json={ + "choices": [{"message": {"content": "hello"}}], + "usage": {"total_tokens": 3}, + }, + ) + + task = _batch_task(["a", "b", "c", "d"], concurrency=1) + transport = _OverlapTransport() + _run(_executor(), task, transport, tmp_path) + assert transport.max_in_flight == 1 + + def test_request_skeleton_constructed_once(self, tmp_path: Path) -> None: + """The request template is built once, not once per row.""" + task = _batch_task(["a", "b", "c"]) + transport = _EchoTransport() + real_build = APIExecutor._build_request_kwargs + + with ( + patch.object( + APIExecutor, + "_get_client", + return_value=httpx.Client(transport=transport), + ), + patch.object( + APIExecutor, "_build_request_kwargs", autospec=True + ) as mock_build, + ): + mock_build.side_effect = lambda *a, **k: real_build(*a, **k) + _executor().run(task, tmp_path) + + assert mock_build.call_count == 1 + assert mock_build.call_args.args[2] is None + + @pytest.mark.parametrize("concurrency", [1, 4, 8]) + def test_client_pool_sized_to_concurrency(self, concurrency: int) -> None: + """The connection pool matches the effective concurrency.""" + APIExecutor.close_all_clients() + try: + client = APIExecutor._get_client( + "https://example.com", + httpx.Timeout(60), + True, + True, + concurrency, + ) + pool = cast(Any, client._transport)._pool + assert pool._max_connections == concurrency + assert pool._max_keepalive_connections == concurrency + finally: + APIExecutor.close_all_clients() + + def test_concurrency_capped_at_max(self, tmp_path: Path) -> None: + """A configured concurrency above the cap is clamped to the cap.""" + task = _batch_task(["a", "b", "c"], concurrency=100) + captured: dict[str, Any] = {} + + def _fake_get_client(*args: Any, **kwargs: Any) -> httpx.Client: + captured["concurrency"] = kwargs.get("concurrency", args[4]) + return httpx.Client(transport=_EchoTransport()) + + with patch.object(APIExecutor, "_get_client", side_effect=_fake_get_client): + _executor().run(task, tmp_path) + + assert captured["concurrency"] == 8 + + @pytest.mark.parametrize("concurrency", [1, 4]) + def test_run_passes_effective_concurrency_to_client( + self, tmp_path: Path, concurrency: int + ) -> None: + """run() forwards the uncapped configured concurrency to the client.""" + task = _batch_task(["a", "b", "c"], concurrency=concurrency) + captured: dict[str, Any] = {} + + def _fake_get_client(*args: Any, **kwargs: Any) -> httpx.Client: + captured["concurrency"] = kwargs.get("concurrency", args[4]) + return httpx.Client(transport=_EchoTransport()) + + with patch.object(APIExecutor, "_get_client", side_effect=_fake_get_client): + _executor().run(task, tmp_path) + + assert captured["concurrency"] == concurrency + + @pytest.mark.parametrize("concurrency", [0, -1]) + def test_concurrency_below_one_rejected( + self, tmp_path: Path, concurrency: int + ) -> None: + """A configured concurrency below 1 is rejected.""" + task = _batch_task(["a", "b", "c"], concurrency=concurrency) + with pytest.raises(ExecutionError, match="spec.api.concurrency must be >= 1"): + _run(_executor(), task, _EchoTransport(), tmp_path) + + def test_client_cache_key_includes_concurrency(self) -> None: + """Pools built for different concurrency values are not shared.""" + APIExecutor.close_all_clients() + try: + c1 = APIExecutor._get_client( + "https://example.com", httpx.Timeout(60), True, True, 1 + ) + c4 = APIExecutor._get_client( + "https://example.com", httpx.Timeout(60), True, True, 4 + ) + assert c1 is not c4 + assert len(APIExecutor._clients) == 2 + finally: + APIExecutor.close_all_clients() + + def test_no_usage_2xx_produces_aligned_item(self, tmp_path: Path) -> None: + """A 2xx response without usage still yields a row-aligned item.""" + + class _NoUsage(httpx.MockTransport): + def __init__(self) -> None: + super().__init__(self._handler) + + def _handler(self, request: httpx.Request) -> httpx.Response: + return httpx.Response( + 200, + json={"choices": [{"message": {"content": "ok"}}]}, + ) + + task = _batch_task(["a", "b"]) + result = _run(_executor(), task, _NoUsage(), tmp_path) + assert len(result.items) == 2 + assert result.items[0].text == "ok" + assert result.items[0].usage is None + + def test_5xx_with_raise_for_status_false_produces_aligned_item( + self, tmp_path: Path + ) -> None: + """A 5xx with raise_for_status false still yields a row-aligned item.""" + + class _ErrorBody(httpx.MockTransport): + def __init__(self) -> None: + super().__init__(self._handler) + + def _handler(self, request: httpx.Request) -> httpx.Response: + return httpx.Response( + 503, + json={"error": {"message": "overloaded"}}, + ) + + task = _batch_task(["a", "b"], response={"raise_for_status": False}) + result = _run(_executor(), task, _ErrorBody(), tmp_path) + assert len(result.items) == 2 + assert result.items[0].status_code == 503 + assert result.items[0].response_json == {"error": {"message": "overloaded"}} + + def test_retryable_503_raises_retryable(self, tmp_path: Path) -> None: + """A 503 is classified retryable even without a success payload.""" + + class _ErrorBody(httpx.MockTransport): + def __init__(self) -> None: + super().__init__(self._handler) + + def _handler(self, request: httpx.Request) -> httpx.Response: + return httpx.Response( + 503, + json={"error": {"message": "overloaded"}}, + ) + + task = _batch_task(["a", "b"], response={"raise_for_status": True}) + with pytest.raises(ExecutionError, match="status 503") as excinfo: + _run(_executor(), task, _ErrorBody(), tmp_path) + assert excinfo.value.retryable is True + + def test_cancel_prevents_queued_rows_from_issuing(self, tmp_path: Path) -> None: + """After cancel, a row that has not started never issues its request.""" + + class _BlockingTransport(httpx.MockTransport): + def __init__(self) -> None: + self.requests: list[httpx.Request] = [] + self.started = threading.Event() + self.release = threading.Event() + super().__init__(self._handler) + + def _handler(self, request: httpx.Request) -> httpx.Response: + self.requests.append(request) + self.started.set() + self.release.wait(timeout=5) + return httpx.Response( + 200, + json={ + "choices": [{"message": {"content": "ok"}}], + "usage": {"total_tokens": 3}, + }, + ) + + executor = _executor() + task = _batch_task(["a", "b", "c", "d"], concurrency=1) + transport = _BlockingTransport() + errors: list[BaseException] = [] + submitted: list[Any] = [] + all_submitted = threading.Event() + real_submit = concurrent.futures.ThreadPoolExecutor.submit + + def _recording_submit(self: Any, fn: Any, *args: Any, **kwargs: Any) -> Any: + submitted.append(fn) + if len(submitted) == 4: + all_submitted.set() + return real_submit(self, fn, *args, **kwargs) + + def _run_in_thread() -> None: + try: + with patch.object( + concurrent.futures.ThreadPoolExecutor, + "submit", + _recording_submit, + ): + _run(executor, task, transport, tmp_path) + except BaseException as exc: # noqa: BLE001 - captured for assertion + errors.append(exc) + + thread = threading.Thread(target=_run_in_thread) + thread.start() + assert transport.started.wait(timeout=5) + assert all_submitted.wait(timeout=5) + executor.cancel("task-api-batch") + transport.release.set() + thread.join(timeout=10) + + assert len(transport.requests) == 1 + assert len(errors) == 1 + assert isinstance(errors[0], TaskCancelledError) + + def test_cancel_before_submission_prevents_any_future(self, tmp_path: Path) -> None: + """Cancelling before futures are submitted surfaces TaskCancelledError + without submitting any future.""" + executor = _executor() + task = _batch_task(["a", "b", "c", "d"]) + transport = _EchoTransport() + errors: list[BaseException] = [] + submitted: list[Any] = [] + real_submit = concurrent.futures.ThreadPoolExecutor.submit + real_base_url = APIExecutor._base_url + + def _recording_submit(self: Any, fn: Any, *args: Any, **kwargs: Any) -> Any: + submitted.append(fn) + return real_submit(self, fn, *args, **kwargs) + + def _run_in_thread() -> None: + def _cancel_then_base_url(url: str) -> str: + executor.cancel("task-api-batch") + return real_base_url(url) + + try: + with ( + patch.object( + concurrent.futures.ThreadPoolExecutor, + "submit", + _recording_submit, + ), + patch.object( + APIExecutor, + "_base_url", + side_effect=_cancel_then_base_url, + ), + ): + _run(executor, task, transport, tmp_path) + except BaseException as exc: # noqa: BLE001 - captured for assertion + errors.append(exc) + + thread = threading.Thread(target=_run_in_thread) + thread.start() + thread.join(timeout=10) + + assert len(errors) == 1 + assert isinstance(errors[0], TaskCancelledError) + assert submitted == [] + + def test_cancel_after_requests_complete_before_collection_not_done( + self, tmp_path: Path + ) -> None: + """A cancel arriving after every request has completed but before results + are collected fails the task rather than returning DONE.""" + + class _RecordingTransport(httpx.MockTransport): + def __init__(self) -> None: + self.requests: list[httpx.Request] = [] + super().__init__(self._handler) + + def _handler(self, request: httpx.Request) -> httpx.Response: + self.requests.append(request) + return httpx.Response( + 200, + json={ + "choices": [{"message": {"content": "ok"}}], + "usage": {"total_tokens": 3}, + }, + ) + + executor = _executor() + task = _batch_task(["a", "b"], concurrency=2) + transport = _RecordingTransport() + errors: list[BaseException] = [] + futures: list[Any] = [] + collect_release = threading.Event() + real_submit = concurrent.futures.ThreadPoolExecutor.submit + real_as_completed = concurrent.futures.as_completed + + def _recording_submit(self: Any, fn: Any, *args: Any, **kwargs: Any) -> Any: + future = real_submit(self, fn, *args, **kwargs) + futures.append(future) + return future + + def _blocking_as_completed(fs: Any, timeout: float | None = None) -> Any: + collect_release.wait(timeout=5) + return real_as_completed(fs, timeout=timeout) + + def _run_in_thread() -> None: + try: + with ( + patch.object( + concurrent.futures.ThreadPoolExecutor, + "submit", + _recording_submit, + ), + patch.object( + api_executor_module, "as_completed", _blocking_as_completed + ), + ): + _run(executor, task, transport, tmp_path) + except BaseException as exc: # noqa: BLE001 - captured for assertion + errors.append(exc) + + thread = threading.Thread(target=_run_in_thread) + thread.start() + assert transport.requests or True + while len(futures) < 2: + time.sleep(0.01) + for future in futures: + assert future.done() + executor.cancel("task-api-batch") + collect_release.set() + thread.join(timeout=10) + + assert len(errors) == 1 + assert isinstance(errors[0], TaskCancelledError) + + def test_cancel_before_run_cancels(self, tmp_path: Path) -> None: + """A cancel that lands before run() starts still cancels the run.""" + executor = _executor() + task = _batch_task(["a", "b"]) + transport = _EchoTransport() + + executor.cancel("task-api-batch") + + with pytest.raises(TaskCancelledError): + _run(executor, task, transport, tmp_path) + assert transport.requests == [] + + def test_cancel_during_run_setup_not_lost(self, tmp_path: Path) -> None: + """A cancel arriving while run() is mid check-and-clear is not dropped. + + run() checks the event, then clears it under the lock; cancel() sets it + under the same lock. Pausing run() inside clear() and firing cancel() + while it is paused must still cancel the run (via the in-flight guards), + not let it complete as if never cancelled.""" + executor = _executor() + task = _batch_task(["a", "b"]) + transport = _EchoTransport() + errors: list[BaseException] = [] + in_clear = threading.Event() + release_clear = threading.Event() + real_clear = executor._cancel_event.clear + + def _blocking_clear() -> None: + in_clear.set() + release_clear.wait(timeout=5) + real_clear() + + executor._cancel_event.clear = _blocking_clear # type: ignore[method-assign] + + def _run_in_thread() -> None: + try: + _run(executor, task, transport, tmp_path) + except BaseException as exc: # noqa: BLE001 - captured for assertion + errors.append(exc) + + thread = threading.Thread(target=_run_in_thread) + thread.start() + assert in_clear.wait(timeout=5) + executor.cancel(task.task_id) + release_clear.set() + thread.join(timeout=10) + + assert len(errors) == 1 + assert isinstance(errors[0], TaskCancelledError) diff --git a/tests/worker/test_api_executor_batch.py b/tests/worker/test_api_executor_batch.py deleted file mode 100644 index 81bae82a7..000000000 --- a/tests/worker/test_api_executor_batch.py +++ /dev/null @@ -1,730 +0,0 @@ -"""Tests for the API executor's batch mode (one task, N row-aligned requests).""" - -import concurrent.futures -import json -import threading -import time -from pathlib import Path -from typing import Any, cast -from unittest.mock import patch - -import httpx -import pytest - -from shared.tasks.worker_message import WorkerTaskMessage -from worker.executors import api_executor as api_executor_module -from worker.executors.api_executor import APIExecutor -from worker.executors.base_executor import ExecutionError, TaskCancelledError - - -def _task_message(**spec_updates: Any) -> WorkerTaskMessage: - payload = { - "task_id": "task-api-batch", - "workflow_id": "wf-1", - "owner_id": "owner", - "assigned_worker": "worker-1", - "dispatched_at": "2026-03-22T00:00:00Z", - "task": { - "apiVersion": "mloc/v1", - "kind": "Task", - "metadata": {"name": "wf:api"}, - "spec": { - "taskType": "api", - "api": { - "method": "POST", - "body": {"messages": [{"role": "user", "content": "{{prompt}}"}]}, - **spec_updates, - }, - }, - }, - } - return WorkerTaskMessage.model_validate(payload) - - -class _RecordingTransport(httpx.MockTransport): - """MockTransport that echoes each row's prompt back with row-specific - status, usage, and headers.""" - - def __init__(self) -> None: - self.requests: list[httpx.Request] = [] - super().__init__(self._handler) - - def _handler(self, request: httpx.Request) -> httpx.Response: - self.requests.append(request) - body = request.read() - prompt = json.loads(body)["messages"][0]["content"] - return httpx.Response( - 200 + len(prompt) % 3, - headers={"X-Row": prompt}, - json={ - "choices": [{"message": {"content": f"echo:{prompt}"}}], - "usage": {"total_tokens": len(prompt)}, - }, - ) - - -def _run( - executor: APIExecutor, - task: WorkerTaskMessage, - transport: httpx.MockTransport, - out_dir: Path, -): - with patch.object( - APIExecutor, "_get_client", return_value=httpx.Client(transport=transport) - ): - return executor.run(task, out_dir) - - -def _executor() -> APIExecutor: - """Construct an APIExecutor without a WorkerConfig, mirroring __init__'s - cancellation state so run()/cancel() work under __new__.""" - executor = APIExecutor.__new__(APIExecutor) - executor._cancel_event = threading.Event() - executor._cancel_task_id = None - executor._cancel_lock = threading.Lock() - return executor - - -def _batch_task(items: list[Any], **api_updates: Any) -> WorkerTaskMessage: - payload = { - "task_id": "task-api-batch", - "workflow_id": "wf-1", - "owner_id": "owner", - "assigned_worker": "worker-1", - "dispatched_at": "2026-03-22T00:00:00Z", - "task": { - "apiVersion": "mloc/v1", - "kind": "Task", - "metadata": {"name": "wf:api"}, - "spec": { - "taskType": "api", - "api": { - "method": "POST", - "body": {"messages": [{"role": "user", "content": "{{prompt}}"}]}, - "response": { - "raise_for_status": False, - "include_headers": True, - }, - **api_updates, - }, - "data": {"type": "list", "items": items}, - }, - }, - } - return WorkerTaskMessage.model_validate(payload) - - -class TestBatch: - @pytest.fixture(autouse=True) - def _nebula_env(self, monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setenv("NEBULA_API_BASE_URL", "https://nebula.example.com") - monkeypatch.setenv("NEBULA_API_TOKEN", "nebula-token") - - def test_issues_one_request_per_row_in_order(self, tmp_path: Path) -> None: - task = _batch_task(["first", "second", "third"]) - transport = _RecordingTransport() - result = _run(_executor(), task, transport, tmp_path) - assert len(transport.requests) == 3 - issued = { - json.loads(req.read())["messages"][0]["content"] - for req in transport.requests - } - assert issued == {"first", "second", "third"} - for idx, prompt in enumerate(["first", "second", "third"]): - item = result.items[idx] - assert item.index == idx - assert item.prompt == prompt - assert item.text == f"echo:{prompt}" - assert item.response_json["choices"][0]["message"]["content"] == ( - f"echo:{prompt}" - ) - assert item.status_code == 200 + len(prompt) % 3 - assert item.usage == {"total_tokens": len(prompt)} - assert item.headers["x-row"] == prompt - - def test_rows_stay_aligned_when_requests_complete_out_of_order( - self, tmp_path: Path - ) -> None: - """Output row i corresponds to input row i even when requests finish - in reverse order.""" - - class _ReverseTransport(httpx.MockTransport): - def __init__(self) -> None: - self.requests: list[httpx.Request] = [] - super().__init__(self._handler) - - def _handler(self, request: httpx.Request) -> httpx.Response: - self.requests.append(request) - prompt = json.loads(request.read())["messages"][0]["content"] - delay = {"a": 0.3, "b": 0.2, "c": 0.1}[prompt] - time.sleep(delay) - return httpx.Response( - 200 + len(prompt) % 3, - headers={"X-Row": prompt}, - json={ - "choices": [{"message": {"content": f"echo:{prompt}"}}], - "usage": {"total_tokens": len(prompt)}, - }, - ) - - task = _batch_task(["a", "b", "c"]) - transport = _ReverseTransport() - result = _run(_executor(), task, transport, tmp_path) - for idx, prompt in enumerate(["a", "b", "c"]): - item = result.items[idx] - assert item.index == idx - assert item.prompt == prompt - assert item.text == f"echo:{prompt}" - assert item.response_json["choices"][0]["message"]["content"] == ( - f"echo:{prompt}" - ) - assert item.status_code == 200 + len(prompt) % 3 - assert item.usage == {"total_tokens": len(prompt)} - assert item.headers["x-row"] == prompt - - def test_single_row_batches_to_one_item(self, tmp_path: Path) -> None: - task = _batch_task(["only"]) - transport = _RecordingTransport() - result = _run(_executor(), task, transport, tmp_path) - assert len(transport.requests) == 1 - assert len(result.items) == 1 - assert result.items[0].index == 0 - assert result.items[0].prompt == "only" - - def test_row_failure_fails_whole_task_without_shifting( - self, tmp_path: Path - ) -> None: - class _FailRow(httpx.MockTransport): - def __init__(self, failing_prompt: str) -> None: - self.failing_prompt = failing_prompt - self.requests: list[httpx.Request] = [] - super().__init__(self._handler) - - def _handler(self, request: httpx.Request) -> httpx.Response: - self.requests.append(request) - prompt = json.loads(request.read())["messages"][0]["content"] - if prompt == self.failing_prompt: - return httpx.Response( - 500, - json={ - "choices": [{"message": {"content": "boom"}}], - "usage": {"total_tokens": 1}, - }, - ) - return httpx.Response( - 200, - json={ - "choices": [{"message": {"content": "ok"}}], - "usage": {"total_tokens": 1}, - }, - ) - - task = _batch_task(["a", "b", "c"], response={"raise_for_status": True}) - transport = _FailRow("b") - with pytest.raises(ExecutionError, match="row 1"): - _run(_executor(), task, transport, tmp_path) - assert len(transport.requests) == 3 - - def test_placeholder_not_required_for_scalar_body(self, tmp_path: Path) -> None: - """A batch task whose body has no placeholder still issues N requests.""" - payload = { - "task_id": "task-api-batch", - "workflow_id": "wf-1", - "owner_id": "owner", - "assigned_worker": "worker-1", - "dispatched_at": "2026-03-22T00:00:00Z", - "task": { - "apiVersion": "mloc/v1", - "kind": "Task", - "metadata": {"name": "wf:api"}, - "spec": { - "taskType": "api", - "api": { - "method": "POST", - "body": {"messages": [{"role": "user", "content": "static"}]}, - }, - "data": {"type": "list", "items": ["a", "b"]}, - }, - }, - } - task = WorkerTaskMessage.model_validate(payload) - transport = _RecordingTransport() - result = _run(_executor(), task, transport, tmp_path) - assert len(transport.requests) == 2 - assert len(result.items) == 2 - - def test_no_rows_raises(self, tmp_path: Path) -> None: - task = _batch_task([]) - with pytest.raises(ExecutionError, match="no rows"): - _run( - _executor(), - task, - _RecordingTransport(), - tmp_path, - ) - - def test_missing_data_raises(self, tmp_path: Path) -> None: - """spec.data is required; an api task without it fails closed.""" - payload = { - "task_id": "task-api-batch", - "workflow_id": "wf-1", - "owner_id": "owner", - "assigned_worker": "worker-1", - "dispatched_at": "2026-03-22T00:00:00Z", - "task": { - "apiVersion": "mloc/v1", - "kind": "Task", - "metadata": {"name": "wf:api"}, - "spec": { - "taskType": "api", - "api": { - "method": "POST", - "body": {"messages": [{"role": "user", "content": "hi"}]}, - }, - }, - }, - } - task = WorkerTaskMessage.model_validate(payload) - with pytest.raises(ExecutionError, match="spec.data is required"): - _run( - _executor(), - task, - _RecordingTransport(), - tmp_path, - ) - - def test_requests_issue_in_parallel(self, tmp_path: Path) -> None: - """N rows take ~one row's latency, not N x, on a network-bound path.""" - - class _SlowTransport(httpx.MockTransport): - def __init__(self) -> None: - self.requests: list[httpx.Request] = [] - super().__init__(self._handler) - - def _handler(self, request: httpx.Request) -> httpx.Response: - self.requests.append(request) - time.sleep(0.2) - return httpx.Response( - 200, - json={ - "choices": [{"message": {"content": "hello"}}], - "usage": {"total_tokens": 3}, - }, - ) - - n_rows = 4 - task = _batch_task([f"row-{i}" for i in range(n_rows)]) - transport = _SlowTransport() - start = time.monotonic() - result = _run(_executor(), task, transport, tmp_path) - elapsed = time.monotonic() - start - - assert len(transport.requests) == n_rows - assert elapsed < 0.2 * n_rows * 0.6 - assert [item.index for item in result.items] == list(range(n_rows)) - - def test_concurrency_one_serializes_requests(self, tmp_path: Path) -> None: - """concurrency: 1 limits the worker pool so requests never overlap.""" - - class _OverlapTransport(httpx.MockTransport): - def __init__(self) -> None: - self.max_in_flight = 0 - self._in_flight = 0 - self._lock = threading.Lock() - super().__init__(self._handler) - - def _handler(self, request: httpx.Request) -> httpx.Response: - with self._lock: - self._in_flight += 1 - self.max_in_flight = max(self.max_in_flight, self._in_flight) - time.sleep(0.05) - with self._lock: - self._in_flight -= 1 - return httpx.Response( - 200, - json={ - "choices": [{"message": {"content": "hello"}}], - "usage": {"total_tokens": 3}, - }, - ) - - task = _batch_task(["a", "b", "c", "d"], concurrency=1) - transport = _OverlapTransport() - _run(_executor(), task, transport, tmp_path) - assert transport.max_in_flight == 1 - - def test_request_skeleton_constructed_once(self, tmp_path: Path) -> None: - """The request template is built once, not once per row.""" - task = _batch_task(["a", "b", "c"]) - transport = _RecordingTransport() - real_build = APIExecutor._build_request_kwargs - - with ( - patch.object( - APIExecutor, - "_get_client", - return_value=httpx.Client(transport=transport), - ), - patch.object( - APIExecutor, "_build_request_kwargs", autospec=True - ) as mock_build, - ): - mock_build.side_effect = lambda *a, **k: real_build(*a, **k) - _executor().run(task, tmp_path) - - assert mock_build.call_count == 1 - assert mock_build.call_args.args[2] is None - - @pytest.mark.parametrize("concurrency", [1, 4, 8]) - def test_client_pool_sized_to_concurrency(self, concurrency: int) -> None: - """The connection pool matches the effective concurrency.""" - APIExecutor.close_all_clients() - try: - client = APIExecutor._get_client( - "https://example.com", - httpx.Timeout(60), - True, - True, - concurrency, - ) - pool = cast(Any, client._transport)._pool - assert pool._max_connections == concurrency - assert pool._max_keepalive_connections == concurrency - finally: - APIExecutor.close_all_clients() - - def test_concurrency_capped_at_max(self, tmp_path: Path) -> None: - """A configured concurrency above the cap is clamped to the cap.""" - task = _batch_task(["a", "b", "c"], concurrency=100) - captured: dict[str, Any] = {} - - def _fake_get_client(*args: Any, **kwargs: Any) -> httpx.Client: - captured["concurrency"] = kwargs.get("concurrency", args[4]) - return httpx.Client(transport=_RecordingTransport()) - - with patch.object(APIExecutor, "_get_client", side_effect=_fake_get_client): - _executor().run(task, tmp_path) - - assert captured["concurrency"] == 8 - - @pytest.mark.parametrize("concurrency", [1, 4]) - def test_run_passes_effective_concurrency_to_client( - self, tmp_path: Path, concurrency: int - ) -> None: - """run() forwards the uncapped configured concurrency to the client.""" - task = _batch_task(["a", "b", "c"], concurrency=concurrency) - captured: dict[str, Any] = {} - - def _fake_get_client(*args: Any, **kwargs: Any) -> httpx.Client: - captured["concurrency"] = kwargs.get("concurrency", args[4]) - return httpx.Client(transport=_RecordingTransport()) - - with patch.object(APIExecutor, "_get_client", side_effect=_fake_get_client): - _executor().run(task, tmp_path) - - assert captured["concurrency"] == concurrency - - @pytest.mark.parametrize("concurrency", [0, -1]) - def test_concurrency_below_one_rejected( - self, tmp_path: Path, concurrency: int - ) -> None: - """A configured concurrency below 1 is rejected.""" - task = _batch_task(["a", "b", "c"], concurrency=concurrency) - with pytest.raises(ExecutionError, match="spec.api.concurrency must be >= 1"): - _run(_executor(), task, _RecordingTransport(), tmp_path) - - def test_client_cache_key_includes_concurrency(self) -> None: - """Pools built for different concurrency values are not shared.""" - APIExecutor.close_all_clients() - try: - c1 = APIExecutor._get_client( - "https://example.com", httpx.Timeout(60), True, True, 1 - ) - c4 = APIExecutor._get_client( - "https://example.com", httpx.Timeout(60), True, True, 4 - ) - assert c1 is not c4 - assert len(APIExecutor._clients) == 2 - finally: - APIExecutor.close_all_clients() - - def test_no_usage_2xx_produces_aligned_item(self, tmp_path: Path) -> None: - """A 2xx response without usage still yields a row-aligned item.""" - - class _NoUsage(httpx.MockTransport): - def __init__(self) -> None: - super().__init__(self._handler) - - def _handler(self, request: httpx.Request) -> httpx.Response: - return httpx.Response( - 200, - json={"choices": [{"message": {"content": "ok"}}]}, - ) - - task = _batch_task(["a", "b"]) - result = _run(_executor(), task, _NoUsage(), tmp_path) - assert len(result.items) == 2 - assert result.items[0].text == "ok" - assert result.items[0].usage is None - - def test_5xx_with_raise_for_status_false_produces_aligned_item( - self, tmp_path: Path - ) -> None: - """A 5xx with raise_for_status false still yields a row-aligned item.""" - - class _ErrorBody(httpx.MockTransport): - def __init__(self) -> None: - super().__init__(self._handler) - - def _handler(self, request: httpx.Request) -> httpx.Response: - return httpx.Response( - 503, - json={"error": {"message": "overloaded"}}, - ) - - task = _batch_task(["a", "b"], response={"raise_for_status": False}) - result = _run(_executor(), task, _ErrorBody(), tmp_path) - assert len(result.items) == 2 - assert result.items[0].status_code == 503 - assert result.items[0].response_json == {"error": {"message": "overloaded"}} - - def test_retryable_503_raises_retryable(self, tmp_path: Path) -> None: - """A 503 is classified retryable even without a success payload.""" - - class _ErrorBody(httpx.MockTransport): - def __init__(self) -> None: - super().__init__(self._handler) - - def _handler(self, request: httpx.Request) -> httpx.Response: - return httpx.Response( - 503, - json={"error": {"message": "overloaded"}}, - ) - - task = _batch_task(["a", "b"], response={"raise_for_status": True}) - with pytest.raises(ExecutionError, match="status 503") as excinfo: - _run(_executor(), task, _ErrorBody(), tmp_path) - assert excinfo.value.retryable is True - - def test_cancel_prevents_queued_rows_from_issuing(self, tmp_path: Path) -> None: - """After cancel, a row that has not started never issues its request.""" - - class _BlockingTransport(httpx.MockTransport): - def __init__(self) -> None: - self.requests: list[httpx.Request] = [] - self.started = threading.Event() - self.release = threading.Event() - super().__init__(self._handler) - - def _handler(self, request: httpx.Request) -> httpx.Response: - self.requests.append(request) - self.started.set() - self.release.wait(timeout=5) - return httpx.Response( - 200, - json={ - "choices": [{"message": {"content": "ok"}}], - "usage": {"total_tokens": 3}, - }, - ) - - executor = _executor() - task = _batch_task(["a", "b", "c", "d"], concurrency=1) - transport = _BlockingTransport() - errors: list[BaseException] = [] - submitted: list[Any] = [] - all_submitted = threading.Event() - real_submit = concurrent.futures.ThreadPoolExecutor.submit - - def _recording_submit(self: Any, fn: Any, *args: Any, **kwargs: Any) -> Any: - submitted.append(fn) - if len(submitted) == 4: - all_submitted.set() - return real_submit(self, fn, *args, **kwargs) - - def _run_in_thread() -> None: - try: - with patch.object( - concurrent.futures.ThreadPoolExecutor, - "submit", - _recording_submit, - ): - _run(executor, task, transport, tmp_path) - except BaseException as exc: # noqa: BLE001 - captured for assertion - errors.append(exc) - - thread = threading.Thread(target=_run_in_thread) - thread.start() - assert transport.started.wait(timeout=5) - assert all_submitted.wait(timeout=5) - executor.cancel("task-api-batch") - transport.release.set() - thread.join(timeout=10) - - assert len(transport.requests) == 1 - assert len(errors) == 1 - assert isinstance(errors[0], TaskCancelledError) - - def test_cancel_before_submission_prevents_any_future(self, tmp_path: Path) -> None: - """Cancelling before futures are submitted surfaces TaskCancelledError - without submitting any future.""" - executor = _executor() - task = _batch_task(["a", "b", "c", "d"]) - transport = _RecordingTransport() - errors: list[BaseException] = [] - submitted: list[Any] = [] - real_submit = concurrent.futures.ThreadPoolExecutor.submit - real_base_url = APIExecutor._base_url - - def _recording_submit(self: Any, fn: Any, *args: Any, **kwargs: Any) -> Any: - submitted.append(fn) - return real_submit(self, fn, *args, **kwargs) - - def _run_in_thread() -> None: - def _cancel_then_base_url(url: str) -> str: - executor.cancel("task-api-batch") - return real_base_url(url) - - try: - with ( - patch.object( - concurrent.futures.ThreadPoolExecutor, - "submit", - _recording_submit, - ), - patch.object( - APIExecutor, - "_base_url", - side_effect=_cancel_then_base_url, - ), - ): - _run(executor, task, transport, tmp_path) - except BaseException as exc: # noqa: BLE001 - captured for assertion - errors.append(exc) - - thread = threading.Thread(target=_run_in_thread) - thread.start() - thread.join(timeout=10) - - assert len(errors) == 1 - assert isinstance(errors[0], TaskCancelledError) - assert submitted == [] - - def test_cancel_after_requests_complete_before_collection_not_done( - self, tmp_path: Path - ) -> None: - """A cancel arriving after every request has completed but before results - are collected fails the task rather than returning DONE.""" - - class _RecordingTransport(httpx.MockTransport): - def __init__(self) -> None: - self.requests: list[httpx.Request] = [] - super().__init__(self._handler) - - def _handler(self, request: httpx.Request) -> httpx.Response: - self.requests.append(request) - return httpx.Response( - 200, - json={ - "choices": [{"message": {"content": "ok"}}], - "usage": {"total_tokens": 3}, - }, - ) - - executor = _executor() - task = _batch_task(["a", "b"], concurrency=2) - transport = _RecordingTransport() - errors: list[BaseException] = [] - futures: list[Any] = [] - collect_release = threading.Event() - real_submit = concurrent.futures.ThreadPoolExecutor.submit - real_as_completed = concurrent.futures.as_completed - - def _recording_submit(self: Any, fn: Any, *args: Any, **kwargs: Any) -> Any: - future = real_submit(self, fn, *args, **kwargs) - futures.append(future) - return future - - def _blocking_as_completed(fs: Any, timeout: float | None = None) -> Any: - collect_release.wait(timeout=5) - return real_as_completed(fs, timeout=timeout) - - def _run_in_thread() -> None: - try: - with ( - patch.object( - concurrent.futures.ThreadPoolExecutor, - "submit", - _recording_submit, - ), - patch.object( - api_executor_module, "as_completed", _blocking_as_completed - ), - ): - _run(executor, task, transport, tmp_path) - except BaseException as exc: # noqa: BLE001 - captured for assertion - errors.append(exc) - - thread = threading.Thread(target=_run_in_thread) - thread.start() - assert transport.requests or True - while len(futures) < 2: - time.sleep(0.01) - for future in futures: - assert future.done() - executor.cancel("task-api-batch") - collect_release.set() - thread.join(timeout=10) - - assert len(errors) == 1 - assert isinstance(errors[0], TaskCancelledError) - - def test_cancel_before_run_cancels(self, tmp_path: Path) -> None: - """A cancel that lands before run() starts still cancels the run.""" - executor = _executor() - task = _batch_task(["a", "b"]) - transport = _RecordingTransport() - - executor.cancel("task-api-batch") - - with pytest.raises(TaskCancelledError): - _run(executor, task, transport, tmp_path) - assert transport.requests == [] - - def test_cancel_during_run_setup_not_lost(self, tmp_path: Path) -> None: - """A cancel arriving while run() is mid check-and-clear is not dropped. - - run() checks the event, then clears it under the lock; cancel() sets it - under the same lock. Pausing run() inside clear() and firing cancel() - while it is paused must still cancel the run (via the in-flight guards), - not let it complete as if never cancelled.""" - executor = _executor() - task = _batch_task(["a", "b"]) - transport = _RecordingTransport() - errors: list[BaseException] = [] - in_clear = threading.Event() - release_clear = threading.Event() - real_clear = executor._cancel_event.clear - - def _blocking_clear() -> None: - in_clear.set() - release_clear.wait(timeout=5) - real_clear() - - executor._cancel_event.clear = _blocking_clear # type: ignore[method-assign] - - def _run_in_thread() -> None: - try: - _run(executor, task, transport, tmp_path) - except BaseException as exc: # noqa: BLE001 - captured for assertion - errors.append(exc) - - thread = threading.Thread(target=_run_in_thread) - thread.start() - assert in_clear.wait(timeout=5) - executor.cancel(task.task_id) - release_clear.set() - thread.join(timeout=10) - - assert len(errors) == 1 - assert isinstance(errors[0], TaskCancelledError) From de41f505cece3d23392cad7d0ae942d157af67bc Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Sat, 19 Sep 2026 13:52:13 +0700 Subject: [PATCH 27/71] refactor: revert gitignore and GPU test, condense added docstrings Signed-off-by: Zhengyuan Su --- .gitignore | 4 +--- src/server/task/n8n_parser.py | 5 +---- src/shared/schemas/result/catalog.py | 6 ++---- src/shared/schemas/result/payloads.py | 6 +++--- src/worker/executors/api_executor.py | 10 ++++------ tests/server/task/test_ssh_result_mounting.py | 7 +------ tests/shared/test_executor_result.py | 6 ++---- tests/worker/test_api_executor.py | 10 ++-------- tests/worker/test_mp_executor_cleanup_gpu.py | 17 ++++------------- 9 files changed, 20 insertions(+), 51 deletions(-) diff --git a/.gitignore b/.gitignore index 5d7ce3094..8a3f0c34b 100644 --- a/.gitignore +++ b/.gitignore @@ -180,6 +180,4 @@ secrets/ dump.rdb plugins/ -plugin-data/ -# Isolated e2e stack env (contains a live credential) -.env.apibatch +plugin-data/ \ No newline at end of file diff --git a/src/server/task/n8n_parser.py b/src/server/task/n8n_parser.py index aa44750fd..a052c0573 100644 --- a/src/server/task/n8n_parser.py +++ b/src/server/task/n8n_parser.py @@ -413,10 +413,7 @@ def _inject_dependency_prompt(prompt_text: str, placeholder: str) -> str: def _dependency_placeholder(dep_name: str, dep_task_type: str) -> str: - """Return the stage reference a dependent node injects for an upstream task. - - APIResult is batch-only, so an api dependency reads the first row's text - (``items.0.text``) rather than a scalar ``text`` field.""" + """Return the stage reference a dependent node injects for an upstream task.""" if dep_task_type == "api": return f"${{{dep_name}.items.0.text}}" if dep_task_type == "inference": diff --git a/src/shared/schemas/result/catalog.py b/src/shared/schemas/result/catalog.py index d7c0aeb19..60c00858b 100644 --- a/src/shared/schemas/result/catalog.py +++ b/src/shared/schemas/result/catalog.py @@ -246,10 +246,8 @@ class EchoResult(StrictExecutorResult): class APIResult(StrictExecutorResult): """HTTP request output. ``response_json``/``usage``/``headers`` are the - upstream API's own payloads and stay open mappings. - - ``items`` carries one entry per row: every API task batches over - ``spec.data``, so even a single request yields one item.""" + upstream API's own payloads and stay open mappings. ``items`` carries one + entry per row.""" task_type: Literal[TaskType.API] = TaskType.API executor: str diff --git a/src/shared/schemas/result/payloads.py b/src/shared/schemas/result/payloads.py index 7769f8926..c53296ad8 100644 --- a/src/shared/schemas/result/payloads.py +++ b/src/shared/schemas/result/payloads.py @@ -212,9 +212,9 @@ class EchoItem(StrictModel): class APIItem(StrictModel): """One row's HTTP response in a batched API task. - Mirrors the per-response fields of :class:`APIResult`; ``response_json`` is - the upstream API's own payload and stays an open mapping. ``populate_by_name`` - lets code construct by field name while the wire key stays ``json``. + ``response_json`` is the upstream API's own payload and stays an open + mapping; ``populate_by_name`` lets code construct by field name while the + wire key stays ``json``. """ model_config = ConfigDict(extra="forbid", populate_by_name=True) diff --git a/src/worker/executors/api_executor.py b/src/worker/executors/api_executor.py index c1fae82e7..69348321c 100644 --- a/src/worker/executors/api_executor.py +++ b/src/worker/executors/api_executor.py @@ -25,8 +25,7 @@ _ClientKey = tuple[str, float, bool, bool, int] -# Worker-side per-row slot in a batched request body. Server-side stage -# references are ${...}; this is a worker-side token, hence {{...}}. +# Worker-side per-row slot; server-side stage references are ${...}. _PROMPT_PLACEHOLDER = "{{prompt}}" # Fixed delay between retry attempts. @@ -104,11 +103,10 @@ def _get_client( follow_redirects: bool, concurrency: int, ) -> httpx.Client: - """Return a cached client or create a new one for the given parameters. + """Return a cached client or create one for the given parameters. - The cache key is ``(base_url, timeout_sec, verify_tls, follow_redirects, - concurrency)``; the pool is sized to ``concurrency`` so parallel row - requests never queue on connections.""" + The pool is sized to ``concurrency`` so parallel row requests never + queue on connections.""" timeout_sec = timeout.connect # all four fields are set to same value if timeout_sec is None: timeout_sec = 0.0 diff --git a/tests/server/task/test_ssh_result_mounting.py b/tests/server/task/test_ssh_result_mounting.py index 1bfe80bb2..c5be4ea01 100644 --- a/tests/server/task/test_ssh_result_mounting.py +++ b/tests/server/task/test_ssh_result_mounting.py @@ -451,12 +451,7 @@ def test_api_dependent_stage_resolves_first_row_text(tmp_path: Path) -> None: def test_translated_n8n_dependent_api_stage_resolves(tmp_path: Path) -> None: - """A translated n8n workflow's dependent API stage resolves end to end. - - The injected ${Upstream.items.0.text} placeholder reads the upstream - stage's first-row text through the dispatcher, covering the n8n - production change through to stage resolution. - """ + """A translated n8n workflow's dependent API stage resolves end to end.""" payload = { "nodes": [ { diff --git a/tests/shared/test_executor_result.py b/tests/shared/test_executor_result.py index 494b78435..d5c0c06d4 100644 --- a/tests/shared/test_executor_result.py +++ b/tests/shared/test_executor_result.py @@ -196,7 +196,7 @@ def test_upstream_results_preserve_subclass_payload_over_the_wire() -> None: def test_api_item_round_trip_construct_serialize_validate() -> None: """An APIItem built by field name round-trips through the worker's serialization and the server's ingest validation.""" - # mypy cannot see populate_by_name; the field's declared name is the json alias. + # mypy cannot see populate_by_name; the declared name is the json alias. item = APIItem( # type: ignore[call-arg] index=0, url="http://example.com/v1/chat/completions", @@ -204,14 +204,12 @@ def test_api_item_round_trip_construct_serialize_validate() -> None: response_json={"choices": [{"message": {"content": "hello"}}]}, text="hello", ) - # Serialize the way the worker does (envelope.model_dump_json, no by_alias). wire = item.model_dump_json() - # Re-validate the way the server does on ingest. reloaded = APIItem.model_validate_json(wire) assert reloaded.index == 0 assert reloaded.response_json["choices"][0]["message"]["content"] == "hello" assert reloaded.text == "hello" - # The wire alias is still accepted on input (backward compatible). + # The wire alias is still accepted on input. by_alias = APIItem.model_validate( {"index": 1, "url": "u", "status_code": 200, "json": {"a": 1}} ) diff --git a/tests/worker/test_api_executor.py b/tests/worker/test_api_executor.py index 274904afb..5fe54d7cc 100644 --- a/tests/worker/test_api_executor.py +++ b/tests/worker/test_api_executor.py @@ -119,8 +119,7 @@ def _run( def _executor() -> APIExecutor: - """Construct an APIExecutor without a WorkerConfig, mirroring __init__'s - cancellation state so run()/cancel() work under __new__.""" + """Build an APIExecutor with cancellation state, without a WorkerConfig.""" executor = APIExecutor.__new__(APIExecutor) executor._cancel_event = threading.Event() executor._cancel_task_id = None @@ -1078,12 +1077,7 @@ def test_cancel_before_run_cancels(self, tmp_path: Path) -> None: assert transport.requests == [] def test_cancel_during_run_setup_not_lost(self, tmp_path: Path) -> None: - """A cancel arriving while run() is mid check-and-clear is not dropped. - - run() checks the event, then clears it under the lock; cancel() sets it - under the same lock. Pausing run() inside clear() and firing cancel() - while it is paused must still cancel the run (via the in-flight guards), - not let it complete as if never cancelled.""" + """A cancel landing mid check-and-clear is not dropped.""" executor = _executor() task = _batch_task(["a", "b"]) transport = _EchoTransport() diff --git a/tests/worker/test_mp_executor_cleanup_gpu.py b/tests/worker/test_mp_executor_cleanup_gpu.py index 64c85c7e6..0c0f409ae 100644 --- a/tests/worker/test_mp_executor_cleanup_gpu.py +++ b/tests/worker/test_mp_executor_cleanup_gpu.py @@ -2,11 +2,10 @@ import tempfile import time import uuid -from collections.abc import Iterator from pathlib import Path import psutil -import pynvml # type: ignore[import-not-found] +import pynvml # type: ignore import pytest from shared.tasks.worker_message import WorkerTaskMessage @@ -14,15 +13,7 @@ from worker.executors.mp_executor import MPExecutor from worker.executors.vllm_executor import VLLMExecutor - -@pytest.fixture(scope="module") -def _nvml() -> Iterator[None]: - try: - pynvml.nvmlInit() - except pynvml.NVMLError: - pytest.skip("NVML unavailable (no GPU)") - yield - pynvml.nvmlShutdown() +pynvml.nvmlInit() def _descendants_of(pid: int) -> set[int]: @@ -34,7 +25,7 @@ def _descendants_of(pid: int) -> set[int]: @pytest.mark.gpu -def test_mp_executor_cleans_up_vllm(caplog, tmp_path: Path, _nvml: None) -> None: +def test_mp_executor_cleans_up_vllm(caplog, tmp_path: Path) -> None: """Start MPExecutor with the real executors, run a minimal task to trigger engine startup, and ensure cleanup removes the worker process and any descendants it spawned. @@ -71,7 +62,7 @@ def total_gpu_used() -> int: "assigned_worker": "test-worker", "dispatched_at": "2026-03-01T00:00:00Z", "task": { - "apiVersion": "flowmesh/v1", + "apiVersion": "mloc/v1", "kind": "InferenceTask", "spec": { "taskType": "inference", From e4991dd4d3f8215aa714cb7e7e5fffcbd2605b39 Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Sat, 19 Sep 2026 14:02:22 +0700 Subject: [PATCH 28/71] fix: use flowmesh/v1 apiVersion in API executor tests Signed-off-by: Zhengyuan Su --- tests/worker/test_api_executor.py | 6 +++--- tests/worker/test_mp_executor_cleanup_gpu.py | 2 +- 2 files changed, 4 insertions(+), 4 deletions(-) diff --git a/tests/worker/test_api_executor.py b/tests/worker/test_api_executor.py index 5fe54d7cc..da2d07417 100644 --- a/tests/worker/test_api_executor.py +++ b/tests/worker/test_api_executor.py @@ -135,7 +135,7 @@ def _batch_task(items: list[Any], **api_updates: Any) -> WorkerTaskMessage: "assigned_worker": "worker-1", "dispatched_at": "2026-03-22T00:00:00Z", "task": { - "apiVersion": "mloc/v1", + "apiVersion": "flowmesh/v1", "kind": "Task", "metadata": {"name": "wf:api"}, "spec": { @@ -619,7 +619,7 @@ def test_placeholder_not_required_for_scalar_body(self, tmp_path: Path) -> None: "assigned_worker": "worker-1", "dispatched_at": "2026-03-22T00:00:00Z", "task": { - "apiVersion": "mloc/v1", + "apiVersion": "flowmesh/v1", "kind": "Task", "metadata": {"name": "wf:api"}, "spec": { @@ -657,7 +657,7 @@ def test_missing_data_raises(self, tmp_path: Path) -> None: "assigned_worker": "worker-1", "dispatched_at": "2026-03-22T00:00:00Z", "task": { - "apiVersion": "mloc/v1", + "apiVersion": "flowmesh/v1", "kind": "Task", "metadata": {"name": "wf:api"}, "spec": { diff --git a/tests/worker/test_mp_executor_cleanup_gpu.py b/tests/worker/test_mp_executor_cleanup_gpu.py index 0c0f409ae..338e5ac56 100644 --- a/tests/worker/test_mp_executor_cleanup_gpu.py +++ b/tests/worker/test_mp_executor_cleanup_gpu.py @@ -62,7 +62,7 @@ def total_gpu_used() -> int: "assigned_worker": "test-worker", "dispatched_at": "2026-03-01T00:00:00Z", "task": { - "apiVersion": "mloc/v1", + "apiVersion": "flowmesh/v1", "kind": "InferenceTask", "spec": { "taskType": "inference", From 2109405461532be385e8665555be724f1b39c4d2 Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Thu, 24 Sep 2026 12:06:02 +0700 Subject: [PATCH 29/71] test: adapt API executor retry tests to the merged batch executor The merged executor runs the atomic cancel check-and-clear under a lock, so the retry tests must build a fully-initialized executor (via _executor()) rather than APIExecutor.__new__, and _run must not reset the cancel event (which would wipe a cancel set before run). Use executor.cancel(task_id) to set the task-scoped cancel in the stop-retrying test. Co-Authored-By: Claude Code Signed-off-by: Zhengyuan Su --- tests/worker/test_api_executor.py | 24 +++++++++--------------- 1 file changed, 9 insertions(+), 15 deletions(-) diff --git a/tests/worker/test_api_executor.py b/tests/worker/test_api_executor.py index da2d07417..b8e1b3432 100644 --- a/tests/worker/test_api_executor.py +++ b/tests/worker/test_api_executor.py @@ -108,10 +108,6 @@ def _run( transport: httpx.MockTransport, out_dir: Path = Path("/tmp/out"), ): - executor._cancel_event = threading.Event() - executor._cancel_lock = threading.Lock() - executor._active_task_id = None - executor._pending_cancelled_ids = set() with patch.object( APIExecutor, "_get_client", return_value=httpx.Client(transport=transport) ): @@ -122,7 +118,8 @@ def _executor() -> APIExecutor: """Build an APIExecutor with cancellation state, without a WorkerConfig.""" executor = APIExecutor.__new__(APIExecutor) executor._cancel_event = threading.Event() - executor._cancel_task_id = None + executor._active_task_id = None + executor._pending_cancelled_ids = set() executor._cancel_lock = threading.Lock() return executor @@ -299,7 +296,7 @@ def test_retry_succeeds_after_transient_failures(self) -> None: transport = _SequenceTransport( [_error_response(504), _error_response(504), _ok_response()] ) - _run(APIExecutor.__new__(APIExecutor), task, transport) + _run(_executor(), task, transport) assert transport.calls == 3 def test_retries_exhausted_still_fails(self) -> None: @@ -309,7 +306,7 @@ def test_retries_exhausted_still_fails(self) -> None: [_error_response(504), _error_response(504), _error_response(504)] ) with pytest.raises(ExecutionError, match="status 504"): - _run(APIExecutor.__new__(APIExecutor), task, transport) + _run(_executor(), task, transport) assert transport.calls == 3 def test_no_retry_by_default(self) -> None: @@ -317,7 +314,7 @@ def test_no_retry_by_default(self) -> None: task = self._task() transport = _SequenceTransport([_error_response(504)]) with pytest.raises(ExecutionError, match="status 504"): - _run(APIExecutor.__new__(APIExecutor), task, transport) + _run(_executor(), task, transport) assert transport.calls == 1 def test_non_retryable_status_not_retried(self) -> None: @@ -325,16 +322,12 @@ def test_non_retryable_status_not_retried(self) -> None: task = self._task(retries=3) transport = _SequenceTransport([_error_response(400)]) with pytest.raises(ExecutionError, match="status 400"): - _run(APIExecutor.__new__(APIExecutor), task, transport) + _run(_executor(), task, transport) assert transport.calls == 1 def test_cancelled_task_stops_retrying(self) -> None: """A cancelled task does not keep retrying.""" - executor = APIExecutor.__new__(APIExecutor) - executor._cancel_event = threading.Event() - executor._cancel_lock = threading.Lock() - executor._active_task_id = None - executor._pending_cancelled_ids = set() + executor = _executor() task = self._task(retries=3) executor.cancel(task.task_id) transport = _SequenceTransport([_error_response(504)]) @@ -350,7 +343,7 @@ def test_invalid_retries_rejected(self) -> None: for bad in (-1, "2", 1.5, True): task = self._task(retries=bad) with pytest.raises(ExecutionError, match="spec.api.retries"): - _run(APIExecutor.__new__(APIExecutor), task, _RecordingTransport()) + _run(_executor(), task, _RecordingTransport()) def test_cancel_previous_task_does_not_cancel_next(self) -> None: """A cancellation left over from a prior task does not cancel the next.""" @@ -499,6 +492,7 @@ def _cancel_after_delay() -> None: assert elapsed < 0.5 assert transport.calls == 1 + class TestBatch: @pytest.fixture(autouse=True) def _nebula_env(self, monkeypatch: pytest.MonkeyPatch) -> None: From 7d2ddb9ae1f34148baec5c06e2bed83c9f81fa88 Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Sun, 27 Sep 2026 11:17:53 +0700 Subject: [PATCH 30/71] fix(worker): substitute a structured prompt object into the API body A data.type: list row that is a chat-message list, used with messages: "{{prompt}}", is now sent as a JSON array of {role, content} dicts instead of a JSON string. A value that is exactly {{prompt}} becomes the prompt object unchanged; a placeholder embedded in a longer string keeps the string for a string prompt and json.dumps for anything else. Signed-off-by: Zhengyuan Su Co-Authored-By: Claude Code --- src/worker/executors/api_executor.py | 13 +++++++--- tests/worker/test_api_executor.py | 38 ++++++++++++++++++++++++++++ 2 files changed, 47 insertions(+), 4 deletions(-) diff --git a/src/worker/executors/api_executor.py b/src/worker/executors/api_executor.py index 69348321c..68d487dfc 100644 --- a/src/worker/executors/api_executor.py +++ b/src/worker/executors/api_executor.py @@ -211,12 +211,17 @@ def _prompt_to_str(prompt: Any) -> str: return json.dumps(prompt) @classmethod - def _substitute_prompt(cls, value: Any, prompt: str) -> Any: - """Replace ``{{prompt}}`` in the request body with a row's prompt.""" + def _substitute_prompt(cls, value: Any, prompt: Any) -> Any: + """Replace ``{{prompt}}`` in the request body with a row's prompt. + + A value that is exactly ``{{prompt}}`` becomes the prompt object + itself, so a chat-message list stays a list of ``{role, content}`` + dicts; a placeholder embedded in a longer string is rendered as text. + """ if isinstance(value, str): if value == _PROMPT_PLACEHOLDER: return prompt - return value.replace(_PROMPT_PLACEHOLDER, prompt) + return value.replace(_PROMPT_PLACEHOLDER, cls._prompt_to_str(prompt)) if isinstance(value, dict): return {k: cls._substitute_prompt(v, prompt) for k, v in value.items()} if isinstance(value, list): @@ -397,7 +402,7 @@ def _issue(idx: int, prompt: Any) -> APIItem: if self._cancel_event.is_set(): raise TaskCancelledError("API task cancelled") prompt_str = self._prompt_to_str(prompt) - kwargs = self._substitute_prompt(request_kwargs, prompt_str) + kwargs = self._substitute_prompt(request_kwargs, prompt) try: resp = self._request_with_retries( client, diff --git a/tests/worker/test_api_executor.py b/tests/worker/test_api_executor.py index b8e1b3432..58badabd1 100644 --- a/tests/worker/test_api_executor.py +++ b/tests/worker/test_api_executor.py @@ -632,6 +632,44 @@ def test_placeholder_not_required_for_scalar_body(self, tmp_path: Path) -> None: assert len(transport.requests) == 2 assert len(result.items) == 2 + def test_message_list_row_sends_json_array(self, tmp_path: Path) -> None: + """A chat-message list row with ``messages: "{{prompt}}"`` is sent as a + JSON array of ``{role, content}`` dicts, not a JSON string.""" + messages = [ + {"role": "system", "content": "rules"}, + {"role": "user", "content": "hi"}, + ] + task = _batch_task( + [messages], + body={"messages": "{{prompt}}"}, + ) + transport = _RecordingTransport() + _run(_executor(), task, transport, tmp_path) + assert transport.request is not None + body = json.loads(transport.request.read()) + assert body["messages"] == messages + + def test_message_list_row_embedded_in_string_sends_json_text( + self, tmp_path: Path + ) -> None: + """A chat-message list row embedded in a longer string is rendered as + valid JSON text inside the surrounding body.""" + messages = [ + {"role": "system", "content": "rules"}, + {"role": "user", "content": "hi"}, + ] + task = _batch_task( + [messages], + body={"messages": [{"role": "user", "content": "context: {{prompt}}"}]}, + ) + transport = _RecordingTransport() + _run(_executor(), task, transport, tmp_path) + assert transport.request is not None + body = json.loads(transport.request.read()) + embedded = body["messages"][0]["content"] + assert embedded == f"context: {json.dumps(messages)}" + assert json.loads(embedded.removeprefix("context: ")) == messages + def test_no_rows_raises(self, tmp_path: Path) -> None: task = _batch_task([]) with pytest.raises(ExecutionError, match="no rows"): From 4a050012334242c8c9fcfc25436d24e0f3a4792e Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Sun, 27 Sep 2026 11:19:54 +0700 Subject: [PATCH 31/71] fix(schema): serialize APIItem's json field under its wire alias APIItem declares the wire alias json but serialized as response_json in both the worker's shared schema and the SDK. Enabling serialize_by_alias emits the json key, and the round-trip tests now assert the serialized key is json and that response_json is absent. Signed-off-by: Zhengyuan Su Co-Authored-By: Claude Code --- sdk/src/flowmesh/models/result/payloads.py | 4 +++- src/shared/schemas/result/payloads.py | 4 +++- tests/sdk/test_models.py | 2 ++ tests/shared/test_executor_result.py | 2 ++ 4 files changed, 10 insertions(+), 2 deletions(-) diff --git a/sdk/src/flowmesh/models/result/payloads.py b/sdk/src/flowmesh/models/result/payloads.py index 9f3eee8b7..cdab96fdc 100644 --- a/sdk/src/flowmesh/models/result/payloads.py +++ b/sdk/src/flowmesh/models/result/payloads.py @@ -161,7 +161,9 @@ class EchoItem(StrictModel): class APIItem(StrictModel): - model_config = ConfigDict(extra="forbid", populate_by_name=True) + model_config = ConfigDict( + extra="forbid", populate_by_name=True, serialize_by_alias=True + ) index: int url: str diff --git a/src/shared/schemas/result/payloads.py b/src/shared/schemas/result/payloads.py index c53296ad8..97a3c3eca 100644 --- a/src/shared/schemas/result/payloads.py +++ b/src/shared/schemas/result/payloads.py @@ -217,7 +217,9 @@ class APIItem(StrictModel): wire key stays ``json``. """ - model_config = ConfigDict(extra="forbid", populate_by_name=True) + model_config = ConfigDict( + extra="forbid", populate_by_name=True, serialize_by_alias=True + ) index: int url: str diff --git a/tests/sdk/test_models.py b/tests/sdk/test_models.py index 505a5a33d..025a26d4f 100644 --- a/tests/sdk/test_models.py +++ b/tests/sdk/test_models.py @@ -507,6 +507,8 @@ def test_construct_by_name_serialize_revalidate(self) -> None: } ) wire = item.model_dump_json() + assert '"json"' in wire + assert '"response_json"' not in wire reloaded = APIItem.model_validate_json(wire) assert reloaded.index == 0 assert reloaded.response_json["choices"][0]["message"]["content"] == "hello" diff --git a/tests/shared/test_executor_result.py b/tests/shared/test_executor_result.py index d5c0c06d4..2df793b09 100644 --- a/tests/shared/test_executor_result.py +++ b/tests/shared/test_executor_result.py @@ -205,6 +205,8 @@ def test_api_item_round_trip_construct_serialize_validate() -> None: text="hello", ) wire = item.model_dump_json() + assert '"json"' in wire + assert '"response_json"' not in wire reloaded = APIItem.model_validate_json(wire) assert reloaded.index == 0 assert reloaded.response_json["choices"][0]["message"]["content"] == "hello" From d11996d44df62c1ebb9d585febc4fe93c1c4c7ed Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Sun, 27 Sep 2026 11:20:56 +0700 Subject: [PATCH 32/71] test: assert completed requests before collection in cancellation test The precondition check asserted transport.requests or True, which is always true. It now asserts the expected completed-request count once both futures are done, so the timing boundary the test targets is actually exercised. Signed-off-by: Zhengyuan Su Co-Authored-By: Claude Code --- tests/worker/test_api_executor.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/worker/test_api_executor.py b/tests/worker/test_api_executor.py index 58badabd1..61de80870 100644 --- a/tests/worker/test_api_executor.py +++ b/tests/worker/test_api_executor.py @@ -1084,11 +1084,11 @@ def _run_in_thread() -> None: thread = threading.Thread(target=_run_in_thread) thread.start() - assert transport.requests or True while len(futures) < 2: time.sleep(0.01) for future in futures: assert future.done() + assert len(transport.requests) == 2 executor.cancel("task-api-batch") collect_release.set() thread.join(timeout=10) From fc8d5c284fa2c774684ea17cd48581f9d04d734d Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Sun, 27 Sep 2026 11:57:08 +0700 Subject: [PATCH 33/71] test: bound the submission wait in the collection-guard cancellation test The wait for both requests to submit had no deadline, so a regression that stopped submission would hang the test instead of failing it. It now times out with a clear error, and a worker-thread error is surfaced before the request-count assertion. Signed-off-by: Zhengyuan Su Co-Authored-By: Claude Code --- tests/worker/test_api_executor.py | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/tests/worker/test_api_executor.py b/tests/worker/test_api_executor.py index 61de80870..4c2ee9299 100644 --- a/tests/worker/test_api_executor.py +++ b/tests/worker/test_api_executor.py @@ -1084,7 +1084,12 @@ def _run_in_thread() -> None: thread = threading.Thread(target=_run_in_thread) thread.start() + deadline = time.monotonic() + 5 while len(futures) < 2: + if errors: + raise errors[0] + if time.monotonic() > deadline: + raise AssertionError("timed out waiting for both requests to submit") time.sleep(0.01) for future in futures: assert future.done() From 821b2d06f47487b7eab04ed6fc2438417c4ec09a Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Thu, 24 Sep 2026 18:51:50 +0700 Subject: [PATCH 34/71] feat: API executor runs DataMixin row-wise and aggregate prompts The batch API executor now serves the same data specs as the local path. A dataframe (row-wise) spec issues one request per row, and a graph_template (aggregate) spec issues one request over all rows. A body value that is exactly {{prompt}} is replaced by the row's prompt object as-is, so a message list stays a list; an embedded {{prompt}} keeps string substitution. Grouped dataframe data comes back as one APIGroupItem {index, rows} per table, the way the vLLM executor regroups with _populate_table; list data stays one item per row. _evaluate_expr maps attribute access over lists of result models and nested lists of groups, and resolves a field by its alias, so items.rows.json and items.json read API content downstream. Signed-off-by: Zhengyuan Su Co-Authored-By: Claude Code --- src/shared/schemas/result/__init__.py | 2 + src/shared/schemas/result/catalog.py | 21 +- src/shared/schemas/result/payloads.py | 12 + src/worker/executors/api_executor.py | 50 ++- src/worker/executors/utils/graph_templates.py | 100 +++--- tests/worker/test_api_executor.py | 288 ++++++++++++++++++ tests/worker/test_graph_templates_expr.py | 46 +++ 7 files changed, 471 insertions(+), 48 deletions(-) create mode 100644 tests/worker/test_graph_templates_expr.py diff --git a/src/shared/schemas/result/__init__.py b/src/shared/schemas/result/__init__.py index 8839005c6..e613b4fe4 100644 --- a/src/shared/schemas/result/__init__.py +++ b/src/shared/schemas/result/__init__.py @@ -39,6 +39,7 @@ AgentItem, AgentMetadata, AgentUsage, + APIGroupItem, APIItem, CostEstimates, DataRetrievalItem, @@ -89,6 +90,7 @@ __all__ = [ "APIItem", + "APIGroupItem", "APIResult", "AgentBatchSummary", "AgentItem", diff --git a/src/shared/schemas/result/catalog.py b/src/shared/schemas/result/catalog.py index 60c00858b..deb73c5df 100644 --- a/src/shared/schemas/result/catalog.py +++ b/src/shared/schemas/result/catalog.py @@ -9,6 +9,7 @@ Field, SerializeAsAny, Tag, + field_validator, ) from shared.tasks.task_type import TaskType @@ -20,6 +21,7 @@ AgentItem, AgentMetadata, AgentUsage, + APIGroupItem, APIItem, CostEstimates, DataRetrievalItem, @@ -259,7 +261,24 @@ class APIResult(StrictExecutorResult): response_json: Any = Field(default=None, alias="json") usage: dict[str, Any] | None = None text: str | None = None - items: list[APIItem] = Field(default_factory=list) + items: list[APIItem | APIGroupItem] = Field(default_factory=list) + + @field_validator("items", mode="before") + @classmethod + def _route_group_items(cls, value: Any) -> Any: + """Route a dict carrying ``rows`` to APIGroupItem before the union runs, + since a group dict also satisfies APIItem's required fields.""" + if not isinstance(value, list): + return value + routed: list[Any] = [] + for item in value: + if isinstance(item, dict) and "rows" in item: + routed.append(APIGroupItem.model_validate(item)) + elif isinstance(item, dict): + routed.append(APIItem.model_validate(item)) + else: + routed.append(item) + return routed class SSHResult(StrictExecutorResult): diff --git a/src/shared/schemas/result/payloads.py b/src/shared/schemas/result/payloads.py index 97a3c3eca..721bf3ca7 100644 --- a/src/shared/schemas/result/payloads.py +++ b/src/shared/schemas/result/payloads.py @@ -230,3 +230,15 @@ class APIItem(StrictModel): usage: dict[str, Any] | None = None text: str | None = None prompt: str | None = None + + +class APIGroupItem(StrictModel): + """One group's row responses in a batched API task over grouped data. + + ``rows`` holds the group's row responses in order. A group is one + dataframe table (one claim), so a downstream column reads ``rows`` as a + per-group list. + """ + + index: int + rows: list[APIItem] diff --git a/src/worker/executors/api_executor.py b/src/worker/executors/api_executor.py index 68d487dfc..bfee64ee7 100644 --- a/src/worker/executors/api_executor.py +++ b/src/worker/executors/api_executor.py @@ -8,7 +8,7 @@ import httpx -from shared.schemas.result import APIItem, APIResult +from shared.schemas.result import APIGroupItem, APIItem, APIResult from shared.tasks.specs import ApiSpecStrict from shared.tasks.task_type import TaskType from shared.utils.redact import is_credential_key @@ -229,7 +229,7 @@ def _substitute_prompt(cls, value: Any, prompt: Any) -> Any: return value def _build_request_kwargs( - self, api_cfg: dict[str, Any], prompt: str | None + self, api_cfg: dict[str, Any], prompt: Any | None ) -> dict[str, Any]: """Build httpx request kwargs from ``spec.api``, substituting the row prompt when batching.""" @@ -454,12 +454,52 @@ def _issue(idx: int, prompt: Any) -> APIItem: items = [results[idx] for idx in range(len(prompts))] + result_items: list[APIItem | APIGroupItem] = [] + if entry.tables: + # Grouped data: one result item per table, holding that group's + # row responses in order (same slicing as DataMixin._populate_table). + grouped: list[APIGroupItem] = [] + cur = 0 + for group_index, df in enumerate(entry.tables): + size = len(df) + grouped.append( + APIGroupItem(index=group_index, rows=items[cur : cur + size]) + ) + cur += size + if cur != len(items): + raise ExecutionError( + f"Output length {len(items)} does not match " + f"the total number of rows {cur} in table stores." + ) + result_items.extend(grouped) + else: + result_items.extend(items) + + if result_items: + first = result_items[0] + if isinstance(first, APIGroupItem): + status_code = first.rows[0].status_code + truncated = any( + r.truncated + for g in result_items + if isinstance(g, APIGroupItem) + for r in g.rows + ) + else: + status_code = first.status_code + truncated = any( + item.truncated for item in result_items if isinstance(item, APIItem) + ) + else: + status_code = 0 + truncated = False + return APIResult( ok=True, executor=self.name, method=method, url=str(url), - status_code=items[0].status_code, - truncated=any(item.truncated for item in items), - items=items, + status_code=status_code, + truncated=truncated, + items=result_items, ) diff --git a/src/worker/executors/utils/graph_templates.py b/src/worker/executors/utils/graph_templates.py index e51aad54d..008701fa7 100644 --- a/src/worker/executors/utils/graph_templates.py +++ b/src/worker/executors/utils/graph_templates.py @@ -602,49 +602,9 @@ def _evaluate_expr(expr: str, context: dict[str, BaseExecutorResult]) -> Any: continue attr, indexes = _split_indexes(token) if attr: - if isinstance(value, dict) and attr in value: - value = value[attr] - elif isinstance(value, list) and all( - isinstance(v, dict) and attr in v for v in value - ): - value = [v[attr] for v in value] - elif isinstance(value, list) and all( - isinstance(v, pd.DataFrame) for v in value - ): - if any(attr not in v.columns for v in value): - raise ExecutionError( - f"{attr} not a valid column in one of the " - f"DataFrames for {token}." - ) - value = [v[attr].tolist() for v in value] - elif isinstance(value, pd.DataFrame): - if attr not in value.columns: - raise ExecutionError( - f"{attr} not a valid column in DataFrame for {token}." - ) - value = value[attr].tolist() - elif isinstance(value, BaseModel): - resolved = getattr(value, attr, _SENTINEL) - if resolved is _SENTINEL: - raise ExecutionError( - f"{attr} not a valid attribute of {type(value).__name__} " - f"for {token}." - ) - value = resolved - else: - raise ExecutionError( - f"{attr} in {parts} is not a valid key - " - f"{type(value).__name__}, {value}" - ) + value = _apply_attr(value, attr, token, parts) for idx in indexes: - if isinstance(value, list) and -len(value) <= idx < len(value): - value = value[idx] - elif isinstance(value, list) and all(isinstance(v, list) for v in value): - value = [v[idx] for v in value] - else: - raise ExecutionError( - f"{idx} not a valid index in {token} - {len(value)}" - ) + value = _apply_index(value, idx, token) # Attempt to deserialize DataFrame if applicable if isinstance(value, dict): value = try_deserialize_dataframe(value) @@ -653,6 +613,62 @@ def _evaluate_expr(expr: str, context: dict[str, BaseExecutorResult]) -> Any: return value +def _apply_attr(value: Any, attr: str, token: str, parts: list[str]) -> Any: + """Resolve an attribute access, mapping over lists of dicts, DataFrames, + or pydantic models (including nested lists).""" + if isinstance(value, dict) and attr in value: + return value[attr] + if isinstance(value, list): + if all(isinstance(v, dict) and attr in v for v in value): + return [v[attr] for v in value] + if all(isinstance(v, pd.DataFrame) for v in value): + if any(attr not in v.columns for v in value): + raise ExecutionError( + f"{attr} not a valid column in one of the " + f"DataFrames for {token}." + ) + return [v[attr].tolist() for v in value] + if all(isinstance(v, BaseModel) for v in value): + return [_model_attr(v, attr, token) for v in value] + if all(isinstance(v, list) for v in value): + return [_apply_attr(v, attr, token, parts) for v in value] + if isinstance(value, pd.DataFrame): + if attr not in value.columns: + raise ExecutionError(f"{attr} not a valid column in DataFrame for {token}.") + return value[attr].tolist() + if isinstance(value, BaseModel): + return _model_attr(value, attr, token) + raise ExecutionError( + f"{attr} in {parts} is not a valid key - " f"{type(value).__name__}, {value}" + ) + + +def _model_attr(value: BaseModel, attr: str, token: str) -> Any: + """Resolve a declared pydantic field by name or alias, never a method.""" + fields = type(value).model_fields + field = fields.get(attr) + if field is not None: + return getattr(value, attr) + for name, f in fields.items(): + if f.alias == attr: + return getattr(value, name) + extras = getattr(value, "__pydantic_extra__", None) + if extras and attr in extras: + return extras[attr] + raise ExecutionError( + f"{attr} not a valid attribute of {type(value).__name__} for {token}." + ) + + +def _apply_index(value: Any, idx: int, token: str) -> Any: + """Index into a list, mapping over a list of lists (one per group).""" + if isinstance(value, list) and all(isinstance(v, list) for v in value): + return [_apply_index(v, idx, token) for v in value] + if isinstance(value, list) and -len(value) <= idx < len(value): + return value[idx] + raise ExecutionError(f"{idx} not a valid index in {token} - {len(value)}") + + def _split_indexes(token: str) -> tuple[str, list[int]]: parts = token.split("[") attr = parts[0] diff --git a/tests/worker/test_api_executor.py b/tests/worker/test_api_executor.py index 4c2ee9299..2e2c925d7 100644 --- a/tests/worker/test_api_executor.py +++ b/tests/worker/test_api_executor.py @@ -12,6 +12,7 @@ import httpx import pytest +from shared.schemas.result import APIGroupItem, APIItem, APIResult from shared.tasks.worker_message import WorkerTaskMessage from worker.executors import api_executor as api_executor_module from worker.executors.api_executor import APIExecutor @@ -121,6 +122,8 @@ def _executor() -> APIExecutor: executor._active_task_id = None executor._pending_cancelled_ids = set() executor._cancel_lock = threading.Lock() + executor._task_id = None + executor._current_batch_id = None return executor @@ -153,6 +156,13 @@ def _batch_task(items: list[Any], **api_updates: Any) -> WorkerTaskMessage: return WorkerTaskMessage.model_validate(payload) +def _api_item(content: str) -> APIItem: + """An upstream APIItem whose response carries the given content.""" + item = APIItem(index=0, url="u", status_code=200) + item.response_json = {"choices": [{"message": {"content": content}}]} + return item + + class TestNebulaPath: def test_no_url_no_header_uses_nebula_url_and_token( self, monkeypatch: pytest.MonkeyPatch @@ -1145,3 +1155,281 @@ def _run_in_thread() -> None: assert len(errors) == 1 assert isinstance(errors[0], TaskCancelledError) + + +class TestPromptSubstitution: + def test_exact_placeholder_substitutes_raw_object(self) -> None: + """A body value that is exactly {{prompt}} is replaced by the prompt + object as-is, so a message list stays a list of dicts.""" + messages = [ + {"role": "system", "content": "sys"}, + {"role": "user", "content": "hi"}, + ] + payload = { + "task_id": "task-api-sub", + "workflow_id": "wf-1", + "owner_id": "owner", + "assigned_worker": "worker-1", + "dispatched_at": "2026-03-22T00:00:00Z", + "task": { + "apiVersion": "flowmesh/v1", + "kind": "Task", + "metadata": {"name": "wf:api"}, + "spec": { + "taskType": "api", + "api": { + "method": "POST", + "url": "https://custom.example.com/v1/chat/completions", + "json": {"messages": "{{prompt}}"}, + }, + "data": {"type": "list", "items": [messages]}, + }, + }, + } + task = WorkerTaskMessage.model_validate(payload) + transport = _EchoTransport() + _run(_executor(), task, transport) + assert len(transport.requests) == 1 + body = json.loads(transport.requests[0].read()) + assert body["messages"] == messages + + def test_embedded_placeholder_substitutes_string(self) -> None: + """An embedded {{prompt}} inside a longer string keeps string + substitution.""" + payload = { + "task_id": "task-api-sub", + "workflow_id": "wf-1", + "owner_id": "owner", + "assigned_worker": "worker-1", + "dispatched_at": "2026-03-22T00:00:00Z", + "task": { + "apiVersion": "flowmesh/v1", + "kind": "Task", + "metadata": {"name": "wf:api"}, + "spec": { + "taskType": "api", + "api": { + "method": "POST", + "url": "https://custom.example.com/v1/chat/completions", + "json": { + "messages": [{"role": "user", "content": "Q: {{prompt}}"}] + }, + }, + "data": {"type": "list", "items": ["hello"]}, + }, + }, + } + task = WorkerTaskMessage.model_validate(payload) + transport = _EchoTransport() + _run(_executor(), task, transport) + body = json.loads(transport.requests[0].read()) + assert body["messages"][0]["content"] == "Q: hello" + + +class TestDataframeRows: + def test_dataframe_column_issues_one_request_per_row(self, tmp_path: Path) -> None: + """A dataframe-spec API task whose column reads an upstream APIResult's + items issues one request per upstream row, each body carrying that + row's messages as a list.""" + upstream = APIResult( + ok=True, + executor="api", + method="POST", + url="https://up.example.com", + status_code=200, + items=[_api_item("c0"), _api_item("c1")], + ) + payload = { + "task_id": "task-api-df", + "workflow_id": "wf-1", + "owner_id": "owner", + "assigned_worker": "worker-1", + "dispatched_at": "2026-03-22T00:00:00Z", + "task": { + "apiVersion": "flowmesh/v1", + "kind": "Task", + "metadata": {"name": "wf:api"}, + "spec": { + "taskType": "api", + "_upstreamResults": {"Up": upstream}, + "api": { + "method": "POST", + "url": "https://custom.example.com/v1/chat/completions", + "json": {"messages": "{{prompt}}"}, + }, + "data": { + "type": "dataframe", + "columns": [ + { + "label": "L", + "node": "Up", + "path": "items.json.choices[0].message.content", + } + ], + "messages": [ + {"role": "user", "content": "row {L}"}, + ], + }, + }, + }, + } + task = WorkerTaskMessage.model_validate(payload) + transport = _EchoTransport() + result = _run(_executor(), task, transport, tmp_path) + assert len(transport.requests) == 2 + issued = { + json.loads(req.read())["messages"][0]["content"] + for req in transport.requests + } + assert issued == {"row c0", "row c1"} + # A single-column dataframe is one table, so one group item holds both rows. + assert len(result.items) == 1 + prompts = [json.loads(r.prompt) for r in result.items[0].rows] + assert prompts == [ + [{"role": "user", "content": "row c0"}], + [{"role": "user", "content": "row c1"}], + ] + + +class TestGraphTemplateAggregate: + def test_graph_template_aggregates_all_rows_into_one_prompt( + self, tmp_path: Path + ) -> None: + """A graph_template aggregate API task over all rows of an upstream + APIResult issues one request whose prompt contains every row.""" + upstream = APIResult( + ok=True, + executor="api", + method="POST", + url="https://up.example.com", + status_code=200, + items=[_api_item("c0"), _api_item("c1")], + ) + payload = { + "task_id": "task-api-gt", + "workflow_id": "wf-1", + "owner_id": "owner", + "assigned_worker": "worker-1", + "dispatched_at": "2026-03-22T00:00:00Z", + "task": { + "apiVersion": "flowmesh/v1", + "kind": "Task", + "metadata": {"name": "wf:api"}, + "spec": { + "taskType": "api", + "_upstreamResults": {"Up": upstream}, + "api": { + "method": "POST", + "url": "https://custom.example.com/v1/chat/completions", + "json": {"messages": "{{prompt}}"}, + }, + "data": { + "type": "graph_template", + "template": { + "name": "format", + "columns": [ + { + "label": "df", + "data": { + "type": "dataframe", + "columns": [ + { + "label": "L", + "node": "Up", + "path": ( + "items.json.choices[0].message.content" + ), + } + ], + }, + } + ], + "options": { + "format": { + "steps": [], + "messages": [ + {"role": "user", "content": "all: {df}"} + ], + } + }, + }, + }, + }, + }, + } + task = WorkerTaskMessage.model_validate(payload) + transport = _EchoTransport() + result = _run(_executor(), task, transport, tmp_path) + assert len(transport.requests) == 1 + body = json.loads(transport.requests[0].read()) + content = body["messages"][0]["content"] + assert "c0" in content and "c1" in content + assert len(result.items) == 1 + # Aggregate result is a plain item (read at items.json...), not a group + # item (read at items.rows.json...). + item = result.items[0] + assert not hasattr(item, "rows") + assert item.response_json["choices"][0]["message"]["content"].startswith( + "echo:all:" + ) + + +class TestGroupedResult: + def test_ragged_groups_return_one_item_per_group(self, tmp_path: Path) -> None: + """A ragged grouped dataframe API task (groups of different sizes) + returns one item per group with the right rows in each.""" + upstream = APIResult( + ok=True, + executor="api", + method="POST", + url="https://up.example.com", + status_code=200, + items=[ + APIGroupItem(index=0, rows=[_api_item("c0"), _api_item("c1")]), + APIGroupItem(index=1, rows=[_api_item("c2")]), + ], + ) + payload = { + "task_id": "task-api-grp", + "workflow_id": "wf-1", + "owner_id": "owner", + "assigned_worker": "worker-1", + "dispatched_at": "2026-03-22T00:00:00Z", + "task": { + "apiVersion": "flowmesh/v1", + "kind": "Task", + "metadata": {"name": "wf:api"}, + "spec": { + "taskType": "api", + "_upstreamResults": {"Up": upstream}, + "api": { + "method": "POST", + "url": "https://custom.example.com/v1/chat/completions", + "json": {"messages": "{{prompt}}"}, + }, + "data": { + "type": "dataframe", + "columns": [ + { + "label": "L", + "node": "Up", + "path": "items.rows.json.choices[0].message.content", + } + ], + "messages": [ + {"role": "user", "content": "row {L}"}, + ], + }, + }, + }, + } + task = WorkerTaskMessage.model_validate(payload) + transport = _EchoTransport() + result = _run(_executor(), task, transport, tmp_path) + assert len(transport.requests) == 3 + assert len(result.items) == 2 + assert [len(item.rows) for item in result.items] == [2, 1] + # Group 0 holds rows c0, c1; group 1 holds c2. + group0 = {json.loads(r.prompt)[0]["content"] for r in result.items[0].rows} + assert group0 == {"row c0", "row c1"} + assert json.loads(result.items[1].rows[0].prompt)[0]["content"] == "row c2" diff --git a/tests/worker/test_graph_templates_expr.py b/tests/worker/test_graph_templates_expr.py new file mode 100644 index 000000000..d50c7b1a5 --- /dev/null +++ b/tests/worker/test_graph_templates_expr.py @@ -0,0 +1,46 @@ +"""Tests for _evaluate_expr attribute/index resolution over pydantic models, +aliases, and nested lists.""" + +from shared.schemas.result import APIGroupItem, APIItem, APIResult +from worker.executors.utils.graph_templates import _evaluate_expr + + +def _item(content: str) -> APIItem: + item = APIItem(index=0, url="u", status_code=200) + item.response_json = {"choices": [{"message": {"content": content}}]} + return item + + +def test_expr_over_list_of_models_with_alias() -> None: + """Attribute access maps over a list of APIItem models, resolving the + ``json`` alias to ``response_json``.""" + upstream = APIResult( + ok=True, + executor="api", + method="POST", + url="https://up.example.com", + status_code=200, + items=[_item("c0"), _item("c1")], + ) + value = _evaluate_expr("Up.items.json.choices[0].message.content", {"Up": upstream}) + assert value == ["c0", "c1"] + + +def test_expr_over_nested_lists_of_models() -> None: + """Attribute access maps through nested lists (a list of groups, each a + list of APIItem models), yielding one inner list per group.""" + upstream = APIResult( + ok=True, + executor="api", + method="POST", + url="https://up.example.com", + status_code=200, + items=[ + APIGroupItem(index=0, rows=[_item("c0"), _item("c1")]), + APIGroupItem(index=1, rows=[_item("c2")]), + ], + ) + value = _evaluate_expr( + "Up.items.rows.json.choices[0].message.content", {"Up": upstream} + ) + assert value == [["c0", "c1"], ["c2"]] From 7ab51b6a96497dcedde1b9534e421a1600c7229a Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Thu, 24 Sep 2026 17:40:03 +0700 Subject: [PATCH 35/71] feat: echo task runs a function over whole upstream lists The echo task gains data type "function": a deterministic Python function, run once in the existing sandbox over the whole resolved arguments ({node, path}, {expr} or {items}), that returns a list of JSON values. Each element becomes one echo item, not flattened, so an element that is a list is one group. Downstream ops read the rows at items.output. A non-list or non-JSON return fails the task. data type "list" keeps its behaviour. Co-Authored-By: Claude Code Signed-off-by: Zhengyuan Su --- src/worker/executors/echo_executor.py | 70 +++++++++-- src/worker/executors/utils/safe_eval.py | 58 +++++++-- tests/worker/test_echo_executor.py | 149 ++++++++++++++++++++++++ tests/worker/test_safe_eval.py | 57 +++++++++ 4 files changed, 313 insertions(+), 21 deletions(-) create mode 100644 tests/worker/test_echo_executor.py create mode 100644 tests/worker/test_safe_eval.py diff --git a/src/worker/executors/echo_executor.py b/src/worker/executors/echo_executor.py index 235bafec9..3c8425caa 100644 --- a/src/worker/executors/echo_executor.py +++ b/src/worker/executors/echo_executor.py @@ -12,6 +12,7 @@ from .mixins.data import DataMixin from .utils.checkpoints import maybe_upload_traces from .utils.graph_templates import _evaluate_expr +from .utils.safe_eval import safe_execute_function, safe_materialize_function logger = logging.getLogger(__name__) @@ -51,6 +52,19 @@ def _resolve_expr_item( ) return resolved + @staticmethod + def _resolve_function_arg( + arg: dict[str, Any], context: dict[str, BaseExecutorResult] + ) -> Any: + if "items" in arg: + items = arg["items"] + if not isinstance(items, list): + raise ExecutionError( + "echo executor function argument 'items' must be a list" + ) + return items + return EchoExecutor._resolve_expr_item(arg, context) + def _resolve_item( self, item: EchoItem, context: dict[str, BaseExecutorResult] ) -> Any: @@ -64,6 +78,48 @@ def _resolve_item( "a string literal or a mapping" ) + def _run_list( + self, data_cfg: dict[str, Any], context: dict[str, BaseExecutorResult] + ) -> list[EchoResultItem]: + items_cfg = data_cfg.get("items") + if not isinstance(items_cfg, list): + raise ExecutionError("echo executor requires spec.data.items to be a list") + merged_items: list[EchoResultItem] = [] + for item in items_cfg: + resolved = self._resolve_item(item, context) + self._append_outputs(merged_items, resolved) + return merged_items + + def _run_function( + self, + data_cfg: dict[str, Any], + context: dict[str, BaseExecutorResult], + task_id: str, + ) -> list[EchoResultItem]: + fn_code = data_cfg.get("function") + if not isinstance(fn_code, str) or not fn_code.strip(): + raise ExecutionError( + f"echo executor task {task_id} requires spec.data.function " + "to be a non-empty string" + ) + args_cfg = data_cfg.get("arguments") + if not isinstance(args_cfg, list): + raise ExecutionError( + f"echo executor task {task_id} requires spec.data.arguments " + "to be a list" + ) + resolved_args = [self._resolve_function_arg(arg, context) for arg in args_cfg] + try: + fn_obj = safe_materialize_function(fn_code) + output = safe_execute_function( + fn_obj, tuple(resolved_args), expect_list=True + ) + except Exception as e: + raise ExecutionError( + f"echo executor task {task_id} function failed: {e}" + ) from e + return [EchoResultItem(output=element) for element in output] + def run(self, task: ExecutorTask, out_dir: Path) -> EchoResult: spec = self.require_spec(task, EchoSpecStrict) task_id = task.task_id.strip() @@ -75,20 +131,16 @@ def run(self, task: ExecutorTask, out_dir: Path) -> EchoResult: if not isinstance(data_cfg, dict): raise ExecutionError("echo executor requires spec.data to be a mapping") - items_cfg = data_cfg.get("items") - if not isinstance(items_cfg, list): - raise ExecutionError( - "echo executor requires spec.data.items to be a list" - ) if not isinstance(context, dict): raise ExecutionError( "echo executor requires spec._upstreamResults to be a mapping" ) - merged_items: list[EchoResultItem] = [] - for item in items_cfg: - resolved = self._resolve_item(item, context) - self._append_outputs(merged_items, resolved) + data_type = data_cfg.get("type") + if data_type == "function": + merged_items = self._run_function(data_cfg, context, task_id) + else: + merged_items = self._run_list(data_cfg, context) result = EchoResult( items=merged_items, diff --git a/src/worker/executors/utils/safe_eval.py b/src/worker/executors/utils/safe_eval.py index b84c1522a..01beca833 100644 --- a/src/worker/executors/utils/safe_eval.py +++ b/src/worker/executors/utils/safe_eval.py @@ -9,7 +9,7 @@ - Two-phase execution: materialize (compile) then execute (run) - Restricted builtins: only safe operations (no open, eval, exec, import, etc.) - Limited module access: json, re, math, numpy, pandas, pyarrow -- Type validation: enforces Callable[[tuple[str, ...]], str] signature +- Type validation: enforces Callable[[tuple[str, ...]], Any] signature - Isolated execution: exec() with explicit safe_globals/safe_locals Typical usage: @@ -69,9 +69,24 @@ } +def _is_json_value(value: Any) -> bool: + """Whether ``value`` is a JSON value (recursively, string keys, finite numbers).""" + if value is None or isinstance(value, (str, bool)): + return True + if isinstance(value, int): + return True + if isinstance(value, float): + return value == value and value not in (float("inf"), float("-inf")) + if isinstance(value, list): + return all(_is_json_value(item) for item in value) + if isinstance(value, dict): + return all(isinstance(k, str) and _is_json_value(v) for k, v in value.items()) + return False + + def safe_materialize_function( fn_code: str, -) -> Callable[[tuple[str | list[dict[str, str]], ...]], str]: +) -> Callable[[tuple[str | list[dict[str, str]], ...]], Any]: """ Compile function source code into a callable object with restricted builtins. @@ -86,8 +101,8 @@ def safe_materialize_function( Type signature enforcement: - Must accept exactly 1 parameter (tuple of strings or list of messages) - - Should return a string (validated at execution time) - - Signature: Callable[[tuple[str, ...]], str] + - Returns a string or a list of JSON values (validated at execution time) + - Signature: Callable[[tuple[str, ...]], Any] Args: fn_code: Python source code defining a function or lambda expression @@ -167,10 +182,11 @@ def safe_materialize_function( def safe_execute_function( - fn_obj: Callable[[tuple[str | list[dict[str, str]], ...]], str], + fn_obj: Callable[[tuple[str | list[dict[str, str]], ...]], Any], args: tuple[str | Sequence[dict[str, str]], ...], allowed_modules: dict[str, Any] | None = None, -) -> str: + expect_list: bool = False, +) -> Any: """ Execute a function in an isolated environment with no access to external state. @@ -185,7 +201,8 @@ def safe_execute_function( 1. Validate input types (args must be tuple of strings or lists) 2. Create isolated globals with SAFE_BUILTINS and SAFE_MODULES 3. Execute function call via exec() in restricted environment - 4. Extract result and validate output type (must be string) + 4. Extract result and validate output type (string, or list of JSON when + ``expect_list`` is set) Args: fn_obj: Compiled function from safe_materialize_function() @@ -193,13 +210,14 @@ def safe_execute_function( allowed_modules: Optional dict of additional modules to allow during execution. If None, uses SAFE_MODULES (json, re, math, numpy, pandas, pyarrow). + expect_list: When True, require the result to be a list of JSON values. Returns: - String result from function execution + Result from function execution Raises: RuntimeError: If function execution fails for any reason - TypeError: If args is not tuple[str, ...] or result is not str + TypeError: If args is not tuple[str, ...] or result is not the expected type Examples: >>> fn = safe_materialize_function("lambda args: args[0].upper()") @@ -213,7 +231,11 @@ def safe_execute_function( if not isinstance(args, tuple): raise TypeError(f"Args must be a tuple, got {type(args).__name__}") - if not all(isinstance(arg, (str, list)) for arg in args): + if expect_list: + if not all(_is_json_value(arg) for arg in args): + arg_types = [type(arg).__name__ for arg in args] + raise TypeError(f"All args must be JSON values, got types: {arg_types}") + elif not all(isinstance(arg, (str, list)) for arg in args): arg_types = [type(arg).__name__ for arg in args] raise TypeError(f"All args must be strings or lists, got types: {arg_types}") @@ -239,8 +261,20 @@ def safe_execute_function( # Extract the result from safe_locals result = safe_locals["__result__"] - # Validate output type: must be a string - if not isinstance(result, str): + # Validate output type + if expect_list: + if not isinstance(result, list): + raise TypeError( + "Function must return a list, but returned " + f"{type(result).__name__}: {result}" + ) + for element in result: + if not _is_json_value(element): + raise TypeError( + "Function must return a list of JSON values, but an element " + f"is {type(element).__name__}: {element}" + ) + elif not isinstance(result, str): raise TypeError( "Function must return a string, but returned " f"{type(result).__name__}: {result}" diff --git a/tests/worker/test_echo_executor.py b/tests/worker/test_echo_executor.py new file mode 100644 index 000000000..8f2e2a520 --- /dev/null +++ b/tests/worker/test_echo_executor.py @@ -0,0 +1,149 @@ +"""Echo executor tests: the literal "list" path and the list-Lambda "function" path.""" + +from pathlib import Path + +import pytest +from pydantic import JsonValue + +from shared.schemas.result import EchoItem, EchoResult +from shared.tasks import TaskType +from worker.executors.base_executor import ExecutionError +from worker.executors.echo_executor import EchoExecutor + +from .factories import make_worker_config, make_worker_task_message + + +def _spec(data: dict, upstream: dict | None = None) -> dict: + spec: dict = {"taskType": "echo", "data": data} + if upstream: + spec["_upstreamResults"] = upstream + return spec + + +def _run( + data: dict, upstream: dict | None = None, tmp_path: Path | None = None +) -> EchoResult: + executor = EchoExecutor(make_worker_config()) + task = make_worker_task_message( + _spec(data, upstream), task_type=TaskType.ECHO, task_id="tsk-echo" + ) + return executor.run(task, tmp_path or Path("/tmp/echo-out")) + + +def _echo_result(*outputs: JsonValue) -> EchoResult: + return EchoResult(items=[EchoItem(output=o) for o in outputs], count=len(outputs)) + + +class TestListPath: + def test_literal_items_are_echoed(self) -> None: + result = _run({"type": "list", "items": ["a", "b", "c"]}) + assert [i.output for i in result.items] == ["a", "b", "c"] + + +class TestFunctionPath: + def test_explode_one_input_list_into_rows(self) -> None: + result = _run( + { + "type": "function", + "function": "lambda args: [x * 2 for x in args[0]]", + "arguments": [{"items": [1, 2, 3]}], + } + ) + assert [i.output for i in result.items] == [2, 4, 6] + + def test_filter_rows(self) -> None: + result = _run( + { + "type": "function", + "function": "lambda args: [x for x in args[0] if x % 2 == 0]", + "arguments": [{"items": [1, 2, 3, 4]}], + } + ) + assert [i.output for i in result.items] == [2, 4] + + def test_split_one_input_into_two_outputs(self) -> None: + result = _run( + { + "type": "function", + "function": "lambda args: [args[0][:2], args[0][2:]]", + "arguments": [{"items": [1, 2, 3, 4]}], + } + ) + assert [i.output for i in result.items] == [[1, 2], [3, 4]] + + def test_cross_product_into_groups_of_different_sizes(self) -> None: + result = _run( + { + "type": "function", + "function": ( + "lambda args: [[[a, b] for b in args[1]] for a in args[0]]" + ), + "arguments": [{"items": [1, 2]}, {"items": ["x", "y", "z"]}], + } + ) + assert [i.output for i in result.items] == [ + [[1, "x"], [1, "y"], [1, "z"]], + [[2, "x"], [2, "y"], [2, "z"]], + ] + + def test_collapse_groups_back_to_one_row_per_group(self) -> None: + result = _run( + { + "type": "function", + "function": "lambda args: [sum(g) for g in args[0]]", + "arguments": [{"items": [[1, 2], [3, 4, 5]]}], + } + ) + assert [i.output for i in result.items] == [3, 12] + + def test_node_path_argument_reads_upstream_echo_result(self) -> None: + upstream = {"echo-a": _echo_result("p", "q")} + result = _run( + { + "type": "function", + "function": "lambda args: [args[0].upper()]", + "arguments": [{"node": "echo-a", "path": "items[0].output"}], + }, + upstream=upstream, + ) + assert [i.output for i in result.items] == ["P"] + + def test_literal_items_argument(self) -> None: + result = _run( + { + "type": "function", + "function": "lambda args: [args[0]]", + "arguments": [{"items": ["a", "b"]}], + } + ) + assert [i.output for i in result.items] == [["a", "b"]] + + def test_non_list_return_raises(self) -> None: + with pytest.raises(ExecutionError, match="must return a list"): + _run( + { + "type": "function", + "function": "lambda args: 'not a list'", + "arguments": [{"items": [1]}], + } + ) + + def test_non_json_element_raises(self) -> None: + with pytest.raises(ExecutionError, match="function failed"): + _run( + { + "type": "function", + "function": "lambda args: [object()]", + "arguments": [{"items": [1]}], + } + ) + + def test_nested_non_json_element_raises(self) -> None: + with pytest.raises(ExecutionError, match="function failed"): + _run( + { + "type": "function", + "function": "lambda args: [{'k': set([1])}]", + "arguments": [{"items": [1]}], + } + ) diff --git a/tests/worker/test_safe_eval.py b/tests/worker/test_safe_eval.py new file mode 100644 index 000000000..4389eb04e --- /dev/null +++ b/tests/worker/test_safe_eval.py @@ -0,0 +1,57 @@ +"""safe_eval tests: the string-result prompt path and the list-result Lambda path.""" + +import pytest + +from worker.executors.utils.safe_eval import ( + safe_execute_function, + safe_materialize_function, +) + + +def _run(fn_code: str, args: tuple, *, expect_list: bool = False): + fn_obj = safe_materialize_function(fn_code) + return safe_execute_function(fn_obj, args, expect_list=expect_list) + + +class TestStringResult: + def test_string_result_is_returned(self) -> None: + assert _run("lambda args: args[0].upper()", ("hello",)) == "HELLO" + + def test_non_string_result_raises(self) -> None: + with pytest.raises(RuntimeError, match="must return a string"): + _run("lambda args: [1, 2]", ([1],)) + + +class TestListResult: + def test_list_of_json_is_returned(self) -> None: + assert _run( + "lambda args: [x * 2 for x in args[0]]", ([1, 2, 3],), expect_list=True + ) == [2, 4, 6] + + def test_groups_are_allowed_as_elements(self) -> None: + assert _run("lambda args: [[1, 2], [3, 4, 5]]", ([],), expect_list=True) == [ + [1, 2], + [3, 4, 5], + ] + + def test_non_list_result_raises(self) -> None: + with pytest.raises(RuntimeError, match="must return a list"): + _run("lambda args: 'nope'", ([],), expect_list=True) + + def test_non_json_element_raises(self) -> None: + with pytest.raises(RuntimeError, match="list of JSON"): + _run("lambda args: [set([1])]", ([],), expect_list=True) + + def test_nested_non_json_element_raises(self) -> None: + with pytest.raises(RuntimeError, match="list of JSON"): + _run("lambda args: [{'k': set([1])}]", ([],), expect_list=True) + + def test_dict_argument_reaches_function(self) -> None: + assert _run("lambda args: [args[0]['a']]", ({"a": 1},), expect_list=True) == [1] + + def test_scalar_argument_reaches_function(self) -> None: + assert _run("lambda args: [args[0]]", (5,), expect_list=True) == [5] + + def test_string_mode_rejects_dict_argument(self) -> None: + with pytest.raises(TypeError, match="strings or lists"): + _run("lambda args: args[0]['a']", ({"a": 1},)) From b4cf53f125f45154eef42c32c043bb87a606af19 Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Thu, 24 Sep 2026 18:54:44 +0700 Subject: [PATCH 36/71] fix: echo function arguments and sandbox sources fail closed A function argument must be exactly {items}, {expr} or {node, path}; a mixed or unknown key set now fails the task instead of silently using items. The sandbox now accepts exactly one top-level statement: a def, or a lambda assigned to one name (the form Lumilake's SDK serialises). It selects that object by name instead of taking the first local. Two defs, extra statements or an import raise, and a syntax error is reported as RuntimeError like other compile failures. Signed-off-by: Zhengyuan Su Co-Authored-By: Claude Code --- src/worker/executors/echo_executor.py | 10 ++++-- src/worker/executors/utils/safe_eval.py | 36 ++++++++++++++++---- tests/worker/test_echo_executor.py | 20 +++++++++++ tests/worker/test_safe_eval.py | 45 +++++++++++++++++++++++++ 4 files changed, 103 insertions(+), 8 deletions(-) diff --git a/src/worker/executors/echo_executor.py b/src/worker/executors/echo_executor.py index 3c8425caa..edc700e60 100644 --- a/src/worker/executors/echo_executor.py +++ b/src/worker/executors/echo_executor.py @@ -56,14 +56,20 @@ def _resolve_expr_item( def _resolve_function_arg( arg: dict[str, Any], context: dict[str, BaseExecutorResult] ) -> Any: - if "items" in arg: + keys = frozenset(arg) + if keys == {"items"}: items = arg["items"] if not isinstance(items, list): raise ExecutionError( "echo executor function argument 'items' must be a list" ) return items - return EchoExecutor._resolve_expr_item(arg, context) + if keys in ({"expr"}, {"node", "path"}): + return EchoExecutor._resolve_expr_item(arg, context) + raise ExecutionError( + "echo executor function argument must have exactly one of " + f"'items', 'expr', or 'node'+'path'; got keys {sorted(keys)}" + ) def _resolve_item( self, item: EchoItem, context: dict[str, BaseExecutorResult] diff --git a/src/worker/executors/utils/safe_eval.py b/src/worker/executors/utils/safe_eval.py index 01beca833..0f976ba96 100644 --- a/src/worker/executors/utils/safe_eval.py +++ b/src/worker/executors/utils/safe_eval.py @@ -17,6 +17,7 @@ result = safe_execute_function(fn_obj, ("hello",)) # Returns "HELLO" """ +import ast import inspect import json import math @@ -137,6 +138,35 @@ def safe_materialize_function( # Case 2: Function definition (use exec - def is a statement) else: + try: + tree = ast.parse(fn_code_stripped) + except SyntaxError as e: + raise RuntimeError( + f"Function definition failed: {e}\nCode: {fn_code}" + ) from e + + if len(tree.body) != 1: + kinds = [type(n).__name__ for n in tree.body] + raise RuntimeError( + "Function source must be a single function definition, " + f"found top-level statements: {kinds}" + ) + stmt = tree.body[0] + if isinstance(stmt, ast.FunctionDef): + fn_name = stmt.name + elif ( + isinstance(stmt, ast.Assign) + and len(stmt.targets) == 1 + and isinstance(stmt.targets[0], ast.Name) + and isinstance(stmt.value, ast.Lambda) + ): + fn_name = stmt.targets[0].id + else: + raise RuntimeError( + "Function source must be a single function definition or " + f"an assignment of a lambda, found: {type(stmt).__name__}" + ) + safe_locals: dict[str, Any] = {} # Execute the function definition (creates function object in locals) @@ -147,12 +177,6 @@ def safe_materialize_function( f"Function definition failed: {e}\nCode: {fn_code}" ) from e - # Find the function object - if not safe_locals: - raise RuntimeError("Function definition did not create any objects") - - # Get the function (usually the first/only item in locals) - fn_name = list(safe_locals.keys())[0] fn_obj = safe_locals[fn_name] if not callable(fn_obj): diff --git a/tests/worker/test_echo_executor.py b/tests/worker/test_echo_executor.py index 8f2e2a520..8cbca28e8 100644 --- a/tests/worker/test_echo_executor.py +++ b/tests/worker/test_echo_executor.py @@ -118,6 +118,26 @@ def test_literal_items_argument(self) -> None: ) assert [i.output for i in result.items] == [["a", "b"]] + def test_mixed_argument_raises(self) -> None: + with pytest.raises(ExecutionError, match="exactly one"): + _run( + { + "type": "function", + "function": "lambda args: [args[0]]", + "arguments": [{"items": [1], "expr": "absent.items"}], + } + ) + + def test_unknown_key_argument_raises(self) -> None: + with pytest.raises(ExecutionError, match="exactly one"): + _run( + { + "type": "function", + "function": "lambda args: [args[0]]", + "arguments": [{"bogus": 1}], + } + ) + def test_non_list_return_raises(self) -> None: with pytest.raises(ExecutionError, match="must return a list"): _run( diff --git a/tests/worker/test_safe_eval.py b/tests/worker/test_safe_eval.py index 4389eb04e..94fd5a03b 100644 --- a/tests/worker/test_safe_eval.py +++ b/tests/worker/test_safe_eval.py @@ -55,3 +55,48 @@ def test_scalar_argument_reaches_function(self) -> None: def test_string_mode_rejects_dict_argument(self) -> None: with pytest.raises(TypeError, match="strings or lists"): _run("lambda args: args[0]['a']", ({"a": 1},)) + + +class TestMaterializeShape: + def test_single_def_works_in_string_mode(self) -> None: + assert _run("def f(args):\n return args[0].upper()", ("hi",)) == "HI" + + def test_single_def_works_in_function_mode(self) -> None: + assert _run( + "def f(args):\n return [x * 2 for x in args[0]]", + ([1, 2],), + expect_list=True, + ) == [2, 4] + + def test_lambda_works_in_both_modes(self) -> None: + assert _run("lambda args: args[0].upper()", ("hi",)) == "HI" + assert _run("lambda args: [args[0]]", (1,), expect_list=True) == [1] + + def test_two_defs_raise(self) -> None: + with pytest.raises(RuntimeError, match="single function definition"): + _run( + "def f(args):\n return args[0]\ndef g(args):\n return args[0]", + (1,), + ) + + def test_def_plus_top_level_statement_raises(self) -> None: + with pytest.raises(RuntimeError, match="single function definition"): + _run("def f(args):\n return args[0]\nx = 1", (1,)) + + def test_assigned_lambda_works_in_string_mode(self) -> None: + assert _run("lam = lambda args: args[0].lower()", ("HI",)) == "hi" + + def test_assigned_lambda_works_in_function_mode(self) -> None: + assert _run( + "lam = lambda args: [x * 2 for x in args[0]]", + ([1, 2],), + expect_list=True, + ) == [2, 4] + + def test_assign_of_non_lambda_raises(self) -> None: + with pytest.raises(RuntimeError, match="single function definition"): + _run("f = 1", (1,)) + + def test_syntax_error_raises_runtime_error(self) -> None: + with pytest.raises(RuntimeError, match="Function definition failed"): + _run("def f(args):\n return )", (1,)) From 1141a9996ffba7e333b2cb63e49536dfe8b98070 Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Thu, 24 Sep 2026 21:08:40 +0700 Subject: [PATCH 37/71] fix(worker): declare tabulate, which graph-template DataFrame rendering needs Signed-off-by: Zhengyuan Su Co-Authored-By: Claude Code --- pyproject.toml | 1 + src/worker/requirements/requirements.txt | 1 + uv.lock | 6 ++++++ 3 files changed, 8 insertions(+) diff --git a/pyproject.toml b/pyproject.toml index 22fbba274..cc2ebe5e0 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -120,6 +120,7 @@ runtime-analytics = [ "pandas>=2.3.3", "psycopg[binary]>=3.2.12", "sqlalchemy>=2.0.44", + "tabulate>=0.9.0", ] runtime-worker-cpu = [ { include-group = "runtime-worker-core" }, diff --git a/src/worker/requirements/requirements.txt b/src/worker/requirements/requirements.txt index 0d504f0b0..63308168a 100644 --- a/src/worker/requirements/requirements.txt +++ b/src/worker/requirements/requirements.txt @@ -51,6 +51,7 @@ qdrant-client==1.15.1 rich==14.2.0 sqlalchemy==2.0.44 sqlmodel==0.0.27 +tabulate==0.9.0 tiktoken==0.12.0 toml==0.10.2 torch==2.13.0 diff --git a/uv.lock b/uv.lock index 9fc3c344e..45582e612 100644 --- a/uv.lock +++ b/uv.lock @@ -1979,6 +1979,7 @@ ci = [ { name = "ruff" }, { name = "sqlalchemy" }, { name = "sqlmodel" }, + { name = "tabulate" }, { name = "tiktoken" }, { name = "toml" }, { name = "torch" }, @@ -2062,6 +2063,7 @@ runtime-analytics = [ { name = "pandas" }, { name = "psycopg", extra = ["binary"] }, { name = "sqlalchemy" }, + { name = "tabulate" }, ] runtime-inference = [ { name = "accelerate" }, @@ -2182,6 +2184,7 @@ runtime-worker-cpu = [ { name = "rich" }, { name = "sqlalchemy" }, { name = "sqlmodel" }, + { name = "tabulate" }, { name = "tiktoken" }, { name = "toml" }, { name = "torch" }, @@ -2290,6 +2293,7 @@ ci = [ { name = "ruff", specifier = ">=0.14.10" }, { name = "sqlalchemy", specifier = ">=2.0.44" }, { name = "sqlmodel", specifier = ">=0.0.27" }, + { name = "tabulate", specifier = ">=0.9.0" }, { name = "tiktoken", specifier = ">=0.12.0" }, { name = "toml", specifier = ">=0.10.2" }, { name = "torch", specifier = ">=2.11.0" }, @@ -2372,6 +2376,7 @@ runtime-analytics = [ { name = "pandas", specifier = ">=2.3.3" }, { name = "psycopg", extras = ["binary"], specifier = ">=3.2.12" }, { name = "sqlalchemy", specifier = ">=2.0.44" }, + { name = "tabulate", specifier = ">=0.9.0" }, ] runtime-inference = [ { name = "accelerate", specifier = ">=1.12.0" }, @@ -2490,6 +2495,7 @@ runtime-worker-cpu = [ { name = "rich", specifier = ">=14.2.0" }, { name = "sqlalchemy", specifier = ">=2.0.44" }, { name = "sqlmodel", specifier = ">=0.0.27" }, + { name = "tabulate", specifier = ">=0.9.0" }, { name = "tiktoken", specifier = ">=0.12.0" }, { name = "toml", specifier = ">=0.10.2" }, { name = "torch", specifier = ">=2.11.0" }, From d1c0e3eb7108def63322ab6b6d6d6d8642264c06 Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Fri, 25 Sep 2026 05:16:41 +0700 Subject: [PATCH 38/71] fix(sdk): mirror APIGroupItem and grouped APIResult.items Signed-off-by: Zhengyuan Su Co-Authored-By: Claude Code --- sdk/src/flowmesh/models/__init__.py | 2 ++ sdk/src/flowmesh/models/result/__init__.py | 2 ++ sdk/src/flowmesh/models/result/catalog.py | 21 +++++++++++- sdk/src/flowmesh/models/result/payloads.py | 12 +++++++ tests/sdk/test_models.py | 40 ++++++++++++++++++++++ tests/sdk/test_schema_compat.py | 17 +++++++++ 6 files changed, 93 insertions(+), 1 deletion(-) diff --git a/sdk/src/flowmesh/models/__init__.py b/sdk/src/flowmesh/models/__init__.py index 69506c497..aba9bc3ee 100644 --- a/sdk/src/flowmesh/models/__init__.py +++ b/sdk/src/flowmesh/models/__init__.py @@ -28,6 +28,7 @@ AgentResult, AgentUsage, AnyExecutorResult, + APIGroupItem, APIItem, APIResult, BaseExecutorResult, @@ -105,6 +106,7 @@ ) __all__ = [ + "APIGroupItem", "APIItem", "APIResult", "ActiveWaitBreakdown", diff --git a/sdk/src/flowmesh/models/result/__init__.py b/sdk/src/flowmesh/models/result/__init__.py index ecd772519..13e0d411b 100644 --- a/sdk/src/flowmesh/models/result/__init__.py +++ b/sdk/src/flowmesh/models/result/__init__.py @@ -39,6 +39,7 @@ AgentItem, AgentMetadata, AgentUsage, + APIGroupItem, APIItem, CostEstimates, DataRetrievalItem, @@ -89,6 +90,7 @@ _model.model_rebuild() __all__ = [ + "APIGroupItem", "APIItem", "APIResult", "AgentBatchSummary", diff --git a/sdk/src/flowmesh/models/result/catalog.py b/sdk/src/flowmesh/models/result/catalog.py index dcdafa54b..5aa8243b2 100644 --- a/sdk/src/flowmesh/models/result/catalog.py +++ b/sdk/src/flowmesh/models/result/catalog.py @@ -9,6 +9,7 @@ Field, SerializeAsAny, Tag, + field_validator, ) from ..artifacts import ArtifactRef @@ -18,6 +19,7 @@ AgentItem, AgentMetadata, AgentUsage, + APIGroupItem, APIItem, CostEstimates, DataRetrievalItem, @@ -213,7 +215,24 @@ class APIResult(StrictExecutorResult): response_json: Any = Field(default=None, alias="json") usage: dict[str, Any] | None = None text: str | None = None - items: list[APIItem] = Field(default_factory=list) + items: list[APIItem | APIGroupItem] = Field(default_factory=list) + + @field_validator("items", mode="before") + @classmethod + def _route_group_items(cls, value: Any) -> Any: + """Route a dict carrying ``rows`` to APIGroupItem before the union runs, + since a group dict also satisfies APIItem's required fields.""" + if not isinstance(value, list): + return value + routed: list[Any] = [] + for item in value: + if isinstance(item, dict) and "rows" in item: + routed.append(APIGroupItem.model_validate(item)) + elif isinstance(item, dict): + routed.append(APIItem.model_validate(item)) + else: + routed.append(item) + return routed class SSHResult(StrictExecutorResult): diff --git a/sdk/src/flowmesh/models/result/payloads.py b/sdk/src/flowmesh/models/result/payloads.py index cdab96fdc..8d9d3d7cc 100644 --- a/sdk/src/flowmesh/models/result/payloads.py +++ b/sdk/src/flowmesh/models/result/payloads.py @@ -174,3 +174,15 @@ class APIItem(StrictModel): usage: dict[str, Any] | None = None text: str | None = None prompt: str | None = None + + +class APIGroupItem(StrictModel): + """One group's row responses in a batched API task over grouped data. + + ``rows`` holds the group's row responses in order. A group is one + dataframe table (one claim), so a downstream column reads ``rows`` as a + per-group list. + """ + + index: int + rows: list[APIItem] diff --git a/tests/sdk/test_models.py b/tests/sdk/test_models.py index 025a26d4f..b2e5ab2cb 100644 --- a/tests/sdk/test_models.py +++ b/tests/sdk/test_models.py @@ -5,7 +5,9 @@ import pytest from flowmesh.models import ( ActiveWaitBreakdown, + APIGroupItem, APIItem, + APIResult, AssetSummary, CriticalPathSummary, E2EBreakdown, @@ -517,3 +519,41 @@ def test_construct_by_name_serialize_revalidate(self) -> None: {"index": 1, "url": "u", "status_code": 200, "json": {"a": 1}} ) assert by_alias.response_json == {"a": 1} + + +class TestAPIGroupResult: + def test_grouped_api_result_round_trip(self) -> None: + """A grouped API result (items carrying ``rows``) validates in the SDK + and round-trips through ``model_dump(by_alias=True)`` then + ``model_validate``.""" + payload = { + "task_type": "api", + "executor": "api", + "method": "POST", + "url": "http://example.com/v1/chat/completions", + "status_code": 200, + "items": [ + { + "index": 0, + "rows": [ + { + "index": 0, + "url": "http://example.com/v1/chat/completions", + "status_code": 200, + "json": {"choices": [{"message": {"content": "hi"}}]}, + "text": "hi", + } + ], + } + ], + } + result = APIResult.model_validate(payload) + assert isinstance(result.items[0], APIGroupItem) + assert ( + result.items[0].rows[0].response_json["choices"][0]["message"]["content"] + == "hi" + ) + wire = result.model_dump(by_alias=True) + reloaded = APIResult.model_validate(wire) + assert isinstance(reloaded.items[0], APIGroupItem) + assert reloaded.items[0].rows[0].text == "hi" diff --git a/tests/sdk/test_schema_compat.py b/tests/sdk/test_schema_compat.py index b30e557f0..ccfce65ca 100644 --- a/tests/sdk/test_schema_compat.py +++ b/tests/sdk/test_schema_compat.py @@ -4,6 +4,8 @@ then compares ``model_fields`` to detect drift. """ +from typing import Any + import pytest from flowmesh import models as sdk_models @@ -153,6 +155,7 @@ "RagQuery", "EchoItem", "APIItem", + "APIGroupItem", ] RESULT_MODEL_PAIRS = [ @@ -250,6 +253,20 @@ def test_worker_register_response_fields() -> None: assert_fields_match(SrvWorkerRegisterResponse, WorkerRegisterResponse) +def _union_member_names(annotation: Any) -> set[str]: + """Names of the members of a ``list[A | B]`` field annotation.""" + list_args = annotation.__args__[0] + return {t.__name__ for t in list_args.__args__} + + +def test_api_result_items_union_matches() -> None: + """APIResult.items must be the same union on both sides; the field-name + check alone misses a type drift (e.g. SDK-only list[APIItem]).""" + srv_items = srv_results.APIResult.model_fields["items"].annotation + sdk_items = sdk_models.APIResult.model_fields["items"].annotation + assert _union_member_names(srv_items) == _union_member_names(sdk_items) + + # ------------------------------------------------------------------ # # Enum compatibility tests # ------------------------------------------------------------------ # From 13a42cf51a0e6f52484c1cea046af93539e9a7f9 Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Fri, 25 Sep 2026 06:34:19 +0700 Subject: [PATCH 39/71] fix(worker): a dataframe over zero rows runs zero rows Signed-off-by: Zhengyuan Su Co-Authored-By: Claude Code --- src/worker/executors/api_executor.py | 4 +- src/worker/executors/mixins/data.py | 2 +- src/worker/executors/utils/graph_templates.py | 2 +- tests/worker/test_api_executor.py | 179 ++++++++++++++++++ tests/worker/test_data_mixin_lineage.py | 38 ++++ 5 files changed, 221 insertions(+), 4 deletions(-) diff --git a/src/worker/executors/api_executor.py b/src/worker/executors/api_executor.py index bfee64ee7..40323ab2b 100644 --- a/src/worker/executors/api_executor.py +++ b/src/worker/executors/api_executor.py @@ -393,7 +393,7 @@ def _run(self, task: ExecutorTask, out_dir: Path) -> APIResult: entry = self._collect_prompts_for_spec(spec, task_id=task.task_id) prompts = entry.prompts - if not prompts: + if not prompts and not entry.tables: raise ExecutionError("spec.data produced no rows") request_kwargs = self._build_request_kwargs(api_cfg, None) @@ -478,7 +478,7 @@ def _issue(idx: int, prompt: Any) -> APIItem: if result_items: first = result_items[0] if isinstance(first, APIGroupItem): - status_code = first.rows[0].status_code + status_code = first.rows[0].status_code if first.rows else 0 truncated = any( r.truncated for g in result_items diff --git a/src/worker/executors/mixins/data.py b/src/worker/executors/mixins/data.py index 5679e55a5..6f4660bce 100644 --- a/src/worker/executors/mixins/data.py +++ b/src/worker/executors/mixins/data.py @@ -612,7 +612,7 @@ def _collect_prompts_for_spec( table_stores_list = [] for group_idx in range(group_count): - max_len = 1 + max_len = 0 raw_group_values: dict[str, list[Any]] = {} for label, groups in grouped_columns.items(): values = groups[group_idx] diff --git a/src/worker/executors/utils/graph_templates.py b/src/worker/executors/utils/graph_templates.py index 008701fa7..fafbab3ad 100644 --- a/src/worker/executors/utils/graph_templates.py +++ b/src/worker/executors/utils/graph_templates.py @@ -215,7 +215,7 @@ def _build_grouped_dataframes(columns: list[dict[str, Any]]) -> list[pd.DataFram dataframes: list[pd.DataFrame] = [] for group_idx in range(group_count): - max_len = 1 + max_len = 0 raw_values: dict[str, list[Any]] = {} for label, groups in grouped_columns.items(): values = groups[group_idx] diff --git a/tests/worker/test_api_executor.py b/tests/worker/test_api_executor.py index 2e2c925d7..fe514df06 100644 --- a/tests/worker/test_api_executor.py +++ b/tests/worker/test_api_executor.py @@ -1290,6 +1290,185 @@ def test_dataframe_column_issues_one_request_per_row(self, tmp_path: Path) -> No [{"role": "user", "content": "row c1"}], ] + def test_all_empty_columns_issue_no_request_and_one_empty_group( + self, tmp_path: Path + ) -> None: + """A dataframe whose columns all resolve to zero rows runs zero rows: + no HTTP request, and one empty group item so downstream paths resolve.""" + upstream = APIResult( + ok=True, + executor="api", + method="POST", + url="https://up.example.com", + status_code=200, + items=[], + ) + payload = { + "task_id": "task-api-df-empty", + "workflow_id": "wf-1", + "owner_id": "owner", + "assigned_worker": "worker-1", + "dispatched_at": "2026-03-22T00:00:00Z", + "task": { + "apiVersion": "flowmesh/v1", + "kind": "Task", + "metadata": {"name": "wf:api"}, + "spec": { + "taskType": "api", + "_upstreamResults": {"Up": upstream}, + "api": { + "method": "POST", + "url": "https://custom.example.com/v1/chat/completions", + "json": {"messages": "{{prompt}}"}, + }, + "data": { + "type": "dataframe", + "columns": [ + { + "label": "L", + "node": "Up", + "path": "items.json.choices[0].message.content", + } + ], + "messages": [ + {"role": "user", "content": "row {L}"}, + ], + }, + }, + }, + } + task = WorkerTaskMessage.model_validate(payload) + transport = _EchoTransport() + result = _run(_executor(), task, transport, tmp_path) + assert transport.requests == [] + assert len(result.items) == 1 + assert isinstance(result.items[0], APIGroupItem) + assert result.items[0].index == 0 + assert result.items[0].rows == [] + assert result.status_code == 0 + + def test_zero_vs_three_rows_still_raises(self, tmp_path: Path) -> None: + """A real mismatch (one column empty, another with rows) still raises.""" + empty = APIResult( + ok=True, + executor="api", + method="POST", + url="https://up.example.com", + status_code=200, + items=[], + ) + full = APIResult( + ok=True, + executor="api", + method="POST", + url="https://up.example.com", + status_code=200, + items=[_api_item("c0"), _api_item("c1"), _api_item("c2")], + ) + payload = { + "task_id": "task-api-df-mismatch", + "workflow_id": "wf-1", + "owner_id": "owner", + "assigned_worker": "worker-1", + "dispatched_at": "2026-03-22T00:00:00Z", + "task": { + "apiVersion": "flowmesh/v1", + "kind": "Task", + "metadata": {"name": "wf:api"}, + "spec": { + "taskType": "api", + "_upstreamResults": {"Empty": empty, "Full": full}, + "api": { + "method": "POST", + "url": "https://custom.example.com/v1/chat/completions", + "json": {"messages": "{{prompt}}"}, + }, + "data": { + "type": "dataframe", + "columns": [ + { + "label": "Empty", + "node": "Empty", + "path": "items.json.choices[0].message.content", + }, + { + "label": "Full", + "node": "Full", + "path": "items.json.choices[0].message.content", + }, + ], + "messages": [ + {"role": "user", "content": "row {Empty} {Full}"}, + ], + }, + }, + }, + } + task = WorkerTaskMessage.model_validate(payload) + with pytest.raises(ExecutionError, match="same number of rows"): + _run(_executor(), task, _EchoTransport(), tmp_path) + + def test_two_vs_three_rows_still_raises(self, tmp_path: Path) -> None: + """A ragged mismatch (2 vs 3 rows) still raises.""" + two = APIResult( + ok=True, + executor="api", + method="POST", + url="https://up.example.com", + status_code=200, + items=[_api_item("c0"), _api_item("c1")], + ) + three = APIResult( + ok=True, + executor="api", + method="POST", + url="https://up.example.com", + status_code=200, + items=[_api_item("c0"), _api_item("c1"), _api_item("c2")], + ) + payload = { + "task_id": "task-api-df-ragged", + "workflow_id": "wf-1", + "owner_id": "owner", + "assigned_worker": "worker-1", + "dispatched_at": "2026-03-22T00:00:00Z", + "task": { + "apiVersion": "flowmesh/v1", + "kind": "Task", + "metadata": {"name": "wf:api"}, + "spec": { + "taskType": "api", + "_upstreamResults": {"Two": two, "Three": three}, + "api": { + "method": "POST", + "url": "https://custom.example.com/v1/chat/completions", + "json": {"messages": "{{prompt}}"}, + }, + "data": { + "type": "dataframe", + "columns": [ + { + "label": "A", + "node": "Two", + "path": "items.json.choices[0].message.content", + }, + { + "label": "B", + "node": "Three", + "path": "items.json.choices[0].message.content", + }, + ], + "messages": [ + {"role": "user", "content": "row {A} {B}"}, + ], + }, + }, + }, + } + task = WorkerTaskMessage.model_validate(payload) + with pytest.raises(ExecutionError, match="same number of rows"): + _run(_executor(), task, _EchoTransport(), tmp_path) + class TestGraphTemplateAggregate: def test_graph_template_aggregates_all_rows_into_one_prompt( diff --git a/tests/worker/test_data_mixin_lineage.py b/tests/worker/test_data_mixin_lineage.py index 80a47f0c1..419831203 100644 --- a/tests/worker/test_data_mixin_lineage.py +++ b/tests/worker/test_data_mixin_lineage.py @@ -470,3 +470,41 @@ def _fake_get( assert len(entry.images) == 1 assert entry.images[0] is not None assert entry.images[0].size == (2, 2) + + +def test_collect_prompts_all_empty_dataframe_columns_yield_zero_rows() -> None: + """A dataframe whose columns all resolve to zero values yields an empty + DataFrame (zero rows) with the column labels, not a row mismatch.""" + mixin = _Mixin() + upstream = BaseExecutorResult.model_validate( + { + "items": [], + "count": 0, + } + ) + spec = cast( + Any, + SimpleNamespace( + data={ + "type": "dataframe", + "columns": [ + {"label": "text", "node": "Up", "path": "items.output.text"}, + { + "label": "statement", + "node": "Up", + "path": "items.output.statement", + }, + ], + "messages": [{"role": "user", "content": "row {text} {statement}"}], + }, + inference={}, + upstreamResults={"Up": upstream}, + ), + ) + + entry = mixin._collect_prompts_for_spec(spec, "tsk-df-empty") + + assert entry.prompts == [] + assert len(entry.tables) == 1 + assert entry.tables[0].empty + assert list(entry.tables[0].columns) == ["text", "statement"] From f86ee964b86c86f0b8cbefa7070c8134e9c6ac4e Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Fri, 25 Sep 2026 06:37:45 +0700 Subject: [PATCH 40/71] fix(worker): expose ValueError to sandboxed functions Signed-off-by: Zhengyuan Su Co-Authored-By: Claude Code --- src/worker/executors/utils/safe_eval.py | 1 + tests/worker/test_safe_eval.py | 14 ++++++++++++++ 2 files changed, 15 insertions(+) diff --git a/src/worker/executors/utils/safe_eval.py b/src/worker/executors/utils/safe_eval.py index 0f976ba96..f181d4a4c 100644 --- a/src/worker/executors/utils/safe_eval.py +++ b/src/worker/executors/utils/safe_eval.py @@ -56,6 +56,7 @@ "all": all, "range": range, "isinstance": isinstance, + "ValueError": ValueError, } # Whitelist of safe modules available during function execution. diff --git a/tests/worker/test_safe_eval.py b/tests/worker/test_safe_eval.py index 94fd5a03b..05c19fc87 100644 --- a/tests/worker/test_safe_eval.py +++ b/tests/worker/test_safe_eval.py @@ -100,3 +100,17 @@ def test_assign_of_non_lambda_raises(self) -> None: def test_syntax_error_raises_runtime_error(self) -> None: with pytest.raises(RuntimeError, match="Function definition failed"): _run("def f(args):\n return )", (1,)) + + +class TestValueErrorPropagation: + def test_value_error_message_survives_sandbox(self) -> None: + """A function raising ValueError fails closed with its message intact.""" + with pytest.raises( + RuntimeError, + match="kept ids not in the candidate table: \\['f6'\\]", + ): + _run( + "def f(args):\n" + " raise ValueError(\"kept ids not in the candidate table: ['f6']\")", + (), + ) From 657ba7b36b3e3a3a4ac6f726293760396fa6b576 Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Fri, 25 Sep 2026 06:49:39 +0700 Subject: [PATCH 41/71] fix(worker): report status from first row across grouped results Signed-off-by: Zhengyuan Su Co-Authored-By: Claude Code --- src/worker/executors/api_executor.py | 10 +++- tests/worker/test_api_executor.py | 59 +++++++++++++++++++++++ tests/worker/test_graph_templates_expr.py | 20 +++++++- 3 files changed, 87 insertions(+), 2 deletions(-) diff --git a/src/worker/executors/api_executor.py b/src/worker/executors/api_executor.py index 40323ab2b..1179a142e 100644 --- a/src/worker/executors/api_executor.py +++ b/src/worker/executors/api_executor.py @@ -478,7 +478,15 @@ def _issue(idx: int, prompt: Any) -> APIItem: if result_items: first = result_items[0] if isinstance(first, APIGroupItem): - status_code = first.rows[0].status_code if first.rows else 0 + status_code = next( + ( + r.status_code + for g in result_items + if isinstance(g, APIGroupItem) + for r in g.rows + ), + 0, + ) truncated = any( r.truncated for g in result_items diff --git a/tests/worker/test_api_executor.py b/tests/worker/test_api_executor.py index fe514df06..faaae5fec 100644 --- a/tests/worker/test_api_executor.py +++ b/tests/worker/test_api_executor.py @@ -1612,3 +1612,62 @@ def test_ragged_groups_return_one_item_per_group(self, tmp_path: Path) -> None: group0 = {json.loads(r.prompt)[0]["content"] for r in result.items[0].rows} assert group0 == {"row c0", "row c1"} assert json.loads(result.items[1].rows[0].prompt)[0]["content"] == "row c2" + + def test_status_code_taken_from_first_row_across_groups( + self, tmp_path: Path + ) -> None: + """A leading empty group must not zero the result status; the first + row across all groups supplies it.""" + upstream = APIResult( + ok=True, + executor="api", + method="POST", + url="https://up.example.com", + status_code=200, + items=[ + APIGroupItem(index=0, rows=[]), + APIGroupItem(index=1, rows=[_api_item("c0")]), + ], + ) + payload = { + "task_id": "task-api-grp-leading-empty", + "workflow_id": "wf-1", + "owner_id": "owner", + "assigned_worker": "worker-1", + "dispatched_at": "2026-03-22T00:00:00Z", + "task": { + "apiVersion": "flowmesh/v1", + "kind": "Task", + "metadata": {"name": "wf:api"}, + "spec": { + "taskType": "api", + "_upstreamResults": {"Up": upstream}, + "api": { + "method": "POST", + "url": "https://custom.example.com/v1/chat/completions", + "json": {"messages": "{{prompt}}"}, + }, + "data": { + "type": "dataframe", + "columns": [ + { + "label": "L", + "node": "Up", + "path": "items.rows.json.choices[0].message.content", + } + ], + "messages": [ + {"role": "user", "content": "row {L}"}, + ], + }, + }, + }, + } + task = WorkerTaskMessage.model_validate(payload) + transport = _EchoTransport() + result = _run(_executor(), task, transport, tmp_path) + assert len(transport.requests) == 1 + assert len(result.items) == 2 + assert result.items[0].rows == [] + assert len(result.items[1].rows) == 1 + assert result.status_code == 200 diff --git a/tests/worker/test_graph_templates_expr.py b/tests/worker/test_graph_templates_expr.py index d50c7b1a5..50b1d0028 100644 --- a/tests/worker/test_graph_templates_expr.py +++ b/tests/worker/test_graph_templates_expr.py @@ -1,8 +1,13 @@ """Tests for _evaluate_expr attribute/index resolution over pydantic models, aliases, and nested lists.""" +import pandas as pd + from shared.schemas.result import APIGroupItem, APIItem, APIResult -from worker.executors.utils.graph_templates import _evaluate_expr +from worker.executors.utils.graph_templates import ( + _build_grouped_dataframes, + _evaluate_expr, +) def _item(content: str) -> APIItem: @@ -44,3 +49,16 @@ def test_expr_over_nested_lists_of_models() -> None: "Up.items.rows.json.choices[0].message.content", {"Up": upstream} ) assert value == [["c0", "c1"], ["c2"]] + + +def test_build_grouped_dataframes_all_empty_columns_yield_zero_rows() -> None: + """All-empty grouped columns yield a zero-row DataFrame, not a mismatch.""" + columns = [ + {"label": "text", "value": []}, + {"label": "statement", "value": []}, + ] + dataframes = _build_grouped_dataframes(columns) + assert len(dataframes) == 1 + assert dataframes[0].empty + assert list(dataframes[0].columns) == ["text", "statement"] + assert isinstance(dataframes[0], pd.DataFrame) From bb949c8223a674e8a494464e1f938c6af40823b9 Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Fri, 25 Sep 2026 09:07:17 +0700 Subject: [PATCH 42/71] fix(worker): decide dataframe grouping from upstream structure, not cell shape A column whose per-row value is a list was misread as grouped values because grouping was inferred from the shape of the resolved value (a non-empty list of lists). The resolver now returns whether the value is grouped, decided from the upstream structure (APIGroupItem.rows or nested lists), and the dataframe and graph-template grouping paths use that flag instead of the shape heuristic. A per-row list stays a cell value. Signed-off-by: Zhengyuan Su Co-Authored-By: Claude Code --- src/worker/executors/echo_executor.py | 2 +- src/worker/executors/mixins/data.py | 19 ++- src/worker/executors/utils/graph_templates.py | 78 ++++++---- .../test_data_mixin_dataframe_grouping.py | 144 ++++++++++++++++++ tests/worker/test_graph_templates_expr.py | 8 +- 5 files changed, 211 insertions(+), 40 deletions(-) create mode 100644 tests/worker/test_data_mixin_dataframe_grouping.py diff --git a/src/worker/executors/echo_executor.py b/src/worker/executors/echo_executor.py index edc700e60..7bbc5f5e9 100644 --- a/src/worker/executors/echo_executor.py +++ b/src/worker/executors/echo_executor.py @@ -45,7 +45,7 @@ def _resolve_expr_item( "echo executor mapping item must contain either 'expr' or " "both 'node' and 'path'" ) - resolved = _evaluate_expr(expr.strip(), context) + resolved, _ = _evaluate_expr(expr.strip(), context) if resolved is None: raise ExecutionError( f"echo executor expression resolved to null: '{expr.strip()}'" diff --git a/src/worker/executors/mixins/data.py b/src/worker/executors/mixins/data.py index 6f4660bce..bf16a9026 100644 --- a/src/worker/executors/mixins/data.py +++ b/src/worker/executors/mixins/data.py @@ -429,7 +429,7 @@ def _collect_prompts_for_spec( if expr: context = self._spec_upstream_results(spec) resolved_expr = expr.strip() - items = _evaluate_expr(resolved_expr, context) + items, _ = _evaluate_expr(resolved_expr, context) root_node = resolved_expr.split(".", 1)[0] or None if not isinstance(items, list): raise ExecutionError( @@ -589,16 +589,15 @@ def _collect_prompts_for_spec( for column in resolved_columns: label = column["label"] value = column["value"] - if ( - isinstance(value, list) - and value - and all(isinstance(v, list) for v in value) - ): + if column.get("grouped"): + if not isinstance(value, list): + raise ExecutionError( + f"Column '{label}' is grouped but did not resolve " + "to a list." + ) groups = value - elif isinstance(value, list): - groups = [value] else: - groups = [[value]] + groups = [value] grouped_columns[label] = groups group_count = max(len(groups) for groups in grouped_columns.values()) @@ -680,7 +679,7 @@ def _collect_prompts_for_spec( resolved_node = node_hint if not resolved_node and isinstance(expr, str): resolved_node = expr.split(".", 1)[0].strip() or None - image_embedding_spec: Any = _evaluate_expr(expr.strip(), context) + image_embedding_spec: Any = _evaluate_expr(expr.strip(), context)[0] artifact_source = maybe_resolve_artifact_ref( image_embedding_spec, context, resolved_node ) diff --git a/src/worker/executors/utils/graph_templates.py b/src/worker/executors/utils/graph_templates.py index fafbab3ad..cb9f6bccb 100644 --- a/src/worker/executors/utils/graph_templates.py +++ b/src/worker/executors/utils/graph_templates.py @@ -7,7 +7,7 @@ import pandas as pd from pydantic import BaseModel -from shared.schemas.result import BaseExecutorResult +from shared.schemas.result import APIGroupItem, BaseExecutorResult from shared.tasks.specs import TaskSpecStrictBase from shared.utils.json import validate_keys @@ -141,10 +141,11 @@ def _resolve_columns( if expr: assert data is None - value = _evaluate_expr(expr.strip(), context) + value, grouped = _evaluate_expr(expr.strip(), context) if value is None: if "default" in raw: value = raw.get("default") + grouped = False else: raise ExecutionError( f"Column '{label}' expression '{expr}' resolved to null." @@ -162,10 +163,12 @@ def _resolve_columns( if not isinstance(items, list): raise ExecutionError("data.items must be a list.") value = items + grouped = False case "dataframe": nested_columns_cfg = data.get("columns") nested_columns = _resolve_columns(nested_columns_cfg, context) value = _build_grouped_dataframes(nested_columns) + grouped = True case _: raise ExecutionError(f"Unsupported column 'data' type: {dtype}") else: @@ -178,6 +181,7 @@ def _resolve_columns( "label": label, "value": value, "expr": expr, + "grouped": grouped, } ) @@ -192,16 +196,14 @@ def _build_grouped_dataframes(columns: list[dict[str, Any]]) -> list[pd.DataFram for column in columns: label = column["label"] value = column["value"] - if ( - isinstance(value, list) - and value - and all(isinstance(v, list) for v in value) - ): + if column.get("grouped"): + if not isinstance(value, list): + raise ExecutionError( + f"Column '{label}' is grouped but did not resolve to a list." + ) groups = value - elif isinstance(value, list): - groups = [value] else: - groups = [[value]] + groups = [value] grouped_columns[label] = groups group_count = max(len(groups) for groups in grouped_columns.values()) @@ -586,58 +588,80 @@ def _format_column_line(label: str, value: str) -> str: return f"• {label}: {indented}" -def _evaluate_expr(expr: str, context: dict[str, BaseExecutorResult]) -> Any: +def _evaluate_expr( + expr: str, context: dict[str, BaseExecutorResult] +) -> tuple[Any, bool]: + """Resolve an expression against upstream results. + + Returns ``(value, grouped)``. ``grouped`` is True only when the resolved + value is a list of groups, decided from the upstream structure (a list of + ``APIGroupItem.rows``, or nested lists) — never from the shape of the cell + values. A per-row list is a cell value, not a group. + """ if not expr: - return None + return None, False parts = expr.split(".") root = parts[0] result = context.get(root) if result is None: - return None + return None, False value: Any = result + grouped = False for token in parts[1:]: if not token: continue attr, indexes = _split_indexes(token) if attr: - value = _apply_attr(value, attr, token, parts) + value, grouped = _apply_attr(value, attr, token, parts, grouped) for idx in indexes: - value = _apply_index(value, idx, token) + value, grouped = _apply_index(value, idx, token, grouped) # Attempt to deserialize DataFrame if applicable if isinstance(value, dict): value = try_deserialize_dataframe(value) elif isinstance(value, list) and all(isinstance(v, dict) for v in value): value = [try_deserialize_dataframe(v) for v in value] - return value + return value, grouped -def _apply_attr(value: Any, attr: str, token: str, parts: list[str]) -> Any: +def _apply_attr( + value: Any, attr: str, token: str, parts: list[str], grouped: bool +) -> tuple[Any, bool]: """Resolve an attribute access, mapping over lists of dicts, DataFrames, or pydantic models (including nested lists).""" if isinstance(value, dict) and attr in value: - return value[attr] + return value[attr], grouped if isinstance(value, list): if all(isinstance(v, dict) and attr in v for v in value): - return [v[attr] for v in value] + return [v[attr] for v in value], grouped if all(isinstance(v, pd.DataFrame) for v in value): if any(attr not in v.columns for v in value): raise ExecutionError( f"{attr} not a valid column in one of the " f"DataFrames for {token}." ) - return [v[attr].tolist() for v in value] + return [v[attr].tolist() for v in value], grouped if all(isinstance(v, BaseModel) for v in value): - return [_model_attr(v, attr, token) for v in value] + is_grouped = attr == "rows" and all( + isinstance(v, APIGroupItem) for v in value + ) + return ( + [_model_attr(v, attr, token) for v in value], + grouped or is_grouped, + ) if all(isinstance(v, list) for v in value): - return [_apply_attr(v, attr, token, parts) for v in value] + mapped: list[Any] = [] + for v in value: + inner, _ = _apply_attr(v, attr, token, parts, grouped) + mapped.append(inner) + return mapped, grouped if isinstance(value, pd.DataFrame): if attr not in value.columns: raise ExecutionError(f"{attr} not a valid column in DataFrame for {token}.") - return value[attr].tolist() + return value[attr].tolist(), grouped if isinstance(value, BaseModel): - return _model_attr(value, attr, token) + return _model_attr(value, attr, token), grouped raise ExecutionError( f"{attr} in {parts} is not a valid key - " f"{type(value).__name__}, {value}" ) @@ -660,12 +684,12 @@ def _model_attr(value: BaseModel, attr: str, token: str) -> Any: ) -def _apply_index(value: Any, idx: int, token: str) -> Any: +def _apply_index(value: Any, idx: int, token: str, grouped: bool) -> tuple[Any, bool]: """Index into a list, mapping over a list of lists (one per group).""" if isinstance(value, list) and all(isinstance(v, list) for v in value): - return [_apply_index(v, idx, token) for v in value] + return [_apply_index(v, idx, token, grouped)[0] for v in value], grouped if isinstance(value, list) and -len(value) <= idx < len(value): - return value[idx] + return value[idx], grouped raise ExecutionError(f"{idx} not a valid index in {token} - {len(value)}") diff --git a/tests/worker/test_data_mixin_dataframe_grouping.py b/tests/worker/test_data_mixin_dataframe_grouping.py new file mode 100644 index 000000000..39f13f86f --- /dev/null +++ b/tests/worker/test_data_mixin_dataframe_grouping.py @@ -0,0 +1,144 @@ +"""DataMixin dataframe grouping: a per-row list column is a cell value, not a +group. Grouping is decided from the upstream structure (APIGroupItem.rows), +never from the shape of the cell values (FM-7).""" + +from types import SimpleNamespace +from typing import Any, cast + +from shared.schemas.result import APIGroupItem, APIItem, APIResult, BaseExecutorResult +from worker.executors.mixins.data import DataMixin + + +class _Mixin(DataMixin): + """Bare-bones DataMixin instance for unit testing.""" + + +def _row(output: dict[str, Any]) -> APIItem: + item = APIItem(index=0, url="u", status_code=200) + item.response_json = {"output": output} + return item + + +def _plain_upstream(rows: list[dict[str, Any]]) -> BaseExecutorResult: + """A non-grouped upstream result: one row dict per item, each carrying an + ``output`` mapping (the live GatherPremises shape).""" + return BaseExecutorResult.model_validate( + { + "items": [{"output": r} for r in rows], + "count": len(rows), + } + ) + + +def _grouped_upstream(groups: list[list[dict[str, Any]]]) -> APIResult: + """A grouped upstream api result: one APIGroupItem per group.""" + return APIResult( + ok=True, + executor="api", + method="POST", + url="https://up.example.com", + status_code=200, + items=[ + APIGroupItem(index=i, rows=[_row(r) for r in group]) + for i, group in enumerate(groups) + ], + ) + + +def _spec( + columns: list[dict[str, Any]], + upstream: Any, + content: str = "row {claim} {premises}", +) -> Any: + return cast( + Any, + SimpleNamespace( + data={ + "type": "dataframe", + "columns": columns, + "messages": [{"role": "user", "content": content}], + }, + inference={}, + upstreamResults={"Up": upstream}, + ), + ) + + +def _collect( + columns: list[dict[str, Any]], + upstream: Any, + content: str = "row {claim} {premises}", +): + return _Mixin()._collect_prompts_for_spec( + _spec(columns, upstream, content), "tsk-fm7" + ) + + +def test_per_row_list_column_is_one_group_with_lists_intact() -> None: + """A column whose per-row value is a list (mixed lengths, incl. empty) is a + single group with one row per upstream row, the list intact in each cell.""" + upstream = _plain_upstream( + [ + {"claim": "c0", "premises": ["p0"]}, + {"claim": "c1", "premises": []}, + {"claim": "c2", "premises": ["p2"]}, + ] + ) + columns = [ + {"label": "claim", "node": "Up", "path": "items.output.claim"}, + {"label": "premises", "node": "Up", "path": "items.output.premises"}, + ] + + entry = _collect(columns, upstream) + + assert len(entry.tables) == 1 + df = entry.tables[0] + assert list(df.columns) == ["claim", "premises"] + assert len(df) == 3 + assert df["claim"].tolist() == ["c0", "c1", "c2"] + assert df["premises"].tolist() == [["p0"], [], ["p2"]] + + +def test_per_row_single_member_list_is_not_repeated() -> None: + """When every per-row list has one member, the column is still one group of + three rows, not three groups each broadcast to all claims.""" + upstream = _plain_upstream( + [ + {"claim": "c0", "premises": ["x0"]}, + {"claim": "c1", "premises": ["x1"]}, + {"claim": "c2", "premises": ["x2"]}, + ] + ) + columns = [ + {"label": "claim", "node": "Up", "path": "items.output.claim"}, + {"label": "premises", "node": "Up", "path": "items.output.premises"}, + ] + + entry = _collect(columns, upstream) + + assert len(entry.tables) == 1 + df = entry.tables[0] + assert len(df) == 3 + assert df["claim"].tolist() == ["c0", "c1", "c2"] + assert df["premises"].tolist() == [["x0"], ["x1"], ["x2"]] + + +def test_genuinely_grouped_upstream_still_groups() -> None: + """A genuinely grouped upstream result (APIGroupItem.rows) still produces + one table per group.""" + upstream = _grouped_upstream( + [ + [{"claim": "c0"}, {"claim": "c1"}], + [{"claim": "c2"}], + ] + ) + columns = [ + {"label": "claim", "node": "Up", "path": "items.rows.json.output.claim"}, + ] + + entry = _collect(columns, upstream, content="row {claim}") + + assert len(entry.tables) == 2 + assert [len(df) for df in entry.tables] == [2, 1] + assert entry.tables[0]["claim"].tolist() == ["c0", "c1"] + assert entry.tables[1]["claim"].tolist() == ["c2"] diff --git a/tests/worker/test_graph_templates_expr.py b/tests/worker/test_graph_templates_expr.py index 50b1d0028..3e999622a 100644 --- a/tests/worker/test_graph_templates_expr.py +++ b/tests/worker/test_graph_templates_expr.py @@ -27,8 +27,11 @@ def test_expr_over_list_of_models_with_alias() -> None: status_code=200, items=[_item("c0"), _item("c1")], ) - value = _evaluate_expr("Up.items.json.choices[0].message.content", {"Up": upstream}) + value, grouped = _evaluate_expr( + "Up.items.json.choices[0].message.content", {"Up": upstream} + ) assert value == ["c0", "c1"] + assert grouped is False def test_expr_over_nested_lists_of_models() -> None: @@ -45,10 +48,11 @@ def test_expr_over_nested_lists_of_models() -> None: APIGroupItem(index=1, rows=[_item("c2")]), ], ) - value = _evaluate_expr( + value, grouped = _evaluate_expr( "Up.items.rows.json.choices[0].message.content", {"Up": upstream} ) assert value == [["c0", "c1"], ["c2"]] + assert grouped is True def test_build_grouped_dataframes_all_empty_columns_yield_zero_rows() -> None: From 5ad3ce1a28e689a98086ceeae291724967cc0dae Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Fri, 25 Sep 2026 09:40:18 +0700 Subject: [PATCH 43/71] fix(worker): group dataframe columns over list-mode lambda output An attribute mapped over a list of lists of records (a list-mode lambda upstream whose items' output is a list of records) is now grouped: each inner list is one group. A per-row list cell inside a record stays a single ungrouped column, and index mapping over a per-item list (e.g. items.json.choices[0]) does not group. Signed-off-by: Zhengyuan Su Co-Authored-By: Claude Code --- src/worker/executors/utils/graph_templates.py | 6 +- .../test_data_mixin_dataframe_grouping.py | 75 +++++++++++++++++++ 2 files changed, 80 insertions(+), 1 deletion(-) diff --git a/src/worker/executors/utils/graph_templates.py b/src/worker/executors/utils/graph_templates.py index cb9f6bccb..4e0ff8af2 100644 --- a/src/worker/executors/utils/graph_templates.py +++ b/src/worker/executors/utils/graph_templates.py @@ -651,11 +651,15 @@ def _apply_attr( grouped or is_grouped, ) if all(isinstance(v, list) for v in value): + # Each inner list of records is one group; a per-row list cell is not. + is_grouped = bool(value) and all( + all(isinstance(r, (dict, BaseModel)) for r in v) for v in value + ) mapped: list[Any] = [] for v in value: inner, _ = _apply_attr(v, attr, token, parts, grouped) mapped.append(inner) - return mapped, grouped + return mapped, grouped or is_grouped if isinstance(value, pd.DataFrame): if attr not in value.columns: raise ExecutionError(f"{attr} not a valid column in DataFrame for {token}.") diff --git a/tests/worker/test_data_mixin_dataframe_grouping.py b/tests/worker/test_data_mixin_dataframe_grouping.py index 39f13f86f..19c69b65c 100644 --- a/tests/worker/test_data_mixin_dataframe_grouping.py +++ b/tests/worker/test_data_mixin_dataframe_grouping.py @@ -142,3 +142,78 @@ def test_genuinely_grouped_upstream_still_groups() -> None: assert [len(df) for df in entry.tables] == [2, 1] assert entry.tables[0]["claim"].tolist() == ["c0", "c1"] assert entry.tables[1]["claim"].tolist() == ["c2"] + + +def _list_mode_lambda_upstream( + groups: list[list[dict[str, Any]]], +) -> BaseExecutorResult: + """A list-mode lambda upstream: one item per group, each item's ``output`` + is itself a list of records (the live PairSources shape).""" + return BaseExecutorResult.model_validate( + { + "items": [{"output": group} for group in groups], + "count": len(groups), + } + ) + + +def test_list_mode_lambda_output_groups_by_item() -> None: + """A list-mode lambda upstream whose items' output is a list of records, + read as ``items.output.`` columns, gives one group per item with the + right row counts (uneven group sizes included).""" + upstream = _list_mode_lambda_upstream( + [ + [{"claim": "c0", "src_text": "s0"}, {"claim": "c0", "src_text": "s1"}], + [{"claim": "c1", "src_text": "s2"}], + [ + {"claim": "c2", "src_text": "s3"}, + {"claim": "c2", "src_text": "s4"}, + {"claim": "c2", "src_text": "s5"}, + ], + ] + ) + columns = [ + {"label": "claim", "node": "Up", "path": "items.output.claim"}, + {"label": "src_text", "node": "Up", "path": "items.output.src_text"}, + ] + + entry = _collect(columns, upstream, content="row {claim} {src_text}") + + assert len(entry.tables) == 3 + assert [len(df) for df in entry.tables] == [2, 1, 3] + assert entry.tables[0]["claim"].tolist() == ["c0", "c0"] + assert entry.tables[0]["src_text"].tolist() == ["s0", "s1"] + assert entry.tables[1]["claim"].tolist() == ["c1"] + assert entry.tables[2]["claim"].tolist() == ["c2", "c2", "c2"] + + +def test_index_mapping_over_nested_list_stays_one_group() -> None: + """``items.json.choices[0].message.content`` over a non-grouped api result + stays one group: the index mapping over a per-item list is not grouping.""" + items: list[APIItem | APIGroupItem] = [] + for i in range(3): + item = APIItem(index=i, url="u", status_code=200) + item.response_json = {"choices": [{"message": {"content": f"c{i}"}}]} + items.append(item) + upstream = APIResult( + ok=True, + executor="api", + method="POST", + url="https://up.example.com", + status_code=200, + items=items, + ) + columns = [ + { + "label": "content", + "node": "Up", + "path": "items.json.choices[0].message.content", + }, + ] + + entry = _collect(columns, upstream, content="row {content}") + + assert len(entry.tables) == 1 + df = entry.tables[0] + assert len(df) == 3 + assert df["content"].tolist() == ["c0", "c1", "c2"] From fc4a960735aa83f96e6e4ffa5474dd14c5b07f0b Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Fri, 25 Sep 2026 14:12:15 +0700 Subject: [PATCH 44/71] test(worker): assert APIResult dumps grouped rows under the json alias A plain model_dump of an APIResult must emit each item's payload under the json wire key, never response_json, including rows nested inside an APIGroupItem, and the aliased payload must revalidate on the server-side result model. Signed-off-by: Zhengyuan Su Co-Authored-By: Claude Code --- tests/shared/test_executor_result.py | 54 ++++++++++++++++++++++++++++ 1 file changed, 54 insertions(+) diff --git a/tests/shared/test_executor_result.py b/tests/shared/test_executor_result.py index 2df793b09..5e95c57e2 100644 --- a/tests/shared/test_executor_result.py +++ b/tests/shared/test_executor_result.py @@ -8,6 +8,7 @@ from shared.schemas.artifact import ArtifactContext, ArtifactRef from shared.schemas.result import ( + APIGroupItem, APIItem, APIResult, BaseExecutorResult, @@ -216,3 +217,56 @@ def test_api_item_round_trip_construct_serialize_validate() -> None: {"index": 1, "url": "u", "status_code": 200, "json": {"a": 1}} ) assert by_alias.response_json == {"a": 1} + + +def test_api_result_payload_uses_wire_alias_json() -> None: + """A plain model_dump emits ``json`` (the wire alias), never ``response_json``.""" + result = APIResult.model_validate( + { + "executor": "api", + "method": "POST", + "url": "http://example.com/v1/chat/completions", + "status_code": 200, + "items": [ + { + "index": 0, + "url": "http://example.com/v1/chat/completions", + "status_code": 200, + "json": {"choices": [{"message": {"content": "hello"}}]}, + "text": "hello", + }, + { + "index": 1, + "rows": [ + { + "index": 0, + "url": "http://example.com/v1/chat/completions", + "status_code": 200, + "json": {"choices": [{"message": {"content": "grouped"}}]}, + } + ], + }, + ], + } + ) + + payload = result.model_dump() + + assert "response_json" not in payload + assert payload["items"][0]["json"] == { + "choices": [{"message": {"content": "hello"}}] + } + assert payload["items"][1]["rows"][0]["json"] == { + "choices": [{"message": {"content": "grouped"}}] + } + + # The server-side result model validates the aliased payload back. + reloaded = APIResult.model_validate(payload) + first = reloaded.items[0] + assert isinstance(first, APIItem) + assert first.response_json["choices"][0]["message"]["content"] == "hello" + grouped = reloaded.items[1] + assert isinstance(grouped, APIGroupItem) + assert grouped.rows[0].response_json["choices"][0]["message"]["content"] == ( + "grouped" + ) From ae0be153818025290a5376970624d9b51ff83e8b Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Fri, 25 Sep 2026 15:10:03 +0700 Subject: [PATCH 45/71] test(sdk): assert grouped APIResult rows dump under the json alias A grouped APIResult must dump each row's payload under the json wire key, never response_json, so a downstream walk of items[].rows[].json succeeds. Signed-off-by: Zhengyuan Su Co-Authored-By: Claude Code --- tests/sdk/test_models.py | 28 ++++++++++++++++++++++++++++ 1 file changed, 28 insertions(+) diff --git a/tests/sdk/test_models.py b/tests/sdk/test_models.py index b2e5ab2cb..202876523 100644 --- a/tests/sdk/test_models.py +++ b/tests/sdk/test_models.py @@ -557,3 +557,31 @@ def test_grouped_api_result_round_trip(self) -> None: reloaded = APIResult.model_validate(wire) assert isinstance(reloaded.items[0], APIGroupItem) assert reloaded.items[0].rows[0].text == "hi" + + def test_grouped_dump_json_uses_json_alias(self) -> None: + """A grouped APIResult dumps each row's payload under the ``json`` wire + key, so a downstream walk of ``items[].rows[].json`` succeeds.""" + payload = { + "ok": True, + "executor": "api", + "method": "POST", + "url": "u", + "status_code": 200, + "items": [ + { + "index": 0, + "rows": [ + { + "index": 0, + "url": "u", + "status_code": 200, + "json": {"a": 1}, + } + ], + } + ], + } + result = APIResult.model_validate(payload) + row = result.model_dump(mode="json")["items"][0]["rows"][0] + assert row["json"] == {"a": 1} + assert "response_json" not in row From 16676b63469ccbdfe1c8e14eefcc08d640fc1f7d Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Fri, 25 Sep 2026 15:50:51 +0700 Subject: [PATCH 46/71] feat(worker): log per-call statistics in the API executor Adds an INFO line per finished call (task, row, attempts, status or exception, wall seconds, chat-completion tokens, finish reason, backend), a 60s heartbeat while calls are outstanding, and a summary line on task finish (calls, failures, retries, wall, latency p50/p95/max, summed tokens, per-backend counts). No request or response content is logged. Signed-off-by: Zhengyuan Su Co-Authored-By: Claude Code --- src/worker/executors/api_executor.py | 227 +++++++++++++++++++++++++-- tests/worker/test_api_executor.py | 106 +++++++++++++ 2 files changed, 319 insertions(+), 14 deletions(-) diff --git a/src/worker/executors/api_executor.py b/src/worker/executors/api_executor.py index 1179a142e..70175b75c 100644 --- a/src/worker/executors/api_executor.py +++ b/src/worker/executors/api_executor.py @@ -1,7 +1,9 @@ import json import logging +import math import os import threading +import time from concurrent.futures import ThreadPoolExecutor, as_completed from pathlib import Path from typing import Any, ClassVar @@ -39,6 +41,53 @@ def _is_retryable_status(status_code: int) -> bool: return status_code >= 500 or status_code in (408, 429) +def _fmt(value: Any) -> str: + """Render a log field, using ``-`` for a missing value.""" + return "-" if value is None else str(value) + + +def _percentile(values: list[float], pct: float) -> float: + """Nearest-rank percentile of a non-empty list of latencies.""" + if not values: + return 0.0 + ordered = sorted(values) + rank = max(1, math.ceil(pct / 100 * len(ordered))) + return ordered[rank - 1] + + +def _chat_completion_stats(body: Any) -> tuple[Any, Any, Any, Any, Any]: + """Extract token and finish-reason fields from a chat-completion body. + + Returns ``(prompt_tokens, completion_tokens, reasoning_tokens, + finish_reason, backend)`` with ``None`` for any missing field. Never + raises on an unexpected body shape. + """ + if not isinstance(body, dict): + return None, None, None, None, None + usage = body.get("usage") + prompt_tokens = completion_tokens = reasoning_tokens = None + if isinstance(usage, dict): + prompt_tokens = usage.get("prompt_tokens") + completion_tokens = usage.get("completion_tokens") + details = usage.get("completion_tokens_details") + if isinstance(details, dict): + reasoning_tokens = details.get("reasoning_tokens") + choices = body.get("choices") + finish_reason = None + if isinstance(choices, list) and choices and isinstance(choices[0], dict): + finish_reason = choices[0].get("finish_reason") + backend = body.get("provider") + if backend is None: + backend = body.get("system_fingerprint") + return ( + prompt_tokens, + completion_tokens, + reasoning_tokens, + finish_reason, + backend, + ) + + class APIExecutor(DataMixin, Executor): """Performs HTTP requests defined by task YAML. @@ -150,12 +199,13 @@ def _request_with_retries( params: dict[str, Any] | None, request_kwargs: dict[str, Any], retries: int, - ) -> httpx.Response: + ) -> tuple[httpx.Response, int]: """Issue the request, retrying transient failures up to ``retries`` times. A retryable failure is a connection error or a transient HTTP status (5xx, 408, 429). Non-retryable failures and a cancelled task stop the loop immediately. The final attempt's failure propagates to the caller. + Returns the response and the number of attempts used. """ attempt = 0 while True: @@ -180,7 +230,7 @@ def _request_with_retries( attempt += 1 self._wait_for_backoff() continue - return resp + return resp, attempt + 1 def _wait_for_backoff(self) -> None: """Wait out the retry backoff, aborting early if the task is cancelled.""" @@ -318,6 +368,42 @@ def _parse_response( return item, body_text + def _log_summary( + self, + task_id: str, + total: int, + failures: int, + total_retries: int, + wall: float, + latencies: list[float], + sum_prompt: int, + sum_completion: int, + sum_reasoning: int, + backend_counts: dict[str, int], + ) -> None: + """Log one summary line for a finished API task.""" + backends = ",".join( + f"{name}={count}" for name, count in sorted(backend_counts.items()) + ) + logger.info( + "api summary task=%s calls=%d failures=%d retries=%d wall=%.3fs " + "latency_p50=%.3fs latency_p95=%.3fs latency_max=%.3fs " + "prompt_tokens=%d completion_tokens=%d reasoning_tokens=%d " + "backends=%s", + task_id, + total, + failures, + total_retries, + wall, + _percentile(latencies, 50), + _percentile(latencies, 95), + max(latencies) if latencies else 0.0, + sum_prompt, + sum_completion, + sum_reasoning, + backends or "-", + ) + def run(self, task: ExecutorTask, out_dir: Path) -> APIResult: with self._cancel_lock: self._active_task_id = task.task_id @@ -398,13 +484,79 @@ def _run(self, task: ExecutorTask, out_dir: Path) -> APIResult: request_kwargs = self._build_request_kwargs(api_cfg, None) + total = len(prompts) + done = 0 + failures = 0 + total_retries = 0 + latencies: list[float] = [] + sum_prompt = 0 + sum_completion = 0 + sum_reasoning = 0 + backend_counts: dict[str, int] = {} + task_start = time.monotonic() + in_flight: dict[int, float] = {} + in_flight_lock = threading.Lock() + + def _record_call( + idx: int, + attempts: int, + status: Any, + start: float, + body: Any, + *, + failed: bool, + ) -> None: + nonlocal done, failures, total_retries + nonlocal sum_prompt, sum_completion, sum_reasoning + wall = time.monotonic() - start + with in_flight_lock: + in_flight.pop(idx, None) + done += 1 + if failed: + failures += 1 + total_retries += max(0, attempts - 1) + latencies.append(wall) + ( + prompt_tokens, + completion_tokens, + reasoning_tokens, + finish_reason, + backend, + ) = _chat_completion_stats(body) + if prompt_tokens is not None: + sum_prompt += prompt_tokens + if completion_tokens is not None: + sum_completion += completion_tokens + if reasoning_tokens is not None: + sum_reasoning += reasoning_tokens + if backend is not None: + backend_counts[backend] = backend_counts.get(backend, 0) + 1 + logger.info( + "api call task=%s row=%d attempts=%d status=%s wall=%.3fs " + "prompt_tokens=%s completion_tokens=%s reasoning_tokens=%s " + "finish_reason=%s backend=%s", + task.task_id, + idx, + attempts, + _fmt(status), + wall, + _fmt(prompt_tokens), + _fmt(completion_tokens), + _fmt(reasoning_tokens), + _fmt(finish_reason), + _fmt(backend), + ) + def _issue(idx: int, prompt: Any) -> APIItem: if self._cancel_event.is_set(): raise TaskCancelledError("API task cancelled") prompt_str = self._prompt_to_str(prompt) kwargs = self._substitute_prompt(request_kwargs, prompt) + start = time.monotonic() + with in_flight_lock: + in_flight[idx] = start try: - resp = self._request_with_retries( + resp, attempts = self._request_with_retries( client, method, str(url), @@ -414,6 +566,7 @@ def _issue(idx: int, prompt: Any) -> APIItem: retries, ) except httpx.RequestError as exc: + _record_call(idx, 0, exc.__class__.__name__, start, None, failed=True) raise ExecutionError( f"API request failed (row {idx}): {exc}", retryable=True ) from exc @@ -424,6 +577,7 @@ def _issue(idx: int, prompt: Any) -> APIItem: if body_text: message = f"{message}: {body_text}" retryable = resp.status_code >= 500 or resp.status_code in (408, 429) + _record_call(idx, attempts, resp.status_code, start, None, failed=True) raise ExecutionError(message, retryable=retryable) item, _ = self._parse_response( @@ -437,20 +591,65 @@ def _issue(idx: int, prompt: Any) -> APIItem: if self._cancel_event.is_set(): raise TaskCancelledError("API task cancelled") + _record_call( + idx, + attempts, + resp.status_code, + start, + item.response_json, + failed=False, + ) return item results: dict[int, APIItem] = {} - with ThreadPoolExecutor(max_workers=concurrency) as pool: - futures = {} - for idx, prompt in enumerate(prompts): - if self._cancel_event.is_set(): - raise TaskCancelledError("API task cancelled") - futures[pool.submit(_issue, idx, prompt)] = idx - for future in as_completed(futures): - idx = futures[future] - if self._cancel_event.is_set(): - raise TaskCancelledError("API task cancelled") - results[idx] = future.result() + heartbeat_stop = threading.Event() + + def _heartbeat() -> None: + while not heartbeat_stop.wait(60): + with in_flight_lock: + outstanding = dict(in_flight) + if not outstanding: + continue + oldest = min(outstanding.values()) + logger.info( + "api heartbeat task=%s done=%d/%d in_flight=%d oldest_age=%.0fs", + task.task_id, + done, + total, + len(outstanding), + time.monotonic() - oldest, + ) + + heartbeat = threading.Thread(target=_heartbeat, daemon=True) + heartbeat.start() + try: + with ThreadPoolExecutor(max_workers=concurrency) as pool: + futures = {} + for idx, prompt in enumerate(prompts): + if self._cancel_event.is_set(): + raise TaskCancelledError("API task cancelled") + futures[pool.submit(_issue, idx, prompt)] = idx + for future in as_completed(futures): + idx = futures[future] + if self._cancel_event.is_set(): + raise TaskCancelledError("API task cancelled") + results[idx] = future.result() + finally: + heartbeat_stop.set() + heartbeat.join(timeout=5) + wall = time.monotonic() - task_start + self._log_summary( + task.task_id, + total, + failures, + total_retries, + wall, + latencies, + sum_prompt, + sum_completion, + sum_reasoning, + backend_counts, + ) items = [results[idx] for idx in range(len(prompts))] diff --git a/tests/worker/test_api_executor.py b/tests/worker/test_api_executor.py index faaae5fec..111e87960 100644 --- a/tests/worker/test_api_executor.py +++ b/tests/worker/test_api_executor.py @@ -3,6 +3,7 @@ import concurrent.futures import json +import logging import threading import time from pathlib import Path @@ -1671,3 +1672,108 @@ def test_status_code_taken_from_first_row_across_groups( assert result.items[0].rows == [] assert len(result.items[1].rows) == 1 assert result.status_code == 200 + + +class TestCallLogging: + @pytest.fixture(autouse=True) + def _nebula_env(self, monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("NEBULA_API_BASE_URL", "https://nebula.example.com") + monkeypatch.setenv("NEBULA_API_TOKEN", "nebula-token") + + @staticmethod + def _records(caplog: pytest.LogCaptureFixture, prefix: str) -> list[str]: + return [ + r.getMessage() for r in caplog.records if r.getMessage().startswith(prefix) + ] + + def test_chat_completion_call_line_has_tokens_and_backend( + self, tmp_path: Path, caplog: pytest.LogCaptureFixture + ) -> None: + """A chat-completion response yields a per-call line with the token and + backend fields.""" + task = _batch_task(["hi"]) + transport = _SequenceTransport( + [ + httpx.Response( + 200, + json={ + "choices": [ + { + "message": {"content": "hello"}, + "finish_reason": "stop", + } + ], + "usage": { + "prompt_tokens": 10, + "completion_tokens": 5, + "completion_tokens_details": {"reasoning_tokens": 2}, + }, + "provider": "nebula", + }, + ) + ] + ) + with caplog.at_level(logging.INFO, logger="worker.executors.api_executor"): + _run(_executor(), task, transport, tmp_path) + call_lines = self._records(caplog, "api call") + assert len(call_lines) == 1 + msg = call_lines[0] + assert "row=0" in msg + assert "attempts=1" in msg + assert "status=200" in msg + assert "prompt_tokens=10" in msg + assert "completion_tokens=5" in msg + assert "reasoning_tokens=2" in msg + assert "finish_reason=stop" in msg + assert "backend=nebula" in msg + + def test_retried_503_then_200_shows_attempts_two( + self, tmp_path: Path, caplog: pytest.LogCaptureFixture + ) -> None: + """A retried 503 then 200 logs attempts=2.""" + task = _batch_task(["hi"], retries=2) + transport = _SequenceTransport([_error_response(503), _ok_response()]) + with caplog.at_level(logging.INFO, logger="worker.executors.api_executor"): + _run(_executor(), task, transport, tmp_path) + call_lines = self._records(caplog, "api call") + assert len(call_lines) == 1 + assert "attempts=2" in call_lines[0] + + def test_non_json_body_logs_dash_without_raising( + self, tmp_path: Path, caplog: pytest.LogCaptureFixture + ) -> None: + """A non-JSON body logs '-' fields without raising.""" + task = _batch_task(["hi"], response={"parse_json": False}) + transport = _SequenceTransport( + [ + httpx.Response( + 200, + text="not json", + headers={"Content-Type": "text/plain"}, + ) + ] + ) + with caplog.at_level(logging.INFO, logger="worker.executors.api_executor"): + result = _run(_executor(), task, transport, tmp_path) + assert result.items[0].text == "not json" + call_lines = self._records(caplog, "api call") + assert len(call_lines) == 1 + msg = call_lines[0] + assert "prompt_tokens=-" in msg + assert "backend=-" in msg + + def test_summary_reports_per_backend_counts( + self, tmp_path: Path, caplog: pytest.LogCaptureFixture + ) -> None: + """The summary line reports per-backend counts.""" + task = _batch_task(["a", "b"]) + transport = _EchoTransport() + with caplog.at_level(logging.INFO, logger="worker.executors.api_executor"): + _run(_executor(), task, transport, tmp_path) + summary = self._records(caplog, "api summary") + assert len(summary) == 1 + msg = summary[0] + assert "calls=2" in msg + assert "failures=0" in msg + assert "retries=0" in msg + assert "backends=-" in msg From 5db7f13c6b75b3abb0db338a06d0bf80ffe6d034 Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Sun, 27 Sep 2026 11:11:31 +0700 Subject: [PATCH 47/71] fix(worker): harden API executor prompt and telemetry handling Embedded {{prompt}} placeholders render a non-string prompt as JSON instead of Python repr, and telemetry skips malformed provider and token values so logging never fails a successful request. A single lock guards all per-call statistics updates and the heartbeat and summary reads. Signed-off-by: Zhengyuan Su Co-Authored-By: Claude Code --- src/worker/executors/api_executor.py | 73 ++++++++++++++++++---------- tests/worker/test_api_executor.py | 65 +++++++++++++++++++++++-- 2 files changed, 108 insertions(+), 30 deletions(-) diff --git a/src/worker/executors/api_executor.py b/src/worker/executors/api_executor.py index 70175b75c..ea060dfaa 100644 --- a/src/worker/executors/api_executor.py +++ b/src/worker/executors/api_executor.py @@ -80,14 +80,22 @@ def _chat_completion_stats(body: Any) -> tuple[Any, Any, Any, Any, Any]: if backend is None: backend = body.get("system_fingerprint") return ( - prompt_tokens, - completion_tokens, - reasoning_tokens, + _as_token_count(prompt_tokens), + _as_token_count(completion_tokens), + _as_token_count(reasoning_tokens), finish_reason, - backend, + backend if isinstance(backend, str) else None, ) +def _as_token_count(value: Any) -> int | None: + """Coerce a token field to an int, skipping non-numeric values (including + bools) so malformed telemetry never changes a request's outcome.""" + if isinstance(value, bool) or not isinstance(value, int): + return None + return value + + class APIExecutor(DataMixin, Executor): """Performs HTTP requests defined by task YAML. @@ -496,6 +504,7 @@ def _run(self, task: ExecutorTask, out_dir: Path) -> APIResult: task_start = time.monotonic() in_flight: dict[int, float] = {} in_flight_lock = threading.Lock() + stats_lock = threading.Lock() def _record_call( idx: int, @@ -511,11 +520,6 @@ def _record_call( wall = time.monotonic() - start with in_flight_lock: in_flight.pop(idx, None) - done += 1 - if failed: - failures += 1 - total_retries += max(0, attempts - 1) - latencies.append(wall) ( prompt_tokens, completion_tokens, @@ -523,14 +527,20 @@ def _record_call( finish_reason, backend, ) = _chat_completion_stats(body) - if prompt_tokens is not None: - sum_prompt += prompt_tokens - if completion_tokens is not None: - sum_completion += completion_tokens - if reasoning_tokens is not None: - sum_reasoning += reasoning_tokens - if backend is not None: - backend_counts[backend] = backend_counts.get(backend, 0) + 1 + with stats_lock: + done += 1 + if failed: + failures += 1 + total_retries += max(0, attempts - 1) + latencies.append(wall) + if prompt_tokens is not None: + sum_prompt += prompt_tokens + if completion_tokens is not None: + sum_completion += completion_tokens + if reasoning_tokens is not None: + sum_reasoning += reasoning_tokens + if backend is not None: + backend_counts[backend] = backend_counts.get(backend, 0) + 1 logger.info( "api call task=%s row=%d attempts=%d status=%s wall=%.3fs " "prompt_tokens=%s completion_tokens=%s reasoning_tokens=%s " @@ -610,11 +620,13 @@ def _heartbeat() -> None: outstanding = dict(in_flight) if not outstanding: continue + with stats_lock: + done_snapshot = done oldest = min(outstanding.values()) logger.info( "api heartbeat task=%s done=%d/%d in_flight=%d oldest_age=%.0fs", task.task_id, - done, + done_snapshot, total, len(outstanding), time.monotonic() - oldest, @@ -638,17 +650,28 @@ def _heartbeat() -> None: heartbeat_stop.set() heartbeat.join(timeout=5) wall = time.monotonic() - task_start + with stats_lock: + summary = ( + done, + failures, + total_retries, + list(latencies), + sum_prompt, + sum_completion, + sum_reasoning, + dict(backend_counts), + ) self._log_summary( task.task_id, total, - failures, - total_retries, + summary[1], + summary[2], wall, - latencies, - sum_prompt, - sum_completion, - sum_reasoning, - backend_counts, + summary[3], + summary[4], + summary[5], + summary[6], + summary[7], ) items = [results[idx] for idx in range(len(prompts))] diff --git a/tests/worker/test_api_executor.py b/tests/worker/test_api_executor.py index 111e87960..f1fb4a929 100644 --- a/tests/worker/test_api_executor.py +++ b/tests/worker/test_api_executor.py @@ -1226,7 +1226,6 @@ def test_embedded_placeholder_substitutes_string(self) -> None: body = json.loads(transport.requests[0].read()) assert body["messages"][0]["content"] == "Q: hello" - class TestDataframeRows: def test_dataframe_column_issues_one_request_per_row(self, tmp_path: Path) -> None: """A dataframe-spec API task whose column reads an upstream APIResult's @@ -1762,18 +1761,74 @@ def test_non_json_body_logs_dash_without_raising( assert "prompt_tokens=-" in msg assert "backend=-" in msg + def test_malformed_telemetry_does_not_fail_request( + self, tmp_path: Path, caplog: pytest.LogCaptureFixture + ) -> None: + """A dict-valued provider and a string token count are skipped in the + stats, so a successful request still succeeds and logs a summary.""" + task = _batch_task(["hi"]) + transport = _SequenceTransport( + [ + httpx.Response( + 200, + json={ + "choices": [{"message": {"content": "ok"}}], + "provider": {"name": "backend-a"}, + "usage": {"prompt_tokens": "ten", "completion_tokens": 5}, + }, + ) + ] + ) + with caplog.at_level(logging.INFO, logger="worker.executors.api_executor"): + result = _run(_executor(), task, transport, tmp_path) + assert result.items[0].status_code == 200 + call_lines = self._records(caplog, "api call") + assert len(call_lines) == 1 + msg = call_lines[0] + assert "prompt_tokens=-" in msg + assert "completion_tokens=5" in msg + assert "backend=-" in msg + summary = self._records(caplog, "api summary") + assert len(summary) == 1 + assert "calls=1" in summary[0] + assert "backends=-" in summary[0] + def test_summary_reports_per_backend_counts( self, tmp_path: Path, caplog: pytest.LogCaptureFixture ) -> None: """The summary line reports per-backend counts.""" - task = _batch_task(["a", "b"]) - transport = _EchoTransport() + task = _batch_task(["a", "b", "c"]) + transport = _SequenceTransport( + [ + httpx.Response( + 200, + json={ + "choices": [{"message": {"content": "a"}}], + "provider": "backend-a", + }, + ), + httpx.Response( + 200, + json={ + "choices": [{"message": {"content": "b"}}], + "provider": "backend-b", + }, + ), + httpx.Response( + 200, + json={ + "choices": [{"message": {"content": "c"}}], + "provider": "backend-a", + }, + ), + ] + ) with caplog.at_level(logging.INFO, logger="worker.executors.api_executor"): _run(_executor(), task, transport, tmp_path) summary = self._records(caplog, "api summary") assert len(summary) == 1 msg = summary[0] - assert "calls=2" in msg + assert "calls=3" in msg assert "failures=0" in msg assert "retries=0" in msg - assert "backends=-" in msg + assert "backends=backend-a=2,backend-b=1" in msg From 77e4bec1ec6f3980be66513b2a44d75d4918647b Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Sun, 27 Sep 2026 11:13:29 +0700 Subject: [PATCH 48/71] docs: document grouped API results and echo function data type Describe APIGroupItem grouped results, how a dataframe spec decides grouping from the upstream structure, empty-group and zero-row semantics, and the downstream items.rows.json / items.json expression shapes. Correct the echo executor entry to cover the sandboxed function data type. Signed-off-by: Zhengyuan Su Co-Authored-By: Claude Code --- docs/EXECUTORS.md | 28 ++++++++++++++++++++++++- docs/WORKFLOWS.md | 53 +++++++++++++++++++++++++++++++++++------------ 2 files changed, 67 insertions(+), 14 deletions(-) diff --git a/docs/EXECUTORS.md b/docs/EXECUTORS.md index 4e04fa2c6..2f158299f 100644 --- a/docs/EXECUTORS.md +++ b/docs/EXECUTORS.md @@ -5,7 +5,7 @@ The worker resolves `spec.taskType` against an executor registry in | `taskType` | Executor | Use case | |-----------|----------|----------| -| `echo` | `EchoExecutor` | Echo input back as result (smoke tests) | +| `echo` | `EchoExecutor` | Echo input back as result, or run a sandboxed function over upstream values | | `inference` | `VLLMExecutor` / `TransformersExecutor` | LLM inference | | `embedding` | `VLLMEmbeddingExecutor` (text, when `model.vllm` is set) / `TransformersExecutor` (visual, `model.transformers.mode: visual-embedding`) | Text / visual embeddings | | `diffusion` | `DiffusersExecutor` | Image / video diffusion models | @@ -64,3 +64,29 @@ Optional, for the search tools: - `SERPER_API_KEY` - `JINA_API_KEY` + +## Echo executor + +`taskType: echo` returns input values back as the result, or runs a sandboxed +Python function over upstream values. It is useful for inspecting and shaping +data between stages. + +`spec.data.type: list` echoes each `spec.data.items` entry. An entry is either +a string literal or a mapping with an expression (`expr`, or both `node` and +`path`) resolved against the upstream results. A resolved list is flattened +into one echo item per element; a scalar becomes a single item. + +`spec.data.type: function` runs a sandboxed source function over upstream +values. `spec.data.function` is the source of a function that takes the +resolved arguments and returns a list of JSON values; each element becomes one +echo item (the list is not flattened further). `spec.data.arguments` is a list +of argument specs, each with exactly one of: + +- `{items: }` — the whole list of upstream values, passed as-is. +- `{expr: }` — a single expression resolved against the upstream + results. +- `{node: , path: }` — a node and path resolved against the + upstream results. + +The function is executed in a sandbox; it must return a list, and each element +is emitted as one echo item. A function failure fails the task. diff --git a/docs/WORKFLOWS.md b/docs/WORKFLOWS.md index 5d1debc61..d6647904b 100644 --- a/docs/WORKFLOWS.md +++ b/docs/WORKFLOWS.md @@ -72,9 +72,9 @@ contract. ## API task `taskType: api` issues one HTTP request per row of `spec.data`, in parallel, -and returns the responses row-aligned in `APIResult.items`. A single request -is a one-row `spec.data`. `spec.data` is required, exactly as for the vLLM -executor; it supports the same data types (`list`, `dataset`, `graph_template`, +and returns the responses in `APIResult.items`. A single request is a one-row +`spec.data`. `spec.data` is required, exactly as for the vLLM executor; it +supports the same data types (`list`, `dataset`, `graph_template`, `dataframe`). By default it routes to the Nebula endpoint and authenticates with the worker's `NEBULA_API_TOKEN`. @@ -105,16 +105,43 @@ spec: ### Batching When `spec.data` is present, the task batches: one request is issued per row, -and the results are returned row-aligned in `APIResult.items`. Each row's -prompt is substituted for the `{{prompt}}` placeholder in the request body. -Server-side stage references are `${...}`; `{{prompt}}` is a worker-side -per-row slot, so it is not touched by server-side resolution. A failure in any -row fails the whole task rather than shifting the remaining rows. -`spec.api.concurrency` bounds the number of in-flight requests and is capped -at 8 (the default); values above 8 are clamped down. Cancelling the task -prevents not-yet-started rows from issuing and marks the task cancelled once -in-flight requests return; a request already inside the HTTP call is not -interrupted. +and the results are returned in `APIResult.items`. Each row's prompt is +substituted for the `{{prompt}}` placeholder in the request body. Server-side +stage references are `${...}`; `{{prompt}}` is a worker-side per-row slot, so +it is not touched by server-side resolution. A failure in any row fails the +whole task rather than shifting the remaining rows. `spec.api.concurrency` +bounds the number of in-flight requests and is capped at 8 (the default); +values above 8 are clamped down. Cancelling the task prevents not-yet-started +rows from issuing and marks the task cancelled once in-flight requests return; +a request already inside the HTTP call is not interrupted. + +A body value that is exactly `{{prompt}}` is replaced by the row's prompt +object as-is (a message list stays a list of `{"role", "content"}` dicts). An +embedded `{{prompt}}` inside a longer string keeps string substitution: a +string prompt is inserted verbatim, and any other prompt value is rendered as +JSON. + +### Grouped results + +When `spec.data` is a `dataframe` whose columns resolve to grouped upstream +values, the result is grouped: `APIResult.items` holds one `APIGroupItem` per +group, each with an `index` and a `rows` list of the group's row responses in +order. A dataframe spec decides grouping from the upstream structure — a list +of `APIGroupItem.rows`, or nested lists — never from the shape of the cell +values; a per-row list is a cell value, not a group. Ungrouped data returns +plain `APIItem`s directly in `items`. + +An empty group or a column that resolves to zero rows yields zero requests for +that group; the group still appears as an `APIGroupItem` with an empty `rows` +list so downstream paths resolve. The result's `status_code` is taken from the +first row across all groups, so a leading empty group does not zero it. + +Downstream stages read grouped responses through the group shape: +`items.rows.json...` addresses a field of each row within a group, while +`items.json...` addresses a field of a plain (ungrouped) item. For example, a +dataframe column that reads a grouped upstream's message content uses +`path: items.rows.json.choices[0].message.content`; an ungrouped upstream uses +`path: items.json.choices[0].message.content`. ```yaml spec: From 608e834983180df634764901fe02ccda32e503fb Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Sun, 27 Sep 2026 11:15:41 +0700 Subject: [PATCH 49/71] refactor(worker): pass summary snapshot as named locals Take the final summary snapshot straight into named locals under the statistics lock and pass those names to the summary logger, instead of packing a tuple and indexing it back out. Signed-off-by: Zhengyuan Su Co-Authored-By: Claude Code --- src/worker/executors/api_executor.py | 31 +++++++++++++--------------- 1 file changed, 14 insertions(+), 17 deletions(-) diff --git a/src/worker/executors/api_executor.py b/src/worker/executors/api_executor.py index ea060dfaa..4cc2abac3 100644 --- a/src/worker/executors/api_executor.py +++ b/src/worker/executors/api_executor.py @@ -651,27 +651,24 @@ def _heartbeat() -> None: heartbeat.join(timeout=5) wall = time.monotonic() - task_start with stats_lock: - summary = ( - done, - failures, - total_retries, - list(latencies), - sum_prompt, - sum_completion, - sum_reasoning, - dict(backend_counts), - ) + failures_snapshot = failures + retries_snapshot = total_retries + latencies_snapshot = list(latencies) + prompt_snapshot = sum_prompt + completion_snapshot = sum_completion + reasoning_snapshot = sum_reasoning + backends_snapshot = dict(backend_counts) self._log_summary( task.task_id, total, - summary[1], - summary[2], + failures_snapshot, + retries_snapshot, wall, - summary[3], - summary[4], - summary[5], - summary[6], - summary[7], + latencies_snapshot, + prompt_snapshot, + completion_snapshot, + reasoning_snapshot, + backends_snapshot, ) items = [results[idx] for idx in range(len(prompts))] From 6bf119240465b5801dc2bd79719daea0812afae7 Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Sun, 27 Sep 2026 12:15:55 +0700 Subject: [PATCH 50/71] fix(worker): cancel a zero-row grouped API task after the pool drains A task whose data resolves to empty grouped tables runs zero requests and never checks the cancel event, so a pending or in-flight cancellation was reported as a successful empty result. A cancel that lands while the final request is in flight was likewise never checked once the collection loop ended. Raise TaskCancelledError after the pool drains and before the result is built, matching the documented behaviour that a cancel marks the task cancelled once in-flight requests return. Signed-off-by: Zhengyuan Su Co-Authored-By: Claude Code --- src/worker/executors/api_executor.py | 3 ++ tests/worker/test_api_executor.py | 52 ++++++++++++++++++++++++++++ 2 files changed, 55 insertions(+) diff --git a/src/worker/executors/api_executor.py b/src/worker/executors/api_executor.py index 4cc2abac3..4adb41090 100644 --- a/src/worker/executors/api_executor.py +++ b/src/worker/executors/api_executor.py @@ -671,6 +671,9 @@ def _heartbeat() -> None: backends_snapshot, ) + if self._cancel_event.is_set(): + raise TaskCancelledError("API task cancelled") + items = [results[idx] for idx in range(len(prompts))] result_items: list[APIItem | APIGroupItem] = [] diff --git a/tests/worker/test_api_executor.py b/tests/worker/test_api_executor.py index f1fb4a929..2e829ba45 100644 --- a/tests/worker/test_api_executor.py +++ b/tests/worker/test_api_executor.py @@ -1226,6 +1226,7 @@ def test_embedded_placeholder_substitutes_string(self) -> None: body = json.loads(transport.requests[0].read()) assert body["messages"][0]["content"] == "Q: hello" + class TestDataframeRows: def test_dataframe_column_issues_one_request_per_row(self, tmp_path: Path) -> None: """A dataframe-spec API task whose column reads an upstream APIResult's @@ -1347,6 +1348,57 @@ def test_all_empty_columns_issue_no_request_and_one_empty_group( assert result.items[0].rows == [] assert result.status_code == 0 + def test_cancelled_zero_row_grouped_task_raises(self, tmp_path: Path) -> None: + """A zero-row grouped task with a pending cancellation for its id is + cancelled, not reported as a successful empty result.""" + upstream = APIResult( + ok=True, + executor="api", + method="POST", + url="https://up.example.com", + status_code=200, + items=[], + ) + payload = { + "task_id": "task-api-df-cancel-empty", + "workflow_id": "wf-1", + "owner_id": "owner", + "assigned_worker": "worker-1", + "dispatched_at": "2026-03-22T00:00:00Z", + "task": { + "apiVersion": "flowmesh/v1", + "kind": "Task", + "metadata": {"name": "wf:api"}, + "spec": { + "taskType": "api", + "_upstreamResults": {"Up": upstream}, + "api": { + "method": "POST", + "url": "https://custom.example.com/v1/chat/completions", + "json": {"messages": "{{prompt}}"}, + }, + "data": { + "type": "dataframe", + "columns": [ + { + "label": "L", + "node": "Up", + "path": "items.json.choices[0].message.content", + } + ], + "messages": [ + {"role": "user", "content": "row {L}"}, + ], + }, + }, + }, + } + task = WorkerTaskMessage.model_validate(payload) + executor = _executor() + executor.cancel(task.task_id) + with pytest.raises(TaskCancelledError): + _run(executor, task, _EchoTransport(), tmp_path) + def test_zero_vs_three_rows_still_raises(self, tmp_path: Path) -> None: """A real mismatch (one column empty, another with rows) still raises.""" empty = APIResult( From 5ef1e96a2f60aaf81de93bb16d65134ac38d8b3d Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Sun, 27 Sep 2026 12:16:23 +0700 Subject: [PATCH 51/71] fix(worker): log the real attempt count for an exhausted connection retry A row whose connection errors exhaust its retries re-raised without carrying the attempt count, so the caller logged attempts=0 and the summary undercounted retries. Carry the attempts used out of the retry loop on that failure path and pass it to the per-call statistics. Signed-off-by: Zhengyuan Su Co-Authored-By: Claude Code --- src/worker/executors/api_executor.py | 28 ++++++++++++++++++++++------ tests/worker/test_api_executor.py | 25 +++++++++++++++++++++++++ 2 files changed, 47 insertions(+), 6 deletions(-) diff --git a/src/worker/executors/api_executor.py b/src/worker/executors/api_executor.py index 4adb41090..b261ec9bf 100644 --- a/src/worker/executors/api_executor.py +++ b/src/worker/executors/api_executor.py @@ -41,6 +41,15 @@ def _is_retryable_status(status_code: int) -> bool: return status_code >= 500 or status_code in (408, 429) +class _RequestFailed(Exception): + """A request that exhausted its retries, carrying the attempts used.""" + + def __init__(self, error: Exception, attempts: int) -> None: + super().__init__(str(error)) + self.error = error + self.attempts = attempts + + def _fmt(value: Any) -> str: """Render a log field, using ``-`` for a missing value.""" return "-" if value is None else str(value) @@ -227,12 +236,12 @@ def _request_with_retries( params=params, **request_kwargs, ) - except httpx.RequestError: + except httpx.RequestError as exc: if attempt < retries: attempt += 1 self._wait_for_backoff() continue - raise + raise _RequestFailed(exc, attempt + 1) from None if resp.is_error and _is_retryable_status(resp.status_code): if attempt < retries: attempt += 1 @@ -575,11 +584,18 @@ def _issue(idx: int, prompt: Any) -> APIItem: kwargs, retries, ) - except httpx.RequestError as exc: - _record_call(idx, 0, exc.__class__.__name__, start, None, failed=True) + except _RequestFailed as exc: + _record_call( + idx, + exc.attempts, + exc.error.__class__.__name__, + start, + None, + failed=True, + ) raise ExecutionError( - f"API request failed (row {idx}): {exc}", retryable=True - ) from exc + f"API request failed (row {idx}): {exc.error}", retryable=True + ) from exc.error if raise_for_status and resp.is_error: message = f"API request returned status {resp.status_code} (row {idx})" diff --git a/tests/worker/test_api_executor.py b/tests/worker/test_api_executor.py index 2e829ba45..1720a6788 100644 --- a/tests/worker/test_api_executor.py +++ b/tests/worker/test_api_executor.py @@ -1790,6 +1790,31 @@ def test_retried_503_then_200_shows_attempts_two( assert len(call_lines) == 1 assert "attempts=2" in call_lines[0] + def test_exhausted_connect_error_logs_attempts_and_retries( + self, tmp_path: Path, caplog: pytest.LogCaptureFixture + ) -> None: + """A row whose connection errors exhaust retries logs the real attempt + count and the retry total, not zero.""" + task = _batch_task(["hi"], retries=2) + + class _RaisingTransport(httpx.MockTransport): + def __init__(self) -> None: + super().__init__(self._handler) + + def _handler(self, request: httpx.Request) -> httpx.Response: + raise httpx.ConnectError("boom", request=request) + + transport = _RaisingTransport() + with caplog.at_level(logging.INFO, logger="worker.executors.api_executor"): + with pytest.raises(ExecutionError, match="API request failed"): + _run(_executor(), task, transport, tmp_path) + call_lines = self._records(caplog, "api call") + assert len(call_lines) == 1 + assert "attempts=3" in call_lines[0] + summary = self._records(caplog, "api summary") + assert len(summary) == 1 + assert "retries=2" in summary[0] + def test_non_json_body_logs_dash_without_raising( self, tmp_path: Path, caplog: pytest.LogCaptureFixture ) -> None: From d8ac38f7605abec9b87058f86ad5d284b5dcf0b1 Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Mon, 28 Sep 2026 15:00:46 +0700 Subject: [PATCH 52/71] Address PR comments Signed-off-by: Zhengyuan Su --- docs/EXECUTORS.md | 2 +- docs/WORKFLOWS.md | 37 ++- examples/templates/api_two_stage.yaml | 23 +- src/shared/tasks/specs/misc.py | 129 +++++++- src/worker/executors/api_executor.py | 149 ++++------ tests/server/task/test_n8n_parser.py | 4 +- tests/server/task/test_ssh_result_mounting.py | 171 +---------- .../task/test_stage_reference_resolution.py | 213 ++++++++++++++ tests/server/test_redact.py | 11 +- tests/shared/test_api_spec.py | 278 ++++++++++++++++++ tests/worker/test_api_executor.py | 88 +++--- 11 files changed, 758 insertions(+), 347 deletions(-) create mode 100644 tests/server/task/test_stage_reference_resolution.py create mode 100644 tests/shared/test_api_spec.py diff --git a/docs/EXECUTORS.md b/docs/EXECUTORS.md index 4e04fa2c6..6d493deaf 100644 --- a/docs/EXECUTORS.md +++ b/docs/EXECUTORS.md @@ -17,7 +17,7 @@ The worker resolves `spec.taskType` against an executor registry in | `data_profiling` | `DataProfilingExecutor` | DataFrame profiling | | `data_retrieval` | `DataRetrievalExecutor` | DataFrame loading from sources (`type: sql`, `type: s3`, `type: lumid` with `mode: sql\|s3\|agent` via lumid-data-app; `type: lumid` (mode `sql`/`s3`/`agent`) requires `lumid_data_token`, the bearer forwarded to lumid-data-app) | | `ssh` | `SSHExecutor` | Interactive SSH session or non-interactive container job | -| `api` | `APIExecutor` | One parallel HTTP request per `spec.data` row (spec.data required) | +| `api` | `APIExecutor` | One parallel HTTP request per `spec.data` row | | `serve` | `VLLMServeExecutor` | Persistent vLLM API server for a single model | Helper utilities live in `src/worker/executors/utils/` (`artifacts`, diff --git a/docs/WORKFLOWS.md b/docs/WORKFLOWS.md index 80ab4196f..a4da7bb5b 100644 --- a/docs/WORKFLOWS.md +++ b/docs/WORKFLOWS.md @@ -71,11 +71,10 @@ contract. ## API task -`taskType: api` issues one HTTP request per row of `spec.data`, in parallel, -and returns the responses row-aligned in `APIResult.items`. A single request -is a one-row `spec.data`. `spec.data` is required, exactly as for the vLLM -executor; it supports the same data types (`list`, `dataset`, `graph_template`, -`dataframe`). +`taskType: api` sends one HTTP request per `spec.data` row, in parallel, and +returns one `APIResult.items` entry per row, in order. `spec.data` is required +and accepts `list`, `dataset`, `graph_template`, and `dataframe`; row metadata +and `dataframe` table grouping are not applied. By default it routes to the Nebula endpoint and authenticates with the worker's `NEBULA_API_TOKEN`. @@ -83,11 +82,15 @@ By default it routes to the Nebula endpoint and authenticates with the worker's Credential handling: a caller-supplied `Authorization` header is always used as-is and never overwritten. With no header, `NEBULA_API_TOKEN` is injected only when the call is on the Nebula url (no custom `spec.api.url`) — the Nebula token is never sent to a custom endpoint. A Nebula-path call with no token available fails closed. -`spec.api.retries` (default `0`, at most `10`) sets how many times a transient failure is retried before the task fails. A transient failure is a connection error or an HTTP status of 5xx, 408, or 429; other 4xx statuses are never retried. Retries back off exponentially: the first waits 1s and each later one doubles, capped at 60s. When a retryable response carries a `Retry-After` header (seconds or an HTTP date), that wait is used instead, also capped at 60s. Each retry logs a warning with the attempt count and the wait. A cancelled task stops retrying immediately. +`spec.api.retries` (default `0`, an integer from 0 to 10; any other value is rejected when the workflow is submitted) sets how many times a transient failure is retried before the task fails. A transient failure is a connection error or an HTTP status of 5xx, 408, or 429; other 4xx statuses are never retried. Retries back off exponentially: the first waits 1s and each later one doubles, capped at 60s. When a retryable response carries a `Retry-After` header (seconds or an HTTP date), that wait is used instead, also capped at 60s. Each retry logs a warning with the attempt count and the wait. A cancelled task stops retrying immediately. ```yaml spec: taskType: api + data: + type: list + items: + - Hello api: method: POST headers: @@ -96,25 +99,19 @@ spec: model: gpt-4o messages: - role: user - content: Hello + content: "{{prompt}}" retries: 3 response: parse_json: true ``` -### Batching - -When `spec.data` is present, the task batches: one request is issued per row, -and the results are returned row-aligned in `APIResult.items`. Each row's -prompt is substituted for the `{{prompt}}` placeholder in the request body. -Server-side stage references are `${...}`; `{{prompt}}` is a worker-side -per-row slot, so it is not touched by server-side resolution. A failure in any -row fails the whole task rather than shifting the remaining rows. -`spec.api.concurrency` bounds the number of in-flight requests and is capped -at 8 (the default); values above 8 are clamped down. Cancelling the task -prevents not-yet-started rows from issuing and marks the task cancelled once -in-flight requests return; a request already inside the HTTP call is not -interrupted. +### Per-row prompts + +Each row's prompt replaces `{{prompt}}` in the request body; a value that is +exactly `{{prompt}}` takes the prompt as-is, so a message-list row fills +`messages`. `spec.api.concurrency` (default 8, an integer from 1 to 8; any other value is rejected when the workflow is submitted) bounds in-flight +requests. Any failed row fails the task. Cancelling the task skips rows that +have not started and marks it cancelled once in-flight requests return. ```yaml spec: diff --git a/examples/templates/api_two_stage.yaml b/examples/templates/api_two_stage.yaml index 136aaad50..6e9e28dbd 100644 --- a/examples/templates/api_two_stage.yaml +++ b/examples/templates/api_two_stage.yaml @@ -1,21 +1,12 @@ -# api_two_stage.yaml +# Two-stage API executor demo. # -# Two-stage batched API executor demo. +# Stage 1 sends one request per spec.data row, filling {{prompt}} in the +# request body. Stage 2 takes the first row's response +# (${stage-1.items.0.text}) and sends it as its own one-row prompt. # -# spec.data is required for an api task and yields one row per request; a -# single request is a one-row list. Each row's prompt is substituted for the -# {{prompt}} placeholder in the request body and one request is issued per -# row, returned row-aligned in APIResult.items. Server-side stage references -# are ${...}; {{prompt}} is a worker-side per-row slot, so it is not touched -# by server-side resolution. -# -# Stage 1 fans out over three rows and calls a chat completion endpoint. -# Stage 2 consumes Stage 1's returned text as the next prompt, so the -# two-stage character is kept while exercising the batch path. A dependent -# stage reads a single upstream value: APIResult is batch-only, so the -# reference addresses one row explicitly (${stage-1.items.0.text} is the -# first row's text), and stage-2's one-row spec.data supplies that value as -# its prompt. +# Stage 1 names a custom endpoint via spec.api.url and supplies its own +# Authorization header. Stage 2 omits both, so it uses the Nebula defaults +# (NEBULA_API_BASE_URL + NEBULA_API_TOKEN). apiVersion: flowmesh/v1 kind: APITask diff --git a/src/shared/tasks/specs/misc.py b/src/shared/tasks/specs/misc.py index 1f8b1963f..72678edc7 100644 --- a/src/shared/tasks/specs/misc.py +++ b/src/shared/tasks/specs/misc.py @@ -1,7 +1,11 @@ -from typing import Any, Literal, Self +from typing import Annotated, Any, Literal, Self + +from pydantic import AfterValidator, AliasChoices, Field, StrictBool, StrictInt from ...utils.pydantic_utils import copy_preserving_fields_set from ...utils.redact import has_redacted_credential_fields, redact_credential_fields +from .._base import StrictBaseModel, TemplateBaseModel +from ..placeholders import PlaceholderString, TemplateBool from ..task_type import TaskType from .common import ( ModelSpecStrict, @@ -10,37 +14,140 @@ TaskSpecTemplateBase, ) +# Upper bounds on spec.api.concurrency and spec.api.retries. +_MAX_CONCURRENCY = 8 +_MAX_RETRIES = 10 + +# Bounds shared by the Strict and Template models. The Template unions each +# bounded literal with a placeholder so a ``${...}`` reference stays accepted. +Retries = Annotated[StrictInt, Field(ge=0, le=_MAX_RETRIES)] +Concurrency = Annotated[StrictInt, Field(ge=1, le=_MAX_CONCURRENCY)] +MaxBodyBytes = Annotated[StrictInt, Field(gt=0)] +TimeoutSec = Annotated[float, Field(gt=0)] + + +def _upper_method(value: str) -> str: + return value.upper() + + +Method = Annotated[str, AfterValidator(_upper_method)] + + +class ApiResponseConfig(StrictBaseModel): + parse_json: StrictBool = True + return_body: StrictBool = True + include_headers: StrictBool = False + max_body_bytes: MaxBodyBytes = 200000 + raise_for_status: StrictBool = True + + +class ApiResponseConfigTemplate(TemplateBaseModel): + parse_json: TemplateBool = True + return_body: TemplateBool = True + include_headers: TemplateBool = False + max_body_bytes: MaxBodyBytes | PlaceholderString = 200000 + raise_for_status: TemplateBool = True + + +class ApiConfig(StrictBaseModel): + url: str | None = None + method: Method = "POST" + headers: dict[str, Any] | None = None + params: dict[str, Any] | None = None + json_body: Any | None = Field( + default=None, + validation_alias=AliasChoices("json_body", "json"), + serialization_alias="json", + ) + body: Any | None = None + data: Any | None = None + timeout_sec: TimeoutSec = 60.0 + verify_tls: StrictBool = True + follow_redirects: StrictBool = True + retries: Retries = 0 + concurrency: Concurrency = _MAX_CONCURRENCY + response: ApiResponseConfig | None = None + + +class ApiConfigTemplate(TemplateBaseModel): + url: str | None = None + method: str = "POST" + headers: dict[str, Any] | None = None + params: dict[str, Any] | None = None + json_body: Any | None = Field( + default=None, + validation_alias=AliasChoices("json_body", "json"), + serialization_alias="json", + ) + body: Any | None = None + data: Any | None = None + timeout_sec: TimeoutSec | PlaceholderString = 60.0 + verify_tls: TemplateBool = True + follow_redirects: TemplateBool = True + retries: Retries | PlaceholderString = 0 + concurrency: Concurrency | PlaceholderString = _MAX_CONCURRENCY + response: ApiResponseConfigTemplate | None = None + + +def _redact_api_config( + api: ApiConfig | ApiConfigTemplate | None, +) -> ApiConfig | ApiConfigTemplate | None: + if api is None: + return None + return copy_preserving_fields_set( + api, + { + "headers": redact_credential_fields(api.headers), + "params": redact_credential_fields(api.params), + "json_body": redact_credential_fields(api.json_body), + "body": redact_credential_fields(api.body), + "data": redact_credential_fields(api.data), + }, + ) + + +def _api_has_redacted_credentials( + api: ApiConfig | ApiConfigTemplate | None, +) -> bool: + if api is None: + return False + return has_redacted_credential_fields( + { + "headers": api.headers, + "params": api.params, + "json_body": api.json_body, + "body": api.body, + "data": api.data, + } + ) + class ApiSpecStrict(TaskSpecStrictBase): taskType: Literal[TaskType.API] - api: dict[str, Any] | None = None + api: ApiConfig | None = None data: dict[str, Any] | None = None def redact_credentials(self) -> Self: spec = super().redact_credentials() - return copy_preserving_fields_set( - spec, {"api": redact_credential_fields(spec.api)} - ) + return copy_preserving_fields_set(spec, {"api": _redact_api_config(spec.api)}) def has_redacted_credentials(self) -> bool: - return super().has_redacted_credentials() or has_redacted_credential_fields( + return super().has_redacted_credentials() or _api_has_redacted_credentials( self.api ) class ApiSpecTemplate(TaskSpecTemplateBase): taskType: Literal[TaskType.API] - api: dict[str, Any] | None = None + api: ApiConfigTemplate | None = None data: dict[str, Any] | None = None def redact_credentials(self) -> Self: spec = super().redact_credentials() - return copy_preserving_fields_set( - spec, {"api": redact_credential_fields(spec.api)} - ) + return copy_preserving_fields_set(spec, {"api": _redact_api_config(spec.api)}) def has_redacted_credentials(self) -> bool: - return super().has_redacted_credentials() or has_redacted_credential_fields( + return super().has_redacted_credentials() or _api_has_redacted_credentials( self.api ) diff --git a/src/worker/executors/api_executor.py b/src/worker/executors/api_executor.py index 5071350ba..c15a3a40d 100644 --- a/src/worker/executors/api_executor.py +++ b/src/worker/executors/api_executor.py @@ -13,6 +13,7 @@ from shared.schemas.result import APIItem, APIResult from shared.tasks.specs import ApiSpecStrict +from shared.tasks.specs.misc import ApiConfig, ApiResponseConfig from shared.tasks.task_type import TaskType from shared.utils.redact import is_credential_key @@ -26,19 +27,16 @@ logger = logging.getLogger(__name__) +# Cache key: (base_url, timeout_seconds, verify_tls, follow_redirects, concurrency) _ClientKey = tuple[str, float, bool, bool, int] # Worker-side per-row slot; server-side stage references are ${...}. _PROMPT_PLACEHOLDER = "{{prompt}}" -_MAX_CONCURRENCY = 8 - # Base delay between retry attempts, doubled each retry. _RETRY_BACKOFF_SEC = 1.0 # Upper bound on any single retry wait. _RETRY_BACKOFF_MAX_SEC = 60.0 -# Upper bound on spec.api.retries. -_MAX_RETRIES = 10 def _is_retryable_status(status_code: int) -> bool: @@ -281,14 +279,11 @@ def _substitute_prompt(cls, value: Any, prompt: Any) -> Any: return [cls._substitute_prompt(v, prompt) for v in value] return value - def _build_request_kwargs( - self, api_cfg: dict[str, Any], prompt: str | None - ) -> dict[str, Any]: - """Build httpx request kwargs from ``spec.api``, substituting the row - prompt when batching.""" - json_payload = api_cfg.get("json") - body = api_cfg.get("body") - data_payload = api_cfg.get("data") + def _build_request_kwargs(self, api_cfg: ApiConfig) -> dict[str, Any]: + """Build httpx request kwargs from ``spec.api``.""" + json_payload = api_cfg.json_body + body = api_cfg.body + data_payload = api_cfg.data if json_payload is not None and body is not None: raise ExecutionError( @@ -297,43 +292,26 @@ def _build_request_kwargs( request_kwargs: dict[str, Any] = {} if json_payload is not None: - request_kwargs["json"] = ( - self._substitute_prompt(json_payload, prompt) - if prompt is not None - else json_payload - ) + request_kwargs["json"] = json_payload elif body is not None: if isinstance(body, (dict, list)): - request_kwargs["json"] = ( - self._substitute_prompt(body, prompt) - if prompt is not None - else body - ) + request_kwargs["json"] = body else: - request_kwargs["content"] = ( - self._substitute_prompt(body, prompt) - if prompt is not None - else body - ) + request_kwargs["content"] = body elif data_payload is not None: - request_kwargs["data"] = ( - self._substitute_prompt(data_payload, prompt) - if prompt is not None - else data_payload - ) + request_kwargs["data"] = data_payload return request_kwargs def _parse_response( self, resp: httpx.Response, *, - response_cfg: dict[str, Any], + response_cfg: ApiResponseConfig, max_body_bytes: int, - ) -> tuple[APIItem, str | None]: - """Turn one HTTP response into an APIItem, applying response config. - - Returns the item and the raw body text (used for error messages). - """ + idx: int, + prompt_str: str, + ) -> APIItem: + """Turn one HTTP response into an APIItem, applying response config.""" body_bytes = resp.content truncated = False if max_body_bytes is not None and len(body_bytes) > max_body_bytes: @@ -341,21 +319,22 @@ def _parse_response( truncated = True item = APIItem( - index=0, + index=idx, url=str(resp.url), status_code=resp.status_code, truncated=truncated, + prompt=prompt_str, ) - if response_cfg.get("include_headers", False): + if response_cfg.include_headers: item.headers = dict(resp.headers) body_text: str | None = None - if response_cfg.get("return_body", True): + if response_cfg.return_body: encoding = resp.encoding or "utf-8" body_text = body_bytes.decode(encoding, errors="replace") - if response_cfg.get("parse_json", True): + if response_cfg.parse_json: item.response_json = resp.json() if not isinstance(item.response_json, dict): raise ExecutionError("Response is not a valid JSON mapping") @@ -366,10 +345,10 @@ def _parse_response( item.text = item.response_json["choices"][0]["message"]["content"] except Exception: item.text = None - elif response_cfg.get("return_body", True): + elif response_cfg.return_body: item.text = body_text - return item, body_text + return item def run(self, task: ExecutorTask, out_dir: Path) -> APIResult: with self._cancel_lock: @@ -386,15 +365,11 @@ def run(self, task: ExecutorTask, out_dir: Path) -> APIResult: def _run(self, task: ExecutorTask, out_dir: Path) -> APIResult: spec = self.require_spec(task, ApiSpecStrict) - api_cfg = spec.api or {} - if not isinstance(api_cfg, dict): - raise ExecutionError("spec.api must be a mapping") + api_cfg = spec.api or ApiConfig.model_validate({}) - url = api_cfg.get("url") - method = str(api_cfg.get("method", "POST")).upper() - headers = api_cfg.get("headers", {}) - if not isinstance(headers, dict): - raise ExecutionError("spec.api.headers must be a mapping") + url = api_cfg.url + method = api_cfg.method + headers = api_cfg.headers or {} if url is None: url = os.getenv("NEBULA_API_BASE_URL") @@ -411,35 +386,19 @@ def _run(self, task: ExecutorTask, out_dir: Path) -> APIResult: ) headers["Authorization"] = f"Bearer {token}" - params = api_cfg.get("params") - if params is not None and not isinstance(params, dict): - raise ExecutionError("spec.api.params must be a mapping") - - timeout_sec = api_cfg.get("timeout_sec", 60) - if not isinstance(timeout_sec, (int, float)): - raise ExecutionError("spec.api.timeout_sec must be a number") - timeout = httpx.Timeout(timeout_sec) - - verify_tls = api_cfg.get("verify_tls", True) - follow_redirects = api_cfg.get("follow_redirects", True) + params = api_cfg.params - response_cfg = api_cfg.get("response") or {} - if response_cfg and not isinstance(response_cfg, dict): - raise ExecutionError("spec.api.response must be a mapping") + timeout = httpx.Timeout(api_cfg.timeout_sec) - max_body_bytes = int(response_cfg.get("max_body_bytes", 200000)) - raise_for_status = bool(response_cfg.get("raise_for_status", True)) + verify_tls = api_cfg.verify_tls + follow_redirects = api_cfg.follow_redirects - retries = api_cfg.get("retries", 0) - if not isinstance(retries, int) or isinstance(retries, bool) or retries < 0: - raise ExecutionError("spec.api.retries must be a non-negative integer") - if retries > _MAX_RETRIES: - raise ExecutionError(f"spec.api.retries must be at most {_MAX_RETRIES}") + response_cfg = api_cfg.response or ApiResponseConfig.model_validate({}) + max_body_bytes = response_cfg.max_body_bytes + raise_for_status = response_cfg.raise_for_status - concurrency = int(api_cfg.get("concurrency", _MAX_CONCURRENCY)) - if concurrency < 1: - raise ExecutionError("spec.api.concurrency must be >= 1") - concurrency = min(concurrency, _MAX_CONCURRENCY) + retries = api_cfg.retries + concurrency = api_cfg.concurrency base = self._base_url(str(url)) client = self._get_client( @@ -451,11 +410,16 @@ def _run(self, task: ExecutorTask, out_dir: Path) -> APIResult: if not prompts: raise ExecutionError("spec.data produced no rows") - request_kwargs = self._build_request_kwargs(api_cfg, None) + request_kwargs = self._build_request_kwargs(api_cfg) + + failed = threading.Event() + first_error: list[BaseException] = [] def _issue(idx: int, prompt: Any) -> APIItem: if self._cancel_event.is_set(): raise TaskCancelledError("API task cancelled") + if failed.is_set(): + raise ExecutionError(f"API task failed on an earlier row (row {idx})") prompt_str = self._prompt_to_str(prompt) kwargs = self._substitute_prompt(request_kwargs, prompt) try: @@ -469,25 +433,31 @@ def _issue(idx: int, prompt: Any) -> APIItem: retries, ) except httpx.RequestError as exc: - raise ExecutionError( + failed.set() + error = ExecutionError( f"API request failed (row {idx}): {exc}", retryable=True - ) from exc + ) + first_error.append(error) + raise error from exc if raise_for_status and resp.is_error: + failed.set() message = f"API request returned status {resp.status_code} (row {idx})" body_text = resp.text[:200] if body_text: message = f"{message}: {body_text}" - retryable = resp.status_code >= 500 or resp.status_code in (408, 429) - raise ExecutionError(message, retryable=retryable) + retryable = _is_retryable_status(resp.status_code) + error = ExecutionError(message, retryable=retryable) + first_error.append(error) + raise error - item, _ = self._parse_response( + item = self._parse_response( resp, response_cfg=response_cfg, max_body_bytes=max_body_bytes, + idx=idx, + prompt_str=prompt_str, ) - item.index = idx - item.prompt = prompt_str if self._cancel_event.is_set(): raise TaskCancelledError("API task cancelled") @@ -500,12 +470,19 @@ def _issue(idx: int, prompt: Any) -> APIItem: for idx, prompt in enumerate(prompts): if self._cancel_event.is_set(): raise TaskCancelledError("API task cancelled") + if failed.is_set(): + break futures[pool.submit(_issue, idx, prompt)] = idx for future in as_completed(futures): idx = futures[future] if self._cancel_event.is_set(): raise TaskCancelledError("API task cancelled") - results[idx] = future.result() + try: + results[idx] = future.result() + except ExecutionError: + pass + if first_error: + raise first_error[0] items = [results[idx] for idx in range(len(prompts))] diff --git a/tests/server/task/test_n8n_parser.py b/tests/server/task/test_n8n_parser.py index dc58f7ba2..33533bd3f 100644 --- a/tests/server/task/test_n8n_parser.py +++ b/tests/server/task/test_n8n_parser.py @@ -1,5 +1,7 @@ """Tests for n8n workflow translation.""" +import base64 + import pytest from server.task.n8n_parser import _decode_secret_part, translate_n8n_workflow @@ -90,8 +92,6 @@ def test_hex_decode(self) -> None: assert _decode_secret_part(encoded) == data def test_base64_decode(self) -> None: - import base64 - data = b"hello world" encoded = base64.b64encode(data).decode() assert _decode_secret_part(encoded) == data diff --git a/tests/server/task/test_ssh_result_mounting.py b/tests/server/task/test_ssh_result_mounting.py index c5be4ea01..1953be090 100644 --- a/tests/server/task/test_ssh_result_mounting.py +++ b/tests/server/task/test_ssh_result_mounting.py @@ -14,9 +14,8 @@ from server.task.models import TaskRecord, TaskStatus from server.task.parser import parse_workflow from server.task.runtime import TaskRuntime -from shared.schemas.result import APIItem, APIResult, ResultEnvelope from shared.tasks import TaskEnvelopeTemplate, TaskType -from shared.tasks.specs import ApiSpecTemplate, SSHSpecStrict +from shared.tasks.specs import SSHSpecStrict class _DummyRuntime: @@ -387,171 +386,3 @@ def test_stage_reference_uses_payload_root_for_local_and_http_results( http_value == "http://flowmesh.example/api/v1/results/task-http/files/final_lora.tar.gz" ) - - -def test_api_dependent_stage_resolves_first_row_text(tmp_path: Path) -> None: - """A dependent stage's ${stage.items.0.text} resolves to the first row's - text of a batch-only APIResult.""" - stage_dir = tmp_path / "task-api" - stage_dir.mkdir() - result = APIResult( - ok=True, - executor="api", - method="POST", - url="https://api.example.com/v1/chat/completions", - status_code=200, - items=[ - APIItem( - index=0, - url="https://api.example.com/v1/chat/completions", - status_code=200, - text="first row text", - prompt="first", - ), - APIItem( - index=1, - url="https://api.example.com/v1/chat/completions", - status_code=200, - text="second row text", - prompt="second", - ), - ], - ) - (stage_dir / "results.json").write_text( - json.dumps( - { - "task_id": "task-api", - "result": json.loads( - ResultEnvelope(task_id="task-api", result=result).model_dump_json() - )["result"], - } - ), - encoding="utf-8", - ) - - upstream = TaskRecord( - task_id="task-api", - workflow_id="wf-1", - owner_id="owner", - source="raw", - task=_task_template(TaskType.API), - status=TaskStatus.DONE, - task_type="api", - local_name="stage", - ) - dispatcher = Dispatcher( - runtime=cast(TaskRuntime, _DummyRuntime({})), - worker_registry=cast(WorkerRegistry, object()), - results_dir=tmp_path, - logger=logging.getLogger("test-api-dependent-stage"), - ) - - value = dispatcher._resolve_reference("stage.items.0.text", {"stage": upstream}) - assert value == "first row text" - - -def test_translated_n8n_dependent_api_stage_resolves(tmp_path: Path) -> None: - """A translated n8n workflow's dependent API stage resolves end to end.""" - payload = { - "nodes": [ - { - "name": "Upstream", - "type": "@n8n/n8n-nodes-langchain.openAi", - "parameters": { - "modelId": {"value": "gpt-4"}, - "responses": {"values": [{"content": "First answer"}]}, - }, - }, - { - "name": "Downstream", - "type": "@n8n/n8n-nodes-langchain.openAi", - "parameters": { - "modelId": {"value": "gpt-4"}, - "responses": {"values": [{"content": "Simplify this"}]}, - }, - }, - ], - "connections": { - "Upstream": {"ai_languageModel": [[{"node": "Downstream"}]]}, - }, - } - parsed = parse_workflow(json.dumps(payload), "n8n") - by_name = {t.graph_node_name: t for t in parsed.tasks} - upstream = by_name["Upstream"] - downstream = by_name["Downstream"] - assert downstream.depends_on == [upstream.task_id] - - stage_dir = tmp_path / upstream.task_id - stage_dir.mkdir() - result = APIResult( - ok=True, - executor="api", - method="POST", - url="https://api.example.com/v1/chat/completions", - status_code=200, - items=[ - APIItem( - index=0, - url="https://api.example.com/v1/chat/completions", - status_code=200, - text="first row text", - prompt="first", - ), - ], - ) - (stage_dir / "results.json").write_text( - json.dumps( - { - "task_id": upstream.task_id, - "result": json.loads( - ResultEnvelope( - task_id=upstream.task_id, result=result - ).model_dump_json() - )["result"], - } - ), - encoding="utf-8", - ) - - upstream_record = TaskRecord( - task_id=upstream.task_id, - workflow_id="wf-1", - owner_id="owner", - source="raw", - task=upstream.task, - status=TaskStatus.DONE, - task_type="api", - graph_node_name="Upstream", - ) - downstream_record = TaskRecord( - task_id=downstream.task_id, - workflow_id="wf-1", - owner_id="owner", - source="raw", - task=downstream.task, - status=TaskStatus.PENDING, - task_type="api", - graph_node_name="Downstream", - ) - dispatcher = Dispatcher( - runtime=cast( - TaskRuntime, - _DummyRuntime( - { - upstream.task_id: upstream_record, - downstream.task_id: downstream_record, - }, - depends_on={downstream.task_id: [upstream.task_id]}, - ), - ), - worker_registry=cast(WorkerRegistry, object()), - results_dir=tmp_path, - logger=logging.getLogger("test-n8n-dependent-stage"), - ) - - context = dispatcher._build_stage_context(downstream_record) - spec = cast(ApiSpecTemplate, downstream.task.spec) - resolved = dispatcher._resolve_placeholders(spec.data, context) - assert resolved["items"][0] == ( - "The previous stage's response is as follows. Simplify this\n" "first row text" - ) diff --git a/tests/server/task/test_stage_reference_resolution.py b/tests/server/task/test_stage_reference_resolution.py new file mode 100644 index 000000000..458582fc2 --- /dev/null +++ b/tests/server/task/test_stage_reference_resolution.py @@ -0,0 +1,213 @@ +"""Dispatcher stage-reference resolution tests.""" + +import json +import logging +from pathlib import Path +from types import SimpleNamespace +from typing import cast + +from server.dispatcher.base import Dispatcher +from server.registries.worker import WorkerRegistry +from server.task.models import TaskRecord, TaskStatus +from server.task.parser import parse_workflow +from server.task.runtime import TaskRuntime +from shared.schemas.result import APIItem, APIResult, ResultEnvelope +from shared.tasks import TaskEnvelopeTemplate, TaskType +from shared.tasks.specs import ApiSpecTemplate + + +class _DummyRuntime: + def __init__( + self, + tasks: dict[str, TaskRecord], + depends_on: dict[str, list[str]] | None = None, + ) -> None: + self.tasks = tasks + self._depends_on = depends_on or {} + + def get_record(self, task_id: str) -> TaskRecord | None: + return self.tasks.get(task_id) + + def describe_task(self, task_id: str) -> SimpleNamespace | None: + record = self.tasks.get(task_id) + if record is None: + return None + return SimpleNamespace(depends_on=list(self._depends_on.get(task_id, []))) + + +def _task_template(task_type: TaskType, **spec_updates: object) -> TaskEnvelopeTemplate: + payload = { + "apiVersion": "flowmesh/v1", + "kind": "Task", + "metadata": {"name": "wf:task"}, + "spec": {"taskType": task_type.value, **spec_updates}, + } + return TaskEnvelopeTemplate.model_validate(payload) + + +def test_api_dependent_stage_resolves_first_row_text(tmp_path: Path) -> None: + """A dependent stage's ${stage.items.0.text} resolves to the first row's + text of a batch-only APIResult.""" + stage_dir = tmp_path / "task-api" + stage_dir.mkdir() + result = APIResult( + ok=True, + executor="api", + method="POST", + url="https://api.example.com/v1/chat/completions", + status_code=200, + items=[ + APIItem( + index=0, + url="https://api.example.com/v1/chat/completions", + status_code=200, + text="first row text", + prompt="first", + ), + APIItem( + index=1, + url="https://api.example.com/v1/chat/completions", + status_code=200, + text="second row text", + prompt="second", + ), + ], + ) + (stage_dir / "results.json").write_text( + json.dumps( + { + "task_id": "task-api", + "result": json.loads( + ResultEnvelope(task_id="task-api", result=result).model_dump_json() + )["result"], + } + ), + encoding="utf-8", + ) + + upstream = TaskRecord( + task_id="task-api", + workflow_id="wf-1", + owner_id="owner", + source="raw", + task=_task_template(TaskType.API), + status=TaskStatus.DONE, + task_type="api", + local_name="stage", + ) + dispatcher = Dispatcher( + runtime=cast(TaskRuntime, _DummyRuntime({})), + worker_registry=cast(WorkerRegistry, object()), + results_dir=tmp_path, + logger=logging.getLogger("test-api-dependent-stage"), + ) + + value = dispatcher._resolve_reference("stage.items.0.text", {"stage": upstream}) + assert value == "first row text" + + +def test_translated_n8n_dependent_api_stage_resolves(tmp_path: Path) -> None: + """A translated n8n workflow's dependent API stage resolves end to end.""" + payload = { + "nodes": [ + { + "name": "Upstream", + "type": "@n8n/n8n-nodes-langchain.openAi", + "parameters": { + "modelId": {"value": "gpt-4"}, + "responses": {"values": [{"content": "First answer"}]}, + }, + }, + { + "name": "Downstream", + "type": "@n8n/n8n-nodes-langchain.openAi", + "parameters": { + "modelId": {"value": "gpt-4"}, + "responses": {"values": [{"content": "Simplify this"}]}, + }, + }, + ], + "connections": { + "Upstream": {"ai_languageModel": [[{"node": "Downstream"}]]}, + }, + } + parsed = parse_workflow(json.dumps(payload), "n8n") + by_name = {t.graph_node_name: t for t in parsed.tasks} + upstream = by_name["Upstream"] + downstream = by_name["Downstream"] + assert downstream.depends_on == [upstream.task_id] + + stage_dir = tmp_path / upstream.task_id + stage_dir.mkdir() + result = APIResult( + ok=True, + executor="api", + method="POST", + url="https://api.example.com/v1/chat/completions", + status_code=200, + items=[ + APIItem( + index=0, + url="https://api.example.com/v1/chat/completions", + status_code=200, + text="first row text", + prompt="first", + ), + ], + ) + (stage_dir / "results.json").write_text( + json.dumps( + { + "task_id": upstream.task_id, + "result": json.loads( + ResultEnvelope( + task_id=upstream.task_id, result=result + ).model_dump_json() + )["result"], + } + ), + encoding="utf-8", + ) + + upstream_record = TaskRecord( + task_id=upstream.task_id, + workflow_id="wf-1", + owner_id="owner", + source="raw", + task=upstream.task, + status=TaskStatus.DONE, + task_type="api", + graph_node_name="Upstream", + ) + downstream_record = TaskRecord( + task_id=downstream.task_id, + workflow_id="wf-1", + owner_id="owner", + source="raw", + task=downstream.task, + status=TaskStatus.PENDING, + task_type="api", + graph_node_name="Downstream", + ) + dispatcher = Dispatcher( + runtime=cast( + TaskRuntime, + _DummyRuntime( + { + upstream.task_id: upstream_record, + downstream.task_id: downstream_record, + }, + depends_on={downstream.task_id: [upstream.task_id]}, + ), + ), + worker_registry=cast(WorkerRegistry, object()), + results_dir=tmp_path, + logger=logging.getLogger("test-n8n-dependent-stage"), + ) + + context = dispatcher._build_stage_context(downstream_record) + spec = cast(ApiSpecTemplate, downstream.task.spec) + resolved = dispatcher._resolve_placeholders(spec.data, context) + assert resolved["items"][0] == ( + "The previous stage's response is as follows. Simplify this\n" "first row text" + ) diff --git a/tests/server/test_redact.py b/tests/server/test_redact.py index cd1c06bf9..3b9b0523e 100644 --- a/tests/server/test_redact.py +++ b/tests/server/test_redact.py @@ -301,7 +301,7 @@ def test_dump_redacts_all_five_locations(self) -> None: "data": {"access_token": "D"}, } rec = _record(api) - dumped = rec.model_dump() + dumped = rec.model_dump(by_alias=True) dumped_api = dumped["task"]["spec"]["api"] for field in ("headers", "params", "body", "json", "data"): assert list(dumped_api[field].values()) == [REDACTED], field @@ -323,7 +323,8 @@ def test_in_memory_keeps_real_credential(self) -> None: rec = _record({"headers": {"Authorization": "Bearer SECRET"}}) assert isinstance(rec.task.spec, ApiSpecTemplate) assert rec.task.spec.api is not None - assert rec.task.spec.api["headers"]["Authorization"] == "Bearer SECRET" + assert rec.task.spec.api.headers is not None + assert rec.task.spec.api.headers["Authorization"] == "Bearer SECRET" def test_dump_excluding_task_does_not_raise(self) -> None: rec = _record({"headers": {"Authorization": "Bearer SECRET"}}) @@ -393,8 +394,10 @@ def test_spec_dump_honors_exclude_defaults(self) -> None: def test_no_credential_unchanged(self) -> None: api = {"url": "http://x", "json": {"model": "gpt"}} rec = _record(api) - dumped = rec.model_dump() - assert dumped["task"]["spec"]["api"] == api + dumped = rec.model_dump(by_alias=True) + assert dumped["task"]["spec"]["api"]["url"] == "http://x" + assert dumped["task"]["spec"]["api"]["json"] == {"model": "gpt"} + assert dumped["task"]["spec"]["api"]["headers"] is None def test_redaction_cached_across_dumps(self) -> None: rec = _record( diff --git a/tests/shared/test_api_spec.py b/tests/shared/test_api_spec.py new file mode 100644 index 000000000..0faf94256 --- /dev/null +++ b/tests/shared/test_api_spec.py @@ -0,0 +1,278 @@ +"""Tests for the typed ``spec.api`` model.""" + +import os +import subprocess +import sys +from typing import Any + +import pytest +from pydantic import ValidationError + +from shared.tasks import TaskEnvelopeStrict, TaskEnvelopeTemplate +from shared.tasks.specs import ApiSpecStrict, ApiSpecTemplate +from shared.tasks.specs.misc import ( + _MAX_CONCURRENCY, + _MAX_RETRIES, + ApiConfig, + ApiConfigTemplate, +) +from shared.utils.redact import REDACTED + + +def _strict(**fields: Any) -> ApiSpecStrict: + return ApiSpecStrict.model_validate({"taskType": "api", **fields}) + + +def _template(**fields: Any) -> ApiSpecTemplate: + return ApiSpecTemplate.model_validate({"taskType": "api", **fields}) + + +def _api(**fields: Any) -> ApiConfig: + api = _strict(api=fields).api + assert api is not None + return api + + +class TestRetries: + def test_maximum_accepted(self) -> None: + assert _api(retries=_MAX_RETRIES).retries == _MAX_RETRIES + + def test_above_maximum_rejected(self) -> None: + with pytest.raises(ValidationError): + _strict(api={"retries": _MAX_RETRIES + 1}) + + def test_zero_accepted(self) -> None: + assert _api(retries=0).retries == 0 + + def test_negative_rejected(self) -> None: + with pytest.raises(ValidationError): + _strict(api={"retries": -1}) + + def test_bool_rejected(self) -> None: + with pytest.raises(ValidationError): + _strict(api={"retries": True}) + + def test_string_rejected(self) -> None: + with pytest.raises(ValidationError): + _strict(api={"retries": "2"}) + + +class TestConcurrency: + def test_maximum_accepted(self) -> None: + assert _api(concurrency=_MAX_CONCURRENCY).concurrency == _MAX_CONCURRENCY + + def test_above_maximum_rejected(self) -> None: + with pytest.raises(ValidationError): + _strict(api={"concurrency": _MAX_CONCURRENCY + 1}) + + def test_one_accepted(self) -> None: + assert _api(concurrency=1).concurrency == 1 + + def test_zero_rejected(self) -> None: + with pytest.raises(ValidationError): + _strict(api={"concurrency": 0}) + + def test_negative_rejected(self) -> None: + with pytest.raises(ValidationError): + _strict(api={"concurrency": -1}) + + def test_bool_rejected(self) -> None: + with pytest.raises(ValidationError): + _strict(api={"concurrency": True}) + + def test_string_rejected(self) -> None: + with pytest.raises(ValidationError): + _strict(api={"concurrency": "2"}) + + +class TestMethod: + def test_upper_cased(self) -> None: + assert _api(method="get").method == "GET" + + def test_default_is_post(self) -> None: + assert ApiConfig.model_validate({}).method == "POST" + + def test_template_keeps_placeholder(self) -> None: + api = _template(api={"method": "${stage.method}"}).api + assert api is not None + assert api.method == "${stage.method}" + + +class TestTemplatePlaceholders: + def test_retries_placeholder_accepted(self) -> None: + api = _template(api={"retries": "${stage.items.0.text}"}).api + assert api is not None + assert api.retries == "${stage.items.0.text}" + + def test_concurrency_placeholder_accepted(self) -> None: + api = _template(api={"concurrency": "${stage.items.0.text}"}).api + assert api is not None + assert api.concurrency == "${stage.items.0.text}" + + def test_timeout_placeholder_accepted(self) -> None: + api = _template(api={"timeout_sec": "${stage.items.0.text}"}).api + assert api is not None + assert api.timeout_sec == "${stage.items.0.text}" + + def test_max_body_bytes_placeholder_accepted(self) -> None: + api = _template( + api={"response": {"max_body_bytes": "${stage.items.0.text}"}} + ).api + assert api is not None + assert api.response is not None + assert api.response.max_body_bytes == "${stage.items.0.text}" + + def test_strict_rejects_placeholder(self) -> None: + with pytest.raises(ValidationError): + _strict(api={"retries": "${stage.items.0.text}"}) + + +class TestTemplateBounds: + def test_retries_above_max_rejected(self) -> None: + with pytest.raises(ValidationError): + _template(api={"retries": _MAX_RETRIES + 1}) + + def test_retries_negative_rejected(self) -> None: + with pytest.raises(ValidationError): + _template(api={"retries": -1}) + + def test_concurrency_above_max_rejected(self) -> None: + with pytest.raises(ValidationError): + _template(api={"concurrency": _MAX_CONCURRENCY + 1}) + + def test_concurrency_zero_rejected(self) -> None: + with pytest.raises(ValidationError): + _template(api={"concurrency": 0}) + + def test_timeout_nonpositive_rejected(self) -> None: + with pytest.raises(ValidationError): + _template(api={"timeout_sec": 0}) + + def test_max_body_bytes_nonpositive_rejected(self) -> None: + with pytest.raises(ValidationError): + _template(api={"response": {"max_body_bytes": 0}}) + + def test_retries_placeholder_accepted(self) -> None: + api = _template(api={"retries": "${stage.items.0.text}"}).api + assert api is not None + assert api.retries == "${stage.items.0.text}" + + +class TestJsonAlias: + def test_json_key_validates(self) -> None: + api = _api(json={"model": "gpt"}) + assert api.json_body == {"model": "gpt"} + + def test_dumps_back_as_json_key(self) -> None: + api = _api(json={"model": "gpt"}) + dumped = api.model_dump(by_alias=True) + assert "json" in dumped + assert dumped["json"] == {"model": "gpt"} + + def test_round_trips_to_equal_model(self) -> None: + api = _api(json={"model": "gpt"}) + dumped = api.model_dump(by_alias=True) + assert ApiConfig.model_validate(dumped) == api + + def test_template_json_key(self) -> None: + api = _template(api={"json": {"model": "gpt"}}).api + assert api is not None + assert api.json_body == {"model": "gpt"} + + +class TestTemplateInstanceConversion: + def test_template_instance_with_json_keeps_value(self) -> None: + template = _template(api={"json": {"model": "gpt"}}).api + assert template is not None + strict = ApiConfig.model_validate(template) + assert strict.json_body == {"model": "gpt"} + + def test_template_instance_without_json_keeps_none(self) -> None: + template = _template(api={}).api + assert template is not None + strict = ApiConfig.model_validate(template) + assert strict.json_body is None + + def test_envelope_conversion_dumps_json_key(self) -> None: + template = TaskEnvelopeTemplate.model_validate( + { + "apiVersion": "flowmesh/v1", + "kind": "APITask", + "spec": { + "taskType": "api", + "data": {"type": "list", "items": ["hi"]}, + "api": {"json": {"model": "gpt"}}, + }, + } + ) + strict = TaskEnvelopeStrict.model_validate(template) + dumped = strict.model_dump_json(by_alias=True) + assert '"json":{"model":"gpt"}' in dumped + + +class TestWarningFreeImport: + def test_import_raises_no_warning(self) -> None: + src = os.path.join(os.path.dirname(__file__), "..", "..", "src") + env = {**os.environ, "PYTHONPATH": os.path.abspath(src)} + result = subprocess.run( + [ + sys.executable, + "-W", + "error::UserWarning", + "-c", + "import shared.tasks.specs.misc", + ], + capture_output=True, + text=True, + env=env, + ) + assert result.returncode == 0, result.stderr + + +class TestRedaction: + def test_redacts_credential_fields(self) -> None: + spec = _strict( + api={ + "headers": {"Authorization": "Bearer H"}, + "params": {"api_key": "P"}, + "body": {"secret": "B"}, + "json": {"token": "J"}, + "data": {"access_token": "D"}, + } + ) + redacted = spec.redact_credentials() + dumped = redacted.model_dump(by_alias=True)["api"] + for field in ("headers", "params", "body", "json", "data"): + assert list(dumped[field].values()) == [REDACTED], field + assert redacted.has_redacted_credentials() + + def test_in_memory_keeps_real_credential(self) -> None: + spec = _strict(api={"headers": {"Authorization": "Bearer SECRET"}}) + assert spec.api is not None + assert spec.api.headers is not None + assert spec.api.headers["Authorization"] == "Bearer SECRET" + assert not spec.has_redacted_credentials() + + def test_template_redacts(self) -> None: + spec = _template(api={"headers": {"Authorization": "Bearer SECRET"}}) + redacted = spec.redact_credentials() + assert redacted.model_dump()["api"]["headers"]["Authorization"] == REDACTED + + +class TestDefaults: + def test_defaults(self) -> None: + api = ApiConfig.model_validate({}) + assert api.method == "POST" + assert api.timeout_sec == 60 + assert api.verify_tls is True + assert api.follow_redirects is True + assert api.retries == 0 + assert api.concurrency == _MAX_CONCURRENCY + assert api.response is None + + def test_template_defaults(self) -> None: + api = ApiConfigTemplate.model_validate({}) + assert api.method == "POST" + assert api.timeout_sec == 60.0 + assert api.retries == 0 + assert api.concurrency == _MAX_CONCURRENCY diff --git a/tests/worker/test_api_executor.py b/tests/worker/test_api_executor.py index af3630650..eb5b04874 100644 --- a/tests/worker/test_api_executor.py +++ b/tests/worker/test_api_executor.py @@ -14,11 +14,12 @@ import httpx import pytest +from pydantic import ValidationError +from shared.tasks.specs.misc import _MAX_CONCURRENCY, _MAX_RETRIES from shared.tasks.worker_message import WorkerTaskMessage from worker.executors import api_executor as api_executor_module from worker.executors.api_executor import ( - _MAX_RETRIES, _RETRY_BACKOFF_MAX_SEC, APIExecutor, ) @@ -367,11 +368,10 @@ def test_cancelled_task_stops_retrying(self) -> None: assert transport.calls == 0 def test_invalid_retries_rejected(self) -> None: - """A negative or non-integer retries value is rejected.""" - for bad in (-1, "2", 1.5, True): - task = self._task(retries=bad) - with pytest.raises(ExecutionError, match="spec.api.retries"): - _run(_executor(), task, _RecordingTransport()) + """A negative, non-integer, or out-of-range retries value is rejected.""" + for bad in (-1, "2", 1.5, True, _MAX_RETRIES + 1): + with pytest.raises(ValidationError): + self._task(retries=bad) def test_cancel_previous_task_does_not_cancel_next(self) -> None: """A cancellation left over from a prior task does not cancel the next.""" @@ -609,12 +609,6 @@ def test_retry_after_is_capped(self) -> None: delays = self._run_recording_delays(task, transport) assert delays == [_RETRY_BACKOFF_MAX_SEC] - def test_retries_above_maximum_rejected(self) -> None: - """A retries value above the maximum is rejected.""" - task = self._task(retries=_MAX_RETRIES + 1) - with pytest.raises(ExecutionError, match=f"at most {_MAX_RETRIES}"): - _run(_executor(), task, _RecordingTransport()) - def test_one_warning_per_retry(self, caplog: pytest.LogCaptureFixture) -> None: """Each retry logs one warning naming the attempt and the delay.""" task = self._task(retries=2) @@ -736,11 +730,43 @@ def _handler(self, request: httpx.Request) -> httpx.Response: }, ) - task = _batch_task(["a", "b", "c"], response={"raise_for_status": True}) + task = _batch_task( + ["a", "b", "c"], + concurrency=1, + response={"raise_for_status": True}, + ) transport = _FailRow("b") with pytest.raises(ExecutionError, match="row 1"): _run(_executor(), task, transport, tmp_path) - assert len(transport.requests) == 3 + assert len(transport.requests) == 2 + + def test_row_failure_stops_later_rows_from_sending(self, tmp_path: Path) -> None: + """With concurrency 1, a failing first row stops later rows from sending.""" + + class _FailFirst(httpx.MockTransport): + def __init__(self) -> None: + self.requests: list[httpx.Request] = [] + super().__init__(self._handler) + + def _handler(self, request: httpx.Request) -> httpx.Response: + self.requests.append(request) + return httpx.Response( + 500, + json={ + "choices": [{"message": {"content": "boom"}}], + "usage": {"total_tokens": 1}, + }, + ) + + task = _batch_task( + ["a", "b", "c"], + concurrency=1, + response={"raise_for_status": True}, + ) + transport = _FailFirst() + with pytest.raises(ExecutionError, match="row 0"): + _run(_executor(), task, transport, tmp_path) + assert len(transport.requests) == 1 def test_placeholder_not_required_for_scalar_body(self, tmp_path: Path) -> None: """A batch task whose body has no placeholder still issues N requests.""" @@ -928,7 +954,7 @@ def test_request_skeleton_constructed_once(self, tmp_path: Path) -> None: _executor().run(task, tmp_path) assert mock_build.call_count == 1 - assert mock_build.call_args.args[2] is None + assert len(mock_build.call_args.args) == 2 @pytest.mark.parametrize("concurrency", [1, 4, 8]) def test_client_pool_sized_to_concurrency(self, concurrency: int) -> None: @@ -948,25 +974,16 @@ def test_client_pool_sized_to_concurrency(self, concurrency: int) -> None: finally: APIExecutor.close_all_clients() - def test_concurrency_capped_at_max(self, tmp_path: Path) -> None: - """A configured concurrency above the cap is clamped to the cap.""" - task = _batch_task(["a", "b", "c"], concurrency=100) - captured: dict[str, Any] = {} - - def _fake_get_client(*args: Any, **kwargs: Any) -> httpx.Client: - captured["concurrency"] = kwargs.get("concurrency", args[4]) - return httpx.Client(transport=_EchoTransport()) - - with patch.object(APIExecutor, "_get_client", side_effect=_fake_get_client): - _executor().run(task, tmp_path) - - assert captured["concurrency"] == 8 + def test_concurrency_above_max_rejected(self) -> None: + """A concurrency above the cap is rejected, not clamped.""" + with pytest.raises(ValidationError): + _batch_task(["a", "b", "c"], concurrency=_MAX_CONCURRENCY + 1) @pytest.mark.parametrize("concurrency", [1, 4]) def test_run_passes_effective_concurrency_to_client( self, tmp_path: Path, concurrency: int ) -> None: - """run() forwards the uncapped configured concurrency to the client.""" + """run() forwards the configured concurrency to the client.""" task = _batch_task(["a", "b", "c"], concurrency=concurrency) captured: dict[str, Any] = {} @@ -979,14 +996,11 @@ def _fake_get_client(*args: Any, **kwargs: Any) -> httpx.Client: assert captured["concurrency"] == concurrency - @pytest.mark.parametrize("concurrency", [0, -1]) - def test_concurrency_below_one_rejected( - self, tmp_path: Path, concurrency: int - ) -> None: - """A configured concurrency below 1 is rejected.""" - task = _batch_task(["a", "b", "c"], concurrency=concurrency) - with pytest.raises(ExecutionError, match="spec.api.concurrency must be >= 1"): - _run(_executor(), task, _EchoTransport(), tmp_path) + @pytest.mark.parametrize("concurrency", [0, -1, "2", 1.5, True]) + def test_invalid_concurrency_rejected(self, concurrency: object) -> None: + """A concurrency below 1, or a non-int or bool, is rejected.""" + with pytest.raises(ValidationError): + _batch_task(["a", "b", "c"], concurrency=concurrency) def test_client_cache_key_includes_concurrency(self) -> None: """Pools built for different concurrency values are not shared.""" From b1fe86684ad85b96ecf5a2216ce66445a4d847ff Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Mon, 28 Sep 2026 22:20:45 +0700 Subject: [PATCH 53/71] Fix PR comments. Signed-off-by: Zhengyuan Su --- docs/WORKFLOWS.md | 4 +- examples/templates/api_two_stage.yaml | 2 +- src/server/task/n8n_parser.py | 1 - src/worker/executors/api_executor.py | 72 +++++++++++++++++++-------- tests/server/task/test_n8n_parser.py | 36 ++++++++++++++ tests/worker/test_api_executor.py | 57 ++++++++++++++++++++- 6 files changed, 147 insertions(+), 25 deletions(-) diff --git a/docs/WORKFLOWS.md b/docs/WORKFLOWS.md index a4da7bb5b..646403fa9 100644 --- a/docs/WORKFLOWS.md +++ b/docs/WORKFLOWS.md @@ -82,7 +82,7 @@ By default it routes to the Nebula endpoint and authenticates with the worker's Credential handling: a caller-supplied `Authorization` header is always used as-is and never overwritten. With no header, `NEBULA_API_TOKEN` is injected only when the call is on the Nebula url (no custom `spec.api.url`) — the Nebula token is never sent to a custom endpoint. A Nebula-path call with no token available fails closed. -`spec.api.retries` (default `0`, an integer from 0 to 10; any other value is rejected when the workflow is submitted) sets how many times a transient failure is retried before the task fails. A transient failure is a connection error or an HTTP status of 5xx, 408, or 429; other 4xx statuses are never retried. Retries back off exponentially: the first waits 1s and each later one doubles, capped at 60s. When a retryable response carries a `Retry-After` header (seconds or an HTTP date), that wait is used instead, also capped at 60s. Each retry logs a warning with the attempt count and the wait. A cancelled task stops retrying immediately. +`spec.api.retries` (default `0`, at most `10`) sets how many times a transient failure is retried before the task fails. A transient failure is a connection error or an HTTP status of 5xx, 408, or 429; other 4xx statuses are never retried. Retries back off exponentially: the first waits 1s and each later one doubles, capped at 60s. When a retryable response carries a `Retry-After` header (seconds or an HTTP date), that wait is used instead, also capped at 60s. Each retry logs a warning with the attempt count and the wait. A cancelled task stops retrying immediately. ```yaml spec: @@ -109,7 +109,7 @@ spec: Each row's prompt replaces `{{prompt}}` in the request body; a value that is exactly `{{prompt}}` takes the prompt as-is, so a message-list row fills -`messages`. `spec.api.concurrency` (default 8, an integer from 1 to 8; any other value is rejected when the workflow is submitted) bounds in-flight +`messages`. `spec.api.concurrency` (default and maximum 8) bounds in-flight requests. Any failed row fails the task. Cancelling the task skips rows that have not started and marks it cancelled once in-flight requests return. diff --git a/examples/templates/api_two_stage.yaml b/examples/templates/api_two_stage.yaml index 6e9e28dbd..365b79df7 100644 --- a/examples/templates/api_two_stage.yaml +++ b/examples/templates/api_two_stage.yaml @@ -37,7 +37,7 @@ spec: - role: user content: "{{prompt}}" response: - parse_json: false + parse_json: true return_body: true raise_for_status: true diff --git a/src/server/task/n8n_parser.py b/src/server/task/n8n_parser.py index a052c0573..f3e8702c9 100644 --- a/src/server/task/n8n_parser.py +++ b/src/server/task/n8n_parser.py @@ -253,7 +253,6 @@ def _build_api_node_spec( }, } if api_key := credential_data.get("api_key"): - api_spec["key"] = api_key headers["Authorization"] = f"Bearer {api_key}" if api_url := credential_data.get("url"): api_spec["url"] = api_url diff --git a/src/worker/executors/api_executor.py b/src/worker/executors/api_executor.py index c15a3a40d..0edf9a686 100644 --- a/src/worker/executors/api_executor.py +++ b/src/worker/executors/api_executor.py @@ -4,6 +4,7 @@ import math import os import threading +import time from concurrent.futures import ThreadPoolExecutor, as_completed from datetime import UTC, datetime from pathlib import Path @@ -155,17 +156,21 @@ def _request_with_retries( params: dict[str, Any] | None, request_kwargs: dict[str, Any], retries: int, + failed: threading.Event, ) -> httpx.Response: """Issue the request, retrying transient failures up to ``retries`` times. A retryable failure is a connection error or a transient HTTP status - (5xx, 408, 429). Non-retryable failures and a cancelled task stop the - loop immediately. The final attempt's failure propagates to the caller. + (5xx, 408, 429). Non-retryable failures, a cancelled task, or a row + already failed elsewhere stop the loop immediately. The final attempt's + failure propagates to the caller. """ attempt = 0 while True: if self._cancel_event.is_set(): raise TaskCancelledError("API request cancelled") + if failed.is_set(): + raise ExecutionError("API task failed on an earlier row") try: resp = client.request( method, @@ -186,7 +191,7 @@ def _request_with_retries( exc, delay, ) - self._wait_for_backoff(delay) + self._wait_for_backoff(delay, failed) continue if resp.is_error and _is_retryable_status(resp.status_code): if attempt >= retries: @@ -200,7 +205,7 @@ def _request_with_retries( retries, delay, ) - self._wait_for_backoff(delay) + self._wait_for_backoff(delay, failed) continue return resp @@ -233,10 +238,19 @@ def _retry_after_seconds(resp: httpx.Response) -> float | None: retry_at = retry_at.replace(tzinfo=UTC) return max(0.0, (retry_at - datetime.now(UTC)).total_seconds()) - def _wait_for_backoff(self, delay: float) -> None: - """Wait out the retry backoff, aborting early if the task is cancelled.""" - if self._cancel_event.wait(delay): - raise TaskCancelledError("API request cancelled") + def _wait_for_backoff(self, delay: float, failed: threading.Event) -> None: + """Wait out the retry backoff, aborting early if the task is cancelled + or another row has already failed.""" + deadline = time.monotonic() + delay + while True: + if self._cancel_event.is_set(): + raise TaskCancelledError("API request cancelled") + if failed.is_set(): + raise ExecutionError("API task failed on an earlier row") + remaining = deadline - time.monotonic() + if remaining <= 0: + return + self._cancel_event.wait(min(remaining, 0.1)) @classmethod def close_all_clients(cls) -> None: @@ -431,33 +445,49 @@ def _issue(idx: int, prompt: Any) -> APIItem: params, kwargs, retries, + failed, ) + except TaskCancelledError: + raise except httpx.RequestError as exc: - failed.set() error = ExecutionError( f"API request failed (row {idx}): {exc}", retryable=True ) - first_error.append(error) + if not first_error: + first_error.append(error) + failed.set() raise error from exc + except BaseException as exc: + if not first_error: + first_error.append(exc) + failed.set() + raise if raise_for_status and resp.is_error: - failed.set() message = f"API request returned status {resp.status_code} (row {idx})" body_text = resp.text[:200] if body_text: message = f"{message}: {body_text}" retryable = _is_retryable_status(resp.status_code) error = ExecutionError(message, retryable=retryable) - first_error.append(error) + if not first_error: + first_error.append(error) + failed.set() raise error - item = self._parse_response( - resp, - response_cfg=response_cfg, - max_body_bytes=max_body_bytes, - idx=idx, - prompt_str=prompt_str, - ) + try: + item = self._parse_response( + resp, + response_cfg=response_cfg, + max_body_bytes=max_body_bytes, + idx=idx, + prompt_str=prompt_str, + ) + except BaseException as exc: + if not first_error: + first_error.append(exc) + failed.set() + raise if self._cancel_event.is_set(): raise TaskCancelledError("API task cancelled") @@ -479,7 +509,9 @@ def _issue(idx: int, prompt: Any) -> APIItem: raise TaskCancelledError("API task cancelled") try: results[idx] = future.result() - except ExecutionError: + except TaskCancelledError: + raise + except BaseException: pass if first_error: raise first_error[0] diff --git a/tests/server/task/test_n8n_parser.py b/tests/server/task/test_n8n_parser.py index 33533bd3f..cc877948a 100644 --- a/tests/server/task/test_n8n_parser.py +++ b/tests/server/task/test_n8n_parser.py @@ -1,11 +1,13 @@ """Tests for n8n workflow translation.""" import base64 +import json import pytest from server.task.n8n_parser import _decode_secret_part, translate_n8n_workflow from server.task.parser import parse_workflow +from shared.tasks.specs import ApiSpecTemplate class TestTranslateN8nWorkflow: @@ -84,6 +86,40 @@ def test_api_dependency_resolves_first_row_text(self) -> None: "${Upstream.items.0.text}" ] + def test_openai_credential_parses_with_header_and_no_key(self) -> None: + """An OpenAI credential yields an Authorization header and no ``key`` + field, so the workflow validates through the submitted path.""" + nodes = [ + { + "name": "Chat", + "type": "@n8n/n8n-nodes-langchain.openAi", + "parameters": { + "modelId": {"value": "gpt-4"}, + "responses": { + "values": [{"content": "Hello, world!"}], + }, + }, + "credentials": { + "openAiApi": { + "data": {"apiKey": "sk-secret"}, + } + }, + } + ] + result = translate_n8n_workflow({"nodes": nodes, "connections": {}}) + api = result["spec"]["api"] + assert api["headers"]["Authorization"] == "Bearer sk-secret" + assert "key" not in api + + parsed = parse_workflow( + json.dumps({"nodes": nodes, "connections": {}}), format="n8n" + ) + spec = parsed.tasks[0].task.spec + assert isinstance(spec, ApiSpecTemplate) + assert spec.api is not None + assert spec.api.headers is not None + assert spec.api.headers["Authorization"] == "Bearer sk-secret" + class TestDecodeSecretPart: def test_hex_decode(self) -> None: diff --git a/tests/worker/test_api_executor.py b/tests/worker/test_api_executor.py index eb5b04874..ab9a6d5a9 100644 --- a/tests/worker/test_api_executor.py +++ b/tests/worker/test_api_executor.py @@ -526,7 +526,7 @@ def _run_recording_delays( """Run a task, recording each backoff delay instead of waiting.""" delays: list[float] = [] - def _record(delay: float) -> None: + def _record(delay: float, failed: threading.Event) -> None: delays.append(delay) executor = APIExecutor.__new__(APIExecutor) @@ -1075,6 +1075,61 @@ def _handler(self, request: httpx.Request) -> httpx.Response: _run(_executor(), task, _ErrorBody(), tmp_path) assert excinfo.value.retryable is True + def test_parse_error_fails_with_mapping_error_not_keyerror( + self, tmp_path: Path + ) -> None: + """A 200 with a non-mapping body and parse_json true fails the task with + the parse error, not a KeyError from the collection loop.""" + + class _ArrayBody(httpx.MockTransport): + def __init__(self) -> None: + super().__init__(self._handler) + + def _handler(self, request: httpx.Request) -> httpx.Response: + return httpx.Response(200, json=[1]) + + task = _batch_task(["a", "b"], response={"parse_json": True}) + with pytest.raises(ExecutionError, match="not a valid JSON mapping"): + _run(_executor(), task, _ArrayBody(), tmp_path) + + def test_retry_stops_when_another_row_fails( + self, tmp_path: Path, monkeypatch: pytest.MonkeyPatch + ) -> None: + """A row already retrying stops as soon as another row fails, issuing no + further requests and returning well under the full backoff time.""" + monkeypatch.setattr("worker.executors.api_executor._RETRY_BACKOFF_SEC", 60.0) + + class _MixedTransport(httpx.MockTransport): + def __init__(self) -> None: + self.requests: list[httpx.Request] = [] + self.b_issued = threading.Event() + super().__init__(self._handler) + + def _handler(self, request: httpx.Request) -> httpx.Response: + self.requests.append(request) + prompt = json.loads(request.read())["messages"][0]["content"] + if prompt == "b": + self.b_issued.set() + return httpx.Response(503, json={"error": "overloaded"}) + # Row a waits until row b has issued its first (retrying) + # request, so row b is mid-backoff when row a fails. + self.b_issued.wait(5.0) + return httpx.Response(400, json={"error": "boom"}) + + task = _batch_task( + ["a", "b"], + concurrency=2, + retries=3, + response={"raise_for_status": True}, + ) + transport = _MixedTransport() + start = time.monotonic() + with pytest.raises(ExecutionError, match="row 0"): + _run(_executor(), task, transport, tmp_path) + elapsed = time.monotonic() - start + assert elapsed < 1.0 + assert len(transport.requests) == 2 + def test_cancel_prevents_queued_rows_from_issuing(self, tmp_path: Path) -> None: """After cancel, a row that has not started never issues its request.""" From 5fee75167c4458bf566f0df43a4966bcf1560020 Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Wed, 30 Sep 2026 15:50:48 +0700 Subject: [PATCH 54/71] refactor: remove echo data.type function mode The echo function mode ran whole-list caller code in-process through safe_eval, duplicating the python task and exposing the same escapable namespace. Remove the executor branch, argument resolution, the expect_list switch and its safe_eval-only changes, their tests, and the docs. Echo data.type list stays. Signed-off-by: Zhengyuan Su Co-Authored-By: Claude Code --- docs/EXECUTORS.md | 22 +-- src/worker/executors/echo_executor.py | 76 ++--------- src/worker/executors/utils/safe_eval.py | 95 +++---------- tests/worker/test_echo_executor.py | 169 ------------------------ tests/worker/test_safe_eval.py | 116 ---------------- 5 files changed, 30 insertions(+), 448 deletions(-) delete mode 100644 tests/worker/test_echo_executor.py delete mode 100644 tests/worker/test_safe_eval.py diff --git a/docs/EXECUTORS.md b/docs/EXECUTORS.md index c1e9c3c8b..1ba70cc97 100644 --- a/docs/EXECUTORS.md +++ b/docs/EXECUTORS.md @@ -5,7 +5,7 @@ The worker resolves `spec.taskType` against an executor registry in | `taskType` | Executor | Use case | |-----------|----------|----------| -| `echo` | `EchoExecutor` | Echo input back as result, or run a sandboxed function over upstream values | +| `echo` | `EchoExecutor` | Echo input back as result (smoke tests) | | `inference` | `VLLMExecutor` / `TransformersExecutor` | LLM inference | | `embedding` | `VLLMEmbeddingExecutor` (text, when `model.vllm` is set) / `TransformersExecutor` (visual, `model.transformers.mode: visual-embedding`) | Text / visual embeddings | | `diffusion` | `DiffusersExecutor` | Image / video diffusion models | @@ -68,26 +68,10 @@ Optional, for the search tools: ## Echo executor -`taskType: echo` returns input values back as the result, or runs a sandboxed -Python function over upstream values. It is useful for inspecting and shaping -data between stages. +`taskType: echo` returns input values back as the result. It is useful for +inspecting and shaping data between stages. `spec.data.type: list` echoes each `spec.data.items` entry. An entry is either a string literal or a mapping with an expression (`expr`, or both `node` and `path`) resolved against the upstream results. A resolved list is flattened into one echo item per element; a scalar becomes a single item. - -`spec.data.type: function` runs a sandboxed source function over upstream -values. `spec.data.function` is the source of a function that takes the -resolved arguments and returns a list of JSON values; each element becomes one -echo item (the list is not flattened further). `spec.data.arguments` is a list -of argument specs, each with exactly one of: - -- `{items: }` — the whole list of upstream values, passed as-is. -- `{expr: }` — a single expression resolved against the upstream - results. -- `{node: , path: }` — a node and path resolved against the - upstream results. - -The function is executed in a sandbox; it must return a list, and each element -is emitted as one echo item. A function failure fails the task. diff --git a/src/worker/executors/echo_executor.py b/src/worker/executors/echo_executor.py index 7bbc5f5e9..ff163e6d7 100644 --- a/src/worker/executors/echo_executor.py +++ b/src/worker/executors/echo_executor.py @@ -12,7 +12,6 @@ from .mixins.data import DataMixin from .utils.checkpoints import maybe_upload_traces from .utils.graph_templates import _evaluate_expr -from .utils.safe_eval import safe_execute_function, safe_materialize_function logger = logging.getLogger(__name__) @@ -52,25 +51,6 @@ def _resolve_expr_item( ) return resolved - @staticmethod - def _resolve_function_arg( - arg: dict[str, Any], context: dict[str, BaseExecutorResult] - ) -> Any: - keys = frozenset(arg) - if keys == {"items"}: - items = arg["items"] - if not isinstance(items, list): - raise ExecutionError( - "echo executor function argument 'items' must be a list" - ) - return items - if keys in ({"expr"}, {"node", "path"}): - return EchoExecutor._resolve_expr_item(arg, context) - raise ExecutionError( - "echo executor function argument must have exactly one of " - f"'items', 'expr', or 'node'+'path'; got keys {sorted(keys)}" - ) - def _resolve_item( self, item: EchoItem, context: dict[str, BaseExecutorResult] ) -> Any: @@ -84,48 +64,6 @@ def _resolve_item( "a string literal or a mapping" ) - def _run_list( - self, data_cfg: dict[str, Any], context: dict[str, BaseExecutorResult] - ) -> list[EchoResultItem]: - items_cfg = data_cfg.get("items") - if not isinstance(items_cfg, list): - raise ExecutionError("echo executor requires spec.data.items to be a list") - merged_items: list[EchoResultItem] = [] - for item in items_cfg: - resolved = self._resolve_item(item, context) - self._append_outputs(merged_items, resolved) - return merged_items - - def _run_function( - self, - data_cfg: dict[str, Any], - context: dict[str, BaseExecutorResult], - task_id: str, - ) -> list[EchoResultItem]: - fn_code = data_cfg.get("function") - if not isinstance(fn_code, str) or not fn_code.strip(): - raise ExecutionError( - f"echo executor task {task_id} requires spec.data.function " - "to be a non-empty string" - ) - args_cfg = data_cfg.get("arguments") - if not isinstance(args_cfg, list): - raise ExecutionError( - f"echo executor task {task_id} requires spec.data.arguments " - "to be a list" - ) - resolved_args = [self._resolve_function_arg(arg, context) for arg in args_cfg] - try: - fn_obj = safe_materialize_function(fn_code) - output = safe_execute_function( - fn_obj, tuple(resolved_args), expect_list=True - ) - except Exception as e: - raise ExecutionError( - f"echo executor task {task_id} function failed: {e}" - ) from e - return [EchoResultItem(output=element) for element in output] - def run(self, task: ExecutorTask, out_dir: Path) -> EchoResult: spec = self.require_spec(task, EchoSpecStrict) task_id = task.task_id.strip() @@ -137,16 +75,20 @@ def run(self, task: ExecutorTask, out_dir: Path) -> EchoResult: if not isinstance(data_cfg, dict): raise ExecutionError("echo executor requires spec.data to be a mapping") + items_cfg = data_cfg.get("items") + if not isinstance(items_cfg, list): + raise ExecutionError( + "echo executor requires spec.data.items to be a list" + ) if not isinstance(context, dict): raise ExecutionError( "echo executor requires spec._upstreamResults to be a mapping" ) - data_type = data_cfg.get("type") - if data_type == "function": - merged_items = self._run_function(data_cfg, context, task_id) - else: - merged_items = self._run_list(data_cfg, context) + merged_items: list[EchoResultItem] = [] + for item in items_cfg: + resolved = self._resolve_item(item, context) + self._append_outputs(merged_items, resolved) result = EchoResult( items=merged_items, diff --git a/src/worker/executors/utils/safe_eval.py b/src/worker/executors/utils/safe_eval.py index f181d4a4c..b84c1522a 100644 --- a/src/worker/executors/utils/safe_eval.py +++ b/src/worker/executors/utils/safe_eval.py @@ -9,7 +9,7 @@ - Two-phase execution: materialize (compile) then execute (run) - Restricted builtins: only safe operations (no open, eval, exec, import, etc.) - Limited module access: json, re, math, numpy, pandas, pyarrow -- Type validation: enforces Callable[[tuple[str, ...]], Any] signature +- Type validation: enforces Callable[[tuple[str, ...]], str] signature - Isolated execution: exec() with explicit safe_globals/safe_locals Typical usage: @@ -17,7 +17,6 @@ result = safe_execute_function(fn_obj, ("hello",)) # Returns "HELLO" """ -import ast import inspect import json import math @@ -56,7 +55,6 @@ "all": all, "range": range, "isinstance": isinstance, - "ValueError": ValueError, } # Whitelist of safe modules available during function execution. @@ -71,24 +69,9 @@ } -def _is_json_value(value: Any) -> bool: - """Whether ``value`` is a JSON value (recursively, string keys, finite numbers).""" - if value is None or isinstance(value, (str, bool)): - return True - if isinstance(value, int): - return True - if isinstance(value, float): - return value == value and value not in (float("inf"), float("-inf")) - if isinstance(value, list): - return all(_is_json_value(item) for item in value) - if isinstance(value, dict): - return all(isinstance(k, str) and _is_json_value(v) for k, v in value.items()) - return False - - def safe_materialize_function( fn_code: str, -) -> Callable[[tuple[str | list[dict[str, str]], ...]], Any]: +) -> Callable[[tuple[str | list[dict[str, str]], ...]], str]: """ Compile function source code into a callable object with restricted builtins. @@ -103,8 +86,8 @@ def safe_materialize_function( Type signature enforcement: - Must accept exactly 1 parameter (tuple of strings or list of messages) - - Returns a string or a list of JSON values (validated at execution time) - - Signature: Callable[[tuple[str, ...]], Any] + - Should return a string (validated at execution time) + - Signature: Callable[[tuple[str, ...]], str] Args: fn_code: Python source code defining a function or lambda expression @@ -139,35 +122,6 @@ def safe_materialize_function( # Case 2: Function definition (use exec - def is a statement) else: - try: - tree = ast.parse(fn_code_stripped) - except SyntaxError as e: - raise RuntimeError( - f"Function definition failed: {e}\nCode: {fn_code}" - ) from e - - if len(tree.body) != 1: - kinds = [type(n).__name__ for n in tree.body] - raise RuntimeError( - "Function source must be a single function definition, " - f"found top-level statements: {kinds}" - ) - stmt = tree.body[0] - if isinstance(stmt, ast.FunctionDef): - fn_name = stmt.name - elif ( - isinstance(stmt, ast.Assign) - and len(stmt.targets) == 1 - and isinstance(stmt.targets[0], ast.Name) - and isinstance(stmt.value, ast.Lambda) - ): - fn_name = stmt.targets[0].id - else: - raise RuntimeError( - "Function source must be a single function definition or " - f"an assignment of a lambda, found: {type(stmt).__name__}" - ) - safe_locals: dict[str, Any] = {} # Execute the function definition (creates function object in locals) @@ -178,6 +132,12 @@ def safe_materialize_function( f"Function definition failed: {e}\nCode: {fn_code}" ) from e + # Find the function object + if not safe_locals: + raise RuntimeError("Function definition did not create any objects") + + # Get the function (usually the first/only item in locals) + fn_name = list(safe_locals.keys())[0] fn_obj = safe_locals[fn_name] if not callable(fn_obj): @@ -207,11 +167,10 @@ def safe_materialize_function( def safe_execute_function( - fn_obj: Callable[[tuple[str | list[dict[str, str]], ...]], Any], + fn_obj: Callable[[tuple[str | list[dict[str, str]], ...]], str], args: tuple[str | Sequence[dict[str, str]], ...], allowed_modules: dict[str, Any] | None = None, - expect_list: bool = False, -) -> Any: +) -> str: """ Execute a function in an isolated environment with no access to external state. @@ -226,8 +185,7 @@ def safe_execute_function( 1. Validate input types (args must be tuple of strings or lists) 2. Create isolated globals with SAFE_BUILTINS and SAFE_MODULES 3. Execute function call via exec() in restricted environment - 4. Extract result and validate output type (string, or list of JSON when - ``expect_list`` is set) + 4. Extract result and validate output type (must be string) Args: fn_obj: Compiled function from safe_materialize_function() @@ -235,14 +193,13 @@ def safe_execute_function( allowed_modules: Optional dict of additional modules to allow during execution. If None, uses SAFE_MODULES (json, re, math, numpy, pandas, pyarrow). - expect_list: When True, require the result to be a list of JSON values. Returns: - Result from function execution + String result from function execution Raises: RuntimeError: If function execution fails for any reason - TypeError: If args is not tuple[str, ...] or result is not the expected type + TypeError: If args is not tuple[str, ...] or result is not str Examples: >>> fn = safe_materialize_function("lambda args: args[0].upper()") @@ -256,11 +213,7 @@ def safe_execute_function( if not isinstance(args, tuple): raise TypeError(f"Args must be a tuple, got {type(args).__name__}") - if expect_list: - if not all(_is_json_value(arg) for arg in args): - arg_types = [type(arg).__name__ for arg in args] - raise TypeError(f"All args must be JSON values, got types: {arg_types}") - elif not all(isinstance(arg, (str, list)) for arg in args): + if not all(isinstance(arg, (str, list)) for arg in args): arg_types = [type(arg).__name__ for arg in args] raise TypeError(f"All args must be strings or lists, got types: {arg_types}") @@ -286,20 +239,8 @@ def safe_execute_function( # Extract the result from safe_locals result = safe_locals["__result__"] - # Validate output type - if expect_list: - if not isinstance(result, list): - raise TypeError( - "Function must return a list, but returned " - f"{type(result).__name__}: {result}" - ) - for element in result: - if not _is_json_value(element): - raise TypeError( - "Function must return a list of JSON values, but an element " - f"is {type(element).__name__}: {element}" - ) - elif not isinstance(result, str): + # Validate output type: must be a string + if not isinstance(result, str): raise TypeError( "Function must return a string, but returned " f"{type(result).__name__}: {result}" diff --git a/tests/worker/test_echo_executor.py b/tests/worker/test_echo_executor.py deleted file mode 100644 index 8cbca28e8..000000000 --- a/tests/worker/test_echo_executor.py +++ /dev/null @@ -1,169 +0,0 @@ -"""Echo executor tests: the literal "list" path and the list-Lambda "function" path.""" - -from pathlib import Path - -import pytest -from pydantic import JsonValue - -from shared.schemas.result import EchoItem, EchoResult -from shared.tasks import TaskType -from worker.executors.base_executor import ExecutionError -from worker.executors.echo_executor import EchoExecutor - -from .factories import make_worker_config, make_worker_task_message - - -def _spec(data: dict, upstream: dict | None = None) -> dict: - spec: dict = {"taskType": "echo", "data": data} - if upstream: - spec["_upstreamResults"] = upstream - return spec - - -def _run( - data: dict, upstream: dict | None = None, tmp_path: Path | None = None -) -> EchoResult: - executor = EchoExecutor(make_worker_config()) - task = make_worker_task_message( - _spec(data, upstream), task_type=TaskType.ECHO, task_id="tsk-echo" - ) - return executor.run(task, tmp_path or Path("/tmp/echo-out")) - - -def _echo_result(*outputs: JsonValue) -> EchoResult: - return EchoResult(items=[EchoItem(output=o) for o in outputs], count=len(outputs)) - - -class TestListPath: - def test_literal_items_are_echoed(self) -> None: - result = _run({"type": "list", "items": ["a", "b", "c"]}) - assert [i.output for i in result.items] == ["a", "b", "c"] - - -class TestFunctionPath: - def test_explode_one_input_list_into_rows(self) -> None: - result = _run( - { - "type": "function", - "function": "lambda args: [x * 2 for x in args[0]]", - "arguments": [{"items": [1, 2, 3]}], - } - ) - assert [i.output for i in result.items] == [2, 4, 6] - - def test_filter_rows(self) -> None: - result = _run( - { - "type": "function", - "function": "lambda args: [x for x in args[0] if x % 2 == 0]", - "arguments": [{"items": [1, 2, 3, 4]}], - } - ) - assert [i.output for i in result.items] == [2, 4] - - def test_split_one_input_into_two_outputs(self) -> None: - result = _run( - { - "type": "function", - "function": "lambda args: [args[0][:2], args[0][2:]]", - "arguments": [{"items": [1, 2, 3, 4]}], - } - ) - assert [i.output for i in result.items] == [[1, 2], [3, 4]] - - def test_cross_product_into_groups_of_different_sizes(self) -> None: - result = _run( - { - "type": "function", - "function": ( - "lambda args: [[[a, b] for b in args[1]] for a in args[0]]" - ), - "arguments": [{"items": [1, 2]}, {"items": ["x", "y", "z"]}], - } - ) - assert [i.output for i in result.items] == [ - [[1, "x"], [1, "y"], [1, "z"]], - [[2, "x"], [2, "y"], [2, "z"]], - ] - - def test_collapse_groups_back_to_one_row_per_group(self) -> None: - result = _run( - { - "type": "function", - "function": "lambda args: [sum(g) for g in args[0]]", - "arguments": [{"items": [[1, 2], [3, 4, 5]]}], - } - ) - assert [i.output for i in result.items] == [3, 12] - - def test_node_path_argument_reads_upstream_echo_result(self) -> None: - upstream = {"echo-a": _echo_result("p", "q")} - result = _run( - { - "type": "function", - "function": "lambda args: [args[0].upper()]", - "arguments": [{"node": "echo-a", "path": "items[0].output"}], - }, - upstream=upstream, - ) - assert [i.output for i in result.items] == ["P"] - - def test_literal_items_argument(self) -> None: - result = _run( - { - "type": "function", - "function": "lambda args: [args[0]]", - "arguments": [{"items": ["a", "b"]}], - } - ) - assert [i.output for i in result.items] == [["a", "b"]] - - def test_mixed_argument_raises(self) -> None: - with pytest.raises(ExecutionError, match="exactly one"): - _run( - { - "type": "function", - "function": "lambda args: [args[0]]", - "arguments": [{"items": [1], "expr": "absent.items"}], - } - ) - - def test_unknown_key_argument_raises(self) -> None: - with pytest.raises(ExecutionError, match="exactly one"): - _run( - { - "type": "function", - "function": "lambda args: [args[0]]", - "arguments": [{"bogus": 1}], - } - ) - - def test_non_list_return_raises(self) -> None: - with pytest.raises(ExecutionError, match="must return a list"): - _run( - { - "type": "function", - "function": "lambda args: 'not a list'", - "arguments": [{"items": [1]}], - } - ) - - def test_non_json_element_raises(self) -> None: - with pytest.raises(ExecutionError, match="function failed"): - _run( - { - "type": "function", - "function": "lambda args: [object()]", - "arguments": [{"items": [1]}], - } - ) - - def test_nested_non_json_element_raises(self) -> None: - with pytest.raises(ExecutionError, match="function failed"): - _run( - { - "type": "function", - "function": "lambda args: [{'k': set([1])}]", - "arguments": [{"items": [1]}], - } - ) diff --git a/tests/worker/test_safe_eval.py b/tests/worker/test_safe_eval.py deleted file mode 100644 index 05c19fc87..000000000 --- a/tests/worker/test_safe_eval.py +++ /dev/null @@ -1,116 +0,0 @@ -"""safe_eval tests: the string-result prompt path and the list-result Lambda path.""" - -import pytest - -from worker.executors.utils.safe_eval import ( - safe_execute_function, - safe_materialize_function, -) - - -def _run(fn_code: str, args: tuple, *, expect_list: bool = False): - fn_obj = safe_materialize_function(fn_code) - return safe_execute_function(fn_obj, args, expect_list=expect_list) - - -class TestStringResult: - def test_string_result_is_returned(self) -> None: - assert _run("lambda args: args[0].upper()", ("hello",)) == "HELLO" - - def test_non_string_result_raises(self) -> None: - with pytest.raises(RuntimeError, match="must return a string"): - _run("lambda args: [1, 2]", ([1],)) - - -class TestListResult: - def test_list_of_json_is_returned(self) -> None: - assert _run( - "lambda args: [x * 2 for x in args[0]]", ([1, 2, 3],), expect_list=True - ) == [2, 4, 6] - - def test_groups_are_allowed_as_elements(self) -> None: - assert _run("lambda args: [[1, 2], [3, 4, 5]]", ([],), expect_list=True) == [ - [1, 2], - [3, 4, 5], - ] - - def test_non_list_result_raises(self) -> None: - with pytest.raises(RuntimeError, match="must return a list"): - _run("lambda args: 'nope'", ([],), expect_list=True) - - def test_non_json_element_raises(self) -> None: - with pytest.raises(RuntimeError, match="list of JSON"): - _run("lambda args: [set([1])]", ([],), expect_list=True) - - def test_nested_non_json_element_raises(self) -> None: - with pytest.raises(RuntimeError, match="list of JSON"): - _run("lambda args: [{'k': set([1])}]", ([],), expect_list=True) - - def test_dict_argument_reaches_function(self) -> None: - assert _run("lambda args: [args[0]['a']]", ({"a": 1},), expect_list=True) == [1] - - def test_scalar_argument_reaches_function(self) -> None: - assert _run("lambda args: [args[0]]", (5,), expect_list=True) == [5] - - def test_string_mode_rejects_dict_argument(self) -> None: - with pytest.raises(TypeError, match="strings or lists"): - _run("lambda args: args[0]['a']", ({"a": 1},)) - - -class TestMaterializeShape: - def test_single_def_works_in_string_mode(self) -> None: - assert _run("def f(args):\n return args[0].upper()", ("hi",)) == "HI" - - def test_single_def_works_in_function_mode(self) -> None: - assert _run( - "def f(args):\n return [x * 2 for x in args[0]]", - ([1, 2],), - expect_list=True, - ) == [2, 4] - - def test_lambda_works_in_both_modes(self) -> None: - assert _run("lambda args: args[0].upper()", ("hi",)) == "HI" - assert _run("lambda args: [args[0]]", (1,), expect_list=True) == [1] - - def test_two_defs_raise(self) -> None: - with pytest.raises(RuntimeError, match="single function definition"): - _run( - "def f(args):\n return args[0]\ndef g(args):\n return args[0]", - (1,), - ) - - def test_def_plus_top_level_statement_raises(self) -> None: - with pytest.raises(RuntimeError, match="single function definition"): - _run("def f(args):\n return args[0]\nx = 1", (1,)) - - def test_assigned_lambda_works_in_string_mode(self) -> None: - assert _run("lam = lambda args: args[0].lower()", ("HI",)) == "hi" - - def test_assigned_lambda_works_in_function_mode(self) -> None: - assert _run( - "lam = lambda args: [x * 2 for x in args[0]]", - ([1, 2],), - expect_list=True, - ) == [2, 4] - - def test_assign_of_non_lambda_raises(self) -> None: - with pytest.raises(RuntimeError, match="single function definition"): - _run("f = 1", (1,)) - - def test_syntax_error_raises_runtime_error(self) -> None: - with pytest.raises(RuntimeError, match="Function definition failed"): - _run("def f(args):\n return )", (1,)) - - -class TestValueErrorPropagation: - def test_value_error_message_survives_sandbox(self) -> None: - """A function raising ValueError fails closed with its message intact.""" - with pytest.raises( - RuntimeError, - match="kept ids not in the candidate table: \\['f6'\\]", - ): - _run( - "def f(args):\n" - " raise ValueError(\"kept ids not in the candidate table: ['f6']\")", - (), - ) From 33f2f9e73050c5d1c96afa3ebfbb73bcd3b5fe9e Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Wed, 30 Sep 2026 15:54:27 +0700 Subject: [PATCH 55/71] test: read a python stage's output in dataframe API tasks A dataframe or graph_template spec reads a python stage's output through the existing {node, path} argument, starting at the result as flowmesh result fetch shows it (value.items.output). _evaluate_expr already resolves a PythonResult and its grouping flag already follows the upstream structure, so these tests cover the behavior without code changes: rows from value.items.output, ragged groups of 3 and 2 into a row-wise dataframe API task, a graph_template aggregate over those groups, and an empty group. Signed-off-by: Zhengyuan Su Co-Authored-By: Claude Code --- tests/worker/test_api_executor.py | 284 +++++++++++++++++++++++++++++- 1 file changed, 283 insertions(+), 1 deletion(-) diff --git a/tests/worker/test_api_executor.py b/tests/worker/test_api_executor.py index 3d87f192d..3cd81fa3e 100644 --- a/tests/worker/test_api_executor.py +++ b/tests/worker/test_api_executor.py @@ -16,7 +16,7 @@ import pytest from pydantic import ValidationError -from shared.schemas.result import APIGroupItem, APIItem, APIResult +from shared.schemas.result import APIGroupItem, APIItem, APIResult, PythonResult from shared.tasks.specs.misc import _MAX_CONCURRENCY, _MAX_RETRIES from shared.tasks.worker_message import WorkerTaskMessage from worker.executors import api_executor as api_executor_module @@ -1931,6 +1931,288 @@ def test_status_code_taken_from_first_row_across_groups( assert result.status_code == 200 +def _python_upstream(value: Any) -> PythonResult: + """A python stage whose return value is the Lumilake shape + ``{"items": [{"output": ...}, ...]}``.""" + return PythonResult(exit_code=0, value=value) + + +class TestPythonStage: + def test_rows_from_python_value_items_output(self, tmp_path: Path) -> None: + """A dataframe API task reads a python stage's rows at + ``value.items.output``, one row per output record.""" + upstream = _python_upstream( + { + "items": [ + {"output": {"claim": "c0", "src": "s0"}}, + {"output": {"claim": "c1", "src": "s1"}}, + {"output": {"claim": "c2", "src": "s2"}}, + ] + } + ) + payload = { + "task_id": "task-api-py", + "workflow_id": "wf-1", + "owner_id": "owner", + "assigned_worker": "worker-1", + "dispatched_at": "2026-03-22T00:00:00Z", + "task": { + "apiVersion": "flowmesh/v1", + "kind": "Task", + "metadata": {"name": "wf:api"}, + "spec": { + "taskType": "api", + "_upstreamResults": {"Py": upstream}, + "api": { + "method": "POST", + "url": "https://custom.example.com/v1/chat/completions", + "json": {"messages": "{{prompt}}"}, + }, + "data": { + "type": "dataframe", + "columns": [ + { + "label": "claim", + "node": "Py", + "path": "value.items.output.claim", + }, + { + "label": "src", + "node": "Py", + "path": "value.items.output.src", + }, + ], + "messages": [ + {"role": "user", "content": "row {claim} {src}"}, + ], + }, + }, + }, + } + task = WorkerTaskMessage.model_validate(payload) + transport = _EchoTransport() + result = _run(_executor(), task, transport, tmp_path) + assert len(transport.requests) == 3 + # A single-column dataframe is one table, so one group item holds all rows. + assert len(result.items) == 1 + prompts = [json.loads(r.prompt)[0]["content"] for r in result.items[0].rows] + assert prompts == ["row c0 s0", "row c1 s1", "row c2 s2"] + + def test_ragged_groups_from_python_value_items_output(self, tmp_path: Path) -> None: + """A python stage whose per-item ``output`` is a list of records gives + one group per item, with uneven group sizes (3 and 2).""" + upstream = _python_upstream( + { + "items": [ + { + "output": [ + {"claim": "c0", "src": "s0"}, + {"claim": "c0", "src": "s1"}, + {"claim": "c0", "src": "s2"}, + ] + }, + { + "output": [ + {"claim": "c1", "src": "s3"}, + {"claim": "c1", "src": "s4"}, + ] + }, + ] + } + ) + payload = { + "task_id": "task-api-py-grp", + "workflow_id": "wf-1", + "owner_id": "owner", + "assigned_worker": "worker-1", + "dispatched_at": "2026-03-22T00:00:00Z", + "task": { + "apiVersion": "flowmesh/v1", + "kind": "Task", + "metadata": {"name": "wf:api"}, + "spec": { + "taskType": "api", + "_upstreamResults": {"Py": upstream}, + "api": { + "method": "POST", + "url": "https://custom.example.com/v1/chat/completions", + "json": {"messages": "{{prompt}}"}, + }, + "data": { + "type": "dataframe", + "columns": [ + { + "label": "claim", + "node": "Py", + "path": "value.items.output.claim", + }, + { + "label": "src", + "node": "Py", + "path": "value.items.output.src", + }, + ], + "messages": [ + {"role": "user", "content": "row {claim} {src}"}, + ], + }, + }, + }, + } + task = WorkerTaskMessage.model_validate(payload) + transport = _EchoTransport() + result = _run(_executor(), task, transport, tmp_path) + assert len(transport.requests) == 5 + assert len(result.items) == 2 + assert [len(item.rows) for item in result.items] == [3, 2] + group0 = {json.loads(r.prompt)[0]["content"] for r in result.items[0].rows} + assert group0 == {"row c0 s0", "row c0 s1", "row c0 s2"} + group1 = {json.loads(r.prompt)[0]["content"] for r in result.items[1].rows} + assert group1 == {"row c1 s3", "row c1 s4"} + + def test_graph_template_aggregate_over_python_groups(self, tmp_path: Path) -> None: + """A graph_template aggregate over a python stage's ragged groups + issues one request per group, each prompt holding all its rows.""" + upstream = _python_upstream( + { + "items": [ + { + "output": [ + {"claim": "c0", "src": "s0"}, + {"claim": "c0", "src": "s1"}, + ] + }, + { + "output": [ + {"claim": "c1", "src": "s2"}, + {"claim": "c1", "src": "s3"}, + {"claim": "c1", "src": "s4"}, + ] + }, + ] + } + ) + payload = { + "task_id": "task-api-py-gt", + "workflow_id": "wf-1", + "owner_id": "owner", + "assigned_worker": "worker-1", + "dispatched_at": "2026-03-22T00:00:00Z", + "task": { + "apiVersion": "flowmesh/v1", + "kind": "Task", + "metadata": {"name": "wf:api"}, + "spec": { + "taskType": "api", + "_upstreamResults": {"Py": upstream}, + "api": { + "method": "POST", + "url": "https://custom.example.com/v1/chat/completions", + "json": {"messages": "{{prompt}}"}, + }, + "data": { + "type": "graph_template", + "template": { + "name": "format", + "columns": [ + { + "label": "df", + "data": { + "type": "dataframe", + "columns": [ + { + "label": "claim", + "node": "Py", + "path": ("value.items.output.claim"), + } + ], + }, + } + ], + "options": { + "format": { + "steps": [], + "messages": [ + {"role": "user", "content": "all: {df}"} + ], + } + }, + }, + }, + }, + }, + } + task = WorkerTaskMessage.model_validate(payload) + transport = _EchoTransport() + result = _run(_executor(), task, transport, tmp_path) + assert len(transport.requests) == 2 + assert len(result.items) == 2 + # Aggregate results are plain items, one per group. + group0 = result.items[0].response_json["choices"][0]["message"]["content"] + group1 = result.items[1].response_json["choices"][0]["message"]["content"] + assert "c0" in group0 and "c1" not in group0 + assert "c1" in group1 and "c0" not in group1 + + def test_empty_group_from_python_value_items_output(self, tmp_path: Path) -> None: + """A python stage with an empty per-item output list yields an empty + group: zero requests for it, an APIGroupItem with no rows.""" + upstream = _python_upstream( + { + "items": [ + {"output": [{"claim": "c0", "src": "s0"}]}, + {"output": []}, + ] + } + ) + payload = { + "task_id": "task-api-py-empty", + "workflow_id": "wf-1", + "owner_id": "owner", + "assigned_worker": "worker-1", + "dispatched_at": "2026-03-22T00:00:00Z", + "task": { + "apiVersion": "flowmesh/v1", + "kind": "Task", + "metadata": {"name": "wf:api"}, + "spec": { + "taskType": "api", + "_upstreamResults": {"Py": upstream}, + "api": { + "method": "POST", + "url": "https://custom.example.com/v1/chat/completions", + "json": {"messages": "{{prompt}}"}, + }, + "data": { + "type": "dataframe", + "columns": [ + { + "label": "claim", + "node": "Py", + "path": "value.items.output.claim", + }, + { + "label": "src", + "node": "Py", + "path": "value.items.output.src", + }, + ], + "messages": [ + {"role": "user", "content": "row {claim} {src}"}, + ], + }, + }, + }, + } + task = WorkerTaskMessage.model_validate(payload) + transport = _EchoTransport() + result = _run(_executor(), task, transport, tmp_path) + assert len(transport.requests) == 1 + assert len(result.items) == 2 + assert len(result.items[0].rows) == 1 + assert result.items[1].rows == [] + assert json.loads(result.items[0].rows[0].prompt)[0]["content"] == ("row c0 s0") + + class TestCallLogging: @pytest.fixture(autouse=True) def _nebula_env(self, monkeypatch: pytest.MonkeyPatch) -> None: From a2be9e5944898f91b27108e8f2993ee42fd43736 Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Wed, 30 Sep 2026 16:36:14 +0700 Subject: [PATCH 56/71] docs: document how a dataframe API task reads a python stage A dataframe column reads a python stage with node and a path that starts at the result as flowmesh result fetch shows it, e.g. value.items.output.q. When each item's output is a list of records, each item is one group and group sizes may differ; a per-row list of scalars stays one cell value. Signed-off-by: Zhengyuan Su Co-Authored-By: Claude Code --- docs/WORKFLOWS.md | 34 ++++++++++++++++++++++++++++++++++ 1 file changed, 34 insertions(+) diff --git a/docs/WORKFLOWS.md b/docs/WORKFLOWS.md index a14566834..96a51907e 100644 --- a/docs/WORKFLOWS.md +++ b/docs/WORKFLOWS.md @@ -148,6 +148,40 @@ dataframe column that reads a grouped upstream's message content uses `path: items.rows.json.choices[0].message.content`; an ungrouped upstream uses `path: items.json.choices[0].message.content`. +A dataframe column reads a python stage with `node: ` and a path that +starts at the result as `flowmesh result fetch` shows it — for a python stage +that returns `{"items": [{"output": [...]}, ...]}`, `path: value.items.output.q` +reads the `q` field of each record. When each item's `output` is a list of +records, each item is one group and group sizes may differ (3 and 2); a per-row +list of scalars stays one cell value. For example, a python stage T0 that +returns two such items feeds a dataframe API task T1 that reads them as two +groups: + +```yaml +spec: + stages: + - name: T0 + spec: + taskType: python + code: | + def main(): + return {"items": [{"output": [{"q": "..."}, {"q": "..."}, {"q": "..."}]}, + {"output": [{"q": "..."}, {"q": "..."}]}]} + - name: T1 + dependsOn: [T0] + spec: + taskType: api + data: + type: dataframe + columns: + - label: Q + node: T0 + path: value.items.output.q + messages: + - role: user + content: "Answer in one word: {Q}" +``` + ```yaml spec: taskType: api From 6e64b31df087192296ed1be932c884153290f4c7 Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Wed, 30 Sep 2026 22:00:04 +0700 Subject: [PATCH 57/71] fix: issue one aggregate prompt per group in graph_template Carry the resolver's grouped flag through to the structural-message renderer so a direct grouped column builds one aggregate prompt per group instead of expanding every inner list into a row request. Signed-off-by: Zhengyuan Su Co-Authored-By: Claude --- src/worker/executors/utils/graph_templates.py | 88 +++++++++++-------- 1 file changed, 51 insertions(+), 37 deletions(-) diff --git a/src/worker/executors/utils/graph_templates.py b/src/worker/executors/utils/graph_templates.py index 4e0ff8af2..01814d5a9 100644 --- a/src/worker/executors/utils/graph_templates.py +++ b/src/worker/executors/utils/graph_templates.py @@ -271,6 +271,7 @@ def __missing__(self, key): # type: ignore[override] def _aggregate_structural_messages( columns: dict[str, Sequence[str | MaterializedMessageOrTable]], msg_options: Sequence[dict[str, str]], + grouped_labels: set[str], ) -> Sequence[MaterializedMessage]: def _is_expandable_group_value(value: Any) -> bool: return isinstance(value, list) and not all( @@ -292,47 +293,52 @@ def _is_expandable_group_value(value: Any) -> bool: group_row_counts: list[int] = [] for group_idx in range(num_groups): row_count = 1 - for values in grouped_columns.values(): + for key, values in grouped_columns.items(): + if key in grouped_labels: + continue group_value = values[group_idx] if _is_expandable_group_value(group_value): row_count = max(row_count, len(group_value)) - group_row_counts.append(row_count) - - columns = {key: [] for key in grouped_columns} - for group_idx, row_count in enumerate(group_row_counts): for key, values in grouped_columns.items(): + if key in grouped_labels: + continue group_value = values[group_idx] - if _is_expandable_group_value(group_value): - value_list = list(group_value) - if len(value_list) == 1 and row_count > 1: - value_list = [value_list[0] for _ in range(row_count)] - elif len(value_list) != row_count: - raise ExecutionError( - "Grouped graph-template values must resolve to the same " - "number of rows per group." - ) - else: - value_list = [group_value for _ in range(row_count)] - columns[key].extend(value_list) # type: ignore - - num_rows = sum(group_row_counts) + if _is_expandable_group_value(group_value) and len(group_value) not in ( + 1, + row_count, + ): + raise ExecutionError( + "Grouped graph-template values must resolve to the same " + "number of rows per group." + ) + group_row_counts.append(row_count) - batch_messages: list[Message] = [[] for _ in range(num_rows)] + batch_messages: list[Message] = [[] for _ in range(num_groups)] class _SafeDict(dict): def __missing__(self, key): # type: ignore[override] return "{" + key + "}" + def _group_value(label: str, group_idx: int, row_idx: int) -> Any: + """The value a column contributes to one rendered row of a group.""" + values = grouped_columns[label] + group_value = values[group_idx] + if label in grouped_labels: + return group_value + if _is_expandable_group_value(group_value): + return list(group_value)[row_idx] + return group_value + for message_metadata in msg_options: if "content" not in message_metadata: raise RuntimeError( f"Each message must have 'content' field. {message_metadata}" ) raw_content: str = message_metadata["content"] - if raw_content in columns: - content = columns[raw_content] # Materialize Message + if raw_content in grouped_columns: + content = grouped_columns[raw_content] # Materialize Message else: - rendered_rows: list[str] = [] + rendered_groups: list[str] = [] # Disable pandas width caps so wide DataFrame cells render in full. with pd.option_context( "display.max_columns", @@ -342,27 +348,32 @@ def __missing__(self, key): # type: ignore[override] "display.max_colwidth", None, ): - for row_idx in range(num_rows): - row_mapping: dict[str, str] = {} - for label, values in columns.items(): - row_value = values[row_idx] - if isinstance(row_value, pd.DataFrame): - row_mapping[label] = row_value.to_markdown(index=False) - else: - row_mapping[label] = _coerce_to_string(row_value) - rendered_rows.append(raw_content.format_map(_SafeDict(row_mapping))) - content = rendered_rows + for group_idx in range(num_groups): + rendered_rows: list[str] = [] + for row_idx in range(group_row_counts[group_idx]): + row_mapping: dict[str, str] = {} + for label in grouped_columns: + row_value = _group_value(label, group_idx, row_idx) + if isinstance(row_value, pd.DataFrame): + row_mapping[label] = row_value.to_markdown(index=False) + else: + row_mapping[label] = _coerce_to_string(row_value) + rendered_rows.append( + raw_content.format_map(_SafeDict(row_mapping)) + ) + rendered_groups.append("\n".join(rendered_rows)) + content = rendered_groups if role := message_metadata.get("role"): assert all(isinstance(prompt, str) for prompt in content), ( content, - columns, + grouped_columns, ) for messages, prompt in zip(batch_messages, content): messages.append({"role": role, "content": prompt}) # type: ignore else: assert all(isinstance(msg, dict) for prompt in content for msg in prompt), ( content, - columns, + grouped_columns, ) for messages, prompt in zip(batch_messages, content): messages.extend(prompt) # type: ignore @@ -441,7 +452,9 @@ def _render_lambda_func( materialized_args: list[Sequence[MaterializedMessage | str]] = [] for arg in fn_args: if not isinstance(arg, str): - materialized_args.append(_aggregate_structural_messages(columns, arg)) + materialized_args.append( + _aggregate_structural_messages(columns, arg, set()) + ) continue if arg in columns: @@ -504,6 +517,7 @@ def _render_structural_messages( ) for column in columns } + grouped_labels = {column["label"] for column in columns if column.get("grouped")} for step_option in format_options.get("steps", []): if "template" in step_option: format_kwargs = { @@ -524,7 +538,7 @@ def _render_structural_messages( ) batch_messages = _aggregate_structural_messages( - formatted_prompts, format_options["messages"] + formatted_prompts, format_options["messages"], grouped_labels ) return batch_messages From 631c8578840a292f306993349f133f31a0e1f257 Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Wed, 30 Sep 2026 22:00:50 +0700 Subject: [PATCH 58/71] test: graph_template direct grouped column issues one request per group Signed-off-by: Zhengyuan Su Co-Authored-By: Claude --- tests/worker/test_api_executor.py | 129 ++++++++++++++++++++++++++++-- 1 file changed, 122 insertions(+), 7 deletions(-) diff --git a/tests/worker/test_api_executor.py b/tests/worker/test_api_executor.py index 3cd81fa3e..91c0c42fc 100644 --- a/tests/worker/test_api_executor.py +++ b/tests/worker/test_api_executor.py @@ -1086,7 +1086,7 @@ def _handler(self, request: httpx.Request) -> httpx.Response: assert excinfo.value.retryable is True def test_parse_error_fails_with_mapping_error_not_keyerror( - self, tmp_path: Path + self, tmp_path: Path, caplog: pytest.LogCaptureFixture ) -> None: """A 200 with a non-mapping body and parse_json true fails the task with the parse error, not a KeyError from the collection loop.""" @@ -1098,9 +1098,25 @@ def __init__(self) -> None: def _handler(self, request: httpx.Request) -> httpx.Response: return httpx.Response(200, json=[1]) - task = _batch_task(["a", "b"], response={"parse_json": True}) - with pytest.raises(ExecutionError, match="not a valid JSON mapping"): - _run(_executor(), task, _ArrayBody(), tmp_path) + task = _batch_task(["a", "b"], response={"parse_json": True}, concurrency=1) + with caplog.at_level(logging.INFO, logger="worker.executors.api_executor"): + with pytest.raises(ExecutionError, match="not a valid JSON mapping"): + _run(_executor(), task, _ArrayBody(), tmp_path) + call_lines = [ + r.getMessage() + for r in caplog.records + if r.getMessage().startswith("api call") + ] + assert len(call_lines) == 1 + assert "status=200" in call_lines[0] + summary = [ + r.getMessage() + for r in caplog.records + if r.getMessage().startswith("api summary") + ] + assert len(summary) == 1 + assert "calls=1" in summary[0] + assert "failures=1" in summary[0] def test_retry_stops_when_another_row_fails( self, tmp_path: Path, monkeypatch: pytest.MonkeyPatch @@ -1802,10 +1818,9 @@ def test_graph_template_aggregates_all_rows_into_one_prompt( content = body["messages"][0]["content"] assert "c0" in content and "c1" in content assert len(result.items) == 1 - # Aggregate result is a plain item (read at items.json...), not a group - # item (read at items.rows.json...). + # Aggregate result is a plain item (read at items.json...), not a group item. item = result.items[0] - assert not hasattr(item, "rows") + assert isinstance(item, APIItem) assert item.response_json["choices"][0]["message"]["content"].startswith( "echo:all:" ) @@ -2153,6 +2168,74 @@ def test_graph_template_aggregate_over_python_groups(self, tmp_path: Path) -> No assert "c0" in group0 and "c1" not in group0 assert "c1" in group1 and "c0" not in group1 + def test_graph_template_direct_grouped_column_issues_one_request_per_group( + self, tmp_path: Path + ) -> None: + """A graph_template whose column reads a python stage's grouped output + directly (not via a nested dataframe) issues one request per group, + each prompt holding all its rows.""" + upstream = _python_upstream( + { + "items": [ + {"output": [{"claim": "c0"}, {"claim": "c1"}]}, + {"output": [{"claim": "c2"}, {"claim": "c3"}, {"claim": "c4"}]}, + ] + } + ) + payload = { + "task_id": "task-api-py-gt-direct", + "workflow_id": "wf-1", + "owner_id": "owner", + "assigned_worker": "worker-1", + "dispatched_at": "2026-03-22T00:00:00Z", + "task": { + "apiVersion": "flowmesh/v1", + "kind": "Task", + "metadata": {"name": "wf:api"}, + "spec": { + "taskType": "api", + "_upstreamResults": {"Py": upstream}, + "api": { + "method": "POST", + "url": "https://custom.example.com/v1/chat/completions", + "json": {"messages": "{{prompt}}"}, + }, + "data": { + "type": "graph_template", + "template": { + "name": "format", + "columns": [ + { + "label": "claim", + "node": "Py", + "path": "value.items.output.claim", + } + ], + "options": { + "format": { + "steps": [], + "messages": [ + {"role": "user", "content": "all: {claim}"} + ], + } + }, + }, + }, + }, + }, + } + task = WorkerTaskMessage.model_validate(payload) + transport = _EchoTransport() + result = _run(_executor(), task, transport, tmp_path) + assert len(transport.requests) == 2 + assert len(result.items) == 2 + group0 = result.items[0].response_json["choices"][0]["message"]["content"] + group1 = result.items[1].response_json["choices"][0]["message"]["content"] + assert "c0" in group0 and "c1" in group0 and "c2" not in group0 + assert ( + "c2" in group1 and "c3" in group1 and "c4" in group1 and "c0" not in group1 + ) + def test_empty_group_from_python_value_items_output(self, tmp_path: Path) -> None: """A python stage with an empty per-item output list yields an empty group: zero requests for it, an APIGroupItem with no rows.""" @@ -2303,6 +2386,38 @@ def _handler(self, request: httpx.Request) -> httpx.Response: assert len(summary) == 1 assert "retries=2" in summary[0] + def test_early_failure_summary_counts_completed_calls( + self, tmp_path: Path, caplog: pytest.LogCaptureFixture + ) -> None: + """With concurrency 1 and the first of three rows failing, the summary + counts the one request actually sent, not the three planned prompts.""" + task = _batch_task( + ["a", "b", "c"], + concurrency=1, + response={"raise_for_status": True}, + ) + + class _FirstFails(httpx.MockTransport): + def __init__(self) -> None: + super().__init__(self._handler) + + def _handler(self, request: httpx.Request) -> httpx.Response: + prompt = json.loads(request.read())["messages"][0]["content"] + if prompt == "a": + return httpx.Response(400, json={"error": "boom"}) + return _ok_response() + + transport = _FirstFails() + with caplog.at_level(logging.INFO, logger="worker.executors.api_executor"): + with pytest.raises(ExecutionError, match="row 0"): + _run(_executor(), task, transport, tmp_path) + call_lines = self._records(caplog, "api call") + assert len(call_lines) == 1 + summary = self._records(caplog, "api summary") + assert len(summary) == 1 + assert "calls=1" in summary[0] + assert "failures=1" in summary[0] + def test_non_json_body_logs_dash_without_raising( self, tmp_path: Path, caplog: pytest.LogCaptureFixture ) -> None: From 73fc8bc0914f60f5e1db170bffb50210faa2eb7f Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Wed, 30 Sep 2026 22:01:31 +0700 Subject: [PATCH 59/71] fix: record api calls that fail parsing or are cancelled A call that fails while parsing a completed response, or is cancelled right after its response, was not recorded. Record it before propagating, and count completed calls (not planned prompts) in the summary. Signed-off-by: Zhengyuan Su Co-Authored-By: Claude --- src/worker/executors/api_executor.py | 12 +++++++++++- 1 file changed, 11 insertions(+), 1 deletion(-) diff --git a/src/worker/executors/api_executor.py b/src/worker/executors/api_executor.py index 52a362bcb..4438d0f56 100644 --- a/src/worker/executors/api_executor.py +++ b/src/worker/executors/api_executor.py @@ -662,12 +662,21 @@ def _issue(idx: int, prompt: Any) -> APIItem: prompt_str=prompt_str, ) except BaseException as exc: + _record_call(idx, attempts, resp.status_code, start, None, failed=True) if not first_error: first_error.append(exc) failed.set() raise if self._cancel_event.is_set(): + _record_call( + idx, + attempts, + resp.status_code, + start, + item.response_json, + failed=True, + ) raise TaskCancelledError("API task cancelled") _record_call( @@ -729,6 +738,7 @@ def _heartbeat() -> None: heartbeat.join(timeout=5) wall = time.monotonic() - task_start with stats_lock: + done_snapshot = done failures_snapshot = failures retries_snapshot = total_retries latencies_snapshot = list(latencies) @@ -738,7 +748,7 @@ def _heartbeat() -> None: backends_snapshot = dict(backend_counts) self._log_summary( task.task_id, - total, + done_snapshot, failures_snapshot, retries_snapshot, wall, From ef976c847050ecef43f8a9f0d2831c24d4abb0d9 Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Wed, 30 Sep 2026 22:03:24 +0700 Subject: [PATCH 60/71] fix: reject echo data.type function with a pointer to python tasks A legacy type: function payload that also carries top-level items was silently run as list mode. Reject it with an error directing the caller to a python task instead. Signed-off-by: Zhengyuan Su Co-Authored-By: Claude --- src/worker/executors/echo_executor.py | 5 +++ tests/worker/test_echo_executor.py | 57 +++++++++++++++++++++++++++ 2 files changed, 62 insertions(+) create mode 100644 tests/worker/test_echo_executor.py diff --git a/src/worker/executors/echo_executor.py b/src/worker/executors/echo_executor.py index ff163e6d7..1aeb5690b 100644 --- a/src/worker/executors/echo_executor.py +++ b/src/worker/executors/echo_executor.py @@ -75,6 +75,11 @@ def run(self, task: ExecutorTask, out_dir: Path) -> EchoResult: if not isinstance(data_cfg, dict): raise ExecutionError("echo executor requires spec.data to be a mapping") + if data_cfg.get("type") == "function": + raise ExecutionError( + "echo executor spec.data.type 'function' is no longer supported; " + "use a python task instead" + ) items_cfg = data_cfg.get("items") if not isinstance(items_cfg, list): raise ExecutionError( diff --git a/tests/worker/test_echo_executor.py b/tests/worker/test_echo_executor.py new file mode 100644 index 000000000..d4aa0ce36 --- /dev/null +++ b/tests/worker/test_echo_executor.py @@ -0,0 +1,57 @@ +"""Echo executor tests: the literal "list" path and rejection of the removed +"function" mode.""" + +from pathlib import Path + +import pytest + +from shared.schemas.result import EchoResult +from shared.tasks import TaskType +from worker.executors.base_executor import ExecutionError +from worker.executors.echo_executor import EchoExecutor + +from .factories import make_worker_config, make_worker_task_message + + +def _spec(data: dict, upstream: dict | None = None) -> dict: + spec: dict = {"taskType": "echo", "data": data} + if upstream: + spec["_upstreamResults"] = upstream + return spec + + +def _run( + data: dict, upstream: dict | None = None, tmp_path: Path | None = None +) -> EchoResult: + executor = EchoExecutor(make_worker_config()) + task = make_worker_task_message( + _spec(data, upstream), task_type=TaskType.ECHO, task_id="tsk-echo" + ) + return executor.run(task, tmp_path or Path("/tmp/echo-out")) + + +class TestListPath: + def test_literal_items_are_echoed(self) -> None: + result = _run({"type": "list", "items": ["a", "b", "c"]}) + assert [i.output for i in result.items] == ["a", "b", "c"] + + +class TestFunctionModeRejected: + def test_type_function_is_rejected(self) -> None: + """A legacy ``type: function`` payload is rejected with a pointer to the + python task, even when it also carries top-level items.""" + with pytest.raises(ExecutionError, match="use a python task instead"): + _run( + { + "type": "function", + "function": "lambda args: args[0]", + "arguments": [{"items": [1, 2, 3]}], + "items": ["ignored"], + } + ) + + def test_type_function_with_items_is_not_run_as_list(self) -> None: + """A ``type: function`` payload that also has top-level items must not + silently fall through to list mode.""" + with pytest.raises(ExecutionError, match="use a python task instead"): + _run({"type": "function", "items": ["a", "b"]}) From 2460ac5d976e52bf356cf10cc75cb7cc3b81dc61 Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Wed, 30 Sep 2026 22:03:53 +0700 Subject: [PATCH 61/71] style: make APIGroupItem docstrings one paragraph Signed-off-by: Zhengyuan Su Co-Authored-By: Claude --- sdk/src/flowmesh/models/result/payloads.py | 8 ++------ src/shared/schemas/result/payloads.py | 8 ++------ 2 files changed, 4 insertions(+), 12 deletions(-) diff --git a/sdk/src/flowmesh/models/result/payloads.py b/sdk/src/flowmesh/models/result/payloads.py index 8d9d3d7cc..411f7bff7 100644 --- a/sdk/src/flowmesh/models/result/payloads.py +++ b/sdk/src/flowmesh/models/result/payloads.py @@ -177,12 +177,8 @@ class APIItem(StrictModel): class APIGroupItem(StrictModel): - """One group's row responses in a batched API task over grouped data. - - ``rows`` holds the group's row responses in order. A group is one - dataframe table (one claim), so a downstream column reads ``rows`` as a - per-group list. - """ + """One group's row responses in a batched API task over grouped data; ``rows`` + holds the group's row responses in order.""" index: int rows: list[APIItem] diff --git a/src/shared/schemas/result/payloads.py b/src/shared/schemas/result/payloads.py index 721bf3ca7..5ada8ce1f 100644 --- a/src/shared/schemas/result/payloads.py +++ b/src/shared/schemas/result/payloads.py @@ -233,12 +233,8 @@ class APIItem(StrictModel): class APIGroupItem(StrictModel): - """One group's row responses in a batched API task over grouped data. - - ``rows`` holds the group's row responses in order. A group is one - dataframe table (one claim), so a downstream column reads ``rows`` as a - per-group list. - """ + """One group's row responses in a batched API task over grouped data; ``rows`` + holds the group's row responses in order.""" index: int rows: list[APIItem] From fb296889aded286080e8d1b6ce8533ba42b72d09 Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Wed, 30 Sep 2026 22:04:13 +0700 Subject: [PATCH 62/71] docs: dataframe api results are always grouped into APIGroupItems The api executor groups every dataframe into one APIGroupItem per table, matching the vLLM executor, so ungrouped data is a single group rather than plain APIItems. Signed-off-by: Zhengyuan Su Co-Authored-By: Claude --- docs/WORKFLOWS.md | 24 +++++++++++------------- 1 file changed, 11 insertions(+), 13 deletions(-) diff --git a/docs/WORKFLOWS.md b/docs/WORKFLOWS.md index 96a51907e..cbbf2892a 100644 --- a/docs/WORKFLOWS.md +++ b/docs/WORKFLOWS.md @@ -128,25 +128,23 @@ JSON. ### Grouped results -When `spec.data` is a `dataframe` whose columns resolve to grouped upstream -values, the result is grouped: `APIResult.items` holds one `APIGroupItem` per -group, each with an `index` and a `rows` list of the group's row responses in -order. A dataframe spec decides grouping from the upstream structure — a list -of `APIGroupItem.rows`, or nested lists — never from the shape of the cell -values; a per-row list is a cell value, not a group. Ungrouped data returns -plain `APIItem`s directly in `items`. +When `spec.data` is a `dataframe`, the result is grouped: `APIResult.items` +holds one `APIGroupItem` per table, each with an `index` and a `rows` list of +the table's row responses in order. A dataframe spec decides grouping from the +upstream structure — a list of `APIGroupItem.rows`, or nested lists — never +from the shape of the cell values; a per-row list is a cell value, not a group. +Ungrouped data is a single table, so it still returns one `APIGroupItem` +holding all rows. An empty group or a column that resolves to zero rows yields zero requests for that group; the group still appears as an `APIGroupItem` with an empty `rows` list so downstream paths resolve. The result's `status_code` is taken from the first row across all groups, so a leading empty group does not zero it. -Downstream stages read grouped responses through the group shape: -`items.rows.json...` addresses a field of each row within a group, while -`items.json...` addresses a field of a plain (ungrouped) item. For example, a -dataframe column that reads a grouped upstream's message content uses -`path: items.rows.json.choices[0].message.content`; an ungrouped upstream uses -`path: items.json.choices[0].message.content`. +Downstream stages read dataframe responses through the group shape: +`items.rows.json...` addresses a field of each row within a group. For example, +a dataframe column that reads an upstream's message content uses +`path: items.rows.json.choices[0].message.content`. A dataframe column reads a python stage with `node: ` and a path that starts at the result as `flowmesh result fetch` shows it — for a python stage From 5217718a84fb5675c87473931f1b208724e0dad8 Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Wed, 30 Sep 2026 22:35:30 +0700 Subject: [PATCH 63/71] fix: keep per-row expansion for ungrouped graph-template columns The previous fix always emitted one aggregate prompt per group, which changed behaviour for every DataMixin caller over grouped data. Restore main's one-message-per-row expansion and only keep a grouped column whole per group, repeated across its rows. Signed-off-by: Zhengyuan Su Co-Authored-By: Claude --- src/worker/executors/utils/graph_templates.py | 79 +++++++++---------- tests/worker/test_graph_templates_expr.py | 52 ++++++++++++ 2 files changed, 89 insertions(+), 42 deletions(-) diff --git a/src/worker/executors/utils/graph_templates.py b/src/worker/executors/utils/graph_templates.py index 01814d5a9..f4bc46acf 100644 --- a/src/worker/executors/utils/graph_templates.py +++ b/src/worker/executors/utils/graph_templates.py @@ -299,46 +299,46 @@ def _is_expandable_group_value(value: Any) -> bool: group_value = values[group_idx] if _is_expandable_group_value(group_value): row_count = max(row_count, len(group_value)) + group_row_counts.append(row_count) + + columns = {key: [] for key in grouped_columns} + for group_idx, row_count in enumerate(group_row_counts): for key, values in grouped_columns.items(): - if key in grouped_labels: - continue group_value = values[group_idx] - if _is_expandable_group_value(group_value) and len(group_value) not in ( - 1, - row_count, - ): - raise ExecutionError( - "Grouped graph-template values must resolve to the same " - "number of rows per group." - ) - group_row_counts.append(row_count) + if key in grouped_labels: + # A grouped column is kept whole per group, repeated across its rows. + columns[key].extend([group_value] * row_count) # type: ignore + elif _is_expandable_group_value(group_value): + value_list = list(group_value) + if len(value_list) == 1 and row_count > 1: + value_list = [value_list[0] for _ in range(row_count)] + elif len(value_list) != row_count: + raise ExecutionError( + "Grouped graph-template values must resolve to the same " + "number of rows per group." + ) + columns[key].extend(value_list) # type: ignore + else: + columns[key].extend([group_value for _ in range(row_count)]) # type: ignore + + num_rows = sum(group_row_counts) - batch_messages: list[Message] = [[] for _ in range(num_groups)] + batch_messages: list[Message] = [[] for _ in range(num_rows)] class _SafeDict(dict): def __missing__(self, key): # type: ignore[override] return "{" + key + "}" - def _group_value(label: str, group_idx: int, row_idx: int) -> Any: - """The value a column contributes to one rendered row of a group.""" - values = grouped_columns[label] - group_value = values[group_idx] - if label in grouped_labels: - return group_value - if _is_expandable_group_value(group_value): - return list(group_value)[row_idx] - return group_value - for message_metadata in msg_options: if "content" not in message_metadata: raise RuntimeError( f"Each message must have 'content' field. {message_metadata}" ) raw_content: str = message_metadata["content"] - if raw_content in grouped_columns: - content = grouped_columns[raw_content] # Materialize Message + if raw_content in columns: + content = columns[raw_content] # Materialize Message else: - rendered_groups: list[str] = [] + rendered_rows: list[str] = [] # Disable pandas width caps so wide DataFrame cells render in full. with pd.option_context( "display.max_columns", @@ -348,32 +348,27 @@ def _group_value(label: str, group_idx: int, row_idx: int) -> Any: "display.max_colwidth", None, ): - for group_idx in range(num_groups): - rendered_rows: list[str] = [] - for row_idx in range(group_row_counts[group_idx]): - row_mapping: dict[str, str] = {} - for label in grouped_columns: - row_value = _group_value(label, group_idx, row_idx) - if isinstance(row_value, pd.DataFrame): - row_mapping[label] = row_value.to_markdown(index=False) - else: - row_mapping[label] = _coerce_to_string(row_value) - rendered_rows.append( - raw_content.format_map(_SafeDict(row_mapping)) - ) - rendered_groups.append("\n".join(rendered_rows)) - content = rendered_groups + for row_idx in range(num_rows): + row_mapping: dict[str, str] = {} + for label, values in columns.items(): + row_value = values[row_idx] + if isinstance(row_value, pd.DataFrame): + row_mapping[label] = row_value.to_markdown(index=False) + else: + row_mapping[label] = _coerce_to_string(row_value) + rendered_rows.append(raw_content.format_map(_SafeDict(row_mapping))) + content = rendered_rows if role := message_metadata.get("role"): assert all(isinstance(prompt, str) for prompt in content), ( content, - grouped_columns, + columns, ) for messages, prompt in zip(batch_messages, content): messages.append({"role": role, "content": prompt}) # type: ignore else: assert all(isinstance(msg, dict) for prompt in content for msg in prompt), ( content, - grouped_columns, + columns, ) for messages, prompt in zip(batch_messages, content): messages.extend(prompt) # type: ignore diff --git a/tests/worker/test_graph_templates_expr.py b/tests/worker/test_graph_templates_expr.py index 3e999622a..7a4cde204 100644 --- a/tests/worker/test_graph_templates_expr.py +++ b/tests/worker/test_graph_templates_expr.py @@ -5,6 +5,7 @@ from shared.schemas.result import APIGroupItem, APIItem, APIResult from worker.executors.utils.graph_templates import ( + _aggregate_structural_messages, _build_grouped_dataframes, _evaluate_expr, ) @@ -16,6 +17,57 @@ def _item(content: str) -> APIItem: return item +def _messages( + columns: dict, grouped_labels: set[str], content: str = "row {L}" +) -> list[str]: + """Render one user message per row over the given columns.""" + batch = _aggregate_structural_messages( + columns, + [{"role": "user", "content": content}], + grouped_labels, + ) + return [m["content"] for message in batch for m in message] + + +def test_aggregate_ungrouped_list_column_expands_per_row() -> None: + """An ungrouped list column over 2 groups of 2 and 3 rows yields 5 messages + in order, one per row.""" + columns = {"L": [["c0", "c1"], ["c2", "c3", "c4"]]} + assert _messages(columns, set()) == [ + "row c0", + "row c1", + "row c2", + "row c3", + "row c4", + ] + + +def test_aggregate_grouped_column_keeps_whole_group() -> None: + """A grouped column over the same data yields 2 messages, each carrying its + whole group.""" + columns = {"L": [["c0", "c1"], ["c2", "c3", "c4"]]} + assert _messages(columns, {"L"}) == [ + 'row ["c0", "c1"]', + 'row ["c2", "c3", "c4"]', + ] + + +def test_aggregate_mixed_grouped_and_ungrouped_columns() -> None: + """A grouped column mixed with an ungrouped list column yields one message + per ungrouped row, each carrying the whole grouped value.""" + columns = { + "G": [["g0", "g1"], ["g2", "g3", "g4"]], + "L": [["c0", "c1"], ["c2", "c3", "c4"]], + } + assert _messages(columns, {"G"}, "row {G} {L}") == [ + 'row ["g0", "g1"] c0', + 'row ["g0", "g1"] c1', + 'row ["g2", "g3", "g4"] c2', + 'row ["g2", "g3", "g4"] c3', + 'row ["g2", "g3", "g4"] c4', + ] + + def test_expr_over_list_of_models_with_alias() -> None: """Attribute access maps over a list of APIItem models, resolving the ``json`` alias to ``response_json``.""" From 82579474f704738e14da9612fd504979a4fdbe6c Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Thu, 1 Oct 2026 11:43:40 +0700 Subject: [PATCH 64/71] fix: address the review round on grouping, echo and API usage - docs: apply the suggested WORKFLOWS.md wording, drop the duplicate api example and the EXECUTORS.md echo section. - echo: drop the function-type rejection branch. - graph templates: a list of DataFrames groups by item, and one grouping rule (mapped items) covers the vLLM and API paths with no APIGroupItem special case; _model_attr reads model_extra. - result catalog: drop the _route_group_items validator in shared and the SDK mirror. - data mixin: the dataframe branch reuses _build_grouped_dataframes. - API executor: APIResult.usage is a typed APIUsage (shared and SDK), and per-call and summary lines log at debug. Co-Authored-By: Claude Opus 5.5 (1M context) Signed-off-by: Zhengyuan Su --- docs/EXECUTORS.md | 10 -- docs/WORKFLOWS.md | 70 +++------ sdk/src/flowmesh/models/result/__init__.py | 2 + sdk/src/flowmesh/models/result/catalog.py | 21 +-- sdk/src/flowmesh/models/result/payloads.py | 11 ++ src/shared/schemas/result/__init__.py | 2 + src/shared/schemas/result/catalog.py | 21 +-- src/shared/schemas/result/payloads.py | 13 ++ src/worker/executors/api_executor.py | 22 ++- src/worker/executors/echo_executor.py | 5 - src/worker/executors/mixins/data.py | 52 +------ src/worker/executors/utils/graph_templates.py | 84 +++++++---- tests/shared/test_executor_result.py | 2 - tests/worker/test_api_executor.py | 137 ++++++++++++++++-- .../test_data_mixin_dataframe_grouping.py | 82 +++++++++++ tests/worker/test_echo_executor.py | 34 +---- tests/worker/test_graph_templates_expr.py | 121 +++++++++++++++- 17 files changed, 463 insertions(+), 226 deletions(-) diff --git a/docs/EXECUTORS.md b/docs/EXECUTORS.md index 1ba70cc97..2f049513b 100644 --- a/docs/EXECUTORS.md +++ b/docs/EXECUTORS.md @@ -65,13 +65,3 @@ Optional, for the search tools: - `SERPER_API_KEY` - `JINA_API_KEY` - -## Echo executor - -`taskType: echo` returns input values back as the result. It is useful for -inspecting and shaping data between stages. - -`spec.data.type: list` echoes each `spec.data.items` entry. An entry is either -a string literal or a mapping with an expression (`expr`, or both `node` and -`path`) resolved against the upstream results. A resolved list is flattened -into one echo item per element; a scalar becomes a single item. diff --git a/docs/WORKFLOWS.md b/docs/WORKFLOWS.md index cbbf2892a..aadff04d6 100644 --- a/docs/WORKFLOWS.md +++ b/docs/WORKFLOWS.md @@ -71,10 +71,9 @@ contract. ## API task -`taskType: api` issues one HTTP request per row of `spec.data`, in parallel, -and returns the responses in `APIResult.items`. A single request is a one-row -`spec.data`. `spec.data` is required, exactly as for the vLLM executor; it -supports the same data types (`list`, `dataset`, `graph_template`, +`taskType: api` sends one HTTP request per `spec.data` prompt, in parallel, and +returns the responses in `APIResult.items`. `spec.data` is required; it supports +the same data types as the vLLM executor (`list`, `dataset`, `graph_template`, `dataframe`). By default it routes to the Nebula endpoint and authenticates with the worker's `NEBULA_API_TOKEN`. @@ -108,17 +107,11 @@ spec: ### Per-row prompts -When `spec.data` is present, the task batches: one request is issued per row, -and the results are returned in `APIResult.items`. Each row's prompt is -substituted for the `{{prompt}}` placeholder in the request body. Server-side -stage references are `${...}`; `{{prompt}}` is a worker-side per-row slot, so -it is not touched by server-side resolution. A failure in any row fails the -whole task rather than shifting the remaining rows. `spec.api.concurrency` -(default 8, an integer from 1 to 8; any other value is rejected when the -workflow is submitted) bounds the number of in-flight requests. Cancelling the -task prevents not-yet-started rows from issuing and marks the task cancelled -once in-flight requests return; a request already inside the HTTP call is not -interrupted. +Each row's prompt replaces `{{prompt}}` in the request body; a value that is +exactly `{{prompt}}` takes the prompt as-is, so a message-list row fills +`messages`. `spec.api.concurrency` (default and maximum 8) bounds in-flight +requests. Any failed row fails the task. Cancelling the task skips rows that +have not started and marks it cancelled once in-flight requests return. A body value that is exactly `{{prompt}}` is replaced by the row's prompt object as-is (a message list stays a list of `{"role", "content"}` dicts). An @@ -128,22 +121,18 @@ JSON. ### Grouped results -When `spec.data` is a `dataframe`, the result is grouped: `APIResult.items` -holds one `APIGroupItem` per table, each with an `index` and a `rows` list of -the table's row responses in order. A dataframe spec decides grouping from the -upstream structure — a list of `APIGroupItem.rows`, or nested lists — never -from the shape of the cell values; a per-row list is a cell value, not a group. -Ungrouped data is a single table, so it still returns one `APIGroupItem` -holding all rows. - -An empty group or a column that resolves to zero rows yields zero requests for -that group; the group still appears as an `APIGroupItem` with an empty `rows` -list so downstream paths resolve. The result's `status_code` is taken from the -first row across all groups, so a leading empty group does not zero it. - -Downstream stages read dataframe responses through the group shape: -`items.rows.json...` addresses a field of each row within a group. For example, -a dataframe column that reads an upstream's message content uses +A `dataframe` spec returns one `APIGroupItem` per table in `APIResult.items`, +with the table's row responses in `rows`; an empty table has empty `rows`. A +`graph_template` spec over grouped columns sends one request per group and +returns one `APIItem` for each. Other specs return one `APIItem` per row. + +A column is grouped when it reads one list of records per upstream item: an +upstream API task's `items.rows`, or a python stage whose items each carry a +list of records in `output` (`path: value.items.output.`). Each list +becomes one table, so group sizes may differ. A per-row list of scalars stays +one cell. + +Downstream stages read a grouped result through `rows`, for example `path: items.rows.json.choices[0].message.content`. A dataframe column reads a python stage with `node: ` and a path that @@ -180,25 +169,6 @@ spec: content: "Answer in one word: {Q}" ``` -```yaml -spec: - taskType: api - data: - type: list - items: - - Explain vector databases - - Explain attention - api: - method: POST - body: - model: gpt-4o - messages: - - role: user - content: "{{prompt}}" - response: - parse_json: true -``` - ## Python task `taskType: python` runs a function from `spec.code` in its own container, and diff --git a/sdk/src/flowmesh/models/result/__init__.py b/sdk/src/flowmesh/models/result/__init__.py index fa7e1d4df..6047335ee 100644 --- a/sdk/src/flowmesh/models/result/__init__.py +++ b/sdk/src/flowmesh/models/result/__init__.py @@ -42,6 +42,7 @@ AgentUsage, APIGroupItem, APIItem, + APIUsage, CostEstimates, DataRetrievalItem, EchoItem, @@ -95,6 +96,7 @@ "APIGroupItem", "APIItem", "APIResult", + "APIUsage", "AgentBatchSummary", "AgentItem", "AgentMetadata", diff --git a/sdk/src/flowmesh/models/result/catalog.py b/sdk/src/flowmesh/models/result/catalog.py index a3f8fbb68..71d79129d 100644 --- a/sdk/src/flowmesh/models/result/catalog.py +++ b/sdk/src/flowmesh/models/result/catalog.py @@ -9,7 +9,6 @@ Field, SerializeAsAny, Tag, - field_validator, ) from ..artifacts import ArtifactRef @@ -21,6 +20,7 @@ AgentUsage, APIGroupItem, APIItem, + APIUsage, CostEstimates, DataRetrievalItem, EchoItem, @@ -213,27 +213,10 @@ class APIResult(StrictExecutorResult): truncated: bool = False headers: dict[str, str] | None = None response_json: Any = Field(default=None, alias="json") - usage: dict[str, Any] | None = None + usage: APIUsage | None = None text: str | None = None items: list[APIItem | APIGroupItem] = Field(default_factory=list) - @field_validator("items", mode="before") - @classmethod - def _route_group_items(cls, value: Any) -> Any: - """Route a dict carrying ``rows`` to APIGroupItem before the union runs, - since a group dict also satisfies APIItem's required fields.""" - if not isinstance(value, list): - return value - routed: list[Any] = [] - for item in value: - if isinstance(item, dict) and "rows" in item: - routed.append(APIGroupItem.model_validate(item)) - elif isinstance(item, dict): - routed.append(APIItem.model_validate(item)) - else: - routed.append(item) - return routed - class SSHResult(StrictExecutorResult): task_type: Literal[TaskType.SSH] = TaskType.SSH diff --git a/sdk/src/flowmesh/models/result/payloads.py b/sdk/src/flowmesh/models/result/payloads.py index 411f7bff7..a53cf9106 100644 --- a/sdk/src/flowmesh/models/result/payloads.py +++ b/sdk/src/flowmesh/models/result/payloads.py @@ -21,6 +21,17 @@ class GenerationUsage(StrictModel): latency_sec: float +class APIUsage(StrictModel): + prompt_tokens: int + completion_tokens: int + reasoning_tokens: int + calls: int + failures: int + retries: int + truncated_calls: int + wall_sec: float + + class EmbeddingUsage(StrictModel): prompt_tokens: int total_tokens: int diff --git a/src/shared/schemas/result/__init__.py b/src/shared/schemas/result/__init__.py index 889651653..a89bd8c36 100644 --- a/src/shared/schemas/result/__init__.py +++ b/src/shared/schemas/result/__init__.py @@ -42,6 +42,7 @@ AgentUsage, APIGroupItem, APIItem, + APIUsage, CostEstimates, DataRetrievalItem, EchoItem, @@ -94,6 +95,7 @@ "APIItem", "APIGroupItem", "APIResult", + "APIUsage", "AgentBatchSummary", "AgentItem", "AgentMetadata", diff --git a/src/shared/schemas/result/catalog.py b/src/shared/schemas/result/catalog.py index 20684e2f5..cfa9234d9 100644 --- a/src/shared/schemas/result/catalog.py +++ b/src/shared/schemas/result/catalog.py @@ -9,7 +9,6 @@ Field, SerializeAsAny, Tag, - field_validator, ) from shared.tasks.task_type import TaskType @@ -23,6 +22,7 @@ AgentUsage, APIGroupItem, APIItem, + APIUsage, CostEstimates, DataRetrievalItem, EchoItem, @@ -259,27 +259,10 @@ class APIResult(StrictExecutorResult): truncated: bool = False headers: dict[str, str] | None = None response_json: Any = Field(default=None, alias="json") - usage: dict[str, Any] | None = None + usage: APIUsage | None = None text: str | None = None items: list[APIItem | APIGroupItem] = Field(default_factory=list) - @field_validator("items", mode="before") - @classmethod - def _route_group_items(cls, value: Any) -> Any: - """Route a dict carrying ``rows`` to APIGroupItem before the union runs, - since a group dict also satisfies APIItem's required fields.""" - if not isinstance(value, list): - return value - routed: list[Any] = [] - for item in value: - if isinstance(item, dict) and "rows" in item: - routed.append(APIGroupItem.model_validate(item)) - elif isinstance(item, dict): - routed.append(APIItem.model_validate(item)) - else: - routed.append(item) - return routed - class SSHResult(StrictExecutorResult): """SSH session output.""" diff --git a/src/shared/schemas/result/payloads.py b/src/shared/schemas/result/payloads.py index 5ada8ce1f..a25fdab2c 100644 --- a/src/shared/schemas/result/payloads.py +++ b/src/shared/schemas/result/payloads.py @@ -19,6 +19,19 @@ class GenerationUsage(StrictModel): latency_sec: float +class APIUsage(StrictModel): + """Token/call accounting for an API task, summed over its requests.""" + + prompt_tokens: int + completion_tokens: int + reasoning_tokens: int + calls: int + failures: int + retries: int + truncated_calls: int + wall_sec: float + + class EmbeddingUsage(StrictModel): """Token/latency accounting for embedding inference.""" diff --git a/src/worker/executors/api_executor.py b/src/worker/executors/api_executor.py index 4438d0f56..3bf1be37f 100644 --- a/src/worker/executors/api_executor.py +++ b/src/worker/executors/api_executor.py @@ -12,7 +12,7 @@ import httpx -from shared.schemas.result import APIGroupItem, APIItem, APIResult +from shared.schemas.result import APIGroupItem, APIItem, APIResult, APIUsage from shared.tasks.specs import ApiSpecStrict from shared.tasks.specs.misc import ApiConfig, ApiResponseConfig from shared.tasks.task_type import TaskType @@ -446,7 +446,7 @@ def _log_summary( backends = ",".join( f"{name}={count}" for name, count in sorted(backend_counts.items()) ) - logger.info( + logger.debug( "api summary task=%s calls=%d failures=%d retries=%d wall=%.3fs " "latency_p50=%.3fs latency_p95=%.3fs latency_max=%.3fs " "prompt_tokens=%d completion_tokens=%d reasoning_tokens=%d " @@ -538,6 +538,7 @@ def _run(self, task: ExecutorTask, out_dir: Path) -> APIResult: sum_prompt = 0 sum_completion = 0 sum_reasoning = 0 + truncated_calls = 0 backend_counts: dict[str, int] = {} task_start = time.monotonic() in_flight: dict[int, float] = {} @@ -554,7 +555,7 @@ def _record_call( failed: bool, ) -> None: nonlocal done, failures, total_retries - nonlocal sum_prompt, sum_completion, sum_reasoning + nonlocal sum_prompt, sum_completion, sum_reasoning, truncated_calls wall = time.monotonic() - start with in_flight_lock: in_flight.pop(idx, None) @@ -577,9 +578,11 @@ def _record_call( sum_completion += completion_tokens if reasoning_tokens is not None: sum_reasoning += reasoning_tokens + if finish_reason == "length": + truncated_calls += 1 if backend is not None: backend_counts[backend] = backend_counts.get(backend, 0) + 1 - logger.info( + logger.debug( "api call task=%s row=%d attempts=%d status=%s wall=%.3fs " "prompt_tokens=%s completion_tokens=%s reasoning_tokens=%s " "finish_reason=%s backend=%s", @@ -745,6 +748,7 @@ def _heartbeat() -> None: prompt_snapshot = sum_prompt completion_snapshot = sum_completion reasoning_snapshot = sum_reasoning + truncated_snapshot = truncated_calls backends_snapshot = dict(backend_counts) self._log_summary( task.task_id, @@ -819,4 +823,14 @@ def _heartbeat() -> None: status_code=status_code, truncated=truncated, items=result_items, + usage=APIUsage( + prompt_tokens=prompt_snapshot, + completion_tokens=completion_snapshot, + reasoning_tokens=reasoning_snapshot, + calls=total, + failures=failures_snapshot, + retries=retries_snapshot, + truncated_calls=truncated_snapshot, + wall_sec=wall, + ), ) diff --git a/src/worker/executors/echo_executor.py b/src/worker/executors/echo_executor.py index 1aeb5690b..ff163e6d7 100644 --- a/src/worker/executors/echo_executor.py +++ b/src/worker/executors/echo_executor.py @@ -75,11 +75,6 @@ def run(self, task: ExecutorTask, out_dir: Path) -> EchoResult: if not isinstance(data_cfg, dict): raise ExecutionError("echo executor requires spec.data to be a mapping") - if data_cfg.get("type") == "function": - raise ExecutionError( - "echo executor spec.data.type 'function' is no longer supported; " - "use a python task instead" - ) items_cfg = data_cfg.get("items") if not isinstance(items_cfg, list): raise ExecutionError( diff --git a/src/worker/executors/mixins/data.py b/src/worker/executors/mixins/data.py index bf16a9026..ba5a2606e 100644 --- a/src/worker/executors/mixins/data.py +++ b/src/worker/executors/mixins/data.py @@ -28,6 +28,7 @@ ) from ..utils.data_utils import normalize_prompt_payload from ..utils.graph_templates import ( + _build_grouped_dataframes, _evaluate_expr, _resolve_columns, build_prompts_from_graph_template, @@ -585,56 +586,9 @@ def _collect_prompts_for_spec( "for type == 'dataframe'." ) - grouped_columns: dict[str, list[list[Any]]] = {} - for column in resolved_columns: - label = column["label"] - value = column["value"] - if column.get("grouped"): - if not isinstance(value, list): - raise ExecutionError( - f"Column '{label}' is grouped but did not resolve " - "to a list." - ) - groups = value - else: - groups = [value] - grouped_columns[label] = groups - - group_count = max(len(groups) for groups in grouped_columns.values()) - for label, groups in list(grouped_columns.items()): - if len(groups) == 1 and group_count > 1: - grouped_columns[label] = groups * group_count - elif len(groups) != group_count: - raise ExecutionError( - "spec.data.columns must resolve to the same number of groups." - ) - - table_stores_list = [] - for group_idx in range(group_count): - max_len = 0 - raw_group_values: dict[str, list[Any]] = {} - for label, groups in grouped_columns.items(): - values = groups[group_idx] - if not isinstance(values, list): - values = [values] - if values: - max_len = max(max_len, len(values)) - raw_group_values[label] = values - - normalized_rows: dict[str, list[Any]] = {} - for label, values in raw_group_values.items(): - if len(values) == 1 and max_len > 1: - values = [values[0] for _ in range(max_len)] - elif len(values) != max_len: - raise ExecutionError( - "spec.data.columns must resolve to " - "the same number of rows per group." - ) - normalized_rows[label] = values - - df = pd.DataFrame(normalized_rows) - table_stores_list.append(df) + table_stores_list = _build_grouped_dataframes(resolved_columns) + for df in table_stores_list: if fetch_images: contents = df.get( "content", pd.Series(["" for _ in range(len(df))]) diff --git a/src/worker/executors/utils/graph_templates.py b/src/worker/executors/utils/graph_templates.py index f4bc46acf..b9a2a9a2d 100644 --- a/src/worker/executors/utils/graph_templates.py +++ b/src/worker/executors/utils/graph_templates.py @@ -7,7 +7,7 @@ import pandas as pd from pydantic import BaseModel -from shared.schemas.result import APIGroupItem, BaseExecutorResult +from shared.schemas.result import BaseExecutorResult from shared.tasks.specs import TaskSpecStrictBase from shared.utils.json import validate_keys @@ -602,10 +602,12 @@ def _evaluate_expr( ) -> tuple[Any, bool]: """Resolve an expression against upstream results. - Returns ``(value, grouped)``. ``grouped`` is True only when the resolved - value is a list of groups, decided from the upstream structure (a list of - ``APIGroupItem.rows``, or nested lists) — never from the shape of the cell - values. A per-row list is a cell value, not a group. + Returns ``(value, grouped)``. ``grouped`` is True when the resolved value + is a list of groups, decided from the upstream structure — the first + attribute access over the items list that yields one list per item (an + upstream API task's ``items.rows``, a python/vLLM stage's ``items.output``, + S3 ``content``) — never from the shape of the cell values. A per-row list + is a cell value, not a group. """ if not expr: return None, False @@ -618,12 +620,15 @@ def _evaluate_expr( value: Any = result grouped = False + mapped_items = False for token in parts[1:]: if not token: continue attr, indexes = _split_indexes(token) if attr: - value, grouped = _apply_attr(value, attr, token, parts, grouped) + value, grouped, mapped_items = _apply_attr( + value, attr, token, parts, grouped, mapped_items + ) for idx in indexes: value, grouped = _apply_index(value, idx, token, grouped) # Attempt to deserialize DataFrame if applicable @@ -635,51 +640,74 @@ def _evaluate_expr( def _apply_attr( - value: Any, attr: str, token: str, parts: list[str], grouped: bool -) -> tuple[Any, bool]: + value: Any, + attr: str, + token: str, + parts: list[str], + grouped: bool, + mapped_items: bool, +) -> tuple[Any, bool, bool]: """Resolve an attribute access, mapping over lists of dicts, DataFrames, - or pydantic models (including nested lists).""" + or pydantic models (including nested lists). Returns ``(value, grouped, + mapped_items)`` where ``mapped_items`` is True once an attribute has been + mapped over the items list; only that first access can set ``grouped``.""" if isinstance(value, dict) and attr in value: - return value[attr], grouped + return value[attr], grouped, mapped_items if isinstance(value, list): if all(isinstance(v, dict) and attr in v for v in value): - return [v[attr] for v in value], grouped + mapped = [v[attr] for v in value] + return ( + mapped, + grouped or _groups_on_first_access(mapped, mapped_items), + True, + ) if all(isinstance(v, pd.DataFrame) for v in value): if any(attr not in v.columns for v in value): raise ExecutionError( f"{attr} not a valid column in one of the " f"DataFrames for {token}." ) - return [v[attr].tolist() for v in value], grouped + return [v[attr].tolist() for v in value], True, True if all(isinstance(v, BaseModel) for v in value): - is_grouped = attr == "rows" and all( - isinstance(v, APIGroupItem) for v in value - ) + mapped = [_model_attr(v, attr, token) for v in value] return ( - [_model_attr(v, attr, token) for v in value], - grouped or is_grouped, + mapped, + grouped or _groups_on_first_access(mapped, mapped_items), + True, ) if all(isinstance(v, list) for v in value): - # Each inner list of records is one group; a per-row list cell is not. - is_grouped = bool(value) and all( - all(isinstance(r, (dict, BaseModel)) for r in v) for v in value - ) - mapped: list[Any] = [] + # Mapping over groups keeps grouped as-is; a raw list of lists + # groups only when its inner lists hold records. + if not grouped: + grouped = bool(value) and all( + all(isinstance(r, (dict, BaseModel)) for r in v) for v in value + ) + mapped_rows: list[Any] = [] for v in value: - inner, _ = _apply_attr(v, attr, token, parts, grouped) - mapped.append(inner) - return mapped, grouped or is_grouped + inner, _, _ = _apply_attr(v, attr, token, parts, grouped, mapped_items) + mapped_rows.append(inner) + return mapped_rows, grouped, mapped_items if isinstance(value, pd.DataFrame): if attr not in value.columns: raise ExecutionError(f"{attr} not a valid column in DataFrame for {token}.") - return value[attr].tolist(), grouped + return value[attr].tolist(), grouped, mapped_items if isinstance(value, BaseModel): - return _model_attr(value, attr, token), grouped + return _model_attr(value, attr, token), grouped, mapped_items raise ExecutionError( f"{attr} in {parts} is not a valid key - " f"{type(value).__name__}, {value}" ) +def _groups_on_first_access(mapped: list[Any], mapped_items: bool) -> bool: + """True when this is the first attribute access over the items list and it + yields one list per item (a list of lists), i.e. the grouping access. A + per-row list of scalars is a later access, which ``mapped_items`` already + rules out.""" + return ( + not mapped_items and bool(mapped) and all(isinstance(v, list) for v in mapped) + ) + + def _model_attr(value: BaseModel, attr: str, token: str) -> Any: """Resolve a declared pydantic field by name or alias, never a method.""" fields = type(value).model_fields @@ -689,7 +717,7 @@ def _model_attr(value: BaseModel, attr: str, token: str) -> Any: for name, f in fields.items(): if f.alias == attr: return getattr(value, name) - extras = getattr(value, "__pydantic_extra__", None) + extras = value.model_extra if extras and attr in extras: return extras[attr] raise ExecutionError( diff --git a/tests/shared/test_executor_result.py b/tests/shared/test_executor_result.py index 5e95c57e2..0e9efee65 100644 --- a/tests/shared/test_executor_result.py +++ b/tests/shared/test_executor_result.py @@ -108,13 +108,11 @@ def test_open_passthrough_nulls_are_preserved() -> None: "url": "http://h", "status_code": 200, "json": {"present": None, "value": 1}, - "usage": {"cost": None}, "text": None, } ) dumped = result.model_dump(by_alias=True) assert dumped["json"] == {"present": None, "value": 1} - assert dumped["usage"] == {"cost": None} assert "text" not in dumped diff --git a/tests/worker/test_api_executor.py b/tests/worker/test_api_executor.py index 91c0c42fc..5795c2190 100644 --- a/tests/worker/test_api_executor.py +++ b/tests/worker/test_api_executor.py @@ -16,7 +16,13 @@ import pytest from pydantic import ValidationError -from shared.schemas.result import APIGroupItem, APIItem, APIResult, PythonResult +from shared.schemas.result import ( + APIGroupItem, + APIItem, + APIResult, + APIUsage, + PythonResult, +) from shared.tasks.specs.misc import _MAX_CONCURRENCY, _MAX_RETRIES from shared.tasks.worker_message import WorkerTaskMessage from worker.executors import api_executor as api_executor_module @@ -1099,7 +1105,7 @@ def _handler(self, request: httpx.Request) -> httpx.Response: return httpx.Response(200, json=[1]) task = _batch_task(["a", "b"], response={"parse_json": True}, concurrency=1) - with caplog.at_level(logging.INFO, logger="worker.executors.api_executor"): + with caplog.at_level(logging.DEBUG, logger="worker.executors.api_executor"): with pytest.raises(ExecutionError, match="not a valid JSON mapping"): _run(_executor(), task, _ArrayBody(), tmp_path) call_lines = [ @@ -2335,7 +2341,7 @@ def test_chat_completion_call_line_has_tokens_and_backend( ) ] ) - with caplog.at_level(logging.INFO, logger="worker.executors.api_executor"): + with caplog.at_level(logging.DEBUG, logger="worker.executors.api_executor"): _run(_executor(), task, transport, tmp_path) call_lines = self._records(caplog, "api call") assert len(call_lines) == 1 @@ -2355,7 +2361,7 @@ def test_retried_503_then_200_shows_attempts_two( """A retried 503 then 200 logs attempts=2.""" task = _batch_task(["hi"], retries=2) transport = _SequenceTransport([_error_response(503), _ok_response()]) - with caplog.at_level(logging.INFO, logger="worker.executors.api_executor"): + with caplog.at_level(logging.DEBUG, logger="worker.executors.api_executor"): _run(_executor(), task, transport, tmp_path) call_lines = self._records(caplog, "api call") assert len(call_lines) == 1 @@ -2376,7 +2382,7 @@ def _handler(self, request: httpx.Request) -> httpx.Response: raise httpx.ConnectError("boom", request=request) transport = _RaisingTransport() - with caplog.at_level(logging.INFO, logger="worker.executors.api_executor"): + with caplog.at_level(logging.DEBUG, logger="worker.executors.api_executor"): with pytest.raises(ExecutionError, match="API request failed"): _run(_executor(), task, transport, tmp_path) call_lines = self._records(caplog, "api call") @@ -2408,7 +2414,7 @@ def _handler(self, request: httpx.Request) -> httpx.Response: return _ok_response() transport = _FirstFails() - with caplog.at_level(logging.INFO, logger="worker.executors.api_executor"): + with caplog.at_level(logging.DEBUG, logger="worker.executors.api_executor"): with pytest.raises(ExecutionError, match="row 0"): _run(_executor(), task, transport, tmp_path) call_lines = self._records(caplog, "api call") @@ -2432,7 +2438,7 @@ def test_non_json_body_logs_dash_without_raising( ) ] ) - with caplog.at_level(logging.INFO, logger="worker.executors.api_executor"): + with caplog.at_level(logging.DEBUG, logger="worker.executors.api_executor"): result = _run(_executor(), task, transport, tmp_path) assert result.items[0].text == "not json" call_lines = self._records(caplog, "api call") @@ -2459,7 +2465,7 @@ def test_malformed_telemetry_does_not_fail_request( ) ] ) - with caplog.at_level(logging.INFO, logger="worker.executors.api_executor"): + with caplog.at_level(logging.DEBUG, logger="worker.executors.api_executor"): result = _run(_executor(), task, transport, tmp_path) assert result.items[0].status_code == 200 call_lines = self._records(caplog, "api call") @@ -2503,7 +2509,7 @@ def test_summary_reports_per_backend_counts( ), ] ) - with caplog.at_level(logging.INFO, logger="worker.executors.api_executor"): + with caplog.at_level(logging.DEBUG, logger="worker.executors.api_executor"): _run(_executor(), task, transport, tmp_path) summary = self._records(caplog, "api summary") assert len(summary) == 1 @@ -2512,3 +2518,116 @@ def test_summary_reports_per_backend_counts( assert "failures=0" in msg assert "retries=0" in msg assert "backends=backend-a=2,backend-b=1" in msg + + +class TestUsage: + @pytest.fixture(autouse=True) + def _nebula_env(self, monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("NEBULA_API_BASE_URL", "https://nebula.example.com") + monkeypatch.setenv("NEBULA_API_TOKEN", "nebula-token") + + def test_usage_totals_across_rows(self, tmp_path: Path) -> None: + """The result's usage sums tokens and calls across every row, counting + a finish_reason=length call as truncated.""" + task = _batch_task(["a", "b", "c"]) + transport = _SequenceTransport( + [ + httpx.Response( + 200, + json={ + "choices": [ + { + "message": {"content": "a"}, + "finish_reason": "stop", + } + ], + "usage": { + "prompt_tokens": 10, + "completion_tokens": 5, + "completion_tokens_details": {"reasoning_tokens": 2}, + }, + }, + ), + httpx.Response( + 200, + json={ + "choices": [ + { + "message": {"content": "b"}, + "finish_reason": "length", + } + ], + "usage": { + "prompt_tokens": 20, + "completion_tokens": 7, + }, + }, + ), + httpx.Response( + 200, + json={"choices": [{"message": {"content": "c"}}]}, + ), + ] + ) + result = _run(_executor(), task, transport, tmp_path) + assert result.usage is not None + assert result.usage.prompt_tokens == 30 + assert result.usage.completion_tokens == 12 + assert result.usage.reasoning_tokens == 2 + assert result.usage.calls == 3 + assert result.usage.failures == 0 + assert result.usage.retries == 0 + assert result.usage.truncated_calls == 1 + assert result.usage.wall_sec >= 0 + + def test_usage_counts_retries(self, tmp_path: Path) -> None: + """A retried 503 then 200 counts the retry in usage.""" + task = _batch_task(["a"], retries=2) + transport = _SequenceTransport( + [ + _error_response(503), + httpx.Response( + 200, + json={ + "choices": [{"message": {"content": "a"}}], + "usage": {"prompt_tokens": 4, "completion_tokens": 1}, + }, + ), + ] + ) + result = _run(_executor(), task, transport, tmp_path) + assert result.usage is not None + assert result.usage.calls == 1 + assert result.usage.retries == 1 + assert result.usage.prompt_tokens == 4 + assert result.usage.completion_tokens == 1 + + def test_usage_absent_when_no_rows(self, tmp_path: Path) -> None: + """A task with no rows raises before producing a usage object.""" + task = _batch_task([]) + with pytest.raises(ExecutionError, match="no rows"): + _run(_executor(), task, _EchoTransport(), tmp_path) + + def test_single_call_usage_shape(self, tmp_path: Path) -> None: + """A one-row task's usage is a typed APIUsage, not a raw provider dict.""" + task = _batch_task(["hi"]) + transport = _SequenceTransport( + [ + httpx.Response( + 200, + json={ + "choices": [{"message": {"content": "hello"}}], + "usage": {"prompt_tokens": 10, "completion_tokens": 5}, + }, + ) + ] + ) + result = _run(_executor(), task, transport, tmp_path) + assert isinstance(result.usage, APIUsage) + assert result.usage.prompt_tokens == 10 + assert result.usage.completion_tokens == 5 + assert result.usage.reasoning_tokens == 0 + assert result.usage.calls == 1 + assert result.usage.failures == 0 + assert result.usage.retries == 0 + assert result.usage.truncated_calls == 0 diff --git a/tests/worker/test_data_mixin_dataframe_grouping.py b/tests/worker/test_data_mixin_dataframe_grouping.py index 19c69b65c..ab511befc 100644 --- a/tests/worker/test_data_mixin_dataframe_grouping.py +++ b/tests/worker/test_data_mixin_dataframe_grouping.py @@ -5,6 +5,8 @@ from types import SimpleNamespace from typing import Any, cast +import pandas as pd + from shared.schemas.result import APIGroupItem, APIItem, APIResult, BaseExecutorResult from worker.executors.mixins.data import DataMixin @@ -217,3 +219,83 @@ def test_index_mapping_over_nested_list_stays_one_group() -> None: df = entry.tables[0] assert len(df) == 3 assert df["content"].tolist() == ["c0", "c1", "c2"] + + +def test_example_table_and_content_shape_yields_three_prompts() -> None: + """The ``data_retrieval_then_inference.yaml`` summarize shape — a dataframe + over ``items.table.`` (a list of DataFrames) and ``items.content`` (S3 + lists) — yields one prompt per article, not one prompt for the whole + table.""" + metadata = BaseExecutorResult.model_validate( + { + "items": [ + { + "table": { + "df": pd.DataFrame( + {"title": ["t0"], "publishedDate": ["d0"]} + ).to_json() + } + }, + { + "table": { + "df": pd.DataFrame( + {"title": ["t1"], "publishedDate": ["d1"]} + ).to_json() + } + }, + { + "table": { + "df": pd.DataFrame( + {"title": ["t2"], "publishedDate": ["d2"]} + ).to_json() + } + }, + ], + "count": 3, + } + ) + html = BaseExecutorResult.model_validate( + { + "items": [ + {"content": ["0"]}, + {"content": ["1"]}, + {"content": ["2"]}, + ], + "count": 3, + } + ) + spec = cast( + Any, + SimpleNamespace( + data={ + "type": "dataframe", + "columns": [ + {"label": "title", "node": "Meta", "path": "items.table.title"}, + { + "label": "publishedDate", + "node": "Meta", + "path": "items.table.publishedDate", + }, + {"label": "html", "node": "Html", "path": "items.content"}, + ], + "messages": [ + { + "role": "user", + "content": "Summarize: {title} {publishedDate} {html}", + } + ], + }, + inference={}, + upstreamResults={"Meta": metadata, "Html": html}, + ), + ) + entry = _Mixin()._collect_prompts_for_spec(spec, "tsk-example") + + assert len(entry.prompts) == 3 + assert [len(df) for df in entry.tables] == [1, 1, 1] + assert [df["title"].tolist() for df in entry.tables] == [["t0"], ["t1"], ["t2"]] + assert [df["html"].tolist() for df in entry.tables] == [ + ["0"], + ["1"], + ["2"], + ] diff --git a/tests/worker/test_echo_executor.py b/tests/worker/test_echo_executor.py index d4aa0ce36..6aab5d133 100644 --- a/tests/worker/test_echo_executor.py +++ b/tests/worker/test_echo_executor.py @@ -1,13 +1,9 @@ -"""Echo executor tests: the literal "list" path and rejection of the removed -"function" mode.""" +"""Echo executor tests: the literal "list" path.""" from pathlib import Path -import pytest - from shared.schemas.result import EchoResult from shared.tasks import TaskType -from worker.executors.base_executor import ExecutionError from worker.executors.echo_executor import EchoExecutor from .factories import make_worker_config, make_worker_task_message @@ -30,28 +26,6 @@ def _run( return executor.run(task, tmp_path or Path("/tmp/echo-out")) -class TestListPath: - def test_literal_items_are_echoed(self) -> None: - result = _run({"type": "list", "items": ["a", "b", "c"]}) - assert [i.output for i in result.items] == ["a", "b", "c"] - - -class TestFunctionModeRejected: - def test_type_function_is_rejected(self) -> None: - """A legacy ``type: function`` payload is rejected with a pointer to the - python task, even when it also carries top-level items.""" - with pytest.raises(ExecutionError, match="use a python task instead"): - _run( - { - "type": "function", - "function": "lambda args: args[0]", - "arguments": [{"items": [1, 2, 3]}], - "items": ["ignored"], - } - ) - - def test_type_function_with_items_is_not_run_as_list(self) -> None: - """A ``type: function`` payload that also has top-level items must not - silently fall through to list mode.""" - with pytest.raises(ExecutionError, match="use a python task instead"): - _run({"type": "function", "items": ["a", "b"]}) +def test_literal_items_are_echoed() -> None: + result = _run({"type": "list", "items": ["a", "b", "c"]}) + assert [i.output for i in result.items] == ["a", "b", "c"] diff --git a/tests/worker/test_graph_templates_expr.py b/tests/worker/test_graph_templates_expr.py index 7a4cde204..98a5ad8ce 100644 --- a/tests/worker/test_graph_templates_expr.py +++ b/tests/worker/test_graph_templates_expr.py @@ -1,9 +1,20 @@ """Tests for _evaluate_expr attribute/index resolution over pydantic models, aliases, and nested lists.""" +from typing import cast + import pandas as pd -from shared.schemas.result import APIGroupItem, APIItem, APIResult +from shared.schemas.result import ( + APIGroupItem, + APIItem, + APIResult, + BaseExecutorResult, + DataRetrievalItem, + DataRetrievalResult, + InferenceItem, + InferenceResult, +) from worker.executors.utils.graph_templates import ( _aggregate_structural_messages, _build_grouped_dataframes, @@ -118,3 +129,111 @@ def test_build_grouped_dataframes_all_empty_columns_yield_zero_rows() -> None: assert dataframes[0].empty assert list(dataframes[0].columns) == ["text", "statement"] assert isinstance(dataframes[0], pd.DataFrame) + + +def _grouped_value(expr: str, upstream: object) -> tuple[list[list[object]], bool]: + value, grouped = _evaluate_expr( + expr, + cast("dict[str, BaseExecutorResult]", {"Up": upstream}), + ) + assert grouped is True + assert isinstance(value, list) and all(isinstance(v, list) for v in value) + return value, grouped + + +def test_api_rows_groups_by_item() -> None: + """An upstream API task's ``items.rows`` is grouped: one list per item.""" + upstream = APIResult( + ok=True, + executor="api", + method="POST", + url="https://up.example.com", + status_code=200, + items=[ + APIGroupItem(index=0, rows=[_item("c0"), _item("c1")]), + APIGroupItem(index=1, rows=[_item("c2")]), + ], + ) + value, _ = _grouped_value("Up.items.rows.json.choices[0].message.content", upstream) + assert value == [["c0", "c1"], ["c2"]] + + +def test_python_output_groups_by_item() -> None: + """A python stage whose items each carry a list of records in ``output`` is + grouped: ``items.output.`` yields one list per item.""" + upstream = { + "items": [ + {"output": [{"q": "a"}, {"q": "b"}]}, + {"output": [{"q": "c"}]}, + ] + } + value, _ = _grouped_value("Up.items.output.q", upstream) + assert value == [["a", "b"], ["c"]] + + +def test_vllm_output_groups_by_item() -> None: + """A vLLM ``_populate_table`` result's ``items.output`` is grouped: one list + of string outputs per item.""" + upstream = InferenceResult( + ok=True, + items=[ + InferenceItem( + index=0, prompt="p0", output=["s0", "s1"], finish_reason=None + ), + InferenceItem(index=1, prompt="p1", output=["s2"], finish_reason=None), + ], + ) + value, _ = _grouped_value("Up.items.output", upstream) + assert value == [["s0", "s1"], ["s2"]] + + +def test_s3_content_groups_by_item() -> None: + """An S3 data-retrieval result's ``items.content`` is grouped: one list per + item.""" + upstream = DataRetrievalResult( + ok=True, + items=[ + DataRetrievalItem(index=0, content=["h0", "h1"]), + DataRetrievalItem(index=1, content=["h2"]), + ], + ) + value, _ = _grouped_value("Up.items.content", upstream) + assert value == [["h0", "h1"], ["h2"]] + + +def test_per_row_scalar_list_stays_one_cell() -> None: + """A per-row list of scalars (e.g. tags) stays one cell: the grouping access + is ``items.output``, and a further attribute over the groups never + re-evaluates grouping from the inner shape.""" + upstream = { + "items": [ + {"output": [{"q": "a", "tags": ["x", "y"]}, {"q": "b", "tags": []}]}, + {"output": [{"q": "c", "tags": ["z"]}]}, + ] + } + value, _ = _grouped_value("Up.items.output.q", upstream) + assert value == [["a", "b"], ["c"]] + tags, grouped = _evaluate_expr( + "Up.items.output.tags", + cast("dict[str, BaseExecutorResult]", {"Up": upstream}), + ) + assert grouped is True + assert tags == [[["x", "y"], []], [["z"]]] + + +def test_list_of_dataframes_groups_by_item() -> None: + """Mapping an attribute over a list of DataFrames (one per item) is grouped, + so ``items.table.`` builds one table per item rather than one row with + list cells.""" + upstream = { + "items": [ + {"table": {"df": pd.DataFrame({"title": ["t0", "t1"]}).to_json()}}, + {"table": {"df": pd.DataFrame({"title": ["t2"]}).to_json()}}, + ] + } + value, grouped = _evaluate_expr( + "Up.items.table.title", + cast("dict[str, BaseExecutorResult]", {"Up": upstream}), + ) + assert grouped is True + assert value == [["t0", "t1"], ["t2"]] From 2bce1a1a9448e0c93b3b945355a89d672073ba0c Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Thu, 1 Oct 2026 12:43:26 +0700 Subject: [PATCH 65/71] feat: report token usage per API task and per workflow - store each task's usage at result ingest (results.py, workflow registry, redis) - GET /workflows/{id} sums stored usage; a merged vLLM parent records its own share so every call counts once; a completed model-calling task with no mappable usage makes the sum null (fail closed); no-model tasks contribute nothing - delete usage keys on unregister; EventMonitor takes the workflow registry - SDK Workflow.usage field Signed-off-by: Zhengyuan Su Co-Authored-By: Claude Code Signed-off-by: Zhengyuan Su --- sdk/src/flowmesh/models/__init__.py | 2 + sdk/src/flowmesh/models/workflows.py | 2 + src/server/clients/redis.py | 4 + src/server/main.py | 1 + src/server/registries/workflow.py | 72 ++++ src/server/routers/v1/results.py | 85 +++++ src/server/routers/v1/workflows.py | 56 ++- src/server/services/monitoring.py | 7 + src/shared/schemas/result/catalog.py | 6 +- .../services/test_monitoring_serve_forward.py | 1 + .../services/test_monitoring_ssh_bind.py | 1 + .../services/test_monitoring_worker_events.py | 1 + tests/server/test_event_monitor_mirror.py | 1 + tests/server/test_workflow_usage.py | 339 ++++++++++++++++++ 14 files changed, 574 insertions(+), 4 deletions(-) create mode 100644 tests/server/test_workflow_usage.py diff --git a/sdk/src/flowmesh/models/__init__.py b/sdk/src/flowmesh/models/__init__.py index 338438b7a..c2a9ee65d 100644 --- a/sdk/src/flowmesh/models/__init__.py +++ b/sdk/src/flowmesh/models/__init__.py @@ -31,6 +31,7 @@ APIGroupItem, APIItem, APIResult, + APIUsage, BaseExecutorResult, CostEstimates, DataProfilingResult, @@ -110,6 +111,7 @@ "APIGroupItem", "APIItem", "APIResult", + "APIUsage", "ActiveWaitBreakdown", "AgentBatchSummary", "AgentItem", diff --git a/sdk/src/flowmesh/models/workflows.py b/sdk/src/flowmesh/models/workflows.py index 005e462aa..5ca345626 100644 --- a/sdk/src/flowmesh/models/workflows.py +++ b/sdk/src/flowmesh/models/workflows.py @@ -3,6 +3,7 @@ from pydantic import BaseModel from .common import TaskStatus, WorkflowStatus +from .result import APIUsage class WorkflowSubmitTaskEntry(BaseModel): @@ -46,3 +47,4 @@ class Workflow(BaseModel): completed_tasks: list[str] failed_tasks: list[str] cancelled_tasks: list[str] + usage: APIUsage | None = None diff --git a/src/server/clients/redis.py b/src/server/clients/redis.py index 41517cb58..eb1112bb0 100644 --- a/src/server/clients/redis.py +++ b/src/server/clients/redis.py @@ -93,6 +93,10 @@ def task_state_key(task_id: str) -> str: return f"task:{task_id}:state" +def task_usage_key(task_id: str) -> str: + return f"task:{task_id}:usage" + + def worker_key(worker_id: str) -> str: return f"worker:{worker_id}" diff --git a/src/server/main.py b/src/server/main.py index eac83cf71..c67aef827 100644 --- a/src/server/main.py +++ b/src/server/main.py @@ -181,6 +181,7 @@ results_dir=RESULTS_DIR, log_stream_ttl_sec=config.log_stream.ttl_sec, server_base_url=config.identity.base_url, + workflow_registry=WORKFLOW_REGISTRY, ) LOG_ARCHIVER = TaskLogArchiver( diff --git a/src/server/registries/workflow.py b/src/server/registries/workflow.py index 200b613e2..593f45e60 100644 --- a/src/server/registries/workflow.py +++ b/src/server/registries/workflow.py @@ -13,10 +13,13 @@ model_serializer, ) +from shared.schemas.result import APIUsage + from ..clients.redis import ( WORKFLOWS_SET_KEY, RedisClient, task_state_key, + task_usage_key, workflow_cancelled_tasks_key, workflow_dispatched_tasks_key, workflow_failed_tasks_key, @@ -28,6 +31,13 @@ from ..utils.time import now_iso +class UnknownUsage: + """Sentinel for a model-calling task whose usage could not be mapped.""" + + +UNKNOWN_USAGE = UnknownUsage() + + class PersistedTask(BaseModel): """A durable per-task snapshot sufficient to rebuild scheduler state.""" @@ -103,6 +113,10 @@ class Workflow(BaseModel): completed_tasks: list[str] = Field(description="Completed task identifiers.") failed_tasks: list[str] = Field(description="Failed task identifiers.") cancelled_tasks: list[str] = Field(description="Cancelled task identifiers.") + usage: APIUsage | None = Field( + default=None, + description="Token/call usage summed over the workflow's finished tasks.", + ) def _create_workflow_record( @@ -175,6 +189,7 @@ def unregister_workflows(self, *workflow_ids: str) -> None: pipe.delete(*(workflow_sched_key(wid) for wid in workflow_ids)) for task_id in task_ids: pipe.delete(task_state_key(task_id)) + pipe.delete(task_usage_key(task_id)) pipe.execute() async def unregister_workflows_async(self, *workflow_ids: str) -> None: @@ -189,6 +204,7 @@ async def unregister_workflows_async(self, *workflow_ids: str) -> None: pipe.delete(*(workflow_sched_key(wid) for wid in workflow_ids)) for task_id in task_ids: pipe.delete(task_state_key(task_id)) + pipe.delete(task_usage_key(task_id)) await pipe.execute() def get_workflow_ids(self) -> set[str]: @@ -332,6 +348,62 @@ async def load_task_states_async( PersistedTask.model_validate_json(blob) if blob else None for blob in blobs ] + def save_task_usage( + self, task_id: str, usage: APIUsage | None | UnknownUsage + ) -> None: + """Synchronous variant of ``save_task_usage_async`` for the event loop.""" + key = task_usage_key(task_id) + if usage is None: + self._rds.sync.set_value(key, "null") + elif isinstance(usage, UnknownUsage): + self._rds.sync.set_value(key, "unknown") + else: + self._rds.sync.set_value(key, usage.model_dump_json()) + + async def save_task_usage_async( + self, task_id: str, usage: APIUsage | None | UnknownUsage + ) -> None: + """Persist one task's usage contribution. + + ``None`` marks a task type that makes no model calls (echo, lambda); + ``UNKNOWN_USAGE`` marks a model-calling task whose usage could not be + mapped. Both are written so a completed task with no entry at all is + detectable as one whose result was never ingested. + """ + key = task_usage_key(task_id) + if usage is None: + await self._rds.asyncio.set_value(key, "null") + elif isinstance(usage, UnknownUsage): + await self._rds.asyncio.set_value(key, "unknown") + else: + await self._rds.asyncio.set_value(key, usage.model_dump_json()) + + async def load_task_usages_async( + self, *task_ids: str + ) -> dict[str, APIUsage | None | UnknownUsage]: + """Load persisted usage for the given tasks. + + ``None`` means the task type makes no model calls; ``UNKNOWN_USAGE`` + means a model-calling task's usage could not be mapped; a task absent + from the result means its usage was never recorded. + """ + if not task_ids: + return {} + blobs = await self._rds.asyncio.mget( + [task_usage_key(task_id) for task_id in task_ids] + ) + result: dict[str, APIUsage | None | UnknownUsage] = {} + for task_id, blob in zip(task_ids, blobs, strict=True): + if blob is None: + continue + if blob == "null": + result[task_id] = None + elif blob == "unknown": + result[task_id] = UNKNOWN_USAGE + else: + result[task_id] = APIUsage.model_validate_json(blob) + return result + async def save_workflow_sched_async( self, workflow_id: str, in_epoch_order: bool, epoch_frontier: int ) -> None: diff --git a/src/server/routers/v1/results.py b/src/server/routers/v1/results.py index 3e01eaeb3..cec2a75be 100644 --- a/src/server/routers/v1/results.py +++ b/src/server/routers/v1/results.py @@ -20,7 +20,16 @@ from shared.schemas.result import ( AnyExecutorResult, + APIResult, + APIUsage, + DataProfilingResult, + DataRetrievalResult, + EchoResult, + GenerationUsage, + InferenceResult, + PythonResult, ResultEnvelope, + SSHResult, read_result, result_file_path, write_result, @@ -32,6 +41,7 @@ get_logger, get_results_dir, get_runtime, + get_workflow_registry, ) from ...auth.security import ( PrincipalContext, @@ -39,6 +49,7 @@ require_permission, ) from ...hooks import ResourceAction, ResourceKind +from ...registries.workflow import UNKNOWN_USAGE, UnknownUsage, WorkflowRegistry from ...schemas.common import PathResponse from ...services.monitoring import EventMonitor from ...task.models import TERMINAL_TASK_STATUSES @@ -52,6 +63,77 @@ router = APIRouter(prefix="/results", tags=["Results"]) +# Result types that make no model calls and therefore never carry usage. +_NO_USAGE_RESULT_TYPES = ( + EchoResult, + SSHResult, + DataProfilingResult, + DataRetrievalResult, + PythonResult, +) + + +def _task_usage_from_envelope( + envelope: ResultEnvelope, +) -> APIUsage | None | UnknownUsage: + """Return a task's usage contribution from its result envelope. + + API tasks carry an ``APIUsage``; vLLM inference tasks map their + ``GenerationUsage`` token counts with reasoning 0, ``calls`` from + ``num_requests``, and ``wall_sec`` from ``latency_sec``. Task types that + make no model calls (echo, ssh, data profiling, data retrieval, python) + contribute nothing (``None``). A model-calling task whose usage cannot be + mapped returns ``UNKNOWN_USAGE`` so the workflow fails closed. + """ + result = envelope.result + if isinstance(result, APIResult): + return result.usage if result.usage is not None else UNKNOWN_USAGE + if isinstance(result, InferenceResult): + if isinstance(result.usage, GenerationUsage): + return _inference_usage(result) + return UNKNOWN_USAGE + if isinstance(result, _NO_USAGE_RESULT_TYPES): + return None + # Any other result type either carries a usage field we do not map to + # APIUsage (embedding, agent, rag) or is an unknown model-calling type. + return UNKNOWN_USAGE + + +def _inference_usage(result: InferenceResult) -> APIUsage: + """Map an inference result's usage, subtracting merged children's shares. + + In a merged dispatch the parent's ``GenerationUsage`` is the whole batch + total, while each child carries its own share in ``result.children``. Each + child is also ingested separately, so the parent must record only its own + share (total minus the children's sum) or the workflow would count every + child's tokens and calls twice. ``wall_sec`` stays per task as reported. + """ + usage = result.usage + assert isinstance(usage, GenerationUsage) + prompt = usage.prompt_tokens + completion = usage.completion_tokens + calls = usage.num_requests + for child in result.children.values(): + if not isinstance(child, InferenceResult): + continue + child_usage = child.usage + if not isinstance(child_usage, GenerationUsage): + continue + prompt -= child_usage.prompt_tokens + completion -= child_usage.completion_tokens + calls -= child_usage.num_requests + return APIUsage( + prompt_tokens=prompt, + completion_tokens=completion, + reasoning_tokens=0, + calls=calls, + failures=0, + retries=0, + truncated_calls=0, + wall_sec=usage.latency_sec, + ) + + def _resolve_artifact_path(filename: str) -> Path: sanitized = Path(filename) if ( @@ -77,6 +159,7 @@ async def ingest_result( runtime: TaskRuntime = Depends(get_runtime), event_monitor: EventMonitor = Depends(get_event_monitor), results_dir: Path = Depends(get_results_dir), + registry: WorkflowRegistry = Depends(get_workflow_registry), logger: logging.Logger = Depends(get_logger), ) -> PathResponse: await require_permission( @@ -97,6 +180,8 @@ async def ingest_result( detail=f"Failed to store result: {exc}", ) from exc + await registry.save_task_usage_async(task_id, _task_usage_from_envelope(envelope)) + expected_artifacts: list[str] = [] record = runtime.get_record(task_id) if record: diff --git a/src/server/routers/v1/workflows.py b/src/server/routers/v1/workflows.py index 523fd37c3..79e94731c 100644 --- a/src/server/routers/v1/workflows.py +++ b/src/server/routers/v1/workflows.py @@ -15,6 +15,7 @@ from fastapi.responses import StreamingResponse from shared.schemas.event import TaskEvent +from shared.schemas.result import APIUsage from ...app_state import ( get_logger, @@ -36,7 +37,7 @@ workflow_log_stream_key, ) from ...hooks import SUBMISSION_GUARDS, ResourceAction, ResourceKind -from ...registries.workflow import Workflow, WorkflowRegistry +from ...registries.workflow import UnknownUsage, Workflow, WorkflowRegistry from ...schemas.logs import LogEntry, LogEvent, LogQueryResponse from ...schemas.workflow import ( WorkflowSubmitResponse, @@ -71,6 +72,57 @@ router = APIRouter(prefix="/workflows", tags=["Workflows"]) +def _sum_usage( + usages: dict[str, APIUsage | None | UnknownUsage], + completed_tasks: list[str], + logger: logging.Logger, +) -> APIUsage | None: + """Sum per-task usage into one workflow figure. + + Fail closed: a completed task whose usage was never recorded (its result + missing or unparsable) or could not be mapped (``UNKNOWN_USAGE``) makes the + whole workflow's usage null, with one log line naming the task. A task type + that makes no model calls (echo, lambda) maps to ``None`` and contributes + nothing, which is not a failure. When every completed task is recorded, the + (possibly all-zero) sum is returned. + """ + totals: dict[str, Any] = { + "prompt_tokens": 0, + "completion_tokens": 0, + "reasoning_tokens": 0, + "calls": 0, + "failures": 0, + "retries": 0, + "truncated_calls": 0, + "wall_sec": 0.0, + } + for task_id in completed_tasks: + if task_id not in usages: + logger.warning( + "workflow usage unavailable: task %s has no recorded usage", + task_id, + ) + return None + usage = usages[task_id] + if usage is None: + continue + if isinstance(usage, UnknownUsage): + logger.warning( + "workflow usage unavailable: task %s has unmappable usage", + task_id, + ) + return None + totals["prompt_tokens"] += usage.prompt_tokens + totals["completion_tokens"] += usage.completion_tokens + totals["reasoning_tokens"] += usage.reasoning_tokens + totals["calls"] += usage.calls + totals["failures"] += usage.failures + totals["retries"] += usage.retries + totals["truncated_calls"] += usage.truncated_calls + totals["wall_sec"] += usage.wall_sec + return APIUsage(**totals) + + def _parse_submission_body(raw_body: bytes, content_type: str) -> str: if "application/json" in content_type: try: @@ -284,6 +336,8 @@ async def get_workflow( status_code=status.HTTP_404_NOT_FOUND, detail=f"Workflow '{workflow_id}' not found", ) + usages = await registry.load_task_usages_async(*workflow.completed_tasks) + workflow.usage = _sum_usage(usages, workflow.completed_tasks, logger) return workflow diff --git a/src/server/services/monitoring.py b/src/server/services/monitoring.py index dc72bb638..60cf5a17b 100644 --- a/src/server/services/monitoring.py +++ b/src/server/services/monitoring.py @@ -51,6 +51,7 @@ ) from ..registries.node import NodeRegistry from ..registries.worker import WorkerRegistry +from ..registries.workflow import WorkflowRegistry from ..schemas.logs import LogEvent from ..task.metadata import extract_model_dataset_names from ..task.models import TaskRecord, TaskStatus, TaskUsage @@ -99,6 +100,7 @@ def __init__( node_registry: NodeRegistry, metrics_recorder: MetricsRecorder, watchdog: WorkerWatchdog, + workflow_registry: WorkflowRegistry, ssh_proxy_enabled: bool = False, serve_proxy_enabled: bool = False, port_forward: PortForwardService | None = None, @@ -119,6 +121,7 @@ def __init__( self._serve_proxy_enabled = serve_proxy_enabled self._port_forward = port_forward self._results_dir = Path(results_dir) + self._workflow_registry = workflow_registry self._log_stream_ttl_sec = max(0, int(log_stream_ttl_sec)) self._server_base_url = self._validate_server_base_url(server_base_url) @@ -277,6 +280,10 @@ def mirror_task_results(self, parent_task_id: str, child_ids: list[str]) -> None if record: expected_artifacts = record.task.spec.get_artifacts() sync_manifest(dst_dir, child_id, expected_artifacts) + # A mirrored child has no result of its own (the executor + # produced no ``children`` entry for it), so the parent's total + # already contains its calls; record the no-usage marker. + self._workflow_registry.save_task_usage(child_id, None) except Exception as exc: self._logger.debug( "Failed to mirror results from %s to %s: %s", diff --git a/src/shared/schemas/result/catalog.py b/src/shared/schemas/result/catalog.py index cfa9234d9..4ba73b77b 100644 --- a/src/shared/schemas/result/catalog.py +++ b/src/shared/schemas/result/catalog.py @@ -247,9 +247,9 @@ class EchoResult(StrictExecutorResult): class APIResult(StrictExecutorResult): - """HTTP request output. ``response_json``/``usage``/``headers`` are the - upstream API's own payloads and stay open mappings. ``items`` carries one - entry per row.""" + """HTTP request output. ``response_json``/``headers`` are the upstream + API's own payloads and stay open mappings; ``usage`` is the task's summed + token/call accounting. ``items`` carries one entry per row.""" task_type: Literal[TaskType.API] = TaskType.API executor: str diff --git a/tests/server/services/test_monitoring_serve_forward.py b/tests/server/services/test_monitoring_serve_forward.py index 55b56865c..df174aaac 100644 --- a/tests/server/services/test_monitoring_serve_forward.py +++ b/tests/server/services/test_monitoring_serve_forward.py @@ -32,6 +32,7 @@ def _make_monitor( serve_proxy_enabled=serve_proxy_enabled, port_forward=port_forward, server_base_url=server_base_url, + workflow_registry=MagicMock(), ) diff --git a/tests/server/services/test_monitoring_ssh_bind.py b/tests/server/services/test_monitoring_ssh_bind.py index 37041ebd3..ab269947e 100644 --- a/tests/server/services/test_monitoring_ssh_bind.py +++ b/tests/server/services/test_monitoring_ssh_bind.py @@ -30,6 +30,7 @@ def _make_monitor( serve_proxy_enabled=False, port_forward=port_forward, server_base_url="http://server.example.com:8000", + workflow_registry=MagicMock(), ) diff --git a/tests/server/services/test_monitoring_worker_events.py b/tests/server/services/test_monitoring_worker_events.py index 45f9f441c..ba252b6d2 100644 --- a/tests/server/services/test_monitoring_worker_events.py +++ b/tests/server/services/test_monitoring_worker_events.py @@ -20,6 +20,7 @@ def _monitor(worker_registry: MagicMock) -> EventMonitor: node_registry=MagicMock(), metrics_recorder=MagicMock(), watchdog=MagicMock(), + workflow_registry=MagicMock(), ) diff --git a/tests/server/test_event_monitor_mirror.py b/tests/server/test_event_monitor_mirror.py index 0fc42f505..74d67db10 100644 --- a/tests/server/test_event_monitor_mirror.py +++ b/tests/server/test_event_monitor_mirror.py @@ -31,6 +31,7 @@ def _make_monitor(results_dir: Path) -> EventMonitor: metrics_recorder=MagicMock(), watchdog=MagicMock(), results_dir=results_dir, + workflow_registry=MagicMock(), ) diff --git a/tests/server/test_workflow_usage.py b/tests/server/test_workflow_usage.py new file mode 100644 index 000000000..a9e594324 --- /dev/null +++ b/tests/server/test_workflow_usage.py @@ -0,0 +1,339 @@ +"""Tests for the workflow-level usage sum on GET /workflows/{id}. + +Usage is captured once at result ingest (``ingest_result``) and stored per task +in the workflow registry's Redis; ``GET /workflows/{id}`` sums the stored +values without opening any result file. +""" + +import logging +from collections.abc import Iterator +from typing import Any, cast +from unittest import mock + +import fakeredis +import pytest +from fastapi import FastAPI +from httpx import ASGITransport, AsyncClient +from lumid_hooks import PrincipalContext, ResourceRef + +from server.app_state import get_logger, get_workflow_registry +from server.auth.security import authenticate_connection +from server.clients.redis import AsyncRedisClient, RedisClient, SyncRedisClient +from server.hooks import PERMISSION_CHECKERS +from server.registries.workflow import ( + UNKNOWN_USAGE, + Workflow, + WorkflowRegistry, + WorkflowStatus, +) +from server.routers.v1 import results as results_router +from server.routers.v1 import workflows as workflows_router +from shared.schemas.result import ( + APIUsage, + GenerationUsage, + InferenceResult, + ResultEnvelope, +) + + +@pytest.fixture +def server() -> fakeredis.FakeServer: + return fakeredis.FakeServer() + + +@pytest.fixture +def registry(server: fakeredis.FakeServer) -> WorkflowRegistry: + sync = SyncRedisClient.__new__(SyncRedisClient) + sync._control = fakeredis.FakeRedis(server=server, decode_responses=True) + async_client = AsyncRedisClient.__new__(AsyncRedisClient) + cast(Any, async_client)._control = fakeredis.FakeAsyncRedis( + server=server, decode_responses=True + ) + client = RedisClient.__new__(RedisClient) + client.sync = sync + client.asyncio = async_client + return WorkflowRegistry(client) + + +def _principal() -> PrincipalContext: + return PrincipalContext( + principal_id="p-1", + org_id="org", + external_id="ext", + principal_type="user", + scopes=[], + ) + + +class _AllowAllChecker: + name = "allow-all" + + async def require( + self, + principal: PrincipalContext, + resource: ResourceRef, + action: str, + logger: logging.Logger, + ) -> None: + return None + + async def accessible_ids( + self, + principal: PrincipalContext, + kind: str, + action: str, + logger: logging.Logger, + ) -> frozenset[str] | None: + return None + + +@pytest.fixture +def allow_all_permissions() -> Iterator[None]: + PERMISSION_CHECKERS.append(_AllowAllChecker()) + try: + yield + finally: + PERMISSION_CHECKERS.clear() + + +def _api_usage(**overrides: Any) -> APIUsage: + base: dict[str, Any] = dict( + prompt_tokens=30, + completion_tokens=12, + reasoning_tokens=2, + calls=3, + failures=0, + retries=1, + truncated_calls=1, + wall_sec=2.5, + ) + base.update(overrides) + return APIUsage(**base) + + +def _inference_usage() -> APIUsage: + result = InferenceResult( + ok=True, + model="m", + items=[], + usage=GenerationUsage( + prompt_tokens=100, + completion_tokens=50, + total_tokens=150, + num_requests=4, + latency_sec=1.0, + ), + ) + usage = results_router._task_usage_from_envelope( + ResultEnvelope(task_id="tsk-inf", result=result) + ) + assert isinstance(usage, APIUsage) + return usage + + +def _merged_parent() -> InferenceResult: + """A merged vLLM parent: batch total usage plus two children's shares.""" + child = lambda pt, ct, nr: InferenceResult( # noqa: E731 + ok=True, + model="m", + items=[], + usage=GenerationUsage( + prompt_tokens=pt, + completion_tokens=ct, + total_tokens=pt + ct, + num_requests=nr, + latency_sec=1.0, + ), + ) + return InferenceResult( + ok=True, + model="m", + items=[], + usage=GenerationUsage( + prompt_tokens=100, + completion_tokens=50, + total_tokens=150, + num_requests=4, + latency_sec=1.0, + ), + children={ + "tsk-child-1": child(30, 20, 2), + "tsk-child-2": child(10, 5, 1), + }, + ) + + +def _workflow(completed: list[str]) -> Workflow: + return Workflow( + workflow_id="wfl-1", + task_ids=completed, + submitted_at="2026-01-01T00:00:00Z", + updated_at="2026-01-01T00:00:00Z", + status=WorkflowStatus.DONE, + dispatched_tasks=[], + completed_tasks=completed, + failed_tasks=[], + cancelled_tasks=[], + ) + + +async def _sum(registry: WorkflowRegistry, completed: list[str]) -> APIUsage | None: + usages = await registry.load_task_usages_async(*completed) + return workflows_router._sum_usage(usages, completed, logging.getLogger("test")) + + +@pytest.mark.anyio +async def test_usage_sums_api_and_inference_tasks(registry: WorkflowRegistry) -> None: + """API and vLLM task usage sum exactly into one workflow figure.""" + await registry.save_task_usage_async("tsk-api", _api_usage()) + await registry.save_task_usage_async("tsk-inf", _inference_usage()) + + usage = await _sum(registry, ["tsk-api", "tsk-inf"]) + assert usage is not None + assert usage.prompt_tokens == 130 + assert usage.completion_tokens == 62 + assert usage.reasoning_tokens == 2 + assert usage.calls == 7 + assert usage.retries == 1 + assert usage.truncated_calls == 1 + assert usage.wall_sec == 3.5 + + +@pytest.mark.anyio +async def test_no_model_task_contributes_nothing(registry: WorkflowRegistry) -> None: + """A task that calls no model contributes nothing and is not a failure.""" + await registry.save_task_usage_async("tsk-api", _api_usage()) + await registry.save_task_usage_async("tsk-echo", None) + + usage = await _sum(registry, ["tsk-api", "tsk-echo"]) + assert usage is not None + assert usage.prompt_tokens == 30 + assert usage.calls == 3 + + +@pytest.mark.anyio +@pytest.mark.parametrize("unmappable", [False, True]) +async def test_missing_or_unmappable_usage_makes_sum_null( + registry: WorkflowRegistry, unmappable: bool +) -> None: + """A completed model-calling task with missing or unmappable usage nulls the sum.""" + await registry.save_task_usage_async("tsk-api", _api_usage()) + if unmappable: + await registry.save_task_usage_async("tsk-inf", UNKNOWN_USAGE) + + assert await _sum(registry, ["tsk-api", "tsk-inf"]) is None + + +@pytest.mark.anyio +@pytest.mark.parametrize( + "order", + [ + ["tsk-parent", "tsk-child-1", "tsk-child-2"], + ["tsk-child-1", "tsk-child-2", "tsk-parent"], + ], +) +async def test_merged_parent_plus_children_sums_to_batch_total( + registry: WorkflowRegistry, order: list[str] +) -> None: + """A merged parent and its children sum to the batch total in either order.""" + parent = _merged_parent() + parent_usage = results_router._task_usage_from_envelope( + ResultEnvelope(task_id="tsk-parent", result=parent) + ) + assert isinstance(parent_usage, APIUsage) + await registry.save_task_usage_async("tsk-parent", parent_usage) + for child_id in ("tsk-child-1", "tsk-child-2"): + child_usage = results_router._task_usage_from_envelope( + ResultEnvelope(task_id=child_id, result=parent.children[child_id]) + ) + assert isinstance(child_usage, APIUsage) + await registry.save_task_usage_async(child_id, child_usage) + + usage = await _sum(registry, order) + assert usage is not None + assert usage.prompt_tokens == 100 + assert usage.completion_tokens == 50 + assert usage.calls == 4 + + +@pytest.mark.anyio +async def test_mirrored_child_records_no_usage( + registry: WorkflowRegistry, tmp_path: Any +) -> None: + """A mirrored child records the no-usage marker (it made no calls).""" + from server.services.monitoring import EventMonitor + + runtime = mock.Mock() + runtime.get_record.return_value = None + monitor = EventMonitor( + redis_client=mock.Mock(), + logger=logging.getLogger("test"), + runtime=runtime, + dispatcher=mock.Mock(), + worker_registry=mock.Mock(), + node_registry=mock.Mock(), + metrics_recorder=mock.Mock(), + watchdog=mock.Mock(), + results_dir=tmp_path, + workflow_registry=registry, + ) + parent_dir = tmp_path / "tsk-parent" + parent_dir.mkdir(parents=True) + (parent_dir / "results.json").write_text("{}", encoding="utf-8") + + monitor.mirror_task_results("tsk-parent", ["tsk-clone"]) + + usages = await registry.load_task_usages_async("tsk-clone") + assert usages["tsk-clone"] is None + + +@pytest.mark.anyio +async def test_unregister_deletes_usage_keys( + registry: WorkflowRegistry, server: fakeredis.FakeServer +) -> None: + """Unregistering a workflow deletes its tasks' usage keys (no leak).""" + rds = fakeredis.FakeRedis(server=server, decode_responses=True) + rds.hset("workflow:wfl-1", "workflow_id", "wfl-1") + rds.hset("workflow:wfl-1", "task_ids", '["tsk-1","tsk-2"]') + await registry.save_task_usage_async("tsk-1", _api_usage()) + await registry.save_task_usage_async("tsk-2", None) + + await registry.unregister_workflows_async("wfl-1") + + assert rds.exists("task:tsk-1:usage") == 0 + assert rds.exists("task:tsk-2:usage") == 0 + + +@pytest.mark.anyio +async def test_get_workflow_does_not_open_result_files( + registry: WorkflowRegistry, allow_all_permissions: None +) -> None: + """GET /workflows/{id} sums stored usage without reading any result file.""" + await registry.save_task_usage_async("tsk-api", _api_usage()) + await registry.save_task_usage_async("tsk-echo", None) + + app = FastAPI() + app.state.logger = logging.getLogger("test.workflow_usage") + app.include_router(workflows_router.router, prefix="/api/v1") + app.dependency_overrides[get_workflow_registry] = lambda: registry + app.dependency_overrides[get_logger] = lambda: logging.getLogger( + "test.workflow_usage" + ) + app.dependency_overrides[authenticate_connection] = lambda: _principal() + + with ( + mock.patch.object( + registry, + "get_workflow_async", + return_value=_workflow(["tsk-api", "tsk-echo"]), + ), + mock.patch("shared.schemas.result.read_result") as read_result, + ): + async with AsyncClient( + transport=ASGITransport(app=app), base_url="http://t" + ) as ac: + resp = await ac.get("/api/v1/workflows/wfl-1") + read_result.assert_not_called() + + assert resp.status_code == 200 + assert resp.json()["usage"]["prompt_tokens"] == 30 From d7af1404e3df6d11f342fb4adff793df3707cf5e Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Thu, 1 Oct 2026 14:14:57 +0700 Subject: [PATCH 66/71] fix: serve tasks contribute no usage to the workflow sum Co-Authored-By: Claude Code Signed-off-by: Zhengyuan Su --- src/server/routers/v1/results.py | 8 +++++--- tests/server/test_workflow_usage.py | 31 ++++++++++++++++++++++++++--- 2 files changed, 33 insertions(+), 6 deletions(-) diff --git a/src/server/routers/v1/results.py b/src/server/routers/v1/results.py index cec2a75be..0b246b4d4 100644 --- a/src/server/routers/v1/results.py +++ b/src/server/routers/v1/results.py @@ -29,6 +29,7 @@ InferenceResult, PythonResult, ResultEnvelope, + ServeResult, SSHResult, read_result, result_file_path, @@ -70,6 +71,7 @@ DataProfilingResult, DataRetrievalResult, PythonResult, + ServeResult, ) @@ -81,9 +83,9 @@ def _task_usage_from_envelope( API tasks carry an ``APIUsage``; vLLM inference tasks map their ``GenerationUsage`` token counts with reasoning 0, ``calls`` from ``num_requests``, and ``wall_sec`` from ``latency_sec``. Task types that - make no model calls (echo, ssh, data profiling, data retrieval, python) - contribute nothing (``None``). A model-calling task whose usage cannot be - mapped returns ``UNKNOWN_USAGE`` so the workflow fails closed. + make no model calls (echo, ssh, serve, data profiling, data retrieval, + python) contribute nothing (``None``). A model-calling task whose usage + cannot be mapped returns ``UNKNOWN_USAGE`` so the workflow fails closed. """ result = envelope.result if isinstance(result, APIResult): diff --git a/tests/server/test_workflow_usage.py b/tests/server/test_workflow_usage.py index a9e594324..f9d325899 100644 --- a/tests/server/test_workflow_usage.py +++ b/tests/server/test_workflow_usage.py @@ -30,9 +30,15 @@ from server.routers.v1 import workflows as workflows_router from shared.schemas.result import ( APIUsage, + DataProfilingResult, + DataRetrievalResult, + EchoResult, GenerationUsage, InferenceResult, + PythonResult, ResultEnvelope, + ServeResult, + SSHResult, ) @@ -200,12 +206,31 @@ async def test_usage_sums_api_and_inference_tasks(registry: WorkflowRegistry) -> @pytest.mark.anyio -async def test_no_model_task_contributes_nothing(registry: WorkflowRegistry) -> None: +@pytest.mark.parametrize( + "result", + [ + EchoResult(ok=True), + SSHResult(ok=True, session_id="s", exit_code=0), + DataProfilingResult(ok=True), + DataRetrievalResult(ok=True), + PythonResult(ok=True, exit_code=0), + ServeResult(ok=True, model="m", port=8000), + ], +) +async def test_no_model_task_contributes_nothing( + registry: WorkflowRegistry, result: Any +) -> None: """A task that calls no model contributes nothing and is not a failure.""" + assert ( + results_router._task_usage_from_envelope( + ResultEnvelope(task_id="tsk-no-model", result=result) + ) + is None + ) await registry.save_task_usage_async("tsk-api", _api_usage()) - await registry.save_task_usage_async("tsk-echo", None) + await registry.save_task_usage_async("tsk-no-model", None) - usage = await _sum(registry, ["tsk-api", "tsk-echo"]) + usage = await _sum(registry, ["tsk-api", "tsk-no-model"]) assert usage is not None assert usage.prompt_tokens == 30 assert usage.calls == 3 From 3ea5ac70523c9283d60df8aa5a48683cfc73d7f0 Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Thu, 1 Oct 2026 14:16:16 +0700 Subject: [PATCH 67/71] docs: shorten the grouped results section Co-Authored-By: Claude Code Signed-off-by: Zhengyuan Su --- docs/WORKFLOWS.md | 58 ++++------------------------------------------- 1 file changed, 5 insertions(+), 53 deletions(-) diff --git a/docs/WORKFLOWS.md b/docs/WORKFLOWS.md index aadff04d6..0f4b13e88 100644 --- a/docs/WORKFLOWS.md +++ b/docs/WORKFLOWS.md @@ -113,61 +113,13 @@ exactly `{{prompt}}` takes the prompt as-is, so a message-list row fills requests. Any failed row fails the task. Cancelling the task skips rows that have not started and marks it cancelled once in-flight requests return. -A body value that is exactly `{{prompt}}` is replaced by the row's prompt -object as-is (a message list stays a list of `{"role", "content"}` dicts). An -embedded `{{prompt}}` inside a longer string keeps string substitution: a -string prompt is inserted verbatim, and any other prompt value is rendered as -JSON. - ### Grouped results -A `dataframe` spec returns one `APIGroupItem` per table in `APIResult.items`, -with the table's row responses in `rows`; an empty table has empty `rows`. A -`graph_template` spec over grouped columns sends one request per group and -returns one `APIItem` for each. Other specs return one `APIItem` per row. - -A column is grouped when it reads one list of records per upstream item: an -upstream API task's `items.rows`, or a python stage whose items each carry a -list of records in `output` (`path: value.items.output.`). Each list -becomes one table, so group sizes may differ. A per-row list of scalars stays -one cell. - -Downstream stages read a grouped result through `rows`, for example -`path: items.rows.json.choices[0].message.content`. - -A dataframe column reads a python stage with `node: ` and a path that -starts at the result as `flowmesh result fetch` shows it — for a python stage -that returns `{"items": [{"output": [...]}, ...]}`, `path: value.items.output.q` -reads the `q` field of each record. When each item's `output` is a list of -records, each item is one group and group sizes may differ (3 and 2); a per-row -list of scalars stays one cell value. For example, a python stage T0 that -returns two such items feeds a dataframe API task T1 that reads them as two -groups: - -```yaml -spec: - stages: - - name: T0 - spec: - taskType: python - code: | - def main(): - return {"items": [{"output": [{"q": "..."}, {"q": "..."}, {"q": "..."}]}, - {"output": [{"q": "..."}, {"q": "..."}]}]} - - name: T1 - dependsOn: [T0] - spec: - taskType: api - data: - type: dataframe - columns: - - label: Q - node: T0 - path: value.items.output.q - messages: - - role: user - content: "Answer in one word: {Q}" -``` +A `dataframe` spec returns one `APIGroupItem` per table, with that table's +responses in `rows`. A column is grouped when it reads one list of records per +upstream item, such as an upstream API task's `items.rows` or a python stage's +`items.output.`; each list becomes one table. Downstream stages read the +responses through `items.rows`. ## Python task From c33e6b2282e5690819e849b7624dcb69b1497996 Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Thu, 1 Oct 2026 14:44:03 +0700 Subject: [PATCH 68/71] docs: restore the embedded {{prompt}} substitution rule Co-Authored-By: Claude Opus 5.5 (1M context) Signed-off-by: Zhengyuan Su --- docs/WORKFLOWS.md | 6 ++++++ 1 file changed, 6 insertions(+) diff --git a/docs/WORKFLOWS.md b/docs/WORKFLOWS.md index 0f4b13e88..4cd074d20 100644 --- a/docs/WORKFLOWS.md +++ b/docs/WORKFLOWS.md @@ -113,6 +113,12 @@ exactly `{{prompt}}` takes the prompt as-is, so a message-list row fills requests. Any failed row fails the task. Cancelling the task skips rows that have not started and marks it cancelled once in-flight requests return. +A body value that is exactly `{{prompt}}` is replaced by the row's prompt +object as-is (a message list stays a list of `{"role", "content"}` dicts). An +embedded `{{prompt}}` inside a longer string keeps string substitution: a +string prompt is inserted verbatim, and any other prompt value is rendered as +JSON. + ### Grouped results A `dataframe` spec returns one `APIGroupItem` per table, with that table's From 43535935f6b6537d1e6ce7ae687605af5aba4e92 Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Thu, 1 Oct 2026 16:14:36 +0700 Subject: [PATCH 69/71] fix: a condition-skipped task records no usage instead of nulling the sum The dispatcher marks a condition-skipped task succeeded without a usage entry, so GET /workflows/{id} treated it as fail-closed and returned null usage for the whole workflow. The skip path now records the no-usage marker before the task is marked succeeded; the workflow registry is a required Dispatcher argument. Co-Authored-By: Claude Sonnet 5.5 Signed-off-by: Zhengyuan Su --- src/server/dispatcher/base.py | 5 + src/server/dispatcher/factory.py | 3 + src/server/main.py | 1 + tests/server/dispatcher/helpers.py | 4 + .../test_merged_child_redaction_dispatch.py | 1 + .../dispatcher/test_python_input_dispatch.py | 1 + .../server/task/test_python_input_mounting.py | 2 + tests/server/task/test_ssh_result_mounting.py | 6 ++ .../task/test_stage_reference_resolution.py | 3 + tests/server/test_workflow_usage.py | 97 +++++++++++++++++++ 10 files changed, 123 insertions(+) diff --git a/src/server/dispatcher/base.py b/src/server/dispatcher/base.py index 740bbfcaf..46a74c19b 100644 --- a/src/server/dispatcher/base.py +++ b/src/server/dispatcher/base.py @@ -35,6 +35,7 @@ from ..clients.redis import REDIS_CONN_ERRORS from ..registries.worker import Worker, WorkerRegistry +from ..registries.workflow import WorkflowRegistry from ..services.metrics import MetricsRecorder from ..task.metadata import extract_model_dataset_names from ..task.models import TaskRecord, TaskStatus @@ -67,6 +68,7 @@ def __init__( worker_registry: WorkerRegistry, results_dir: Path, logger: logging.Logger, + workflow_registry: WorkflowRegistry, worker_selection_strategy: str = DEFAULT_WORKER_SELECTION, enable_context_reuse: bool = True, enable_task_merge: bool = True, @@ -82,6 +84,7 @@ def __init__( self._worker_registry = worker_registry self._logger = logger self._results_dir = Path(results_dir) + self._workflow_registry = workflow_registry self._worker_selection_strategy = worker_selection_strategy self._context_reuse_enabled = enable_context_reuse self._task_merge_enabled = enable_task_merge @@ -1256,6 +1259,8 @@ def _evaluate_condition_skip( write_result(self._results_dir, skip_envelope) self._runtime.release_merge(task_id) ts = now_iso() + # A skipped task made no model calls; record the no-usage marker. + self._workflow_registry.save_task_usage(task_id, None) self._runtime.mark_succeeded( task_id, worker_id=None, diff --git a/src/server/dispatcher/factory.py b/src/server/dispatcher/factory.py index de73df041..615276197 100644 --- a/src/server/dispatcher/factory.py +++ b/src/server/dispatcher/factory.py @@ -4,6 +4,7 @@ from ..config import DispatchConfig from ..dispatcher import Dispatcher from ..registries.worker import WorkerRegistry +from ..registries.workflow import WorkflowRegistry from ..services.metrics import MetricsRecorder from ..task.runtime import TaskRuntime @@ -20,6 +21,7 @@ def create_dispatcher( worker_registry: WorkerRegistry, results_dir: Path, logger: logging.Logger, + workflow_registry: WorkflowRegistry, metrics_recorder: MetricsRecorder | None = None, ) -> Dispatcher: """ @@ -59,4 +61,5 @@ def create_dispatcher( enable_stage_weight_stickiness=config.enable_stage_weight_stickiness, no_worker_grace_sec=config.no_worker_grace_sec, metrics_recorder=metrics_recorder, + workflow_registry=workflow_registry, ) diff --git a/src/server/main.py b/src/server/main.py index c67aef827..bf93fbff0 100644 --- a/src/server/main.py +++ b/src/server/main.py @@ -127,6 +127,7 @@ RESULTS_DIR, logger=logger, metrics_recorder=METRICS_RECORDER, + workflow_registry=WORKFLOW_REGISTRY, ) _pf_cfg = config.port_forward diff --git a/tests/server/dispatcher/helpers.py b/tests/server/dispatcher/helpers.py index 7216b398d..29954e660 100644 --- a/tests/server/dispatcher/helpers.py +++ b/tests/server/dispatcher/helpers.py @@ -56,6 +56,7 @@ def make_capturing_dispatcher( idle_ids: list[str] | None = None, satisfying_ids: list[str] | None = None, grace_sec: int = 60, + workflow_registry: Any = None, ) -> CapturingDispatcher: """Build a CapturingDispatcher whose registry returns the given worker ids.""" registry = mock.Mock() @@ -71,4 +72,7 @@ def make_capturing_dispatcher( results_dir=Path(tempfile.gettempdir()), logger=logging.getLogger("dispatcher-test"), no_worker_grace_sec=grace_sec, + workflow_registry=( + workflow_registry if workflow_registry is not None else mock.Mock() + ), ) diff --git a/tests/server/dispatcher/test_merged_child_redaction_dispatch.py b/tests/server/dispatcher/test_merged_child_redaction_dispatch.py index d1de71137..4224f7ee0 100644 --- a/tests/server/dispatcher/test_merged_child_redaction_dispatch.py +++ b/tests/server/dispatcher/test_merged_child_redaction_dispatch.py @@ -94,6 +94,7 @@ def test_dispatch_fails_redacted_merged_child_and_carries_survivors() -> None: enable_context_reuse=False, enable_task_merge=True, task_merge_max_batch_size=4, + workflow_registry=mock.Mock(), ) assert disp.dispatch_once(parent) is True diff --git a/tests/server/dispatcher/test_python_input_dispatch.py b/tests/server/dispatcher/test_python_input_dispatch.py index c6d75d748..2b28a5455 100644 --- a/tests/server/dispatcher/test_python_input_dispatch.py +++ b/tests/server/dispatcher/test_python_input_dispatch.py @@ -58,6 +58,7 @@ def test_unknown_input_stage_fails_the_task(tmp_path: Path) -> None: logger=logging.getLogger("dispatch-python-inputs"), worker_selection_strategy="first_fit", enable_context_reuse=False, + workflow_registry=mock.Mock(), ) assert disp.dispatch_once(nodes["score"]) is True diff --git a/tests/server/task/test_python_input_mounting.py b/tests/server/task/test_python_input_mounting.py index 1c0a02120..9dab6b713 100644 --- a/tests/server/task/test_python_input_mounting.py +++ b/tests/server/task/test_python_input_mounting.py @@ -3,6 +3,7 @@ import logging from pathlib import Path from typing import cast +from unittest import mock import pytest @@ -47,6 +48,7 @@ def _dispatcher( worker_registry=cast(WorkerRegistry, object()), results_dir=results_dir, logger=logging.getLogger("test-python-inputs"), + workflow_registry=mock.Mock(), ) diff --git a/tests/server/task/test_ssh_result_mounting.py b/tests/server/task/test_ssh_result_mounting.py index 1953be090..2240dd2ea 100644 --- a/tests/server/task/test_ssh_result_mounting.py +++ b/tests/server/task/test_ssh_result_mounting.py @@ -6,6 +6,7 @@ from pathlib import Path from types import SimpleNamespace from typing import cast +from unittest import mock import pytest @@ -109,6 +110,7 @@ def test_dispatcher_resolves_ssh_input_stage_names_from_local_stage_names() -> N worker_registry=cast(WorkerRegistry, object()), results_dir=Path("/tmp"), logger=logging.getLogger("test-ssh-phase2"), + workflow_registry=mock.Mock(), ) spec = SSHSpecStrict.model_validate(current.task.spec.model_dump()) @@ -149,6 +151,7 @@ def test_dispatcher_requeues_when_ssh_input_stage_not_done() -> None: worker_registry=cast(WorkerRegistry, object()), results_dir=Path("/tmp"), logger=logging.getLogger("test-ssh-phase2"), + workflow_registry=mock.Mock(), ) spec = SSHSpecStrict.model_validate(current.task.spec.model_dump()) @@ -219,6 +222,7 @@ def test_build_stage_context_includes_only_transitive_dependencies() -> None: worker_registry=cast(WorkerRegistry, object()), results_dir=Path("/tmp"), logger=logging.getLogger("test-stage-context"), + workflow_registry=mock.Mock(), ) context = dispatcher._build_stage_context(current) @@ -298,6 +302,7 @@ def test_collect_upstream_results_excludes_unrelated_completed_stages( worker_registry=cast(WorkerRegistry, object()), results_dir=tmp_path, logger=logging.getLogger("test-stage-results"), + workflow_registry=mock.Mock(), ) context = dispatcher._build_stage_context(current) @@ -372,6 +377,7 @@ def test_stage_reference_uses_payload_root_for_local_and_http_results( worker_registry=cast(WorkerRegistry, object()), results_dir=tmp_path, logger=logging.getLogger("test-stage-reference-root"), + workflow_registry=mock.Mock(), ) local_value = dispatcher._resolve_reference( diff --git a/tests/server/task/test_stage_reference_resolution.py b/tests/server/task/test_stage_reference_resolution.py index 458582fc2..21a2c4dd9 100644 --- a/tests/server/task/test_stage_reference_resolution.py +++ b/tests/server/task/test_stage_reference_resolution.py @@ -5,6 +5,7 @@ from pathlib import Path from types import SimpleNamespace from typing import cast +from unittest import mock from server.dispatcher.base import Dispatcher from server.registries.worker import WorkerRegistry @@ -100,6 +101,7 @@ def test_api_dependent_stage_resolves_first_row_text(tmp_path: Path) -> None: worker_registry=cast(WorkerRegistry, object()), results_dir=tmp_path, logger=logging.getLogger("test-api-dependent-stage"), + workflow_registry=mock.Mock(), ) value = dispatcher._resolve_reference("stage.items.0.text", {"stage": upstream}) @@ -203,6 +205,7 @@ def test_translated_n8n_dependent_api_stage_resolves(tmp_path: Path) -> None: worker_registry=cast(WorkerRegistry, object()), results_dir=tmp_path, logger=logging.getLogger("test-n8n-dependent-stage"), + workflow_registry=mock.Mock(), ) context = dispatcher._build_stage_context(downstream_record) diff --git a/tests/server/test_workflow_usage.py b/tests/server/test_workflow_usage.py index f9d325899..04490d2e5 100644 --- a/tests/server/test_workflow_usage.py +++ b/tests/server/test_workflow_usage.py @@ -7,6 +7,7 @@ import logging from collections.abc import Iterator +from pathlib import Path from typing import Any, cast from unittest import mock @@ -19,6 +20,7 @@ from server.app_state import get_logger, get_workflow_registry from server.auth.security import authenticate_connection from server.clients.redis import AsyncRedisClient, RedisClient, SyncRedisClient +from server.dispatcher.base import Dispatcher from server.hooks import PERMISSION_CHECKERS from server.registries.workflow import ( UNKNOWN_USAGE, @@ -28,7 +30,10 @@ ) from server.routers.v1 import results as results_router from server.routers.v1 import workflows as workflows_router +from server.task.runtime import TaskRuntime from shared.schemas.result import ( + APIItem, + APIResult, APIUsage, DataProfilingResult, DataRetrievalResult, @@ -39,7 +44,9 @@ ResultEnvelope, ServeResult, SSHResult, + write_result, ) +from shared.tasks import TaskEnvelopeStrict @pytest.fixture @@ -362,3 +369,93 @@ async def test_get_workflow_does_not_open_result_files( assert resp.status_code == 200 assert resp.json()["usage"]["prompt_tokens"] == 30 + + +_SKIP_WORKFLOW = """ +apiVersion: flowmesh/v1 +kind: Workflow +metadata: + name: skip-usage +spec: + graph: + nodes: + - name: judge + spec: + taskType: api + - name: refine + dependsOn: [judge] + spec: + taskType: api + condition: + node: judge + field: items.0.text + equals: insufficient +""" + + +@pytest.mark.anyio +async def test_condition_skipped_task_records_no_usage_and_keeps_sum( + registry: WorkflowRegistry, allow_all_permissions: None, tmp_path: Path +) -> None: + """A condition-skipped task records the no-usage marker, so the workflow + usage stays the sum of the tasks that made calls.""" + logger = logging.getLogger("test.workflow_usage") + runtime = TaskRuntime(registry, mock.Mock(), logger) + workflow_id, parsed = await runtime.register("owner", "org", _SKIP_WORKFLOW) + ids = {str(p.graph_node_name): p.task_id for p in parsed} + + write_result( + tmp_path, + ResultEnvelope( + task_id=ids["judge"], + result=APIResult( + ok=True, + executor="api", + method="POST", + url="https://api.example.com/v1/chat/completions", + status_code=200, + items=[ + APIItem( + index=0, + url="https://api.example.com/v1/chat/completions", + status_code=200, + text="sufficient", + ) + ], + ), + ), + ) + await registry.save_task_usage_async(ids["judge"], _api_usage(calls=7)) + runtime.mark_succeeded(ids["judge"], None, {}, "2026-01-01T00:00:00Z") + + dispatcher = Dispatcher( + runtime=runtime, + worker_registry=mock.Mock(), + results_dir=tmp_path, + logger=logger, + workflow_registry=registry, + ) + record = runtime.get_record(ids["refine"]) + assert record is not None + skipped = dispatcher._evaluate_condition_skip( + ids["refine"], TaskEnvelopeStrict.model_validate(record.task), record + ) + + assert skipped is True + usages = await registry.load_task_usages_async(ids["refine"]) + assert usages == {ids["refine"]: None} + + app = FastAPI() + app.state.logger = logger + app.include_router(workflows_router.router, prefix="/api/v1") + app.dependency_overrides[get_workflow_registry] = lambda: registry + app.dependency_overrides[get_logger] = lambda: logger + app.dependency_overrides[authenticate_connection] = lambda: _principal() + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://t") as ac: + resp = await ac.get(f"/api/v1/workflows/{workflow_id}") + + assert resp.status_code == 200 + body = resp.json() + assert body["usage"] is not None + assert body["usage"]["calls"] == 7 + assert body["usage"]["prompt_tokens"] == 30 From 01ee51cc0e5a1c27c4986ba56961ee020f73d5c7 Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Thu, 1 Oct 2026 16:14:42 +0700 Subject: [PATCH 70/71] test: guard APIUsage against server and SDK drift Co-Authored-By: Claude Sonnet 5.5 Signed-off-by: Zhengyuan Su --- tests/sdk/test_schema_compat.py | 1 + 1 file changed, 1 insertion(+) diff --git a/tests/sdk/test_schema_compat.py b/tests/sdk/test_schema_compat.py index ccfce65ca..1290f5849 100644 --- a/tests/sdk/test_schema_compat.py +++ b/tests/sdk/test_schema_compat.py @@ -156,6 +156,7 @@ "EchoItem", "APIItem", "APIGroupItem", + "APIUsage", ] RESULT_MODEL_PAIRS = [ From 349dfaeb651e705db53cf0c6c033bacf7ed4c163 Mon Sep 17 00:00:00 2001 From: Zhengyuan Su Date: Fri, 2 Oct 2026 18:31:43 +0700 Subject: [PATCH 71/71] refactor: drop workflow-level usage; clients sum per-task usage Remove the per-workflow usage store and sum: the usage keys in the Redis client, the save/load/delete and UnknownUsage sentinel in the workflow registry, the usage ingest and task-type mapping in the results router, _sum_usage and the usage line in GET /workflows/{id}, and the registry wiring threaded through main, monitoring, and the dispatcher. The SDK Workflow.usage field and its test file go too. Per-task APIResult.usage stays: clients sum each task's result themselves, which also survives a Redis restart. Co-Authored-By: Claude Opus 5.5 (1M context) Signed-off-by: Zhengyuan Su --- sdk/src/flowmesh/models/workflows.py | 2 - src/server/clients/redis.py | 4 - src/server/dispatcher/base.py | 5 - src/server/dispatcher/factory.py | 3 - src/server/main.py | 2 - src/server/registries/workflow.py | 72 --- src/server/routers/v1/results.py | 87 ---- src/server/routers/v1/workflows.py | 56 +-- src/server/services/monitoring.py | 7 - tests/server/dispatcher/helpers.py | 4 - .../test_merged_child_redaction_dispatch.py | 1 - .../dispatcher/test_python_input_dispatch.py | 1 - .../services/test_monitoring_serve_forward.py | 1 - .../services/test_monitoring_ssh_bind.py | 1 - .../services/test_monitoring_worker_events.py | 1 - .../server/task/test_python_input_mounting.py | 2 - tests/server/task/test_ssh_result_mounting.py | 6 - .../task/test_stage_reference_resolution.py | 3 - tests/server/test_event_monitor_mirror.py | 1 - tests/server/test_workflow_usage.py | 461 ------------------ 20 files changed, 1 insertion(+), 719 deletions(-) delete mode 100644 tests/server/test_workflow_usage.py diff --git a/sdk/src/flowmesh/models/workflows.py b/sdk/src/flowmesh/models/workflows.py index 5ca345626..005e462aa 100644 --- a/sdk/src/flowmesh/models/workflows.py +++ b/sdk/src/flowmesh/models/workflows.py @@ -3,7 +3,6 @@ from pydantic import BaseModel from .common import TaskStatus, WorkflowStatus -from .result import APIUsage class WorkflowSubmitTaskEntry(BaseModel): @@ -47,4 +46,3 @@ class Workflow(BaseModel): completed_tasks: list[str] failed_tasks: list[str] cancelled_tasks: list[str] - usage: APIUsage | None = None diff --git a/src/server/clients/redis.py b/src/server/clients/redis.py index eb1112bb0..41517cb58 100644 --- a/src/server/clients/redis.py +++ b/src/server/clients/redis.py @@ -93,10 +93,6 @@ def task_state_key(task_id: str) -> str: return f"task:{task_id}:state" -def task_usage_key(task_id: str) -> str: - return f"task:{task_id}:usage" - - def worker_key(worker_id: str) -> str: return f"worker:{worker_id}" diff --git a/src/server/dispatcher/base.py b/src/server/dispatcher/base.py index 46a74c19b..740bbfcaf 100644 --- a/src/server/dispatcher/base.py +++ b/src/server/dispatcher/base.py @@ -35,7 +35,6 @@ from ..clients.redis import REDIS_CONN_ERRORS from ..registries.worker import Worker, WorkerRegistry -from ..registries.workflow import WorkflowRegistry from ..services.metrics import MetricsRecorder from ..task.metadata import extract_model_dataset_names from ..task.models import TaskRecord, TaskStatus @@ -68,7 +67,6 @@ def __init__( worker_registry: WorkerRegistry, results_dir: Path, logger: logging.Logger, - workflow_registry: WorkflowRegistry, worker_selection_strategy: str = DEFAULT_WORKER_SELECTION, enable_context_reuse: bool = True, enable_task_merge: bool = True, @@ -84,7 +82,6 @@ def __init__( self._worker_registry = worker_registry self._logger = logger self._results_dir = Path(results_dir) - self._workflow_registry = workflow_registry self._worker_selection_strategy = worker_selection_strategy self._context_reuse_enabled = enable_context_reuse self._task_merge_enabled = enable_task_merge @@ -1259,8 +1256,6 @@ def _evaluate_condition_skip( write_result(self._results_dir, skip_envelope) self._runtime.release_merge(task_id) ts = now_iso() - # A skipped task made no model calls; record the no-usage marker. - self._workflow_registry.save_task_usage(task_id, None) self._runtime.mark_succeeded( task_id, worker_id=None, diff --git a/src/server/dispatcher/factory.py b/src/server/dispatcher/factory.py index 615276197..de73df041 100644 --- a/src/server/dispatcher/factory.py +++ b/src/server/dispatcher/factory.py @@ -4,7 +4,6 @@ from ..config import DispatchConfig from ..dispatcher import Dispatcher from ..registries.worker import WorkerRegistry -from ..registries.workflow import WorkflowRegistry from ..services.metrics import MetricsRecorder from ..task.runtime import TaskRuntime @@ -21,7 +20,6 @@ def create_dispatcher( worker_registry: WorkerRegistry, results_dir: Path, logger: logging.Logger, - workflow_registry: WorkflowRegistry, metrics_recorder: MetricsRecorder | None = None, ) -> Dispatcher: """ @@ -61,5 +59,4 @@ def create_dispatcher( enable_stage_weight_stickiness=config.enable_stage_weight_stickiness, no_worker_grace_sec=config.no_worker_grace_sec, metrics_recorder=metrics_recorder, - workflow_registry=workflow_registry, ) diff --git a/src/server/main.py b/src/server/main.py index bf93fbff0..eac83cf71 100644 --- a/src/server/main.py +++ b/src/server/main.py @@ -127,7 +127,6 @@ RESULTS_DIR, logger=logger, metrics_recorder=METRICS_RECORDER, - workflow_registry=WORKFLOW_REGISTRY, ) _pf_cfg = config.port_forward @@ -182,7 +181,6 @@ results_dir=RESULTS_DIR, log_stream_ttl_sec=config.log_stream.ttl_sec, server_base_url=config.identity.base_url, - workflow_registry=WORKFLOW_REGISTRY, ) LOG_ARCHIVER = TaskLogArchiver( diff --git a/src/server/registries/workflow.py b/src/server/registries/workflow.py index 593f45e60..200b613e2 100644 --- a/src/server/registries/workflow.py +++ b/src/server/registries/workflow.py @@ -13,13 +13,10 @@ model_serializer, ) -from shared.schemas.result import APIUsage - from ..clients.redis import ( WORKFLOWS_SET_KEY, RedisClient, task_state_key, - task_usage_key, workflow_cancelled_tasks_key, workflow_dispatched_tasks_key, workflow_failed_tasks_key, @@ -31,13 +28,6 @@ from ..utils.time import now_iso -class UnknownUsage: - """Sentinel for a model-calling task whose usage could not be mapped.""" - - -UNKNOWN_USAGE = UnknownUsage() - - class PersistedTask(BaseModel): """A durable per-task snapshot sufficient to rebuild scheduler state.""" @@ -113,10 +103,6 @@ class Workflow(BaseModel): completed_tasks: list[str] = Field(description="Completed task identifiers.") failed_tasks: list[str] = Field(description="Failed task identifiers.") cancelled_tasks: list[str] = Field(description="Cancelled task identifiers.") - usage: APIUsage | None = Field( - default=None, - description="Token/call usage summed over the workflow's finished tasks.", - ) def _create_workflow_record( @@ -189,7 +175,6 @@ def unregister_workflows(self, *workflow_ids: str) -> None: pipe.delete(*(workflow_sched_key(wid) for wid in workflow_ids)) for task_id in task_ids: pipe.delete(task_state_key(task_id)) - pipe.delete(task_usage_key(task_id)) pipe.execute() async def unregister_workflows_async(self, *workflow_ids: str) -> None: @@ -204,7 +189,6 @@ async def unregister_workflows_async(self, *workflow_ids: str) -> None: pipe.delete(*(workflow_sched_key(wid) for wid in workflow_ids)) for task_id in task_ids: pipe.delete(task_state_key(task_id)) - pipe.delete(task_usage_key(task_id)) await pipe.execute() def get_workflow_ids(self) -> set[str]: @@ -348,62 +332,6 @@ async def load_task_states_async( PersistedTask.model_validate_json(blob) if blob else None for blob in blobs ] - def save_task_usage( - self, task_id: str, usage: APIUsage | None | UnknownUsage - ) -> None: - """Synchronous variant of ``save_task_usage_async`` for the event loop.""" - key = task_usage_key(task_id) - if usage is None: - self._rds.sync.set_value(key, "null") - elif isinstance(usage, UnknownUsage): - self._rds.sync.set_value(key, "unknown") - else: - self._rds.sync.set_value(key, usage.model_dump_json()) - - async def save_task_usage_async( - self, task_id: str, usage: APIUsage | None | UnknownUsage - ) -> None: - """Persist one task's usage contribution. - - ``None`` marks a task type that makes no model calls (echo, lambda); - ``UNKNOWN_USAGE`` marks a model-calling task whose usage could not be - mapped. Both are written so a completed task with no entry at all is - detectable as one whose result was never ingested. - """ - key = task_usage_key(task_id) - if usage is None: - await self._rds.asyncio.set_value(key, "null") - elif isinstance(usage, UnknownUsage): - await self._rds.asyncio.set_value(key, "unknown") - else: - await self._rds.asyncio.set_value(key, usage.model_dump_json()) - - async def load_task_usages_async( - self, *task_ids: str - ) -> dict[str, APIUsage | None | UnknownUsage]: - """Load persisted usage for the given tasks. - - ``None`` means the task type makes no model calls; ``UNKNOWN_USAGE`` - means a model-calling task's usage could not be mapped; a task absent - from the result means its usage was never recorded. - """ - if not task_ids: - return {} - blobs = await self._rds.asyncio.mget( - [task_usage_key(task_id) for task_id in task_ids] - ) - result: dict[str, APIUsage | None | UnknownUsage] = {} - for task_id, blob in zip(task_ids, blobs, strict=True): - if blob is None: - continue - if blob == "null": - result[task_id] = None - elif blob == "unknown": - result[task_id] = UNKNOWN_USAGE - else: - result[task_id] = APIUsage.model_validate_json(blob) - return result - async def save_workflow_sched_async( self, workflow_id: str, in_epoch_order: bool, epoch_frontier: int ) -> None: diff --git a/src/server/routers/v1/results.py b/src/server/routers/v1/results.py index 0b246b4d4..3e01eaeb3 100644 --- a/src/server/routers/v1/results.py +++ b/src/server/routers/v1/results.py @@ -20,17 +20,7 @@ from shared.schemas.result import ( AnyExecutorResult, - APIResult, - APIUsage, - DataProfilingResult, - DataRetrievalResult, - EchoResult, - GenerationUsage, - InferenceResult, - PythonResult, ResultEnvelope, - ServeResult, - SSHResult, read_result, result_file_path, write_result, @@ -42,7 +32,6 @@ get_logger, get_results_dir, get_runtime, - get_workflow_registry, ) from ...auth.security import ( PrincipalContext, @@ -50,7 +39,6 @@ require_permission, ) from ...hooks import ResourceAction, ResourceKind -from ...registries.workflow import UNKNOWN_USAGE, UnknownUsage, WorkflowRegistry from ...schemas.common import PathResponse from ...services.monitoring import EventMonitor from ...task.models import TERMINAL_TASK_STATUSES @@ -64,78 +52,6 @@ router = APIRouter(prefix="/results", tags=["Results"]) -# Result types that make no model calls and therefore never carry usage. -_NO_USAGE_RESULT_TYPES = ( - EchoResult, - SSHResult, - DataProfilingResult, - DataRetrievalResult, - PythonResult, - ServeResult, -) - - -def _task_usage_from_envelope( - envelope: ResultEnvelope, -) -> APIUsage | None | UnknownUsage: - """Return a task's usage contribution from its result envelope. - - API tasks carry an ``APIUsage``; vLLM inference tasks map their - ``GenerationUsage`` token counts with reasoning 0, ``calls`` from - ``num_requests``, and ``wall_sec`` from ``latency_sec``. Task types that - make no model calls (echo, ssh, serve, data profiling, data retrieval, - python) contribute nothing (``None``). A model-calling task whose usage - cannot be mapped returns ``UNKNOWN_USAGE`` so the workflow fails closed. - """ - result = envelope.result - if isinstance(result, APIResult): - return result.usage if result.usage is not None else UNKNOWN_USAGE - if isinstance(result, InferenceResult): - if isinstance(result.usage, GenerationUsage): - return _inference_usage(result) - return UNKNOWN_USAGE - if isinstance(result, _NO_USAGE_RESULT_TYPES): - return None - # Any other result type either carries a usage field we do not map to - # APIUsage (embedding, agent, rag) or is an unknown model-calling type. - return UNKNOWN_USAGE - - -def _inference_usage(result: InferenceResult) -> APIUsage: - """Map an inference result's usage, subtracting merged children's shares. - - In a merged dispatch the parent's ``GenerationUsage`` is the whole batch - total, while each child carries its own share in ``result.children``. Each - child is also ingested separately, so the parent must record only its own - share (total minus the children's sum) or the workflow would count every - child's tokens and calls twice. ``wall_sec`` stays per task as reported. - """ - usage = result.usage - assert isinstance(usage, GenerationUsage) - prompt = usage.prompt_tokens - completion = usage.completion_tokens - calls = usage.num_requests - for child in result.children.values(): - if not isinstance(child, InferenceResult): - continue - child_usage = child.usage - if not isinstance(child_usage, GenerationUsage): - continue - prompt -= child_usage.prompt_tokens - completion -= child_usage.completion_tokens - calls -= child_usage.num_requests - return APIUsage( - prompt_tokens=prompt, - completion_tokens=completion, - reasoning_tokens=0, - calls=calls, - failures=0, - retries=0, - truncated_calls=0, - wall_sec=usage.latency_sec, - ) - - def _resolve_artifact_path(filename: str) -> Path: sanitized = Path(filename) if ( @@ -161,7 +77,6 @@ async def ingest_result( runtime: TaskRuntime = Depends(get_runtime), event_monitor: EventMonitor = Depends(get_event_monitor), results_dir: Path = Depends(get_results_dir), - registry: WorkflowRegistry = Depends(get_workflow_registry), logger: logging.Logger = Depends(get_logger), ) -> PathResponse: await require_permission( @@ -182,8 +97,6 @@ async def ingest_result( detail=f"Failed to store result: {exc}", ) from exc - await registry.save_task_usage_async(task_id, _task_usage_from_envelope(envelope)) - expected_artifacts: list[str] = [] record = runtime.get_record(task_id) if record: diff --git a/src/server/routers/v1/workflows.py b/src/server/routers/v1/workflows.py index 79e94731c..523fd37c3 100644 --- a/src/server/routers/v1/workflows.py +++ b/src/server/routers/v1/workflows.py @@ -15,7 +15,6 @@ from fastapi.responses import StreamingResponse from shared.schemas.event import TaskEvent -from shared.schemas.result import APIUsage from ...app_state import ( get_logger, @@ -37,7 +36,7 @@ workflow_log_stream_key, ) from ...hooks import SUBMISSION_GUARDS, ResourceAction, ResourceKind -from ...registries.workflow import UnknownUsage, Workflow, WorkflowRegistry +from ...registries.workflow import Workflow, WorkflowRegistry from ...schemas.logs import LogEntry, LogEvent, LogQueryResponse from ...schemas.workflow import ( WorkflowSubmitResponse, @@ -72,57 +71,6 @@ router = APIRouter(prefix="/workflows", tags=["Workflows"]) -def _sum_usage( - usages: dict[str, APIUsage | None | UnknownUsage], - completed_tasks: list[str], - logger: logging.Logger, -) -> APIUsage | None: - """Sum per-task usage into one workflow figure. - - Fail closed: a completed task whose usage was never recorded (its result - missing or unparsable) or could not be mapped (``UNKNOWN_USAGE``) makes the - whole workflow's usage null, with one log line naming the task. A task type - that makes no model calls (echo, lambda) maps to ``None`` and contributes - nothing, which is not a failure. When every completed task is recorded, the - (possibly all-zero) sum is returned. - """ - totals: dict[str, Any] = { - "prompt_tokens": 0, - "completion_tokens": 0, - "reasoning_tokens": 0, - "calls": 0, - "failures": 0, - "retries": 0, - "truncated_calls": 0, - "wall_sec": 0.0, - } - for task_id in completed_tasks: - if task_id not in usages: - logger.warning( - "workflow usage unavailable: task %s has no recorded usage", - task_id, - ) - return None - usage = usages[task_id] - if usage is None: - continue - if isinstance(usage, UnknownUsage): - logger.warning( - "workflow usage unavailable: task %s has unmappable usage", - task_id, - ) - return None - totals["prompt_tokens"] += usage.prompt_tokens - totals["completion_tokens"] += usage.completion_tokens - totals["reasoning_tokens"] += usage.reasoning_tokens - totals["calls"] += usage.calls - totals["failures"] += usage.failures - totals["retries"] += usage.retries - totals["truncated_calls"] += usage.truncated_calls - totals["wall_sec"] += usage.wall_sec - return APIUsage(**totals) - - def _parse_submission_body(raw_body: bytes, content_type: str) -> str: if "application/json" in content_type: try: @@ -336,8 +284,6 @@ async def get_workflow( status_code=status.HTTP_404_NOT_FOUND, detail=f"Workflow '{workflow_id}' not found", ) - usages = await registry.load_task_usages_async(*workflow.completed_tasks) - workflow.usage = _sum_usage(usages, workflow.completed_tasks, logger) return workflow diff --git a/src/server/services/monitoring.py b/src/server/services/monitoring.py index 60cf5a17b..dc72bb638 100644 --- a/src/server/services/monitoring.py +++ b/src/server/services/monitoring.py @@ -51,7 +51,6 @@ ) from ..registries.node import NodeRegistry from ..registries.worker import WorkerRegistry -from ..registries.workflow import WorkflowRegistry from ..schemas.logs import LogEvent from ..task.metadata import extract_model_dataset_names from ..task.models import TaskRecord, TaskStatus, TaskUsage @@ -100,7 +99,6 @@ def __init__( node_registry: NodeRegistry, metrics_recorder: MetricsRecorder, watchdog: WorkerWatchdog, - workflow_registry: WorkflowRegistry, ssh_proxy_enabled: bool = False, serve_proxy_enabled: bool = False, port_forward: PortForwardService | None = None, @@ -121,7 +119,6 @@ def __init__( self._serve_proxy_enabled = serve_proxy_enabled self._port_forward = port_forward self._results_dir = Path(results_dir) - self._workflow_registry = workflow_registry self._log_stream_ttl_sec = max(0, int(log_stream_ttl_sec)) self._server_base_url = self._validate_server_base_url(server_base_url) @@ -280,10 +277,6 @@ def mirror_task_results(self, parent_task_id: str, child_ids: list[str]) -> None if record: expected_artifacts = record.task.spec.get_artifacts() sync_manifest(dst_dir, child_id, expected_artifacts) - # A mirrored child has no result of its own (the executor - # produced no ``children`` entry for it), so the parent's total - # already contains its calls; record the no-usage marker. - self._workflow_registry.save_task_usage(child_id, None) except Exception as exc: self._logger.debug( "Failed to mirror results from %s to %s: %s", diff --git a/tests/server/dispatcher/helpers.py b/tests/server/dispatcher/helpers.py index 29954e660..7216b398d 100644 --- a/tests/server/dispatcher/helpers.py +++ b/tests/server/dispatcher/helpers.py @@ -56,7 +56,6 @@ def make_capturing_dispatcher( idle_ids: list[str] | None = None, satisfying_ids: list[str] | None = None, grace_sec: int = 60, - workflow_registry: Any = None, ) -> CapturingDispatcher: """Build a CapturingDispatcher whose registry returns the given worker ids.""" registry = mock.Mock() @@ -72,7 +71,4 @@ def make_capturing_dispatcher( results_dir=Path(tempfile.gettempdir()), logger=logging.getLogger("dispatcher-test"), no_worker_grace_sec=grace_sec, - workflow_registry=( - workflow_registry if workflow_registry is not None else mock.Mock() - ), ) diff --git a/tests/server/dispatcher/test_merged_child_redaction_dispatch.py b/tests/server/dispatcher/test_merged_child_redaction_dispatch.py index 4224f7ee0..d1de71137 100644 --- a/tests/server/dispatcher/test_merged_child_redaction_dispatch.py +++ b/tests/server/dispatcher/test_merged_child_redaction_dispatch.py @@ -94,7 +94,6 @@ def test_dispatch_fails_redacted_merged_child_and_carries_survivors() -> None: enable_context_reuse=False, enable_task_merge=True, task_merge_max_batch_size=4, - workflow_registry=mock.Mock(), ) assert disp.dispatch_once(parent) is True diff --git a/tests/server/dispatcher/test_python_input_dispatch.py b/tests/server/dispatcher/test_python_input_dispatch.py index 2b28a5455..c6d75d748 100644 --- a/tests/server/dispatcher/test_python_input_dispatch.py +++ b/tests/server/dispatcher/test_python_input_dispatch.py @@ -58,7 +58,6 @@ def test_unknown_input_stage_fails_the_task(tmp_path: Path) -> None: logger=logging.getLogger("dispatch-python-inputs"), worker_selection_strategy="first_fit", enable_context_reuse=False, - workflow_registry=mock.Mock(), ) assert disp.dispatch_once(nodes["score"]) is True diff --git a/tests/server/services/test_monitoring_serve_forward.py b/tests/server/services/test_monitoring_serve_forward.py index df174aaac..55b56865c 100644 --- a/tests/server/services/test_monitoring_serve_forward.py +++ b/tests/server/services/test_monitoring_serve_forward.py @@ -32,7 +32,6 @@ def _make_monitor( serve_proxy_enabled=serve_proxy_enabled, port_forward=port_forward, server_base_url=server_base_url, - workflow_registry=MagicMock(), ) diff --git a/tests/server/services/test_monitoring_ssh_bind.py b/tests/server/services/test_monitoring_ssh_bind.py index ab269947e..37041ebd3 100644 --- a/tests/server/services/test_monitoring_ssh_bind.py +++ b/tests/server/services/test_monitoring_ssh_bind.py @@ -30,7 +30,6 @@ def _make_monitor( serve_proxy_enabled=False, port_forward=port_forward, server_base_url="http://server.example.com:8000", - workflow_registry=MagicMock(), ) diff --git a/tests/server/services/test_monitoring_worker_events.py b/tests/server/services/test_monitoring_worker_events.py index ba252b6d2..45f9f441c 100644 --- a/tests/server/services/test_monitoring_worker_events.py +++ b/tests/server/services/test_monitoring_worker_events.py @@ -20,7 +20,6 @@ def _monitor(worker_registry: MagicMock) -> EventMonitor: node_registry=MagicMock(), metrics_recorder=MagicMock(), watchdog=MagicMock(), - workflow_registry=MagicMock(), ) diff --git a/tests/server/task/test_python_input_mounting.py b/tests/server/task/test_python_input_mounting.py index 9dab6b713..1c0a02120 100644 --- a/tests/server/task/test_python_input_mounting.py +++ b/tests/server/task/test_python_input_mounting.py @@ -3,7 +3,6 @@ import logging from pathlib import Path from typing import cast -from unittest import mock import pytest @@ -48,7 +47,6 @@ def _dispatcher( worker_registry=cast(WorkerRegistry, object()), results_dir=results_dir, logger=logging.getLogger("test-python-inputs"), - workflow_registry=mock.Mock(), ) diff --git a/tests/server/task/test_ssh_result_mounting.py b/tests/server/task/test_ssh_result_mounting.py index 2240dd2ea..1953be090 100644 --- a/tests/server/task/test_ssh_result_mounting.py +++ b/tests/server/task/test_ssh_result_mounting.py @@ -6,7 +6,6 @@ from pathlib import Path from types import SimpleNamespace from typing import cast -from unittest import mock import pytest @@ -110,7 +109,6 @@ def test_dispatcher_resolves_ssh_input_stage_names_from_local_stage_names() -> N worker_registry=cast(WorkerRegistry, object()), results_dir=Path("/tmp"), logger=logging.getLogger("test-ssh-phase2"), - workflow_registry=mock.Mock(), ) spec = SSHSpecStrict.model_validate(current.task.spec.model_dump()) @@ -151,7 +149,6 @@ def test_dispatcher_requeues_when_ssh_input_stage_not_done() -> None: worker_registry=cast(WorkerRegistry, object()), results_dir=Path("/tmp"), logger=logging.getLogger("test-ssh-phase2"), - workflow_registry=mock.Mock(), ) spec = SSHSpecStrict.model_validate(current.task.spec.model_dump()) @@ -222,7 +219,6 @@ def test_build_stage_context_includes_only_transitive_dependencies() -> None: worker_registry=cast(WorkerRegistry, object()), results_dir=Path("/tmp"), logger=logging.getLogger("test-stage-context"), - workflow_registry=mock.Mock(), ) context = dispatcher._build_stage_context(current) @@ -302,7 +298,6 @@ def test_collect_upstream_results_excludes_unrelated_completed_stages( worker_registry=cast(WorkerRegistry, object()), results_dir=tmp_path, logger=logging.getLogger("test-stage-results"), - workflow_registry=mock.Mock(), ) context = dispatcher._build_stage_context(current) @@ -377,7 +372,6 @@ def test_stage_reference_uses_payload_root_for_local_and_http_results( worker_registry=cast(WorkerRegistry, object()), results_dir=tmp_path, logger=logging.getLogger("test-stage-reference-root"), - workflow_registry=mock.Mock(), ) local_value = dispatcher._resolve_reference( diff --git a/tests/server/task/test_stage_reference_resolution.py b/tests/server/task/test_stage_reference_resolution.py index 21a2c4dd9..458582fc2 100644 --- a/tests/server/task/test_stage_reference_resolution.py +++ b/tests/server/task/test_stage_reference_resolution.py @@ -5,7 +5,6 @@ from pathlib import Path from types import SimpleNamespace from typing import cast -from unittest import mock from server.dispatcher.base import Dispatcher from server.registries.worker import WorkerRegistry @@ -101,7 +100,6 @@ def test_api_dependent_stage_resolves_first_row_text(tmp_path: Path) -> None: worker_registry=cast(WorkerRegistry, object()), results_dir=tmp_path, logger=logging.getLogger("test-api-dependent-stage"), - workflow_registry=mock.Mock(), ) value = dispatcher._resolve_reference("stage.items.0.text", {"stage": upstream}) @@ -205,7 +203,6 @@ def test_translated_n8n_dependent_api_stage_resolves(tmp_path: Path) -> None: worker_registry=cast(WorkerRegistry, object()), results_dir=tmp_path, logger=logging.getLogger("test-n8n-dependent-stage"), - workflow_registry=mock.Mock(), ) context = dispatcher._build_stage_context(downstream_record) diff --git a/tests/server/test_event_monitor_mirror.py b/tests/server/test_event_monitor_mirror.py index 74d67db10..0fc42f505 100644 --- a/tests/server/test_event_monitor_mirror.py +++ b/tests/server/test_event_monitor_mirror.py @@ -31,7 +31,6 @@ def _make_monitor(results_dir: Path) -> EventMonitor: metrics_recorder=MagicMock(), watchdog=MagicMock(), results_dir=results_dir, - workflow_registry=MagicMock(), ) diff --git a/tests/server/test_workflow_usage.py b/tests/server/test_workflow_usage.py deleted file mode 100644 index 04490d2e5..000000000 --- a/tests/server/test_workflow_usage.py +++ /dev/null @@ -1,461 +0,0 @@ -"""Tests for the workflow-level usage sum on GET /workflows/{id}. - -Usage is captured once at result ingest (``ingest_result``) and stored per task -in the workflow registry's Redis; ``GET /workflows/{id}`` sums the stored -values without opening any result file. -""" - -import logging -from collections.abc import Iterator -from pathlib import Path -from typing import Any, cast -from unittest import mock - -import fakeredis -import pytest -from fastapi import FastAPI -from httpx import ASGITransport, AsyncClient -from lumid_hooks import PrincipalContext, ResourceRef - -from server.app_state import get_logger, get_workflow_registry -from server.auth.security import authenticate_connection -from server.clients.redis import AsyncRedisClient, RedisClient, SyncRedisClient -from server.dispatcher.base import Dispatcher -from server.hooks import PERMISSION_CHECKERS -from server.registries.workflow import ( - UNKNOWN_USAGE, - Workflow, - WorkflowRegistry, - WorkflowStatus, -) -from server.routers.v1 import results as results_router -from server.routers.v1 import workflows as workflows_router -from server.task.runtime import TaskRuntime -from shared.schemas.result import ( - APIItem, - APIResult, - APIUsage, - DataProfilingResult, - DataRetrievalResult, - EchoResult, - GenerationUsage, - InferenceResult, - PythonResult, - ResultEnvelope, - ServeResult, - SSHResult, - write_result, -) -from shared.tasks import TaskEnvelopeStrict - - -@pytest.fixture -def server() -> fakeredis.FakeServer: - return fakeredis.FakeServer() - - -@pytest.fixture -def registry(server: fakeredis.FakeServer) -> WorkflowRegistry: - sync = SyncRedisClient.__new__(SyncRedisClient) - sync._control = fakeredis.FakeRedis(server=server, decode_responses=True) - async_client = AsyncRedisClient.__new__(AsyncRedisClient) - cast(Any, async_client)._control = fakeredis.FakeAsyncRedis( - server=server, decode_responses=True - ) - client = RedisClient.__new__(RedisClient) - client.sync = sync - client.asyncio = async_client - return WorkflowRegistry(client) - - -def _principal() -> PrincipalContext: - return PrincipalContext( - principal_id="p-1", - org_id="org", - external_id="ext", - principal_type="user", - scopes=[], - ) - - -class _AllowAllChecker: - name = "allow-all" - - async def require( - self, - principal: PrincipalContext, - resource: ResourceRef, - action: str, - logger: logging.Logger, - ) -> None: - return None - - async def accessible_ids( - self, - principal: PrincipalContext, - kind: str, - action: str, - logger: logging.Logger, - ) -> frozenset[str] | None: - return None - - -@pytest.fixture -def allow_all_permissions() -> Iterator[None]: - PERMISSION_CHECKERS.append(_AllowAllChecker()) - try: - yield - finally: - PERMISSION_CHECKERS.clear() - - -def _api_usage(**overrides: Any) -> APIUsage: - base: dict[str, Any] = dict( - prompt_tokens=30, - completion_tokens=12, - reasoning_tokens=2, - calls=3, - failures=0, - retries=1, - truncated_calls=1, - wall_sec=2.5, - ) - base.update(overrides) - return APIUsage(**base) - - -def _inference_usage() -> APIUsage: - result = InferenceResult( - ok=True, - model="m", - items=[], - usage=GenerationUsage( - prompt_tokens=100, - completion_tokens=50, - total_tokens=150, - num_requests=4, - latency_sec=1.0, - ), - ) - usage = results_router._task_usage_from_envelope( - ResultEnvelope(task_id="tsk-inf", result=result) - ) - assert isinstance(usage, APIUsage) - return usage - - -def _merged_parent() -> InferenceResult: - """A merged vLLM parent: batch total usage plus two children's shares.""" - child = lambda pt, ct, nr: InferenceResult( # noqa: E731 - ok=True, - model="m", - items=[], - usage=GenerationUsage( - prompt_tokens=pt, - completion_tokens=ct, - total_tokens=pt + ct, - num_requests=nr, - latency_sec=1.0, - ), - ) - return InferenceResult( - ok=True, - model="m", - items=[], - usage=GenerationUsage( - prompt_tokens=100, - completion_tokens=50, - total_tokens=150, - num_requests=4, - latency_sec=1.0, - ), - children={ - "tsk-child-1": child(30, 20, 2), - "tsk-child-2": child(10, 5, 1), - }, - ) - - -def _workflow(completed: list[str]) -> Workflow: - return Workflow( - workflow_id="wfl-1", - task_ids=completed, - submitted_at="2026-01-01T00:00:00Z", - updated_at="2026-01-01T00:00:00Z", - status=WorkflowStatus.DONE, - dispatched_tasks=[], - completed_tasks=completed, - failed_tasks=[], - cancelled_tasks=[], - ) - - -async def _sum(registry: WorkflowRegistry, completed: list[str]) -> APIUsage | None: - usages = await registry.load_task_usages_async(*completed) - return workflows_router._sum_usage(usages, completed, logging.getLogger("test")) - - -@pytest.mark.anyio -async def test_usage_sums_api_and_inference_tasks(registry: WorkflowRegistry) -> None: - """API and vLLM task usage sum exactly into one workflow figure.""" - await registry.save_task_usage_async("tsk-api", _api_usage()) - await registry.save_task_usage_async("tsk-inf", _inference_usage()) - - usage = await _sum(registry, ["tsk-api", "tsk-inf"]) - assert usage is not None - assert usage.prompt_tokens == 130 - assert usage.completion_tokens == 62 - assert usage.reasoning_tokens == 2 - assert usage.calls == 7 - assert usage.retries == 1 - assert usage.truncated_calls == 1 - assert usage.wall_sec == 3.5 - - -@pytest.mark.anyio -@pytest.mark.parametrize( - "result", - [ - EchoResult(ok=True), - SSHResult(ok=True, session_id="s", exit_code=0), - DataProfilingResult(ok=True), - DataRetrievalResult(ok=True), - PythonResult(ok=True, exit_code=0), - ServeResult(ok=True, model="m", port=8000), - ], -) -async def test_no_model_task_contributes_nothing( - registry: WorkflowRegistry, result: Any -) -> None: - """A task that calls no model contributes nothing and is not a failure.""" - assert ( - results_router._task_usage_from_envelope( - ResultEnvelope(task_id="tsk-no-model", result=result) - ) - is None - ) - await registry.save_task_usage_async("tsk-api", _api_usage()) - await registry.save_task_usage_async("tsk-no-model", None) - - usage = await _sum(registry, ["tsk-api", "tsk-no-model"]) - assert usage is not None - assert usage.prompt_tokens == 30 - assert usage.calls == 3 - - -@pytest.mark.anyio -@pytest.mark.parametrize("unmappable", [False, True]) -async def test_missing_or_unmappable_usage_makes_sum_null( - registry: WorkflowRegistry, unmappable: bool -) -> None: - """A completed model-calling task with missing or unmappable usage nulls the sum.""" - await registry.save_task_usage_async("tsk-api", _api_usage()) - if unmappable: - await registry.save_task_usage_async("tsk-inf", UNKNOWN_USAGE) - - assert await _sum(registry, ["tsk-api", "tsk-inf"]) is None - - -@pytest.mark.anyio -@pytest.mark.parametrize( - "order", - [ - ["tsk-parent", "tsk-child-1", "tsk-child-2"], - ["tsk-child-1", "tsk-child-2", "tsk-parent"], - ], -) -async def test_merged_parent_plus_children_sums_to_batch_total( - registry: WorkflowRegistry, order: list[str] -) -> None: - """A merged parent and its children sum to the batch total in either order.""" - parent = _merged_parent() - parent_usage = results_router._task_usage_from_envelope( - ResultEnvelope(task_id="tsk-parent", result=parent) - ) - assert isinstance(parent_usage, APIUsage) - await registry.save_task_usage_async("tsk-parent", parent_usage) - for child_id in ("tsk-child-1", "tsk-child-2"): - child_usage = results_router._task_usage_from_envelope( - ResultEnvelope(task_id=child_id, result=parent.children[child_id]) - ) - assert isinstance(child_usage, APIUsage) - await registry.save_task_usage_async(child_id, child_usage) - - usage = await _sum(registry, order) - assert usage is not None - assert usage.prompt_tokens == 100 - assert usage.completion_tokens == 50 - assert usage.calls == 4 - - -@pytest.mark.anyio -async def test_mirrored_child_records_no_usage( - registry: WorkflowRegistry, tmp_path: Any -) -> None: - """A mirrored child records the no-usage marker (it made no calls).""" - from server.services.monitoring import EventMonitor - - runtime = mock.Mock() - runtime.get_record.return_value = None - monitor = EventMonitor( - redis_client=mock.Mock(), - logger=logging.getLogger("test"), - runtime=runtime, - dispatcher=mock.Mock(), - worker_registry=mock.Mock(), - node_registry=mock.Mock(), - metrics_recorder=mock.Mock(), - watchdog=mock.Mock(), - results_dir=tmp_path, - workflow_registry=registry, - ) - parent_dir = tmp_path / "tsk-parent" - parent_dir.mkdir(parents=True) - (parent_dir / "results.json").write_text("{}", encoding="utf-8") - - monitor.mirror_task_results("tsk-parent", ["tsk-clone"]) - - usages = await registry.load_task_usages_async("tsk-clone") - assert usages["tsk-clone"] is None - - -@pytest.mark.anyio -async def test_unregister_deletes_usage_keys( - registry: WorkflowRegistry, server: fakeredis.FakeServer -) -> None: - """Unregistering a workflow deletes its tasks' usage keys (no leak).""" - rds = fakeredis.FakeRedis(server=server, decode_responses=True) - rds.hset("workflow:wfl-1", "workflow_id", "wfl-1") - rds.hset("workflow:wfl-1", "task_ids", '["tsk-1","tsk-2"]') - await registry.save_task_usage_async("tsk-1", _api_usage()) - await registry.save_task_usage_async("tsk-2", None) - - await registry.unregister_workflows_async("wfl-1") - - assert rds.exists("task:tsk-1:usage") == 0 - assert rds.exists("task:tsk-2:usage") == 0 - - -@pytest.mark.anyio -async def test_get_workflow_does_not_open_result_files( - registry: WorkflowRegistry, allow_all_permissions: None -) -> None: - """GET /workflows/{id} sums stored usage without reading any result file.""" - await registry.save_task_usage_async("tsk-api", _api_usage()) - await registry.save_task_usage_async("tsk-echo", None) - - app = FastAPI() - app.state.logger = logging.getLogger("test.workflow_usage") - app.include_router(workflows_router.router, prefix="/api/v1") - app.dependency_overrides[get_workflow_registry] = lambda: registry - app.dependency_overrides[get_logger] = lambda: logging.getLogger( - "test.workflow_usage" - ) - app.dependency_overrides[authenticate_connection] = lambda: _principal() - - with ( - mock.patch.object( - registry, - "get_workflow_async", - return_value=_workflow(["tsk-api", "tsk-echo"]), - ), - mock.patch("shared.schemas.result.read_result") as read_result, - ): - async with AsyncClient( - transport=ASGITransport(app=app), base_url="http://t" - ) as ac: - resp = await ac.get("/api/v1/workflows/wfl-1") - read_result.assert_not_called() - - assert resp.status_code == 200 - assert resp.json()["usage"]["prompt_tokens"] == 30 - - -_SKIP_WORKFLOW = """ -apiVersion: flowmesh/v1 -kind: Workflow -metadata: - name: skip-usage -spec: - graph: - nodes: - - name: judge - spec: - taskType: api - - name: refine - dependsOn: [judge] - spec: - taskType: api - condition: - node: judge - field: items.0.text - equals: insufficient -""" - - -@pytest.mark.anyio -async def test_condition_skipped_task_records_no_usage_and_keeps_sum( - registry: WorkflowRegistry, allow_all_permissions: None, tmp_path: Path -) -> None: - """A condition-skipped task records the no-usage marker, so the workflow - usage stays the sum of the tasks that made calls.""" - logger = logging.getLogger("test.workflow_usage") - runtime = TaskRuntime(registry, mock.Mock(), logger) - workflow_id, parsed = await runtime.register("owner", "org", _SKIP_WORKFLOW) - ids = {str(p.graph_node_name): p.task_id for p in parsed} - - write_result( - tmp_path, - ResultEnvelope( - task_id=ids["judge"], - result=APIResult( - ok=True, - executor="api", - method="POST", - url="https://api.example.com/v1/chat/completions", - status_code=200, - items=[ - APIItem( - index=0, - url="https://api.example.com/v1/chat/completions", - status_code=200, - text="sufficient", - ) - ], - ), - ), - ) - await registry.save_task_usage_async(ids["judge"], _api_usage(calls=7)) - runtime.mark_succeeded(ids["judge"], None, {}, "2026-01-01T00:00:00Z") - - dispatcher = Dispatcher( - runtime=runtime, - worker_registry=mock.Mock(), - results_dir=tmp_path, - logger=logger, - workflow_registry=registry, - ) - record = runtime.get_record(ids["refine"]) - assert record is not None - skipped = dispatcher._evaluate_condition_skip( - ids["refine"], TaskEnvelopeStrict.model_validate(record.task), record - ) - - assert skipped is True - usages = await registry.load_task_usages_async(ids["refine"]) - assert usages == {ids["refine"]: None} - - app = FastAPI() - app.state.logger = logger - app.include_router(workflows_router.router, prefix="/api/v1") - app.dependency_overrides[get_workflow_registry] = lambda: registry - app.dependency_overrides[get_logger] = lambda: logger - app.dependency_overrides[authenticate_connection] = lambda: _principal() - async with AsyncClient(transport=ASGITransport(app=app), base_url="http://t") as ac: - resp = await ac.get(f"/api/v1/workflows/{workflow_id}") - - assert resp.status_code == 200 - body = resp.json() - assert body["usage"] is not None - assert body["usage"]["calls"] == 7 - assert body["usage"]["prompt_tokens"] == 30