diff --git a/docs/WORKFLOWS.md b/docs/WORKFLOWS.md index 272ed2347..4cd074d20 100644 --- a/docs/WORKFLOWS.md +++ b/docs/WORKFLOWS.md @@ -71,10 +71,10 @@ contract. ## API task -`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. +`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`. @@ -113,24 +113,19 @@ 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. -```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 -``` +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 +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 diff --git a/pyproject.toml b/pyproject.toml index 5a6a21c14..90cd96e9a 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/sdk/src/flowmesh/models/__init__.py b/sdk/src/flowmesh/models/__init__.py index 6ccdd211d..c2a9ee65d 100644 --- a/sdk/src/flowmesh/models/__init__.py +++ b/sdk/src/flowmesh/models/__init__.py @@ -28,8 +28,10 @@ AgentResult, AgentUsage, AnyExecutorResult, + APIGroupItem, APIItem, APIResult, + APIUsage, BaseExecutorResult, CostEstimates, DataProfilingResult, @@ -106,8 +108,10 @@ ) __all__ = [ + "APIGroupItem", "APIItem", "APIResult", + "APIUsage", "ActiveWaitBreakdown", "AgentBatchSummary", "AgentItem", diff --git a/sdk/src/flowmesh/models/result/__init__.py b/sdk/src/flowmesh/models/result/__init__.py index 346d26092..6047335ee 100644 --- a/sdk/src/flowmesh/models/result/__init__.py +++ b/sdk/src/flowmesh/models/result/__init__.py @@ -40,7 +40,9 @@ AgentItem, AgentMetadata, AgentUsage, + APIGroupItem, APIItem, + APIUsage, CostEstimates, DataRetrievalItem, EchoItem, @@ -91,8 +93,10 @@ _model.model_rebuild() __all__ = [ + "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 974b22d3b..71d79129d 100644 --- a/sdk/src/flowmesh/models/result/catalog.py +++ b/sdk/src/flowmesh/models/result/catalog.py @@ -18,7 +18,9 @@ AgentItem, AgentMetadata, AgentUsage, + APIGroupItem, APIItem, + APIUsage, CostEstimates, DataRetrievalItem, EchoItem, @@ -211,9 +213,9 @@ 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] = Field(default_factory=list) + items: list[APIItem | APIGroupItem] = 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 cdab96fdc..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 @@ -174,3 +185,11 @@ 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.""" + + index: int + rows: list[APIItem] diff --git a/src/shared/schemas/result/__init__.py b/src/shared/schemas/result/__init__.py index c86e46ff5..a89bd8c36 100644 --- a/src/shared/schemas/result/__init__.py +++ b/src/shared/schemas/result/__init__.py @@ -40,7 +40,9 @@ AgentItem, AgentMetadata, AgentUsage, + APIGroupItem, APIItem, + APIUsage, CostEstimates, DataRetrievalItem, EchoItem, @@ -91,7 +93,9 @@ __all__ = [ "APIItem", + "APIGroupItem", "APIResult", + "APIUsage", "AgentBatchSummary", "AgentItem", "AgentMetadata", diff --git a/src/shared/schemas/result/catalog.py b/src/shared/schemas/result/catalog.py index 93532e2fc..4ba73b77b 100644 --- a/src/shared/schemas/result/catalog.py +++ b/src/shared/schemas/result/catalog.py @@ -20,7 +20,9 @@ AgentItem, AgentMetadata, AgentUsage, + APIGroupItem, APIItem, + APIUsage, CostEstimates, DataRetrievalItem, EchoItem, @@ -245,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 @@ -257,9 +259,9 @@ 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] = Field(default_factory=list) + items: list[APIItem | APIGroupItem] = Field(default_factory=list) class SSHResult(StrictExecutorResult): diff --git a/src/shared/schemas/result/payloads.py b/src/shared/schemas/result/payloads.py index 97a3c3eca..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.""" @@ -230,3 +243,11 @@ 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.""" + + index: int + rows: list[APIItem] diff --git a/src/worker/executors/api_executor.py b/src/worker/executors/api_executor.py index 0edf9a686..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 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 @@ -45,6 +45,70 @@ 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) + + +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 ( + _as_token_count(prompt_tokens), + _as_token_count(completion_tokens), + _as_token_count(reasoning_tokens), + finish_reason, + 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. @@ -157,13 +221,14 @@ def _request_with_retries( request_kwargs: dict[str, Any], retries: int, failed: threading.Event, - ) -> 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, a cancelled task, or a row already failed elsewhere stop the loop immediately. The final attempt's - failure propagates to the caller. + failure propagates to the caller. Returns the response and the number + of attempts used. """ attempt = 0 while True: @@ -181,7 +246,7 @@ def _request_with_retries( ) except httpx.RequestError as exc: if attempt >= retries: - raise + raise _RequestFailed(exc, attempt + 1) from None attempt += 1 delay = self._backoff_delay(attempt) logger.warning( @@ -195,7 +260,7 @@ def _request_with_retries( continue if resp.is_error and _is_retryable_status(resp.status_code): if attempt >= retries: - return resp + return resp, attempt + 1 attempt += 1 delay = self._backoff_delay(attempt, resp) logger.warning( @@ -207,7 +272,7 @@ def _request_with_retries( ) self._wait_for_backoff(delay, failed) continue - return resp + return resp, attempt + 1 def _backoff_delay(self, attempt: int, resp: httpx.Response | None = None) -> float: """Return the wait before the next attempt, honouring Retry-After.""" @@ -364,6 +429,42 @@ def _parse_response( return item + 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.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 " + "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 @@ -421,7 +522,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) @@ -429,6 +530,74 @@ def _run(self, task: ExecutorTask, out_dir: Path) -> APIResult: failed = threading.Event() first_error: list[BaseException] = [] + total = len(prompts) + done = 0 + failures = 0 + total_retries = 0 + latencies: list[float] = [] + 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] = {} + in_flight_lock = threading.Lock() + stats_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, truncated_calls + wall = time.monotonic() - start + with in_flight_lock: + in_flight.pop(idx, None) + ( + prompt_tokens, + completion_tokens, + reasoning_tokens, + finish_reason, + backend, + ) = _chat_completion_stats(body) + 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 finish_reason == "length": + truncated_calls += 1 + if backend is not None: + backend_counts[backend] = backend_counts.get(backend, 0) + 1 + 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", + 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") @@ -436,8 +605,11 @@ def _issue(idx: int, prompt: Any) -> APIItem: 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) + 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), @@ -447,16 +619,24 @@ def _issue(idx: int, prompt: Any) -> APIItem: retries, failed, ) - except TaskCancelledError: - raise - except httpx.RequestError as exc: + except _RequestFailed as exc: + _record_call( + idx, + exc.attempts, + exc.error.__class__.__name__, + start, + None, + failed=True, + ) error = ExecutionError( - f"API request failed (row {idx}): {exc}", retryable=True + f"API request failed (row {idx}): {exc.error}", retryable=True ) if not first_error: first_error.append(error) failed.set() - raise error from exc + raise error from exc.error + except TaskCancelledError: + raise except BaseException as exc: if not first_error: first_error.append(exc) @@ -469,6 +649,7 @@ def _issue(idx: int, prompt: Any) -> APIItem: if body_text: message = f"{message}: {body_text}" retryable = _is_retryable_status(resp.status_code) + _record_call(idx, attempts, resp.status_code, start, None, failed=True) error = ExecutionError(message, retryable=retryable) if not first_error: first_error.append(error) @@ -484,46 +665,172 @@ 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( + 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") - 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") - try: - results[idx] = future.result() - except TaskCancelledError: - raise - except BaseException: - pass - if first_error: - raise first_error[0] + 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 + 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_snapshot, + 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") + 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") + try: + results[idx] = future.result() + except TaskCancelledError: + raise + except BaseException: + pass + if first_error: + raise first_error[0] + finally: + heartbeat_stop.set() + 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) + 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, + done_snapshot, + failures_snapshot, + retries_snapshot, + wall, + latencies_snapshot, + prompt_snapshot, + completion_snapshot, + reasoning_snapshot, + 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] = [] + if entry.tables: + # One result item per table, sliced as in 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 = 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 + 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, + 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 235bafec9..ff163e6d7 100644 --- a/src/worker/executors/echo_executor.py +++ b/src/worker/executors/echo_executor.py @@ -44,7 +44,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 5679e55a5..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, @@ -429,7 +430,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( @@ -585,57 +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 ( - isinstance(value, list) - and value - and all(isinstance(v, list) for v in value) - ): - groups = value - elif isinstance(value, 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 = 1 - 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))]) @@ -680,7 +633,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 e51aad54d..b9a2a9a2d 100644 --- a/src/worker/executors/utils/graph_templates.py +++ b/src/worker/executors/utils/graph_templates.py @@ -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()) @@ -215,7 +217,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] @@ -269,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( @@ -290,7 +293,9 @@ 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)) @@ -300,7 +305,10 @@ def _is_expandable_group_value(value: Any) -> bool: for group_idx, row_count in enumerate(group_row_counts): for key, values in grouped_columns.items(): group_value = values[group_idx] - if _is_expandable_group_value(group_value): + 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)] @@ -309,9 +317,9 @@ def _is_expandable_group_value(value: Any) -> bool: "Grouped graph-template values must resolve to the same " "number of rows per group." ) + columns[key].extend(value_list) # type: ignore else: - value_list = [group_value for _ in range(row_count)] - columns[key].extend(value_list) # type: ignore + columns[key].extend([group_value for _ in range(row_count)]) # type: ignore num_rows = sum(group_row_counts) @@ -439,7 +447,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: @@ -502,6 +512,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 = { @@ -522,7 +533,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 @@ -586,71 +597,141 @@ 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 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 + 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 + mapped_items = False for token in parts[1:]: if not token: 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, grouped, mapped_items = _apply_attr( + value, attr, token, parts, grouped, mapped_items + ) 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, 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], + 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). 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, mapped_items + if isinstance(value, list): + if all(isinstance(v, dict) and attr in v for v in value): + 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], True, True + if all(isinstance(v, BaseModel) for v in value): + mapped = [_model_attr(v, attr, token) for v in value] + return ( + mapped, + grouped or _groups_on_first_access(mapped, mapped_items), + True, + ) + if all(isinstance(v, list) for v in value): + # 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_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, mapped_items + if isinstance(value, BaseModel): + 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 + 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 = value.model_extra + 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, 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, grouped)[0] for v in value], grouped + if isinstance(value, list) and -len(value) <= idx < len(value): + return value[idx], grouped + raise ExecutionError(f"{idx} not a valid index in {token} - {len(value)}") def _split_indexes(token: str) -> tuple[str, list[int]]: 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/tests/sdk/test_models.py b/tests/sdk/test_models.py index 720083efd..72b854cb0 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, @@ -525,3 +527,69 @@ 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" + + 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 diff --git a/tests/sdk/test_schema_compat.py b/tests/sdk/test_schema_compat.py index b30e557f0..1290f5849 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,8 @@ "RagQuery", "EchoItem", "APIItem", + "APIGroupItem", + "APIUsage", ] RESULT_MODEL_PAIRS = [ @@ -250,6 +254,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 # ------------------------------------------------------------------ # diff --git a/tests/shared/test_executor_result.py b/tests/shared/test_executor_result.py index 2df793b09..0e9efee65 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, @@ -107,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 @@ -216,3 +215,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" + ) diff --git a/tests/worker/test_api_executor.py b/tests/worker/test_api_executor.py index ab9a6d5a9..5795c2190 100644 --- a/tests/worker/test_api_executor.py +++ b/tests/worker/test_api_executor.py @@ -16,6 +16,13 @@ import pytest from pydantic import ValidationError +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 @@ -144,6 +151,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 @@ -176,6 +185,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 @@ -1076,7 +1092,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.""" @@ -1088,9 +1104,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.DEBUG, 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 @@ -1352,3 +1384,1250 @@ 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"}], + ] + + 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_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( + 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( + 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. + item = result.items[0] + assert isinstance(item, APIItem) + 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" + + 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 + + +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_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.""" + 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: + 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.DEBUG, 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.DEBUG, 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_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.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") + 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_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.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") + 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: + """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.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") + assert len(call_lines) == 1 + msg = call_lines[0] + 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.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") + 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", "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.DEBUG, 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=3" in msg + 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 new file mode 100644 index 000000000..ab511befc --- /dev/null +++ b/tests/worker/test_data_mixin_dataframe_grouping.py @@ -0,0 +1,301 @@ +"""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 + +import pandas as pd + +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"] + + +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"] + + +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_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"] diff --git a/tests/worker/test_echo_executor.py b/tests/worker/test_echo_executor.py new file mode 100644 index 000000000..6aab5d133 --- /dev/null +++ b/tests/worker/test_echo_executor.py @@ -0,0 +1,31 @@ +"""Echo executor tests: the literal "list" path.""" + +from pathlib import Path + +from shared.schemas.result import EchoResult +from shared.tasks import TaskType +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 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 new file mode 100644 index 000000000..98a5ad8ce --- /dev/null +++ b/tests/worker/test_graph_templates_expr.py @@ -0,0 +1,239 @@ +"""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, + BaseExecutorResult, + DataRetrievalItem, + DataRetrievalResult, + InferenceItem, + InferenceResult, +) +from worker.executors.utils.graph_templates import ( + _aggregate_structural_messages, + _build_grouped_dataframes, + _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 _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``.""" + upstream = APIResult( + ok=True, + executor="api", + method="POST", + url="https://up.example.com", + status_code=200, + items=[_item("c0"), _item("c1")], + ) + 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: + """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, 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: + """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) + + +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"]] diff --git a/uv.lock b/uv.lock index d7aa99faa..2b5b94448 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" },