diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml
index 0b40c22..4bf1858 100644
--- a/.github/workflows/ci.yml
+++ b/.github/workflows/ci.yml
@@ -5,9 +5,19 @@ on:
branches: [main]
pull_request:
+# Least privilege: the CI job only reads the repository.
+permissions:
+ contents: read
+
+# Cancel superseded runs for the same ref (e.g. force-pushes to a PR).
+concurrency:
+ group: ci-${{ github.workflow }}-${{ github.ref }}
+ cancel-in-progress: true
+
jobs:
ci:
runs-on: ubuntu-latest
+ timeout-minutes: 15
strategy:
fail-fast: false
matrix:
@@ -20,18 +30,20 @@ jobs:
uses: actions/setup-python@v5
with:
python-version: ${{ matrix.python-version }}
+ cache: pip
+ cache-dependency-path: pyproject.toml
- name: Install dependencies
run: pip install -e ".[dev]"
- name: Ruff check
- run: ruff check taskmaestro/ tests/
+ run: ruff check taskmaestro/ tests/ examples/
- name: Ruff format check
- run: ruff format --check taskmaestro/ tests/
+ run: ruff format --check taskmaestro/ tests/ examples/
- name: Mypy type check
run: mypy taskmaestro
- name: Run tests with coverage
- run: pytest --cov=taskmaestro --cov-report=term-missing
+ run: pytest --cov=taskmaestro --cov-report=term-missing --cov-fail-under=100
diff --git a/.github/workflows/publish.yml b/.github/workflows/publish.yml
index 29ba318..f743daf 100644
--- a/.github/workflows/publish.yml
+++ b/.github/workflows/publish.yml
@@ -4,10 +4,15 @@ on:
push:
tags: ["v*"]
+# Default to read-only; the publish job opts into id-token below.
+permissions:
+ contents: read
+
jobs:
build:
name: Build distributions
runs-on: ubuntu-latest
+ timeout-minutes: 10
steps:
- uses: actions/checkout@v4
@@ -15,6 +20,8 @@ jobs:
uses: actions/setup-python@v5
with:
python-version: "3.12"
+ cache: pip
+ cache-dependency-path: pyproject.toml
- name: Install build
run: pip install build
@@ -32,6 +39,7 @@ jobs:
name: Publish to PyPI
needs: build
runs-on: ubuntu-latest
+ timeout-minutes: 10
environment: pypi
permissions:
id-token: write
@@ -43,4 +51,6 @@ jobs:
path: dist/
- name: Publish via trusted publishing
- uses: pypa/gh-action-pypi-publish@release/v1
+ # Third-party action pinned to a full commit SHA (tag v1.14.2); the
+ # `release/v1` branch is mutable and could be repointed.
+ uses: pypa/gh-action-pypi-publish@dc37677b2e1c63e2034f94d8a5b11f265b73ba33 # v1.14.2
diff --git a/CLAUDE.md b/CLAUDE.md
index 69ecebd..0439165 100644
--- a/CLAUDE.md
+++ b/CLAUDE.md
@@ -33,7 +33,8 @@ mypy taskmaestro # type check (strict mode)
- **Type introspection**: Walk MRO via `__orig_bases__` + `typing.get_args()` to extract concrete `I`/`O` types
- **Fan-in**: Downstream task input model fields mapped to upstream outputs via `model_fields` (Pydantic v2)
- **Timeouts**: `signal.alarm` (Unix only, main thread); gracefully warns if unavailable
-- **Hook error swallowing**: `_emit()` wraps each hook call in try/except, reports via `warnings.warn()`
+- **Hook error swallowing**: `_emit()` wraps each hook call in try/except, reports via `warnings.warn(..., HookError, source=exc)` — message includes `repr(exc)`; `HookError` subclasses `UserWarning` so it can be filtered or escalated
+- **Inner-workflow failures**: `workflow_task` raises `WorkflowTaskError` (a `TaskExecutionError`) carrying `inner_job` and chaining the original exception via `__cause__`; `Job.exception` keeps the raw exception alongside `Job.error`
- **Validation order**: unique names → acyclic (DFS) → type chain → result task detection
## Testing Conventions
diff --git a/README.md b/README.md
index c0186e0..61a542b 100644
--- a/README.md
+++ b/README.md
@@ -110,7 +110,7 @@ You define **Tasks** (typed units of work), compose them into a **Workflow** (li
| Concept | Description |
|---|---|
| **Task** | Subclass `Task[I, O]` with Pydantic models for input and output, then implement `run(input, ctx)`. Each task can declare an optional `timeout_seconds`. For tasks with multiple named outputs, use inline `Inputs`/`Outputs` classes inside the task body. |
-| **Workflow** | Build a linear pipeline with `Workflow(tasks=[...])` or a DAG with `Workflow.builder()`. The builder accepts `depends_on` for single dependencies, fan-in dicts (`{"field": UpstreamTask}`), and `(Task, "field")` tuples for output field routing. Use `config_fields` to declare which input fields come from `JobConfiguration`. Workflows are validated at build time for cycles, type compatibility, and input completeness. |
+| **Workflow** | Build a linear pipeline with `Workflow(tasks=[...])` or a DAG with `Workflow.builder()`. The builder accepts `depends_on` for single dependencies, fan-in dicts (`{"field": UpstreamTask}`), `(Task, "field")` tuples for output field routing, `collect()` for gathering outputs into collection fields, and `mapped_over=TaskMap(...)` for sequential expansion over configured mappings. Use `config_fields` to declare which input fields come from `JobConfiguration`. Workflows are validated at build time for cycles, type compatibility, and input completeness. |
| **Job** | Binds a Workflow to a typed config (the root task's input). Tracks `status` (`pending` → `running` → `completed`/`failed`), the final `result`, any `error`, and per-task `task_results`. Optionally accepts a `JobConfiguration` for per-task static config values. A job can only be run once. |
| **Runner** | Executes tasks in topological order, stopping on the first failure (fail-fast). Supports per-task and per-job timeouts via `signal.alarm` (Unix only). Dispatches lifecycle events to registered hooks. |
| **ExecutionContext** | Passed to every `run()` call. Provides a `logger`, an auto-generated `correlation_id` (UUID), a `scratch_dir` (temporary directory), and a service registry (`register()`/`resolve()`) for injecting shared resources like DB connections. |
@@ -231,6 +231,163 @@ workflow = (
)
```
+## Collecting Multiple Outputs
+
+Use `collect()` when several task outputs should populate one `list[T]` or
+`dict[str, T]` field. Positional members preserve declaration order:
+
+```python
+from taskmaestro import collect
+
+class GridInput(BaseModel):
+ surfaces: list[Surface]
+
+workflow = (
+ Workflow.builder("create_grid")
+ .add_task(LoadSurface, name="top")
+ .add_task(GenerateSurface, name="middle")
+ .add_task(LoadSurface, name="base")
+ .add_task(
+ CreateGrid,
+ depends_on={"surfaces": collect("top", "middle", "base")},
+ )
+ .build()
+)
+```
+
+Use a mapping to preserve aliases in a `dict[str, T]`, and use `(task, "field")`
+to collect a specific output field:
+
+```python
+depends_on={
+ "surfaces": collect({
+ "top": ("top_loader", "surface"),
+ "base": ("base_loader", "surface"),
+ })
+}
+```
+
+The equivalent YAML forms are:
+
+```yaml
+depends_on:
+ surfaces:
+ collect:
+ - top
+ - [middle, generated_surface]
+ - base
+```
+
+```yaml
+depends_on:
+ surfaces:
+ collect:
+ top: [top_loader, surface]
+ base: [base_loader, surface]
+```
+
+Every member is checked against the field's element type when the workflow is
+built. Subtypes are accepted. `collect()` and `collect({})` explicitly create
+empty list and dictionary inputs, respectively.
+
+## Mapped Tasks
+
+A mapped task invokes one task declaration for every entry in a configured
+mapping. Mapped items execute sequentially in mapping declaration order.
+Each item gets a fresh task instance and child `ExecutionContext`.
+
+```python
+from taskmaestro import TaskMap
+
+workflow = (
+ Workflow.builder("create_grid")
+ .add_task(ConnectToResInsight)
+ .add_task(
+ LoadRegularSurface,
+ name="load_surfaces",
+ depends_on={"resinsight": ConnectToResInsight},
+ config_fields=["unit"],
+ mapped_over=TaskMap(
+ over="surfaces",
+ key_as="surface_name",
+ value_as="path",
+ error_mode="fail_fast",
+ ),
+ )
+ .add_task(
+ CreateGrid,
+ depends_on={"surfaces": ("load_surfaces", "root")},
+ )
+ .build()
+)
+```
+
+The mapped task's input model contains the injected key and value fields, not
+the source mapping:
+
+```python
+class LoadSurfaceInput(BaseModel):
+ resinsight: RipsInstance
+ unit: str
+ surface_name: str # key_as
+ path: str # value_as
+```
+
+Configure the source through `JobConfiguration`:
+
+```python
+job_configuration = JobConfiguration({
+ "load_surfaces": {
+ "unit": "meters",
+ "surfaces": {
+ "top": "/data/top.irap",
+ "base": "/data/base.irap",
+ },
+ },
+})
+```
+
+The logical output is a `MappedOutput[O]` Pydantic root model containing an
+insertion-ordered `dict[str, O]`, where `O` is the task's declared output type.
+Routing its `root` field lets a downstream input consume the dictionary:
+
+```python
+class CreateGridInput(BaseModel):
+ surfaces: dict[str, RegularSurface]
+```
+
+The equivalent YAML task declaration is:
+
+```yaml
+- task: resinsight.load_regular_surface
+ name: load_surfaces
+ map:
+ over: surfaces
+ key_as: surface_name
+ value_as: path
+ error_mode: fail_fast
+ depends_on:
+ resinsight: resinsight.connect
+ config_fields: [unit]
+```
+
+Input YAML:
+
+```yaml
+load_surfaces:
+ unit: meters
+ surfaces:
+ top: /data/top.irap
+ base: /data/base.irap
+```
+
+`fail_fast` stops at the first failed item. `collect_all` attempts every item
+and reports an aggregate `MappedTaskExecutionError`. An empty mapping succeeds
+with `MappedOutput(root={})`. Per-item records are available in
+`job.mapped_item_results`, and
+built-in logging, timing, and persistence hooks observe individual items.
+Concurrent mapped execution is intentionally deferred.
+
## ObjectModel
`ObjectModel[T]` wraps arbitrary (non-Pydantic) objects so they can flow through workflows. Use it as a type alias for simple wrappers, or subclass it to add extra fields:
@@ -343,6 +500,16 @@ result = run_workflow_from_yaml("workflow.yaml", "input.yaml")
YAML supports named task instances (`name:`), per-task input config (keyed by task name in the input file), fan-in dicts, and output field routing via `[task, field]` lists.
+How `input.yaml` is read is controlled by `workflow.input_mode`:
+
+| `input_mode` | Meaning |
+|---|---|
+| `auto` (default) | Per-task if every top-level key is a task name whose value is a mapping (or null); otherwise flat. If the file is *also* a valid input for the root task, loading fails and asks you to pick explicitly. |
+| `flat` | The whole mapping is the root task's input model. |
+| `per_task` | Top-level keys must be task names; unknown keys or non-mapping values are errors. |
+
+When the same task class (or the same inner YAML file) appears more than once under different `name:`s, `depends_on` and `result_task` must use the instance name — referencing the class path is rejected as ambiguous.
+
Use `workflow:` instead of `task:` to compose another YAML workflow. Paths are resolved
relative to the containing workflow file, and `workflow_input:` optionally supplies the
inner workflow's per-task configuration:
@@ -405,6 +572,7 @@ WorkflowRunnerError (base)
├── JobStateError # e.g., re-running a completed job
├── ConfigLoadError # YAML config loading failure
└── TaskExecutionError # Runtime task failure
+ ├── MappedTaskExecutionError # One or more mapped items failed
├── TaskOutputTypeError # Output type mismatch
└── TaskTimeoutError # Task exceeded timeout
```
diff --git a/pyproject.toml b/pyproject.toml
index 2e4a407..5e5dff1 100644
--- a/pyproject.toml
+++ b/pyproject.toml
@@ -19,6 +19,7 @@ classifiers = [
"Programming Language :: Python :: 3",
"Programming Language :: Python :: 3.12",
"Programming Language :: Python :: 3.13",
+ "Programming Language :: Python :: 3.14",
"Topic :: Software Development :: Libraries :: Python Modules",
"Typing :: Typed",
]
diff --git a/taskmaestro/__init__.py b/taskmaestro/__init__.py
index f613347..03e0253 100644
--- a/taskmaestro/__init__.py
+++ b/taskmaestro/__init__.py
@@ -3,6 +3,7 @@
__version__ = "0.2.0"
from taskmaestro.context import ExecutionContext
+from taskmaestro.dependencies import collect
from taskmaestro.discovery import (
TASK_ENTRY_POINT_GROUP,
WORKFLOW_ENTRY_POINT_GROUP,
@@ -18,16 +19,19 @@
CycleDetectedError,
IncompleteInputError,
JobStateError,
+ MappedTaskExecutionError,
PluginLoadError,
TaskExecutionError,
TaskOutputTypeError,
TaskTimeoutError,
WorkflowDefinitionError,
WorkflowRunnerError,
+ WorkflowTaskError,
)
from taskmaestro.job import EmptyConfig, Job, JobConfiguration, JobStatus, TaskResult, TaskStatus
+from taskmaestro.mapping import MappedOutput, TaskMap
from taskmaestro.object_model import ObjectModel
-from taskmaestro.runner import Runner
+from taskmaestro.runner import HookError, Runner
from taskmaestro.task import Task
from taskmaestro.visualization import to_mermaid
from taskmaestro.workflow import Workflow, WorkflowBuilder
@@ -45,17 +49,21 @@
"CycleDetectedError",
"EmptyConfig",
"ExecutionContext",
+ "HookError",
"IncompleteInputError",
"Job",
"JobConfiguration",
"JobStateError",
"JobStatus",
"LoadedWorkflow",
+ "MappedOutput",
+ "MappedTaskExecutionError",
"ObjectModel",
"PluginLoadError",
"Runner",
"Task",
"TaskExecutionError",
+ "TaskMap",
"TaskOutputTypeError",
"TaskResult",
"TaskStatus",
@@ -64,6 +72,8 @@
"WorkflowBuilder",
"WorkflowDefinitionError",
"WorkflowRunnerError",
+ "WorkflowTaskError",
+ "collect",
"get_registered_task",
"get_registered_workflow",
"load_workflow_from_yaml",
diff --git a/taskmaestro/context.py b/taskmaestro/context.py
index 2f5fcc6..80a9ff3 100644
--- a/taskmaestro/context.py
+++ b/taskmaestro/context.py
@@ -2,7 +2,9 @@
from __future__ import annotations
+import hashlib
import logging
+import re
import tempfile
import uuid
from pathlib import Path
@@ -22,8 +24,11 @@ def __init__(
correlation_id: str | None = None,
logger: logging.Logger | None = None,
scratch_dir: Path | None = None,
+ *,
+ parent_correlation_id: str | None = None,
) -> None:
self.correlation_id = correlation_id or str(uuid.uuid4())
+ self.parent_correlation_id = parent_correlation_id
self.logger = logger or logging.getLogger("taskmaestro")
self.scratch_dir = scratch_dir or Path(tempfile.gettempdir()) / self.correlation_id
self._registry: dict[str, Any] = {}
@@ -35,3 +40,18 @@ def register(self, key: str, service: Any) -> None:
def resolve(self, key: str) -> Any:
"""Retrieve a registered service. Raises KeyError if not found."""
return self._registry[key]
+
+ def child(self, *, task_name: str, item_key: str) -> ExecutionContext:
+ """Create a mapped-item context sharing this context's services."""
+ raw_suffix = f"{task_name}:{item_key}"
+ safe_suffix = re.sub(r"[^A-Za-z0-9_.-]+", "_", raw_suffix).strip("_") or "item"
+ digest = hashlib.sha256(raw_suffix.encode()).hexdigest()[:8]
+ child_id = f"{self.correlation_id}:{safe_suffix}:{digest}"
+ child = ExecutionContext(
+ correlation_id=child_id,
+ logger=self.logger,
+ scratch_dir=self.scratch_dir / f"{safe_suffix}-{digest}",
+ parent_correlation_id=self.correlation_id,
+ )
+ child._registry = self._registry
+ return child
diff --git a/taskmaestro/dependencies.py b/taskmaestro/dependencies.py
new file mode 100644
index 0000000..6f379d0
--- /dev/null
+++ b/taskmaestro/dependencies.py
@@ -0,0 +1,80 @@
+"""Dependency references used to collect multiple task outputs."""
+
+from __future__ import annotations
+
+from collections.abc import Mapping
+from dataclasses import dataclass
+from typing import Any, Literal, overload
+
+from taskmaestro.task import Task
+
+type TaskReference = type[Task[Any, Any]] | str
+type OutputReference = TaskReference | tuple[TaskReference, str]
+
+
+@dataclass(frozen=True)
+class CollectionDependency:
+ """Unresolved collection declared through :func:`collect`."""
+
+ kind: Literal["positional", "keyed"]
+ positional_members: tuple[OutputReference, ...] = ()
+ keyed_members: tuple[tuple[str, OutputReference], ...] = ()
+
+
+@dataclass(frozen=True)
+class OutputRef:
+ """A resolved reference to a task output or one of its fields."""
+
+ task_name: str
+ output_field: str | None = None
+
+
+@dataclass(frozen=True)
+class CollectionRef:
+ """A collection dependency whose task references have been resolved."""
+
+ kind: Literal["positional", "keyed"]
+ positional_members: tuple[OutputRef, ...] = ()
+ keyed_members: tuple[tuple[str, OutputRef], ...] = ()
+
+ def output_refs(self) -> tuple[OutputRef, ...]:
+ """Return all output references in declaration order."""
+ if self.kind == "positional":
+ return self.positional_members
+ return tuple(ref for _key, ref in self.keyed_members)
+
+
+@overload
+def collect() -> CollectionDependency: ...
+
+
+@overload
+def collect(*members: OutputReference) -> CollectionDependency: ...
+
+
+@overload
+def collect(members: Mapping[str, OutputReference], /) -> CollectionDependency: ...
+
+
+def collect(
+ *members: OutputReference | Mapping[str, OutputReference],
+) -> CollectionDependency:
+ """Collect several upstream outputs into one list or dictionary input field.
+
+ Positional members target ``list[T]`` fields. A single mapping argument
+ targets ``dict[str, T]`` fields. Members may be task classes, registered
+ task names, or ``(task, output_field)`` references.
+ """
+ if len(members) == 1 and isinstance(members[0], Mapping):
+ mapping = members[0]
+ if not all(isinstance(key, str) for key in mapping):
+ raise TypeError("collect() dictionary keys must be strings")
+ return CollectionDependency("keyed", keyed_members=tuple(mapping.items()))
+
+ if any(isinstance(member, Mapping) for member in members):
+ raise TypeError("collect() accepts either positional members or one mapping")
+
+ return CollectionDependency(
+ "positional",
+ positional_members=tuple(members), # type: ignore[arg-type]
+ )
diff --git a/taskmaestro/exceptions.py b/taskmaestro/exceptions.py
index c5c47b1..bf6006e 100644
--- a/taskmaestro/exceptions.py
+++ b/taskmaestro/exceptions.py
@@ -1,5 +1,7 @@
"""Exception hierarchy for the workflow runner library."""
+from typing import Any
+
class WorkflowRunnerError(Exception):
"""Base exception for all workflow runner errors."""
@@ -25,6 +27,16 @@ class TaskExecutionError(WorkflowRunnerError):
"""Raised during task execution."""
+class MappedTaskExecutionError(TaskExecutionError):
+ """One or more invocations of a mapped task failed."""
+
+ def __init__(self, task_name: str, errors: dict[str, Exception]) -> None:
+ self.task_name = task_name
+ self.errors = errors
+ details = "; ".join(f"{key}: {error}" for key, error in errors.items())
+ super().__init__(f"Mapped task '{task_name}' failed: {details}")
+
+
class TaskOutputTypeError(TaskExecutionError):
"""Task returned an output whose type doesn't match the declared output type."""
@@ -33,6 +45,24 @@ class TaskTimeoutError(TaskExecutionError):
"""Raised when a task exceeds its timeout_seconds."""
+class WorkflowTaskError(TaskExecutionError):
+ """An inner workflow wrapped by ``workflow_task`` failed.
+
+ Carries the completed inner :class:`~taskmaestro.job.Job` so callers can
+ inspect ``inner_job.task_results``, ``inner_job.failed_task`` and the
+ per-item results of mapped tasks. The original exception raised by the
+ failing inner task is attached as ``__cause__`` when it is available.
+ """
+
+ def __init__(self, workflow_name: str, inner_job: Any) -> None:
+ self.workflow_name = workflow_name
+ self.inner_job = inner_job
+ super().__init__(
+ f"Inner workflow '{workflow_name}' failed at task "
+ f"'{inner_job.failed_task}': {inner_job.error}"
+ )
+
+
class ConfigLoadError(WorkflowRunnerError):
"""Raised when YAML config loading fails (parse errors, import failures, validation)."""
diff --git a/taskmaestro/hooks/base.py b/taskmaestro/hooks/base.py
index 54d1312..7c854e6 100644
--- a/taskmaestro/hooks/base.py
+++ b/taskmaestro/hooks/base.py
@@ -21,6 +21,9 @@ class Event(StrEnum):
TASK_START = "task_start"
TASK_COMPLETE = "task_complete"
TASK_FAIL = "task_fail"
+ MAP_ITEM_START = "map_item_start"
+ MAP_ITEM_COMPLETE = "map_item_complete"
+ MAP_ITEM_FAIL = "map_item_fail"
@runtime_checkable
@@ -33,6 +36,13 @@ def on_job_fail(self, job: Job[Any]) -> None: ...
def on_task_start(self, job: Job[Any], task: Task[Any, Any]) -> None: ...
def on_task_complete(self, job: Job[Any], task: Task[Any, Any], output: BaseModel) -> None: ...
def on_task_fail(self, job: Job[Any], task: Task[Any, Any], error: Exception) -> None: ...
+ def on_map_item_start(self, job: Job[Any], task: Task[Any, Any], key: str) -> None: ...
+ def on_map_item_complete(
+ self, job: Job[Any], task: Task[Any, Any], key: str, output: BaseModel
+ ) -> None: ...
+ def on_map_item_fail(
+ self, job: Job[Any], task: Task[Any, Any], key: str, error: Exception
+ ) -> None: ...
class BaseHook:
@@ -55,3 +65,16 @@ def on_task_complete(self, job: Job[Any], task: Task[Any, Any], output: BaseMode
def on_task_fail(self, job: Job[Any], task: Task[Any, Any], error: Exception) -> None:
pass
+
+ def on_map_item_start(self, job: Job[Any], task: Task[Any, Any], key: str) -> None:
+ pass
+
+ def on_map_item_complete(
+ self, job: Job[Any], task: Task[Any, Any], key: str, output: BaseModel
+ ) -> None:
+ pass
+
+ def on_map_item_fail(
+ self, job: Job[Any], task: Task[Any, Any], key: str, error: Exception
+ ) -> None:
+ pass
diff --git a/taskmaestro/hooks/logging.py b/taskmaestro/hooks/logging.py
index c10052b..0a51253 100644
--- a/taskmaestro/hooks/logging.py
+++ b/taskmaestro/hooks/logging.py
@@ -42,3 +42,16 @@ def on_task_complete(self, job: Job[Any], task: Task[Any, Any], output: BaseMode
def on_task_fail(self, job: Job[Any], task: Task[Any, Any], error: Exception) -> None:
self._logger.log(self._level, "Task failed: %s, error=%s", task.name, error)
+
+ def on_map_item_start(self, job: Job[Any], task: Task[Any, Any], key: str) -> None:
+ self._logger.log(self._level, "Map item started: %s[%s]", task.name, key)
+
+ def on_map_item_complete(
+ self, job: Job[Any], task: Task[Any, Any], key: str, output: BaseModel
+ ) -> None:
+ self._logger.log(self._level, "Map item completed: %s[%s]", task.name, key)
+
+ def on_map_item_fail(
+ self, job: Job[Any], task: Task[Any, Any], key: str, error: Exception
+ ) -> None:
+ self._logger.log(self._level, "Map item failed: %s[%s], error=%s", task.name, key, error)
diff --git a/taskmaestro/hooks/persistence.py b/taskmaestro/hooks/persistence.py
index 6ee8006..f73b326 100644
--- a/taskmaestro/hooks/persistence.py
+++ b/taskmaestro/hooks/persistence.py
@@ -4,6 +4,7 @@
from pathlib import Path
from typing import Any
+from urllib.parse import quote
from pydantic import BaseModel
@@ -13,12 +14,31 @@
class ResultPersistenceHook(BaseHook):
- """Writes {task_name}.json per completed task to an output directory."""
+ """Writes {task_name}.json per completed task to an output directory.
+
+ Task names and mapped-item keys are percent-encoded so that a name such
+ as ``../evil`` or ``a/b`` can never escape ``output_dir``.
+ """
def __init__(self, output_dir: Path) -> None:
self.output_dir = output_dir
def on_task_complete(self, job: Job[Any], task: Task[Any, Any], output: BaseModel) -> None:
+ self._write(f"{_safe(task.name)}.json", output)
+
+ def on_map_item_complete(
+ self, job: Job[Any], task: Task[Any, Any], key: str, output: BaseModel
+ ) -> None:
+ self._write(f"{_safe(task.name)}[{_safe(key)}].json", output)
+
+ def _write(self, filename: str, output: BaseModel) -> None:
self.output_dir.mkdir(parents=True, exist_ok=True)
- output_path = self.output_dir / f"{task.name}.json"
- output_path.write_text(output.model_dump_json(indent=2))
+ (self.output_dir / filename).write_text(output.model_dump_json(indent=2))
+
+
+def _safe(component: str) -> str:
+ """Return a single, traversal-free filename component.
+
+ Escapes ``%`` too, so distinct inputs cannot collapse onto the same name.
+ """
+ return quote(component, safe="")
diff --git a/taskmaestro/hooks/timing.py b/taskmaestro/hooks/timing.py
index d0acb27..8b426b2 100644
--- a/taskmaestro/hooks/timing.py
+++ b/taskmaestro/hooks/timing.py
@@ -19,7 +19,9 @@ def __init__(self) -> None:
self.job_duration: float | None = None
self.task_timings: dict[str, float] = {}
self._job_start: float | None = None
+ self.mapped_item_timings: dict[str, dict[str, float]] = {}
self._task_starts: dict[str, float] = {}
+ self._map_item_starts: dict[tuple[str, str], float] = {}
def on_job_start(self, job: Job[Any]) -> None:
self._job_start = time.monotonic()
@@ -44,3 +46,21 @@ def on_task_fail(self, job: Job[Any], task: Task[Any, Any], error: Exception) ->
start = self._task_starts.get(task.name)
if start is not None:
self.task_timings[task.name] = time.monotonic() - start
+
+ def on_map_item_start(self, job: Job[Any], task: Task[Any, Any], key: str) -> None:
+ self._map_item_starts[(task.name, key)] = time.monotonic()
+
+ def on_map_item_complete(
+ self, job: Job[Any], task: Task[Any, Any], key: str, output: BaseModel
+ ) -> None:
+ self._record_map_item(task.name, key)
+
+ def on_map_item_fail(
+ self, job: Job[Any], task: Task[Any, Any], key: str, error: Exception
+ ) -> None:
+ self._record_map_item(task.name, key)
+
+ def _record_map_item(self, task_name: str, key: str) -> None:
+ start = self._map_item_starts.get((task_name, key))
+ if start is not None:
+ self.mapped_item_timings.setdefault(task_name, {})[key] = time.monotonic() - start
diff --git a/taskmaestro/job.py b/taskmaestro/job.py
index 549a3c4..4a13f29 100644
--- a/taskmaestro/job.py
+++ b/taskmaestro/job.py
@@ -2,12 +2,13 @@
from __future__ import annotations
+from collections.abc import Mapping
from dataclasses import dataclass
from datetime import datetime
from enum import StrEnum
from typing import Any, Generic, TypeVar
-from pydantic import BaseModel
+from pydantic import BaseModel, TypeAdapter, ValidationError
from taskmaestro.exceptions import WorkflowDefinitionError
from taskmaestro.task import get_input_type
@@ -92,20 +93,23 @@ def __init__(
self.status: JobStatus = JobStatus.PENDING
self.result: BaseModel | None = None
self.error: str | None = None
+ self.exception: Exception | None = None
self.failed_task: str | None = None
self.started_at: datetime | None = None
self.completed_at: datetime | None = None
self.task_results: list[TaskResult] = []
+ self.mapped_item_results: dict[str, list[TaskResult]] = {}
self._validate_root_task_inputs(config)
+ self._validate_task_maps()
def _validate_root_task_inputs(self, config: C) -> None:
"""Validate that config type matches the input type of all root tasks."""
for task_name, deps in self.workflow._dependencies.items():
if deps is None:
- # Skip validation for root tasks that have config_fields
+ # Configured and mapped roots do not consume job.config directly.
config_fields = self.workflow.get_config_fields(task_name)
- if config_fields:
+ if config_fields or self.workflow.is_mapped_task(task_name):
continue
task_cls = self.workflow._tasks[task_name]
expected_input = get_input_type(task_cls)
@@ -114,3 +118,46 @@ def _validate_root_task_inputs(self, config: C) -> None:
f"Root task '{task_name}' expects input type "
f"{expected_input.__name__} but got {type(config).__name__}"
)
+
+ def _validate_task_maps(self) -> None:
+ """Validate configured map sources and their key/value types."""
+ for task_name, task_cls in self.workflow._tasks.items():
+ task_map = self.workflow.get_task_map(task_name)
+ if task_map is None:
+ continue
+ if self.job_configuration is None:
+ raise WorkflowDefinitionError(
+ f"Mapped task '{task_name}' requires JobConfiguration"
+ )
+ task_config = self.job_configuration.get_config_for_task(task_name)
+ if task_map.over not in task_config:
+ raise WorkflowDefinitionError(
+ f"Mapped task '{task_name}' requires configuration field '{task_map.over}'"
+ )
+ source = task_config[task_map.over]
+ if not isinstance(source, Mapping):
+ raise WorkflowDefinitionError(
+ f"Configuration field '{task_name}.{task_map.over}' must be a mapping"
+ )
+
+ input_type = get_input_type(task_cls)
+ input_fields = input_type.model_fields
+ # A tuple adapter carries the input model's config while allowing
+ # nested BaseModels to retain their own config. Keep field metadata too.
+ key_type = input_fields[task_map.key_as].rebuild_annotation()
+ value_type = input_fields[task_map.value_as].rebuild_annotation()
+ item_adapter: TypeAdapter[Any] = TypeAdapter(
+ tuple[key_type, value_type], # type: ignore[valid-type]
+ config=input_type.model_config,
+ )
+ for key, value in source.items():
+ if not isinstance(key, str):
+ raise WorkflowDefinitionError(
+ f"Mapping keys for task '{task_name}' must be strings"
+ )
+ try:
+ item_adapter.validate_python((key, value))
+ except ValidationError as exc:
+ raise WorkflowDefinitionError(
+ f"Invalid mapping item '{key}' for task '{task_name}': {exc}"
+ ) from exc
diff --git a/taskmaestro/mapping.py b/taskmaestro/mapping.py
new file mode 100644
index 0000000..851dc86
--- /dev/null
+++ b/taskmaestro/mapping.py
@@ -0,0 +1,37 @@
+"""Configuration for expanding one task over a configured mapping."""
+
+from __future__ import annotations
+
+from dataclasses import dataclass
+from typing import Generic, Literal, TypeVar
+
+from pydantic import BaseModel, RootModel
+
+O = TypeVar("O", bound=BaseModel)
+
+
+class MappedOutput(RootModel[dict[str, O]], Generic[O]):
+ """Typed aggregate output produced by a mapped workflow task."""
+
+
+@dataclass(frozen=True)
+class TaskMap:
+ """Describe how mapping keys and values populate task input fields."""
+
+ over: str
+ key_as: str
+ value_as: str
+ error_mode: Literal["fail_fast", "collect_all"] = "fail_fast"
+
+ def __post_init__(self) -> None:
+ for field_name, value in (
+ ("over", self.over),
+ ("key_as", self.key_as),
+ ("value_as", self.value_as),
+ ):
+ if not value:
+ raise ValueError(f"TaskMap.{field_name} must be a non-empty string")
+ if self.key_as == self.value_as:
+ raise ValueError("TaskMap.key_as and TaskMap.value_as must be different")
+ if self.error_mode not in ("fail_fast", "collect_all"):
+ raise ValueError("TaskMap.error_mode must be 'fail_fast' or 'collect_all'")
diff --git a/taskmaestro/runner.py b/taskmaestro/runner.py
index 65e7012..34ea08d 100644
--- a/taskmaestro/runner.py
+++ b/taskmaestro/runner.py
@@ -3,21 +3,70 @@
from __future__ import annotations
import signal
+import time
import warnings
+from collections.abc import Mapping
+from contextlib import suppress
+from dataclasses import dataclass, field
from datetime import datetime
from typing import Any
from pydantic import BaseModel
from taskmaestro.context import ExecutionContext
+from taskmaestro.dependencies import CollectionRef, OutputRef
from taskmaestro.exceptions import (
JobStateError,
+ MappedTaskExecutionError,
TaskOutputTypeError,
TaskTimeoutError,
)
from taskmaestro.hooks.base import BaseHook, Event
from taskmaestro.job import Job, JobStatus, TaskResult, TaskStatus
-from taskmaestro.task import get_input_type, get_output_type
+from taskmaestro.mapping import MappedOutput, TaskMap
+from taskmaestro.task import Task, get_input_type, get_output_type
+
+
+class _JobTimeoutError(TaskTimeoutError):
+ """A job deadline must abort even when mapped items collect failures."""
+
+
+class HookError(UserWarning):
+ """Warning category used when a lifecycle hook raises.
+
+ Subclasses :class:`UserWarning` so existing ``pytest.warns(UserWarning)``
+ and ``-W error::UserWarning`` configurations keep working, while allowing
+ callers to filter hook failures specifically.
+ """
+
+
+@dataclass
+class _Deadline:
+ """Per-run timer state shared by the job and its tasks.
+
+ There is only one ``SIGALRM`` per process, so the job deadline is kept as an
+ absolute ``time.monotonic()`` timestamp and folded into every task or item
+ alarm. Whichever deadline is nearer wins, and the job deadline is
+ re-checked before each unit of work so an inner alarm can never cancel it.
+ """
+
+ job_timeout: float | None = None
+ job_deadline: float | None = None
+ warned: bool = False
+ previous_handler: Any = field(default=None, repr=False)
+ handler_installed: bool = False
+
+ def remaining(self) -> float | None:
+ """Seconds left until the job deadline, or ``None`` if there is none."""
+ if self.job_deadline is None:
+ return None
+ return self.job_deadline - time.monotonic()
+
+ def check(self) -> None:
+ """Raise :class:`_JobTimeoutError` if the job deadline has passed."""
+ remaining = self.remaining()
+ if remaining is not None and remaining <= 0:
+ raise _JobTimeoutError(f"Job timed out after {self.job_timeout}s")
class Runner:
@@ -47,10 +96,11 @@ def run(
job.status = JobStatus.RUNNING
job.started_at = datetime.now()
- # Set up job-level timeout
- job_alarm_set = False
+ # Job-level timeout is tracked as an absolute deadline and folded into
+ # every task/item alarm; see _Deadline.
+ deadline = _Deadline(job_timeout=timeout_seconds)
if timeout_seconds is not None:
- job_alarm_set = self._set_alarm(timeout_seconds, "Job")
+ deadline.job_deadline = time.monotonic() + timeout_seconds
outputs: dict[str, BaseModel] = {}
job_config = job.job_configuration
@@ -61,14 +111,25 @@ def run(
task.name = task_name # instance-level override for named instances
deps = workflow.get_dependencies(task_name)
config_fields = workflow.get_config_fields(task_name)
- config_values = (
+ task_map = workflow.get_task_map(task_name)
+ all_config_values = (
job_config.get_config_for_task(task_name)
- if job_config and config_fields
+ if job_config and (config_fields or task_map is not None)
else {}
)
+ # Mapped tasks consume the map source themselves; every other
+ # configured value is passed through to the input model as before.
+ config_values = {
+ key: value
+ for key, value in all_config_values.items()
+ if task_map is None or key != task_map.over
+ }
- # Assemble input based on dependency type
- if deps is None:
+ # Assemble input based on dependency type. Mapped tasks build
+ # one validated input per configured item below.
+ if task_map is not None:
+ task_input: Any = None
+ elif deps is None:
if config_values:
# Root task with config: build input from config values
input_type = get_input_type(task_cls)
@@ -79,7 +140,9 @@ def run(
if config_values:
# Single dep with config: decompose upstream, merge with config
input_type = get_input_type(task_cls)
- upstream_data = outputs[deps].model_dump()
+ upstream_output = outputs[deps]
+ assert isinstance(upstream_output, BaseModel)
+ upstream_data = upstream_output.model_dump()
down_fields = input_type.model_fields
merged: dict[str, object] = {
k: v for k, v in upstream_data.items() if k in down_fields
@@ -95,7 +158,9 @@ def run(
input_type = get_input_type(task_cls)
field_values: dict[str, object] = {}
for fname, upstream_ref in deps.items():
- if isinstance(upstream_ref, tuple):
+ if isinstance(upstream_ref, CollectionRef):
+ field_values[fname] = self._resolve_collection(upstream_ref, outputs)
+ elif isinstance(upstream_ref, tuple):
up_name, up_field = upstream_ref
field_values[fname] = getattr(outputs[up_name], up_field)
else:
@@ -109,21 +174,35 @@ def run(
task_started = datetime.now()
self._emit(Event.TASK_START, job, task)
- # Set up per-task timeout
- task_alarm_set = False
- if task.timeout_seconds is not None:
- task_alarm_set = self._set_alarm(task.timeout_seconds, task.name)
-
try:
- output = task.run(task_input, ctx)
-
- # Validate output matches declared type
- expected_output_type = get_output_type(task_cls)
- if not isinstance(output, expected_output_type):
- raise TaskOutputTypeError(
- f"Task '{task.name}' returned {type(output).__name__}, "
- f"expected {expected_output_type.__name__}"
+ # Arming happens inside the guarded block so that an expired
+ # job deadline or an unusable timer is recorded as a task
+ # failure rather than escaping with the job left RUNNING.
+ deadline.check()
+ if task_map is not None:
+ output = self._run_mapped_task(
+ job,
+ task_cls,
+ task,
+ task_map,
+ deps,
+ config_values,
+ all_config_values,
+ outputs,
+ ctx,
+ deadline,
)
+ else:
+ self._arm(task.timeout_seconds, task.name, deadline)
+ output = task.run(task_input, ctx)
+
+ # Validate output matches declared type
+ expected_output_type = get_output_type(task_cls)
+ if not isinstance(output, expected_output_type):
+ raise TaskOutputTypeError(
+ f"Task '{task.name}' returned {type(output).__name__}, "
+ f"expected {expected_output_type.__name__}"
+ )
duration = (datetime.now() - task_started).total_seconds()
outputs[task.name] = output
@@ -141,6 +220,7 @@ def run(
duration = (datetime.now() - task_started).total_seconds()
job.status = JobStatus.FAILED
job.error = str(exc)
+ job.exception = exc
job.failed_task = task.name
job.completed_at = datetime.now()
job.task_results.append(
@@ -157,11 +237,10 @@ def run(
self._emit(Event.JOB_FAIL, job)
return job
finally:
- if task_alarm_set:
- signal.alarm(0)
+ self._disarm(deadline)
finally:
- if job_alarm_set:
- signal.alarm(0)
+ self._disarm(deadline)
+ self._restore_handler(deadline)
job.status = JobStatus.COMPLETED
job.result = outputs[workflow.result_task_name]
@@ -169,33 +248,229 @@ def run(
self._emit(Event.JOB_COMPLETE, job)
return job
- def _set_alarm(self, seconds: float, label: str) -> bool:
- """Set a signal.alarm for timeout. Returns True if alarm was set."""
- try:
+ def _run_mapped_task(
+ self,
+ job: Job[Any],
+ task_cls: type[Task[Any, Any]],
+ parent_task: Task[Any, Any],
+ task_map: TaskMap,
+ deps: Any,
+ config_values: dict[str, Any],
+ all_config_values: dict[str, Any],
+ outputs: dict[str, BaseModel],
+ ctx: ExecutionContext,
+ deadline: _Deadline,
+ ) -> BaseModel:
+ """Run all configured items for one mapped workflow node."""
+ source = all_config_values[task_map.over]
+ assert isinstance(source, Mapping) # validated when the Job was created
+ shared_values = self._mapped_shared_values(deps, config_values, outputs)
+ expected_output_type = get_output_type(task_cls)
+ collected: dict[str, BaseModel] = {}
+ errors: dict[str, Exception] = {}
+ item_results = job.mapped_item_results.setdefault(parent_task.name, [])
+
+ for key, value in source.items():
+ assert isinstance(key, str) # validated when the Job was created
+ item_task = task_cls()
+ item_task.name = parent_task.name
+ item_input_values = dict(shared_values)
+ item_input_values[task_map.key_as] = key
+ item_input_values[task_map.value_as] = value
+ item_ctx = ctx.child(task_name=parent_task.name, item_key=key)
+ item_started = datetime.now()
+ self._emit(Event.MAP_ITEM_START, job, item_task, key)
+ try:
+ deadline.check()
+ self._arm(item_task.timeout_seconds, f"{parent_task.name}[{key}]", deadline)
+ input_type = get_input_type(task_cls)
+ item_input = input_type.model_validate(item_input_values)
+ output = item_task.run(item_input, item_ctx)
+ if not isinstance(output, expected_output_type):
+ raise TaskOutputTypeError(
+ f"Task '{parent_task.name}[{key}]' returned "
+ f"{type(output).__name__}, expected {expected_output_type.__name__}"
+ )
+ collected[key] = output
+ item_results.append(
+ TaskResult(
+ task_name=f"{parent_task.name}[{key}]",
+ status=TaskStatus.COMPLETED,
+ output=output,
+ started_at=item_started,
+ duration_seconds=(datetime.now() - item_started).total_seconds(),
+ )
+ )
+ self._emit(Event.MAP_ITEM_COMPLETE, job, item_task, key, output)
+ except Exception as exc:
+ errors[key] = exc
+ item_results.append(
+ TaskResult(
+ task_name=f"{parent_task.name}[{key}]",
+ status=TaskStatus.FAILED,
+ output=None,
+ started_at=item_started,
+ duration_seconds=(datetime.now() - item_started).total_seconds(),
+ error=str(exc),
+ )
+ )
+ self._emit(Event.MAP_ITEM_FAIL, job, item_task, key, exc)
+ if isinstance(exc, _JobTimeoutError):
+ raise
+ if task_map.error_mode == "fail_fast":
+ raise MappedTaskExecutionError(parent_task.name, errors) from exc
+ finally:
+ self._disarm(deadline)
+
+ if errors:
+ raise MappedTaskExecutionError(parent_task.name, errors)
+ mapped_output_type = MappedOutput[expected_output_type] # type: ignore[valid-type]
+ return mapped_output_type(root=collected)
+
+ def _mapped_shared_values(
+ self,
+ deps: Any,
+ config_values: dict[str, Any],
+ outputs: dict[str, BaseModel],
+ ) -> dict[str, object]:
+ """Resolve fields shared by every invocation of a mapped task."""
+ values: dict[str, object] = {}
+ if isinstance(deps, dict):
+ for field_name, ref in deps.items():
+ if isinstance(ref, CollectionRef):
+ values[field_name] = self._resolve_collection(ref, outputs)
+ elif isinstance(ref, tuple):
+ upstream_name, output_field = ref
+ values[field_name] = getattr(outputs[upstream_name], output_field)
+ else:
+ values[field_name] = outputs[ref]
+ values.update(config_values)
+ return values
+
+ @staticmethod
+ def _resolve_output_ref(
+ ref: OutputRef,
+ outputs: dict[str, BaseModel],
+ ) -> object:
+ """Resolve one task output or output field from completed outputs."""
+ output = outputs[ref.task_name]
+ if ref.output_field is None:
+ return output
+ return getattr(output, ref.output_field)
+
+ def _resolve_collection(
+ self,
+ collection: CollectionRef,
+ outputs: dict[str, BaseModel],
+ ) -> object:
+ """Resolve a collection while preserving its declaration order."""
+ if collection.kind == "positional":
+ return [
+ self._resolve_output_ref(ref, outputs) for ref in collection.positional_members
+ ]
+ return {
+ key: self._resolve_output_ref(ref, outputs) for key, ref in collection.keyed_members
+ }
+
+ def _arm(self, task_timeout: float | None, label: str, deadline: _Deadline) -> None:
+ """Arm the timer for one unit of work.
- def _handler(signum: int, frame: Any) -> None:
- raise TaskTimeoutError(f"{label} timed out after {seconds}s")
+ The nearer of the task's own timeout and the remaining job time wins.
+ Raises :class:`_JobTimeoutError` immediately if the job deadline has
+ already passed.
+ """
+ remaining = deadline.remaining()
+ if remaining is not None and remaining <= 0:
+ raise _JobTimeoutError(f"Job timed out after {deadline.job_timeout}s")
- signal.signal(signal.SIGALRM, _handler)
- signal.alarm(int(seconds) if seconds >= 1 else 1)
+ if task_timeout is not None and (remaining is None or task_timeout <= remaining):
+ self._set_alarm(task_timeout, label, deadline=deadline)
+ elif remaining is not None:
+ self._set_alarm(remaining, "Job", deadline=deadline, job_timeout=True)
+
+ def _set_alarm(
+ self,
+ seconds: float,
+ label: str,
+ *,
+ deadline: _Deadline | None = None,
+ job_timeout: bool = False,
+ ) -> bool:
+ """Install a SIGALRM handler and start a one-shot timer.
+
+ Uses ``signal.setitimer`` for sub-second precision, falling back to
+ ``signal.alarm`` where unavailable. Returns True if the timer was set.
+ On platforms or threads where signals cannot be used, a single warning
+ is issued per run and the timeout is not enforced.
+ """
+ if job_timeout and deadline is not None:
+ message = f"Job timed out after {deadline.job_timeout}s"
+ else:
+ message = f"{label} timed out after {seconds}s"
+ error_type: type[TaskTimeoutError] = _JobTimeoutError if job_timeout else TaskTimeoutError
+
+ def _handler(signum: int, frame: Any) -> None:
+ raise error_type(message)
+
+ try:
+ previous = signal.signal(signal.SIGALRM, _handler)
+ if deadline is not None and not deadline.handler_installed:
+ deadline.previous_handler = previous
+ deadline.handler_installed = True
+ setitimer = getattr(signal, "setitimer", None)
+ if setitimer is not None:
+ setitimer(signal.ITIMER_REAL, max(seconds, 1e-6))
+ else: # pragma: no cover - every SIGALRM platform has setitimer
+ signal.alarm(max(1, int(seconds + 0.999999)))
return True
- except (AttributeError, OSError):
- warnings.warn(
- f"signal.alarm not available on this platform; "
- f"timeout for {label} will not be enforced",
- stacklevel=2,
- )
+ except (AttributeError, OSError, ValueError):
+ # ValueError: signal.signal() called outside the main thread.
+ if deadline is None or not deadline.warned:
+ if deadline is not None:
+ deadline.warned = True
+ warnings.warn(
+ f"signal.alarm not available on this platform or thread; "
+ f"timeout for {label} will not be enforced",
+ stacklevel=2,
+ )
return False
+ @staticmethod
+ def _disarm(deadline: _Deadline) -> None:
+ """Cancel any pending timer without touching the handler."""
+ if not deadline.handler_installed:
+ return
+ setitimer = getattr(signal, "setitimer", None)
+ if setitimer is not None:
+ setitimer(signal.ITIMER_REAL, 0)
+ else: # pragma: no cover
+ signal.alarm(0)
+
+ @staticmethod
+ def _restore_handler(deadline: _Deadline) -> None:
+ """Put back the SIGALRM handler that was installed before this run."""
+ if not deadline.handler_installed:
+ return
+ with suppress(AttributeError, OSError, ValueError, TypeError): # pragma: no cover
+ signal.signal(signal.SIGALRM, deadline.previous_handler)
+ deadline.handler_installed = False
+
def _emit(self, event: Event, *args: object) -> None:
- """Dispatch event to all hooks, swallowing any hook errors."""
+ """Dispatch event to all hooks, swallowing any hook errors.
+
+ A failing hook must not abort the workflow, but its error should not
+ vanish either: the warning carries the exception and the original
+ traceback is attached via ``source`` for ``-W error`` / logging capture.
+ """
for hook in self.hooks:
handler = getattr(hook, f"on_{event}", None)
if handler is not None:
try:
handler(*args)
- except Exception:
+ except Exception as exc:
warnings.warn(
- f"Hook {type(hook).__name__} raised during {event}",
+ f"Hook {type(hook).__name__} raised during {event}: {exc!r}",
+ HookError,
stacklevel=2,
+ source=exc,
)
diff --git a/taskmaestro/visualization.py b/taskmaestro/visualization.py
index 6e2c096..04c680b 100644
--- a/taskmaestro/visualization.py
+++ b/taskmaestro/visualization.py
@@ -3,22 +3,28 @@
from __future__ import annotations
import sys
-from typing import TYPE_CHECKING
+from typing import TYPE_CHECKING, Any, get_args, get_origin
-from taskmaestro.task import get_input_type, get_output_type
+from taskmaestro.dependencies import CollectionRef, OutputRef
+from taskmaestro.task import get_input_type
if TYPE_CHECKING:
from taskmaestro.job import JobConfiguration
from taskmaestro.workflow import Workflow
-def _safe_type_name(tp: type, context_cls: type | None = None) -> str:
+def _safe_type_name(tp: Any, context_cls: type | None = None) -> str:
"""Return a Mermaid-safe type name, resolving module-level aliases.
When *context_cls* is provided, its module namespace is scanned for a
variable that refers to *tp*, so that ``GridCase = ObjectModel[X]``
renders as ``GridCase`` instead of ``ObjectModel[X]``.
"""
+ origin = get_origin(tp)
+ if origin is not None:
+ origin_name = getattr(origin, "__name__", str(origin))
+ args = ", ".join(_safe_type_name(arg, context_cls) for arg in get_args(tp))
+ return f"{origin_name}‹{args}›"
name = tp.__name__ if hasattr(tp, "__name__") else str(tp)
if "[" not in name:
return name
@@ -32,16 +38,29 @@ def _safe_type_name(tp: type, context_cls: type | None = None) -> str:
return name.replace("[", "‹").replace("]", "›")
-def _field_type_label(task_by_name: dict[str, type], upstream_name: str, field_name: str) -> str:
+def _field_type_label(
+ workflow: Workflow,
+ task_by_name: dict[str, type],
+ upstream_name: str,
+ field_name: str,
+) -> str:
"""Return ``'.field: FieldType'`` for a field-ref edge."""
upstream_cls = task_by_name[upstream_name]
- output_model = get_output_type(upstream_cls)
+ output_model = workflow.get_output_annotation(upstream_name)
field_info = output_model.model_fields[field_name]
annotation = field_info.annotation
type_label = _safe_type_name(annotation, upstream_cls) if annotation is not None else "Any"
return f".{field_name}: {type_label}"
+def _output_ref_label(workflow: Workflow, task_by_name: dict[str, type], ref: OutputRef) -> str:
+ """Return the type label for a resolved output reference."""
+ task_cls = task_by_name[ref.task_name]
+ if ref.output_field is None:
+ return _safe_type_name(workflow.get_output_annotation(ref.task_name), task_cls)
+ return _field_type_label(workflow, task_by_name, ref.task_name, ref.output_field)
+
+
def _apply_redirect(name: str, redirect: dict[str, str]) -> str:
"""Replace *name* with its redirect target if one exists."""
return redirect.get(name, name)
@@ -72,24 +91,50 @@ def _emit_edges(
elif isinstance(deps, str):
upstream_src = _apply_redirect(deps, source_redirect)
upstream_cls = task_by_name[deps]
- output_name = _safe_type_name(get_output_type(upstream_cls), upstream_cls)
+ output_name = _safe_type_name(workflow.get_output_annotation(deps), upstream_cls)
lines.append(f"{indent}{upstream_src} -->|{output_name}| {tgt_name}")
elif isinstance(deps, tuple):
upstream_name, field_name = deps
upstream_src = _apply_redirect(upstream_name, source_redirect)
- label = _field_type_label(task_by_name, upstream_name, field_name)
+ label = _field_type_label(workflow, task_by_name, upstream_name, field_name)
lines.append(f"{indent}{upstream_src} -->|{label}| {tgt_name}")
elif isinstance(deps, dict):
for down_field, upstream_ref in sorted(deps.items()):
- if isinstance(upstream_ref, tuple):
+ if isinstance(upstream_ref, CollectionRef):
+ collection_node = f"_collect_{tgt_name}_{down_field}_"
+ lines.append(f'{indent}{collection_node}{{{{"collect {down_field}"}}}}')
+ if upstream_ref.kind == "positional":
+ members = [
+ (str(index), ref)
+ for index, ref in enumerate(upstream_ref.positional_members)
+ ]
+ else:
+ members = list(upstream_ref.keyed_members)
+ for member_label, ref in members:
+ upstream_src = _apply_redirect(ref.task_name, source_redirect)
+ label = _output_ref_label(workflow, task_by_name, ref)
+ lines.append(
+ f"{indent}{upstream_src} -->|{member_label}: {label}| "
+ f"{collection_node}"
+ )
+ input_model = get_input_type(task_cls)
+ annotation = input_model.model_fields[down_field].annotation
+ collection_type = _safe_type_name(annotation, task_cls)
+ lines.append(
+ f"{indent}{collection_node} -->|{down_field}: {collection_type}| "
+ f"{tgt_name}"
+ )
+ elif isinstance(upstream_ref, tuple):
upstream_name, up_field = upstream_ref
upstream_src = _apply_redirect(upstream_name, source_redirect)
- label = _field_type_label(task_by_name, upstream_name, up_field)
+ label = _field_type_label(workflow, task_by_name, upstream_name, up_field)
lines.append(f"{indent}{upstream_src} -->|{down_field}: {label}| {tgt_name}")
else:
upstream_src = _apply_redirect(upstream_ref, source_redirect)
up_cls = task_by_name[upstream_ref]
- output_name = _safe_type_name(get_output_type(up_cls), up_cls)
+ output_name = _safe_type_name(
+ workflow.get_output_annotation(upstream_ref), up_cls
+ )
lines.append(
f"{indent}{upstream_src} -->|{down_field}: {output_name}| {tgt_name}"
)
@@ -126,7 +171,11 @@ def to_mermaid(
sinks = [(name, cls) for name, cls in tasks if name not in has_dependents]
# Collect tasks with config_fields
- configured_tasks = {name for name, _cls in tasks if workflow.get_config_fields(name)}
+ configured_tasks = {
+ name
+ for name, _cls in tasks
+ if workflow.get_config_fields(name) or workflow.is_mapped_task(name)
+ }
# Detect workflow_task nodes and build redirect maps
source_redirect: dict[str, str] = {}
@@ -193,7 +242,11 @@ def to_mermaid(
lines.append(" end")
else:
- lines.append(f' {task_name}["{task_name}"]')
+ task_map = workflow.get_task_map(task_name)
+ label = (
+ f"{task_name}
map over: {task_map.over}" if task_map is not None else task_name
+ )
+ lines.append(f' {task_name}["{label}"]')
# Outer edge definitions
_emit_edges(
@@ -210,14 +263,17 @@ def to_mermaid(
# JobConfiguration dashed edges to configured tasks
if configured_tasks:
for task_name in sorted(configured_tasks):
- cf = workflow.get_config_fields(task_name)
+ cf = set(workflow.get_config_fields(task_name))
+ task_map = workflow.get_task_map(task_name)
+ if task_map is not None:
+ cf.add(task_map.over)
label = ", ".join(sorted(cf))
lines.append(f" _job_config_ -.->|{label}| {task_name}")
# Sink tasks: edge to end, labeled with output type
for task_name, task_cls in sinks:
src = _apply_redirect(task_name, source_redirect)
- output_name = _safe_type_name(get_output_type(task_cls), task_cls)
+ output_name = _safe_type_name(workflow.get_output_annotation(task_name), task_cls)
lines.append(f" {src} -->|{output_name}| _end_")
return "\n".join(lines) + "\n"
diff --git a/taskmaestro/workflow.py b/taskmaestro/workflow.py
index 4b9761e..392e7d8 100644
--- a/taskmaestro/workflow.py
+++ b/taskmaestro/workflow.py
@@ -2,18 +2,27 @@
from __future__ import annotations
-from typing import TYPE_CHECKING, Any, Union
+import types
+import typing
+from typing import TYPE_CHECKING, Any, get_args, get_origin
from pydantic import BaseModel
if TYPE_CHECKING:
from taskmaestro.job import JobConfiguration
+from taskmaestro.dependencies import (
+ CollectionDependency,
+ CollectionRef,
+ OutputRef,
+ OutputReference,
+)
from taskmaestro.exceptions import (
CycleDetectedError,
IncompleteInputError,
WorkflowDefinitionError,
)
+from taskmaestro.mapping import MappedOutput, TaskMap
from taskmaestro.task import Task, get_input_type, get_output_type
# Stored dependency types after name resolution:
@@ -21,8 +30,9 @@
# str — single upstream (whole output)
# tuple[str, str] — single upstream, specific field
# dict[str, str | tuple[str, str]] — fan-in (values may be field refs)
-DepValue = Union[str, "tuple[str, str]"]
-StoredDeps = Union[dict[str, DepValue], str, "tuple[str, str]", None]
+type DepValue = str | tuple[str, str]
+type FanInValue = DepValue | CollectionRef
+type StoredDeps = dict[str, FanInValue] | str | tuple[str, str] | None
def _extract_upstream_names(deps: StoredDeps) -> set[str]:
@@ -33,16 +43,94 @@ def _extract_upstream_names(deps: StoredDeps) -> set[str]:
return {deps}
if isinstance(deps, tuple):
return {deps[0]}
- # dict
names: set[str] = set()
- for v in deps.values():
- if isinstance(v, tuple):
- names.add(v[0])
+ for value in deps.values():
+ if isinstance(value, CollectionRef):
+ names.update(ref.task_name for ref in value.output_refs())
+ elif isinstance(value, tuple):
+ names.add(value[0])
else:
- names.add(v)
+ names.add(value)
return names
+def _is_union(annotation: Any) -> bool:
+ """Return whether *annotation* is either spelling of a union."""
+ return get_origin(annotation) in (typing.Union, types.UnionType)
+
+
+def _is_type_compatible(produced: Any, expected: Any) -> bool:
+ """Return whether a produced annotation can be assigned to an expected one.
+
+ Parameterized annotations are compared recursively. This deliberately
+ treats their arguments covariantly: task outputs are validated by Pydantic
+ before they cross an edge, so the question here is whether every produced
+ value is accepted by the downstream annotation rather than whether a
+ mutable container may safely be shared between arbitrary Python callers.
+ """
+ if expected is Any or produced is Any or produced == expected:
+ return True
+
+ # Every possible produced value must be accepted. Conversely, an expected
+ # union only needs one arm which accepts the produced annotation.
+ if _is_union(produced):
+ return all(_is_type_compatible(option, expected) for option in get_args(produced))
+ if _is_union(expected):
+ return any(_is_type_compatible(produced, option) for option in get_args(expected))
+
+ produced_origin = get_origin(produced)
+ expected_origin = get_origin(expected)
+ if produced_origin is not None or expected_origin is not None:
+ produced_base = produced_origin or produced
+ expected_base = expected_origin or expected
+ if not isinstance(produced_base, type) or not isinstance(expected_base, type):
+ return False
+ if not issubclass(produced_base, expected_base):
+ return False
+
+ produced_args = get_args(produced)
+ expected_args = get_args(expected)
+ if not expected_args:
+ return True
+ if not produced_args:
+ return False
+
+ # A fixed-length tuple can be assigned to tuple[T, ...] when each of
+ # its elements can be assigned to T.
+ if expected_base is tuple and len(expected_args) == 2 and expected_args[1] is Ellipsis:
+ if len(produced_args) == 2 and produced_args[1] is Ellipsis:
+ return _is_type_compatible(produced_args[0], expected_args[0])
+ return all(_is_type_compatible(arg, expected_args[0]) for arg in produced_args)
+
+ if len(produced_args) != len(expected_args):
+ return False
+ return all(
+ _is_type_compatible(produced_arg, expected_arg)
+ for produced_arg, expected_arg in zip(produced_args, expected_args, strict=True)
+ )
+
+ if isinstance(produced, type) and isinstance(expected, type):
+ return issubclass(produced, expected)
+ return False
+
+
+def _type_name(annotation: Any) -> str:
+ """Return a readable, complete name for a runtime or typing annotation."""
+ if annotation is Any:
+ return "Any"
+ if annotation is None or annotation is type(None):
+ return "None"
+ if annotation is Ellipsis:
+ return "..."
+ if _is_union(annotation):
+ return " | ".join(_type_name(arg) for arg in get_args(annotation))
+ origin = get_origin(annotation)
+ if origin is not None:
+ origin_name = getattr(origin, "__name__", str(origin).removeprefix("typing."))
+ return f"{origin_name}[{', '.join(_type_name(arg) for arg in get_args(annotation))}]"
+ return getattr(annotation, "__name__", str(annotation))
+
+
class Workflow:
"""A DAG of tasks. Linear pipelines are a special case."""
@@ -61,20 +149,27 @@ def __init__(
self._tasks: dict[str, type[Task[Any, Any]]] = {}
self._dependencies: dict[str, StoredDeps] = {}
self._config_fields: dict[str, set[str]] = {}
+ self._task_maps: dict[str, TaskMap] = {}
self._result_task_name: str | None = None
- if tasks:
+ if tasks is not None:
+ if not tasks:
+ raise WorkflowDefinitionError(
+ f"Workflow '{name}' was given an empty task list; "
+ "pass at least one task or use Workflow.builder()"
+ )
for i, task_cls in enumerate(tasks):
+ if task_cls.name in self._tasks:
+ raise WorkflowDefinitionError(f"Duplicate task name '{task_cls.name}'")
self._tasks[task_cls.name] = task_cls
if i == 0:
self._dependencies[task_cls.name] = None
else:
prev = tasks[i - 1]
self._dependencies[task_cls.name] = prev.name
- self._result_task_name = tasks[-1].name
+ self._result_task_name = result_task.name if result_task is not None else None
self._validate()
-
- if result_task is not None:
+ elif result_task is not None:
self._result_task_name = result_task.name
@classmethod
@@ -133,18 +228,37 @@ def get_config_fields(self, task_name: str) -> set[str]:
"""Return the set of config field names for a task, or empty set."""
return self._config_fields.get(task_name, set())
+ def get_task_map(self, task_name: str) -> TaskMap | None:
+ """Return the mapping declaration for a task, if it is mapped."""
+ return self._task_maps.get(task_name)
+
+ def is_mapped_task(self, task_name: str) -> bool:
+ """Return whether a registered task expands over configured items."""
+ return task_name in self._task_maps
+
+ def get_output_annotation(self, task_name: str) -> Any:
+ """Return a task instance's effective output annotation."""
+ output_type = get_output_type(self._tasks[task_name])
+ if self.is_mapped_task(task_name):
+ return MappedOutput[output_type] # type: ignore[valid-type]
+ return output_type
+
def _validate(self) -> None:
self._validate_unique_names()
self._validate_references()
self._validate_acyclic()
+ self._validate_task_maps()
self._validate_types()
self._validate_result_task()
def _validate_unique_names(self) -> None:
- """Raise WorkflowDefinitionError on duplicate task names."""
- # Already handled by dict keys in _tasks; duplicates would overwrite.
- # For linear shorthand, check the input list explicitly.
- pass
+ """Duplicate task names are rejected at registration time.
+
+ Both ``Workflow(tasks=[...])`` and ``WorkflowBuilder.add_task`` check
+ before inserting into ``_tasks``, so by the time validation runs the
+ mapping is guaranteed to be unique. Kept as an explicit step so the
+ validation order documented in CLAUDE.md remains visible here.
+ """
def _validate_references(self) -> None:
"""Ensure all dependency references point to registered task names."""
@@ -183,6 +297,55 @@ def dfs(node: str) -> None:
if color[node] == WHITE:
dfs(node)
+ def _validate_task_maps(self) -> None:
+ """Validate mapped input fields and their sources."""
+ for task_name, task_map in self._task_maps.items():
+ input_type = get_input_type(self._tasks[task_name])
+ fields = input_type.model_fields
+ for map_field in (task_map.key_as, task_map.value_as):
+ if map_field not in fields:
+ raise WorkflowDefinitionError(
+ f"Map field '{map_field}' not found on {input_type.__name__} "
+ f"(input of '{task_name}')"
+ )
+ if not _is_type_compatible(str, fields[task_map.key_as].annotation):
+ raise WorkflowDefinitionError(
+ f"Map key field '{task_name}.{task_map.key_as}' must accept strings"
+ )
+
+ config_fields = self.get_config_fields(task_name)
+ reserved = {task_map.key_as, task_map.value_as, task_map.over}
+ overlap = reserved & config_fields
+ if overlap:
+ raise WorkflowDefinitionError(
+ f"Mapped task '{task_name}' fields {sorted(overlap)} cannot also be "
+ "config_fields"
+ )
+
+ deps = self._dependencies[task_name]
+ if deps is None:
+ dependency_fields: set[str] = set()
+ elif isinstance(deps, dict):
+ dependency_fields = set(deps)
+ else:
+ raise WorkflowDefinitionError(
+ f"Mapped task '{task_name}' requires named field dependencies"
+ )
+ injected = {task_map.key_as, task_map.value_as}
+ overlap = injected & dependency_fields
+ if overlap:
+ raise WorkflowDefinitionError(
+ f"Mapped task '{task_name}' fields {sorted(overlap)} cannot also be "
+ "dependencies"
+ )
+ covered = dependency_fields | config_fields | injected
+ for field_name, field_info in fields.items():
+ if field_name not in covered and field_info.is_required():
+ raise IncompleteInputError(
+ f"Required field '{field_name}' on {input_type.__name__} is not "
+ f"covered for mapped task '{task_name}'"
+ )
+
def _validate_types(self) -> None:
"""Validate type compatibility for all edges."""
for name, deps in self._dependencies.items():
@@ -194,8 +357,12 @@ def _validate_types(self) -> None:
downstream_input = get_input_type(task_cls)
model_fields = downstream_input.model_fields
# Validate config_fields cover all required input fields
+ task_map = self.get_task_map(name)
+ map_fields = (
+ {task_map.key_as, task_map.value_as} if task_map is not None else set()
+ )
for field_name, field_info in model_fields.items():
- if field_name not in cf and field_info.is_required():
+ if field_name not in cf | map_fields and field_info.is_required():
raise IncompleteInputError(
f"Required field '{field_name}' on "
f"{downstream_input.__name__} is not covered by "
@@ -211,10 +378,14 @@ def _validate_types(self) -> None:
continue
elif isinstance(deps, str):
# Single dependency (whole output)
- upstream_cls = self._tasks[deps]
- upstream_output = get_output_type(upstream_cls)
+ upstream_output = self.get_output_annotation(deps)
downstream_input = get_input_type(task_cls)
if cf:
+ if self.is_mapped_task(deps):
+ raise WorkflowDefinitionError(
+ f"Mapped upstream task '{deps}' must be connected through "
+ "a named input field"
+ )
# With config_fields: check upstream output fields exist in
# downstream input with compatible types, and that upstream
# fields + config_fields cover all required fields
@@ -235,12 +406,12 @@ def _validate_types(self) -> None:
if (
up_annotation is not None
and down_annotation is not None
- and not issubclass(up_annotation, down_annotation)
+ and not _is_type_compatible(up_annotation, down_annotation)
):
raise WorkflowDefinitionError(
f"Type mismatch: {deps}.{field_name} is "
- f"{up_annotation.__name__} but {name}.{field_name} "
- f"expects {down_annotation.__name__}"
+ f"{_type_name(up_annotation)} but {name}.{field_name} "
+ f"expects {_type_name(down_annotation)}"
)
# Check all required fields are covered by upstream or config
covered = set(up_fields.keys()) | cf
@@ -252,17 +423,16 @@ def _validate_types(self) -> None:
f"upstream output or config_fields"
)
else:
- if upstream_output is not downstream_input:
+ if not _is_type_compatible(upstream_output, downstream_input):
raise WorkflowDefinitionError(
f"Type mismatch: {deps} outputs "
- f"{upstream_output.__name__} but {name} expects "
- f"{downstream_input.__name__}"
+ f"{_type_name(upstream_output)} but {name} expects "
+ f"{_type_name(downstream_input)}"
)
elif isinstance(deps, tuple):
# Single dependency, specific output field
upstream_name, field_name = deps
- upstream_cls = self._tasks[upstream_name]
- upstream_output = get_output_type(upstream_cls)
+ upstream_output = self.get_output_annotation(upstream_name)
upstream_fields = upstream_output.model_fields
if field_name not in upstream_fields:
raise WorkflowDefinitionError(
@@ -271,11 +441,13 @@ def _validate_types(self) -> None:
)
field_annotation = upstream_fields[field_name].annotation
downstream_input = get_input_type(task_cls)
- if field_annotation is not None and downstream_input is not field_annotation:
+ if field_annotation is not None and not _is_type_compatible(
+ field_annotation, downstream_input
+ ):
raise WorkflowDefinitionError(
f"Type mismatch: {upstream_name}.{field_name} is "
- f"{field_annotation.__name__} but {name} expects "
- f"{downstream_input.__name__}"
+ f"{_type_name(field_annotation)} but {name} expects "
+ f"{_type_name(downstream_input)}"
)
elif isinstance(deps, dict):
# Fan-in: validate each field
@@ -288,10 +460,23 @@ def _validate_types(self) -> None:
raise WorkflowDefinitionError(
f"Fan-in field '{field_name}' not found on {downstream_input.__name__}"
)
+ field_annotation = model_fields[field_name].annotation
+ if isinstance(upstream_ref, CollectionRef):
+ if field_name in cf:
+ raise WorkflowDefinitionError(
+ f"Field '{field_name}' on task '{name}' is supplied by both "
+ "a collection dependency and config_fields"
+ )
+ self._validate_collection(
+ name,
+ field_name,
+ field_annotation,
+ upstream_ref,
+ )
+ continue
if isinstance(upstream_ref, tuple):
up_name, up_field = upstream_ref
- up_cls = self._tasks[up_name]
- up_output = get_output_type(up_cls)
+ up_output = self.get_output_annotation(up_name)
up_fields = up_output.model_fields
if up_field not in up_fields:
raise WorkflowDefinitionError(
@@ -300,19 +485,17 @@ def _validate_types(self) -> None:
)
resolved_type = up_fields[up_field].annotation
else:
- up_cls = self._tasks[upstream_ref]
- resolved_type = get_output_type(up_cls)
- field_annotation = model_fields[field_name].annotation
+ resolved_type = self.get_output_annotation(upstream_ref)
if (
field_annotation is not None
and resolved_type is not None
- and not issubclass(resolved_type, field_annotation)
+ and not _is_type_compatible(resolved_type, field_annotation)
):
raise WorkflowDefinitionError(
f"Fan-in type mismatch: {upstream_ref} outputs "
- f"{resolved_type.__name__} but field '{field_name}' "
+ f"{_type_name(resolved_type)} but field '{field_name}' "
f"on {downstream_input.__name__} expects "
- f"{field_annotation.__name__}"
+ f"{_type_name(field_annotation)}"
)
# Validate config field names exist on the model
for field_name in cf:
@@ -322,7 +505,11 @@ def _validate_types(self) -> None:
f"{downstream_input.__name__} (input of '{name}')"
)
# Check all required fields are covered by deps or config_fields
- covered = set(deps.keys()) | cf
+ task_map = self.get_task_map(name)
+ map_fields = (
+ {task_map.key_as, task_map.value_as} if task_map is not None else set()
+ )
+ covered = set(deps.keys()) | cf | map_fields
for field_name, field_info in model_fields.items():
if field_name not in covered and field_info.is_required():
raise IncompleteInputError(
@@ -331,17 +518,77 @@ def _validate_types(self) -> None:
f"upstream task"
)
+ def _resolve_output_ref_type(self, ref: OutputRef) -> Any:
+ """Resolve the type produced by an output reference."""
+ output_type = self.get_output_annotation(ref.task_name)
+ if ref.output_field is None:
+ return output_type
+ if ref.output_field not in output_type.model_fields:
+ raise WorkflowDefinitionError(
+ f"Field '{ref.output_field}' not found on {output_type.__name__} "
+ f"(output of {ref.task_name})"
+ )
+ return output_type.model_fields[ref.output_field].annotation
+
+ def _validate_collection(
+ self,
+ task_name: str,
+ field_name: str,
+ field_annotation: Any,
+ collection: CollectionRef,
+ ) -> None:
+ """Validate a collection dependency against its destination field."""
+ origin = get_origin(field_annotation)
+ args = get_args(field_annotation)
+ if collection.kind == "positional":
+ if origin is not list or len(args) != 1:
+ raise WorkflowDefinitionError(
+ f"Positional collection for '{task_name}.{field_name}' requires "
+ f"a list[T] field, got {_type_name(field_annotation)}"
+ )
+ expected_type = args[0]
+ members = [
+ (str(index), ref) for index, ref in enumerate(collection.positional_members)
+ ]
+ else:
+ if origin is not dict or len(args) != 2 or args[0] is not str:
+ raise WorkflowDefinitionError(
+ f"Keyed collection for '{task_name}.{field_name}' requires "
+ f"a dict[str, T] field, got {_type_name(field_annotation)}"
+ )
+ expected_type = args[1]
+ members = list(collection.keyed_members)
+
+ for member_label, ref in members:
+ produced_type = self._resolve_output_ref_type(ref)
+ if produced_type is not None and not _is_type_compatible(produced_type, expected_type):
+ source = ref.task_name
+ if ref.output_field is not None:
+ source += f".{ref.output_field}"
+ raise WorkflowDefinitionError(
+ f"Collection type mismatch for "
+ f"'{task_name}.{field_name}[{member_label}]': '{source}' produces "
+ f"{_type_name(produced_type)}, but collection element type is "
+ f"{_type_name(expected_type)}"
+ )
+
def _validate_result_task(self) -> None:
- """Ensure result_task is set. Default to sole sink; raise if ambiguous."""
- sinks = self._find_sinks()
- if self._result_task_name is None:
- if len(sinks) == 1:
- self._result_task_name = sinks[0]
- else:
+ """Ensure result_task is set and registered. Default to sole sink; raise if ambiguous."""
+ if self._result_task_name is not None:
+ if self._result_task_name not in self._tasks:
raise WorkflowDefinitionError(
- f"Workflow '{self.name}' has {len(sinks)} sink tasks "
- f"({sinks}); specify result_task explicitly"
+ f"result_task '{self._result_task_name}' is not registered in "
+ f"workflow '{self.name}' (known tasks: {sorted(self._tasks)})"
)
+ return
+ sinks = self._find_sinks()
+ if len(sinks) == 1:
+ self._result_task_name = sinks[0]
+ else:
+ raise WorkflowDefinitionError(
+ f"Workflow '{self.name}' has {len(sinks)} sink tasks "
+ f"({sinks}); specify result_task explicitly"
+ )
def as_task(
self,
@@ -387,6 +634,7 @@ def __init__(
self._workflow._tasks = {}
self._workflow._dependencies = {}
self._workflow._config_fields = {}
+ self._workflow._task_maps = {}
self._workflow._result_task_name = None
# Store the raw result_task ref for resolution at build() time
self._result_task_ref: type[Task[Any, Any]] | str | None = result_task
@@ -425,6 +673,29 @@ def _resolve_dep_ref(
return dep
return self._resolve_dep_name(dep)
+ def _resolve_output_reference(self, ref: OutputReference) -> OutputRef:
+ """Resolve a public task/output-field reference."""
+ if isinstance(ref, tuple):
+ task_ref, output_field = ref
+ return OutputRef(self._resolve_dep_ref(task_ref), output_field)
+ return OutputRef(self._resolve_dep_ref(ref))
+
+ def _resolve_collection(self, collection: CollectionDependency) -> CollectionRef:
+ """Resolve every task reference in a collection dependency."""
+ if collection.kind == "positional":
+ return CollectionRef(
+ "positional",
+ positional_members=tuple(
+ self._resolve_output_reference(ref) for ref in collection.positional_members
+ ),
+ )
+ return CollectionRef(
+ "keyed",
+ keyed_members=tuple(
+ (key, self._resolve_output_reference(ref)) for key, ref in collection.keyed_members
+ ),
+ )
+
def add_task(
self,
task_cls: type[Task[Any, Any]],
@@ -434,10 +705,17 @@ def add_task(
type[Task[Any, Any]]
| str
| tuple[type[Task[Any, Any]] | str, str]
- | dict[str, type[Task[Any, Any]] | str | tuple[type[Task[Any, Any]] | str, str]]
+ | dict[
+ str,
+ type[Task[Any, Any]]
+ | str
+ | tuple[type[Task[Any, Any]] | str, str]
+ | CollectionDependency,
+ ]
| None
) = None,
config_fields: list[str] | None = None,
+ mapped_over: TaskMap | None = None,
) -> WorkflowBuilder:
"""Add a task to the DAG. Returns self for chaining.
@@ -451,6 +729,9 @@ def add_task(
- ``(TaskClass | "name", "field")`` — single upstream, specific output field
- ``{"field": TaskClass | "name", ...}`` — fan-in, whole outputs
- ``{"field": (TaskClass | "name", "f"), ...}`` — fan-in with field routing
+ - ``{"field": collect(...), ...}`` — collect outputs into a list or dictionary
+
+ ``mapped_over`` expands this logical task over a configured mapping.
"""
wf = self._workflow
task_name = name if name is not None else task_cls.name
@@ -465,9 +746,11 @@ def add_task(
resolved_name = self._resolve_dep_ref(dep_ref)
wf._dependencies[task_name] = (resolved_name, field)
elif isinstance(depends_on, dict):
- resolved: dict[str, DepValue] = {}
+ resolved: dict[str, FanInValue] = {}
for field, dep in depends_on.items():
- if isinstance(dep, tuple):
+ if isinstance(dep, CollectionDependency):
+ resolved[field] = self._resolve_collection(dep)
+ elif isinstance(dep, tuple):
dep_ref, dep_field = dep
resolved_name = self._resolve_dep_ref(dep_ref)
resolved[field] = (resolved_name, dep_field)
@@ -484,6 +767,8 @@ def add_task(
if config_fields is not None:
wf._config_fields[task_name] = set(config_fields)
+ if mapped_over is not None:
+ wf._task_maps[task_name] = mapped_over
return self
diff --git a/taskmaestro/workflow_task.py b/taskmaestro/workflow_task.py
index 6a94e6e..09fac92 100644
--- a/taskmaestro/workflow_task.py
+++ b/taskmaestro/workflow_task.py
@@ -5,10 +5,10 @@
from typing import Any
from taskmaestro.context import ExecutionContext
-from taskmaestro.exceptions import WorkflowDefinitionError
+from taskmaestro.exceptions import WorkflowDefinitionError, WorkflowTaskError
from taskmaestro.job import EmptyConfig, Job, JobConfiguration, JobStatus
from taskmaestro.runner import Runner
-from taskmaestro.task import Task, get_input_type, get_output_type
+from taskmaestro.task import Task, get_input_type
from taskmaestro.workflow import Workflow
@@ -39,13 +39,18 @@ def workflow_task(
WorkflowDefinitionError: If the inner workflow does not have exactly one
root task without config_fields (unless all roots are covered by
job_configuration).
+
+ At run time, a failure inside the inner workflow surfaces as
+ :class:`~taskmaestro.exceptions.WorkflowTaskError`, which carries the
+ completed inner :class:`~taskmaestro.job.Job` and chains the original
+ exception as ``__cause__``.
"""
# Find root tasks: tasks with deps=None and no config_fields
roots: list[tuple[str, type[Task[Any, Any]]]] = []
for task_name, deps in workflow._dependencies.items():
if deps is None:
config_fields = workflow.get_config_fields(task_name)
- if not config_fields:
+ if not config_fields and not workflow.is_mapped_task(task_name):
roots.append((task_name, workflow._tasks[task_name]))
all_roots_configured = False
@@ -70,8 +75,7 @@ def workflow_task(
else:
input_type = get_input_type(roots[0][1])
- result_task_cls = workflow.result_task
- output_type = get_output_type(result_task_cls)
+ output_type = workflow.get_output_annotation(workflow.result_task_name)
resolved_name = name if name is not None else workflow.name
inner_wf = workflow
@@ -87,10 +91,7 @@ def run(self, input: Any, ctx: ExecutionContext) -> Any:
job = Job(workflow=inner_wf, config=cfg, job_configuration=inner_jc)
result_job = Runner().run(job, ctx=ctx)
if result_job.status == JobStatus.FAILED:
- raise RuntimeError(
- f"Inner workflow '{inner_wf.name}' failed at task "
- f"'{result_job.failed_task}': {result_job.error}"
- )
+ raise WorkflowTaskError(inner_wf.name, result_job) from result_job.exception
return result_job.result
_WorkflowTask.__name__ = f"WorkflowTask_{resolved_name}"
diff --git a/taskmaestro/yaml_config.py b/taskmaestro/yaml_config.py
index 3f0b644..f9c4595 100644
--- a/taskmaestro/yaml_config.py
+++ b/taskmaestro/yaml_config.py
@@ -12,10 +12,12 @@
from pydantic import BaseModel, Field, ValidationError, model_validator
from taskmaestro.context import ExecutionContext
+from taskmaestro.dependencies import CollectionDependency, OutputReference, collect
from taskmaestro.discovery import get_registered_task, registered_task_names
-from taskmaestro.exceptions import ConfigLoadError, PluginLoadError
+from taskmaestro.exceptions import ConfigLoadError, PluginLoadError, WorkflowDefinitionError
from taskmaestro.hooks.base import BaseHook
from taskmaestro.job import EmptyConfig, Job, JobConfiguration
+from taskmaestro.mapping import TaskMap
from taskmaestro.runner import Runner
from taskmaestro.task import Task, get_input_type
from taskmaestro.workflow import Workflow, WorkflowBuilder
@@ -23,6 +25,15 @@
# --- Pydantic schema models for YAML validation ---
+class TaskMapConfig(BaseModel):
+ """Mapped execution settings for a YAML task entry."""
+
+ over: str
+ key_as: str
+ value_as: str
+ error_mode: typing.Literal["fail_fast", "collect_all"] = "fail_fast"
+
+
class TaskConfig(BaseModel):
"""A single task entry in the YAML workflow config."""
@@ -32,6 +43,7 @@ class TaskConfig(BaseModel):
name: str | None = None
depends_on: str | list[str] | dict[str, Any] | None = None
config_fields: list[str] | None = None
+ map: TaskMapConfig | None = None
@model_validator(mode="after")
def _check_task_or_workflow(self) -> TaskConfig:
@@ -71,6 +83,7 @@ class WorkflowSectionConfig(BaseModel):
name: str
result_task: str | None = None
+ input_mode: typing.Literal["auto", "flat", "per_task"] = "auto"
tasks: list[TaskConfig] = Field(min_length=1)
@@ -85,6 +98,42 @@ class YamlWorkflowConfig(BaseModel):
# --- Utilities ---
+class _UniqueKeyLoader(yaml.SafeLoader):
+ """Safe YAML loader that rejects duplicate mapping keys."""
+
+ def __init__(self, stream: str) -> None:
+ super().__init__(stream)
+ self._checked_mappings: set[yaml.nodes.MappingNode] = set()
+
+ def flatten_mapping(self, node: yaml.nodes.MappingNode) -> None:
+ # Check declarations before merges add inherited keys. Anchors can reuse
+ # already-flattened nodes, whose override keys are legitimately repeated.
+ if node in self._checked_mappings:
+ return
+ self._checked_mappings.add(node)
+ keys: set[Any] = set()
+ for key_node, _value_node in node.value:
+ key = (
+ "<<"
+ if key_node.tag == "tag:yaml.org,2002:merge"
+ else self.construct_object(key_node)
+ )
+ if key in keys:
+ raise yaml.constructor.ConstructorError(
+ "while constructing a mapping",
+ node.start_mark,
+ f"found duplicate key {key!r}",
+ key_node.start_mark,
+ )
+ keys.add(key)
+ super().flatten_mapping(node)
+
+
+def _yaml_load(text: str) -> Any:
+ """Safely parse YAML while rejecting duplicate mapping keys."""
+ return yaml.load(text, Loader=_UniqueKeyLoader)
+
+
def import_class(dotted_path: str) -> type[Any]:
"""Import a class from a dotted path like 'pkg.mod.ClassName'.
@@ -123,6 +172,87 @@ def _coerce_hook_params(hook_cls: type[Any], params: dict[str, Any]) -> dict[str
# --- LoadedWorkflow ---
+@dataclass(frozen=True)
+class _LinearDep:
+ """An already-registered upstream name produced by linear-mode chaining."""
+
+ upstream: str
+
+
+@dataclass(frozen=True)
+class _Entry:
+ """One resolved ``tasks:`` entry, kept positionally."""
+
+ config: TaskConfig
+ cls: type[Task[Any, Any]]
+ key: str # import path or inner-workflow path as written in YAML
+ registered_name: str
+
+
+def _decide_input_mode(
+ mode: typing.Literal["auto", "flat", "per_task"],
+ raw_input: dict[str, Any],
+ entries: list[_Entry],
+ *,
+ linear: bool,
+) -> bool:
+ """Return True when *raw_input* is per-task configuration.
+
+ See the comment at the call site for the three modes. Raises
+ ConfigLoadError for an explicit ``per_task`` file with unknown keys or
+ non-mapping values, and for an ``auto`` file that reads validly both ways.
+ """
+ names = {entry.registered_name for entry in entries}
+
+ if mode == "flat":
+ return False
+
+ if mode == "per_task":
+ for key, value in raw_input.items():
+ if key not in names:
+ raise ConfigLoadError(
+ f"input_mode is 'per_task' but top-level key '{key}' is not a task "
+ f"name (known tasks: {sorted(names)})"
+ )
+ if value is not None and not isinstance(value, dict):
+ raise ConfigLoadError(
+ f"input_mode is 'per_task' but the value for task '{key}' is not a "
+ f"mapping (got {type(value).__name__})"
+ )
+ return True
+
+ # auto
+ looks_per_task = bool(raw_input) and all(
+ key in names and isinstance(raw_input[key], (dict, type(None))) for key in raw_input
+ )
+ if not looks_per_task:
+ return False
+
+ # The heuristic fired. If the same mapping is also a valid input for the
+ # sole unconfigured root task, both readings are plausible: refuse to guess.
+ roots = [
+ entry
+ for index, entry in enumerate(entries)
+ if (entry.config.depends_on is None if not linear else index == 0)
+ and not entry.config.config_fields
+ and entry.config.map is None
+ ]
+ if len(roots) == 1:
+ try:
+ get_input_type(roots[0].cls).model_validate(raw_input)
+ except ValidationError:
+ pass
+ else:
+ raise ConfigLoadError(
+ f"Input file is ambiguous: its top-level keys {sorted(raw_input)} are task "
+ f"names, but the mapping is also a valid "
+ f"{get_input_type(roots[0].cls).__name__} for root task "
+ f"'{roots[0].registered_name}'. Set workflow.input_mode to 'flat' or "
+ f"'per_task'."
+ )
+ return True
+
+
@dataclass(frozen=True)
class LoadedWorkflow:
"""A fully resolved workflow ready to execute."""
@@ -153,19 +283,30 @@ def _timeout_seconds(self) -> float | None:
def _load_workflow_only(
workflow_path: Path,
input_path: Path | None = None,
+ *,
+ _ancestors: frozenset[Path] = frozenset(),
) -> tuple[Workflow, JobConfiguration | None]:
"""Build a Workflow and optional JobConfiguration from YAML files.
This is the core logic shared by ``load_workflow_from_yaml`` and
- recursive ``workflow:`` references in YAML configs.
+ recursive ``workflow:`` references in YAML configs. ``_ancestors`` holds
+ the resolved paths of every enclosing workflow file so that a self- or
+ mutually-referencing ``workflow:`` entry is rejected instead of recursing
+ without bound.
Returns (workflow, job_configuration).
"""
from taskmaestro.workflow_task import workflow_task as _workflow_task
+ resolved_path = workflow_path.resolve()
+ if resolved_path in _ancestors:
+ chain = " -> ".join(str(p) for p in (*sorted(_ancestors), resolved_path))
+ raise ConfigLoadError(f"Recursive workflow reference: {chain}")
+ ancestors = _ancestors | {resolved_path}
+
# 1. Parse workflow YAML
try:
- raw = yaml.safe_load(workflow_path.read_text())
+ raw = _yaml_load(workflow_path.read_text())
except yaml.YAMLError as exc:
raise ConfigLoadError(f"YAML parse error: {exc}") from exc
except OSError as exc:
@@ -178,7 +319,7 @@ def _load_workflow_only(
raw_input: dict[str, Any] = {}
if input_path is not None:
try:
- raw_input = yaml.safe_load(input_path.read_text())
+ raw_input = _yaml_load(input_path.read_text())
except yaml.YAMLError as exc:
raise ConfigLoadError(f"Input YAML parse error: {exc}") from exc
except OSError as exc:
@@ -193,9 +334,11 @@ def _load_workflow_only(
except ValidationError as exc:
raise ConfigLoadError(f"YAML schema validation error: {exc}") from exc
- # 4. Resolve task import paths (handles both task: and workflow: entries)
+ # 4. Resolve task import paths (handles both task: and workflow: entries).
+ # Entries are kept positionally: the same class path or inner YAML file
+ # may legitimately appear more than once under different ``name:``s.
base_dir = workflow_path.parent
- task_classes: dict[str, type[Task[Any, Any]]] = {}
+ entries: list[_Entry] = []
installed_task_names = registered_task_names()
for task_config in config.workflow.tasks:
if task_config.workflow:
@@ -204,11 +347,14 @@ def _load_workflow_only(
inner_input_path = (
base_dir / task_config.workflow_input if task_config.workflow_input else None
)
- inner_wf, inner_jc = _load_workflow_only(inner_wf_path, inner_input_path)
+ inner_wf, inner_jc = _load_workflow_only(
+ inner_wf_path, inner_input_path, _ancestors=ancestors
+ )
inner_name = task_config.name if task_config.name else inner_wf.name
- wrapped_cls = _workflow_task(inner_wf, name=inner_name, job_configuration=inner_jc)
- # Use a synthetic key for this entry (the workflow path)
- task_classes[task_config.workflow] = wrapped_cls
+ cls: type[Task[Any, Any]] = _workflow_task(
+ inner_wf, name=inner_name, job_configuration=inner_jc
+ )
+ key = task_config.workflow
else:
assert task_config.task is not None
try:
@@ -220,40 +366,85 @@ def _load_workflow_only(
raise ConfigLoadError(str(exc)) from exc
if not (isinstance(cls, type) and issubclass(cls, Task)):
raise ConfigLoadError(f"'{task_config.task}' is not a Task subclass")
- task_classes[task_config.task] = cls
-
- # Helper to get the lookup key for a task config entry
- def _task_key(tc: TaskConfig) -> str:
- return tc.workflow if tc.workflow else tc.task # type: ignore[return-value]
+ key = task_config.task
+ registered_name = task_config.name if task_config.name else cls.name
+ entries.append(_Entry(task_config, cls, key, registered_name))
# 5. Build a lookup from instance names and import paths to registered names.
- name_lookup: dict[str, str] = {}
- for task_config in config.workflow.tasks:
- key = _task_key(task_config)
- registered_name = task_config.name if task_config.name else task_classes[key].name
- name_lookup[key] = registered_name
- if task_config.name:
- name_lookup[task_config.name] = registered_name
-
- def _resolve_yaml_dep(dep_str: str, context_task: str) -> str:
+ # A key that maps to more than one registered name is ambiguous and is
+ # rejected when used, with the candidates listed.
+ candidates: dict[str, set[str]] = {}
+ for entry in entries:
+ candidates.setdefault(entry.key, set()).add(entry.registered_name)
+ if entry.config.name:
+ candidates.setdefault(entry.config.name, set()).add(entry.registered_name)
+
+ def _resolve_yaml_dep(dep_str: str, context_task: str, *, what: str = "Dependency") -> str:
"""Resolve a YAML dependency string to a registered task name."""
- if dep_str in name_lookup:
- return name_lookup[dep_str]
- raise ConfigLoadError(f"Dependency '{dep_str}' for task '{context_task}' not found")
+ where = f" for task '{context_task}'" if context_task else ""
+ names = candidates.get(dep_str)
+ if names is None:
+ raise ConfigLoadError(f"{what} '{dep_str}'{where} not found")
+ if len(names) > 1:
+ raise ConfigLoadError(
+ f"{what} '{dep_str}'{where} is ambiguous; it matches "
+ f"{sorted(names)}. Use the instance name."
+ )
+ return next(iter(names))
+
+ def _resolve_yaml_output_ref(raw_ref: Any, context_task: str) -> OutputReference:
+ """Resolve a YAML task or ``[task, field]`` output reference."""
+ if isinstance(raw_ref, str):
+ return _resolve_yaml_dep(raw_ref, context_task)
+ if isinstance(raw_ref, list):
+ if len(raw_ref) != 2 or not all(isinstance(item, str) for item in raw_ref):
+ raise ConfigLoadError(
+ f"Collection member must be a task name or [task, field], "
+ f"got {raw_ref!r} for task '{context_task}'"
+ )
+ return (_resolve_yaml_dep(raw_ref[0], context_task), raw_ref[1])
+ raise ConfigLoadError(
+ f"Collection member must be a task name or [task, field], "
+ f"got {raw_ref!r} for task '{context_task}'"
+ )
- # 6. Detect linear vs DAG mode
- has_depends_on = any(tc.depends_on is not None for tc in config.workflow.tasks)
+ def _resolve_yaml_collection(raw_collection: Any, context_task: str) -> CollectionDependency:
+ """Resolve a YAML collect list or mapping."""
+ if isinstance(raw_collection, list):
+ return collect(
+ *(_resolve_yaml_output_ref(member, context_task) for member in raw_collection)
+ )
+ if isinstance(raw_collection, dict):
+ if not all(isinstance(key, str) for key in raw_collection):
+ raise ConfigLoadError(f"Collection keys must be strings for task '{context_task}'")
+ return collect(
+ {
+ key: _resolve_yaml_output_ref(member, context_task)
+ for key, member in raw_collection.items()
+ }
+ )
+ raise ConfigLoadError(
+ f"'collect' must contain a list or mapping for task '{context_task}'"
+ )
- # 6b. Detect per-task config format early (before building workflow)
- all_registered_names: set[str] = set()
- for task_config in config.workflow.tasks:
- key = _task_key(task_config)
- registered_name = task_config.name if task_config.name else task_classes[key].name
- all_registered_names.add(registered_name)
+ # 6. Detect linear vs DAG mode
+ has_depends_on = any(
+ tc.depends_on is not None or tc.map is not None for tc in config.workflow.tasks
+ )
- is_per_task_config = bool(raw_input) and all(
- key in all_registered_names and isinstance(raw_input[key], (dict, type(None)))
- for key in raw_input
+ # 6b. Decide how the input YAML is interpreted.
+ #
+ # flat — the mapping is the single unconfigured root task's input model
+ # per_task — top-level keys are task names, values are per-task config
+ # auto — infer; every key must name a task with a mapping/null value.
+ # If the flat reading is *also* valid the file is ambiguous
+ # and the user must set ``input_mode`` explicitly.
+ entry_by_name = {entry.registered_name: entry for entry in entries}
+ is_per_task_config = _decide_input_mode(
+ config.workflow.input_mode,
+ raw_input,
+ entries,
+ linear=not has_depends_on,
)
per_task_data: dict[str, dict[str, Any]] = {}
@@ -262,51 +453,63 @@ def _resolve_yaml_dep(dep_str: str, context_task: str) -> str:
for task_name, task_values in raw_input.items():
per_task_data[task_name] = dict(task_values) if task_values else {}
if task_values:
- per_task_cfg_fields[task_name] = list(task_values.keys())
+ task_config = entry_by_name[task_name].config
+ map_source = task_config.map.over if task_config.map is not None else None
+ per_task_cfg_fields[task_name] = [
+ field_name for field_name in task_values if field_name != map_source
+ ]
# 7. Resolve result_task
result_task_name: str | None = None
if config.workflow.result_task:
- if config.workflow.result_task in name_lookup:
- result_task_name = name_lookup[config.workflow.result_task]
- else:
- raise ConfigLoadError(f"result_task '{config.workflow.result_task}' not found")
-
- # 8. Build Workflow
- if not has_depends_on:
- task_list = [task_classes[_task_key(tc)] for tc in config.workflow.tasks]
- result_task_cls = (
- task_classes[config.workflow.result_task] if config.workflow.result_task else None
- )
- workflow = Workflow(
- name=config.workflow.name,
- tasks=task_list,
- result_task=result_task_cls,
- )
- for task_config in config.workflow.tasks:
- key = _task_key(task_config)
- registered_name = task_config.name if task_config.name else task_classes[key].name
- cfg = task_config.config_fields or per_task_cfg_fields.get(registered_name)
- if cfg:
- workflow._config_fields[registered_name] = set(cfg)
- else:
- builder = WorkflowBuilder(
- config.workflow.name,
- result_task=result_task_name,
+ result_task_name = _resolve_yaml_dep(config.workflow.result_task, "", what="result_task")
+
+ # 8. Build Workflow. Both linear and DAG configs go through the builder so
+ # that ``name:`` overrides, config_fields and validation behave the same.
+ # In linear mode each task depends on the whole output of the previous one.
+ builder = WorkflowBuilder(config.workflow.name, result_task=result_task_name)
+ previous_registered: str | None = None
+ for entry in entries:
+ task_config = entry.config
+ key = entry.key
+ cls = entry.cls
+ registered_name = entry.registered_name
+ instance_name = task_config.name
+ cfg_fields = task_config.config_fields or per_task_cfg_fields.get(registered_name)
+ mapped_over = (
+ TaskMap(**task_config.map.model_dump()) if task_config.map is not None else None
)
- for task_config in config.workflow.tasks:
- key = _task_key(task_config)
- cls = task_classes[key]
- deps = task_config.depends_on
- registered_name = name_lookup[key]
- instance_name = task_config.name
- cfg_fields = task_config.config_fields or per_task_cfg_fields.get(registered_name)
+ deps: str | list[str] | dict[str, Any] | _LinearDep | None = task_config.depends_on
+ if not has_depends_on and previous_registered is not None:
+ # Linear mode: chain on the previous task's registered name, which
+ # the builder accepts verbatim as a string dependency.
+ deps = _LinearDep(previous_registered)
+ previous_registered = registered_name
+
+ try:
if deps is None:
- builder.add_task(cls, name=instance_name, config_fields=cfg_fields)
+ builder.add_task(
+ cls,
+ name=instance_name,
+ config_fields=cfg_fields,
+ mapped_over=mapped_over,
+ )
+ elif isinstance(deps, _LinearDep):
+ builder.add_task(
+ cls,
+ name=instance_name,
+ depends_on=deps.upstream,
+ config_fields=cfg_fields,
+ mapped_over=mapped_over,
+ )
elif isinstance(deps, str):
resolved_dep = _resolve_yaml_dep(deps, key)
builder.add_task(
- cls, name=instance_name, depends_on=resolved_dep, config_fields=cfg_fields
+ cls,
+ name=instance_name,
+ depends_on=resolved_dep,
+ config_fields=cfg_fields,
+ mapped_over=mapped_over,
)
elif isinstance(deps, list):
if len(deps) != 2 or not all(isinstance(e, str) for e in deps):
@@ -321,13 +524,20 @@ def _resolve_yaml_dep(dep_str: str, context_task: str) -> str:
name=instance_name,
depends_on=(resolved_dep, field_name),
config_fields=cfg_fields,
+ mapped_over=mapped_over,
)
elif isinstance(deps, dict):
fan_in: dict[
- str, type[Task[Any, Any]] | str | tuple[type[Task[Any, Any]] | str, str]
+ str,
+ type[Task[Any, Any]]
+ | str
+ | tuple[type[Task[Any, Any]] | str, str]
+ | CollectionDependency,
] = {}
for field_name, upstream_ref in deps.items():
- if isinstance(upstream_ref, list):
+ if isinstance(upstream_ref, dict) and set(upstream_ref) == {"collect"}:
+ fan_in[field_name] = _resolve_yaml_collection(upstream_ref["collect"], key)
+ elif isinstance(upstream_ref, list):
if len(upstream_ref) != 2 or not all(
isinstance(e, str) for e in upstream_ref
):
@@ -339,16 +549,27 @@ def _resolve_yaml_dep(dep_str: str, context_task: str) -> str:
up_path, up_field = upstream_ref
resolved_dep = _resolve_yaml_dep(up_path, key)
fan_in[field_name] = (resolved_dep, up_field)
- else:
+ elif isinstance(upstream_ref, str):
resolved_dep = _resolve_yaml_dep(upstream_ref, key)
fan_in[field_name] = resolved_dep
+ else:
+ raise ConfigLoadError(
+ f"Invalid dependency {upstream_ref!r} for field "
+ f"'{field_name}' on task '{key}'"
+ )
builder.add_task(
- cls, name=instance_name, depends_on=fan_in, config_fields=cfg_fields
+ cls,
+ name=instance_name,
+ depends_on=fan_in,
+ config_fields=cfg_fields,
+ mapped_over=mapped_over,
)
- try:
- workflow = builder.build()
- except Exception as exc:
+ except WorkflowDefinitionError as exc:
raise ConfigLoadError(f"Workflow validation failed: {exc}") from exc
+ try:
+ workflow = builder.build()
+ except Exception as exc:
+ raise ConfigLoadError(f"Workflow validation failed: {exc}") from exc
# 9. Build JobConfiguration if per-task config detected
job_configuration: JobConfiguration | None = None
@@ -371,7 +592,7 @@ def load_workflow_from_yaml(workflow_path: str | Path, input_path: str | Path) -
# 1. Parse workflow YAML (needed for runner/context config)
try:
- raw = yaml.safe_load(workflow_path.read_text())
+ raw = _yaml_load(workflow_path.read_text())
except yaml.YAMLError as exc:
raise ConfigLoadError(f"YAML parse error: {exc}") from exc
except OSError as exc:
@@ -382,7 +603,7 @@ def load_workflow_from_yaml(workflow_path: str | Path, input_path: str | Path) -
# 2. Parse input YAML
try:
- raw_input = yaml.safe_load(input_path.read_text())
+ raw_input = _yaml_load(input_path.read_text())
except yaml.YAMLError as exc:
raise ConfigLoadError(f"Input YAML parse error: {exc}") from exc
except OSError as exc:
@@ -402,26 +623,36 @@ def load_workflow_from_yaml(workflow_path: str | Path, input_path: str | Path) -
# 5. Validate input and build Job
job: Job[Any]
- if job_configuration is not None:
- job = Job(workflow, EmptyConfig(), job_configuration=job_configuration)
- else:
- # Flat config mode: find root tasks from the built workflow
- root_task_classes = [
- workflow._tasks[task_name]
- for task_name, deps in workflow._dependencies.items()
- if deps is None and not workflow.get_config_fields(task_name)
- ]
- assert root_task_classes, (
- "job_configuration is None yet no roots found without config_fields"
- )
-
- input_type = get_input_type(root_task_classes[0])
- try:
- validated_input = input_type.model_validate(raw_input)
- except ValidationError as exc:
- raise ConfigLoadError(f"Input validation error: {exc}") from exc
-
- job = Job(workflow, validated_input)
+ try:
+ if job_configuration is not None:
+ job = Job(workflow, EmptyConfig(), job_configuration=job_configuration)
+ else:
+ # Flat config mode: find root tasks from the built workflow
+ root_task_classes = [
+ workflow._tasks[task_name]
+ for task_name, deps in workflow._dependencies.items()
+ if deps is None and not workflow.get_config_fields(task_name)
+ ]
+ if not root_task_classes and all(
+ deps is not None for deps in workflow._dependencies.values()
+ ):
+ # The workflow is self-contained, for example a task fed only by
+ # explicitly empty collection dependencies.
+ job = Job(workflow, EmptyConfig())
+ elif not root_task_classes:
+ raise ConfigLoadError(
+ "Workflow has configured root tasks but no per-task input configuration"
+ )
+ else:
+ input_type = get_input_type(root_task_classes[0])
+ try:
+ validated_input = input_type.model_validate(raw_input)
+ except ValidationError as exc:
+ raise ConfigLoadError(f"Input validation error: {exc}") from exc
+
+ job = Job(workflow, validated_input)
+ except WorkflowDefinitionError as exc:
+ raise ConfigLoadError(f"Job validation failed: {exc}") from exc
# 6. Instantiate hooks
hooks: list[BaseHook] = []
diff --git a/tests/test_collections.py b/tests/test_collections.py
new file mode 100644
index 0000000..30e2520
--- /dev/null
+++ b/tests/test_collections.py
@@ -0,0 +1,573 @@
+"""Tests for collecting multiple task outputs into one input field."""
+
+from __future__ import annotations
+
+from pathlib import Path
+from typing import Any, Literal
+
+import pytest
+from pydantic import BaseModel
+
+from taskmaestro import (
+ ConfigLoadError,
+ EmptyConfig,
+ ExecutionContext,
+ Job,
+ JobStatus,
+ Runner,
+ Task,
+ Workflow,
+ WorkflowDefinitionError,
+ collect,
+ load_workflow_from_yaml,
+)
+from taskmaestro.workflow import _is_type_compatible, _type_name
+from tests.conftest import NumberInput
+
+
+class Surface(BaseModel):
+ name: str
+
+
+class RegularSurface(Surface):
+ source: str = "generated"
+
+
+class ProduceSurface(Task[NumberInput, RegularSurface]):
+ name = "produce_surface"
+
+ def run(self, input: NumberInput, ctx: ExecutionContext) -> RegularSurface:
+ return RegularSurface(name=f"{self.name}-{input.value}")
+
+
+class SurfaceEnvelope(BaseModel):
+ surface: RegularSurface
+ ignored: str
+
+
+class ProduceEnvelope(Task[NumberInput, SurfaceEnvelope]):
+ name = "produce_envelope"
+
+ def run(self, input: NumberInput, ctx: ExecutionContext) -> SurfaceEnvelope:
+ return SurfaceEnvelope(
+ surface=RegularSurface(name=f"{self.name}-{input.value}"),
+ ignored="ignored",
+ )
+
+
+class SurfaceListInput(BaseModel):
+ surfaces: list[Surface]
+
+
+class SurfaceNames(BaseModel):
+ names: list[str]
+
+
+class CollectSurfaceList(Task[SurfaceListInput, SurfaceNames]):
+ name = "collect_surface_list"
+
+ def run(self, input: SurfaceListInput, ctx: ExecutionContext) -> SurfaceNames:
+ return SurfaceNames(names=[surface.name for surface in input.surfaces])
+
+
+class SurfaceDictInput(BaseModel):
+ surfaces: dict[str, Surface]
+
+
+class CollectSurfaceDict(Task[SurfaceDictInput, SurfaceNames]):
+ name = "collect_surface_dict"
+
+ def run(self, input: SurfaceDictInput, ctx: ExecutionContext) -> SurfaceNames:
+ return SurfaceNames(
+ names=[f"{key}:{surface.name}" for key, surface in input.surfaces.items()]
+ )
+
+
+class TextOutput(BaseModel):
+ text: str
+
+
+class GenericSurfaceOutput(BaseModel):
+ surfaces: dict[str, RegularSurface]
+
+
+class BadGenericSurfaceOutput(BaseModel):
+ surfaces: dict[int, RegularSurface]
+
+
+class GenericSurfaceInput(BaseModel):
+ surfaces: dict[str, Surface | None]
+
+
+class ProduceGenericSurfaces(Task[NumberInput, GenericSurfaceOutput]):
+ name = "produce_generic_surfaces"
+
+ def run(self, input: NumberInput, ctx: ExecutionContext) -> GenericSurfaceOutput:
+ return GenericSurfaceOutput(surfaces={})
+
+
+class ProduceBadGenericSurfaces(Task[NumberInput, BadGenericSurfaceOutput]):
+ name = "produce_bad_generic_surfaces"
+
+ def run(self, input: NumberInput, ctx: ExecutionContext) -> BadGenericSurfaceOutput:
+ return BadGenericSurfaceOutput(surfaces={})
+
+
+class ConsumeGenericSurfaces(Task[GenericSurfaceInput, SurfaceNames]):
+ name = "consume_generic_surfaces"
+
+ def run(self, input: GenericSurfaceInput, ctx: ExecutionContext) -> SurfaceNames:
+ return SurfaceNames(names=[])
+
+
+class ProduceText(Task[NumberInput, TextOutput]):
+ name = "produce_text"
+
+ def run(self, input: NumberInput, ctx: ExecutionContext) -> TextOutput:
+ return TextOutput(text=str(input.value))
+
+
+class TestCollectDeclaration:
+ def test_dictionary_keys_must_be_strings(self) -> None:
+ with pytest.raises(TypeError, match="keys must be strings"):
+ collect({1: ProduceSurface}) # type: ignore[dict-item]
+
+ def test_mapping_cannot_be_mixed_with_positional_members(self) -> None:
+ with pytest.raises(TypeError, match="either positional members or one mapping"):
+ collect(ProduceSurface, {"other": ProduceSurface}) # type: ignore[call-overload]
+
+
+class TestCollectionWorkflow:
+ def test_list_collects_outputs_in_declaration_order(self) -> None:
+ workflow = (
+ Workflow.builder("surface_list")
+ .add_task(ProduceSurface, name="second")
+ .add_task(ProduceSurface, name="first")
+ .add_task(
+ CollectSurfaceList,
+ depends_on={"surfaces": collect("first", "second")},
+ )
+ .build()
+ )
+
+ result = Runner().run(Job(workflow, NumberInput(value=7)))
+
+ assert result.status == JobStatus.COMPLETED
+ assert result.result == SurfaceNames(names=["first-7", "second-7"])
+ collection = workflow.get_dependencies("collect_surface_list")
+ assert collection is not None
+
+ def test_collects_whole_outputs_and_routed_fields(self) -> None:
+ workflow = (
+ Workflow.builder("routed_collection")
+ .add_task(ProduceSurface)
+ .add_task(ProduceEnvelope)
+ .add_task(
+ CollectSurfaceList,
+ depends_on={
+ "surfaces": collect(
+ ProduceSurface,
+ (ProduceEnvelope, "surface"),
+ )
+ },
+ )
+ .build()
+ )
+
+ result = Runner().run(Job(workflow, NumberInput(value=3)))
+
+ assert result.status == JobStatus.COMPLETED
+ assert result.result == SurfaceNames(names=["produce_surface-3", "produce_envelope-3"])
+
+ def test_collects_keyed_outputs_in_declaration_order(self) -> None:
+ workflow = (
+ Workflow.builder("surface_dict")
+ .add_task(ProduceSurface, name="top_task")
+ .add_task(ProduceEnvelope, name="base_task")
+ .add_task(
+ CollectSurfaceDict,
+ depends_on={
+ "surfaces": collect(
+ {
+ "top": "top_task",
+ "base": ("base_task", "surface"),
+ }
+ )
+ },
+ )
+ .build()
+ )
+
+ result = Runner().run(Job(workflow, NumberInput(value=4)))
+
+ assert result.status == JobStatus.COMPLETED
+ assert result.result == SurfaceNames(names=["top:top_task-4", "base:base_task-4"])
+
+ def test_empty_list_collection(self) -> None:
+ workflow = (
+ Workflow.builder("empty_collection")
+ .add_task(
+ CollectSurfaceList,
+ depends_on={"surfaces": collect()},
+ )
+ .build()
+ )
+
+ result = Runner().run(Job(workflow, EmptyConfig()))
+
+ assert result.status == JobStatus.COMPLETED
+ assert result.result == SurfaceNames(names=[])
+
+ def test_empty_dictionary_collection(self) -> None:
+ workflow = (
+ Workflow.builder("empty_dictionary")
+ .add_task(
+ CollectSurfaceDict,
+ depends_on={"surfaces": collect({})},
+ )
+ .build()
+ )
+
+ result = Runner().run(Job(workflow, EmptyConfig()))
+
+ assert result.status == JobStatus.COMPLETED
+ assert result.result == SurfaceNames(names=[])
+
+ def test_incompatible_member_is_rejected(self) -> None:
+ with pytest.raises(WorkflowDefinitionError, match="Collection type mismatch"):
+ (
+ Workflow.builder("bad_collection")
+ .add_task(ProduceText)
+ .add_task(
+ CollectSurfaceList,
+ depends_on={"surfaces": collect(ProduceText)},
+ )
+ .build()
+ )
+
+ def test_missing_routed_output_field_is_rejected(self) -> None:
+ with pytest.raises(WorkflowDefinitionError, match="Field 'missing' not found"):
+ (
+ Workflow.builder("missing_field")
+ .add_task(ProduceEnvelope)
+ .add_task(
+ CollectSurfaceList,
+ depends_on={"surfaces": collect((ProduceEnvelope, "missing"))},
+ )
+ .build()
+ )
+
+ def test_routed_output_field_type_mismatch_names_source_field(self) -> None:
+ with pytest.raises(
+ WorkflowDefinitionError,
+ match=r"produce_envelope\.ignored.*collection element type is Surface",
+ ):
+ (
+ Workflow.builder("bad_field_type")
+ .add_task(ProduceEnvelope)
+ .add_task(
+ CollectSurfaceList,
+ depends_on={"surfaces": collect((ProduceEnvelope, "ignored"))},
+ )
+ .build()
+ )
+
+ def test_collection_shape_must_match_field(self) -> None:
+ with pytest.raises(WorkflowDefinitionError, match=r"requires a list\[T\] field"):
+ (
+ Workflow.builder("bad_shape")
+ .add_task(ProduceSurface)
+ .add_task(
+ CollectSurfaceDict,
+ depends_on={"surfaces": collect(ProduceSurface)},
+ )
+ .build()
+ )
+
+ def test_keyed_collection_requires_dictionary_field(self) -> None:
+ with pytest.raises(WorkflowDefinitionError, match=r"requires a dict\[str, T\] field"):
+ (
+ Workflow.builder("bad_keyed_shape")
+ .add_task(ProduceSurface)
+ .add_task(
+ CollectSurfaceList,
+ depends_on={"surfaces": collect({"surface": ProduceSurface})},
+ )
+ .build()
+ )
+
+ def test_type_compatibility_handles_unions_and_parameterized_types(self) -> None:
+ assert _is_type_compatible(RegularSurface, Surface | TextOutput)
+ assert _is_type_compatible(dict[str, RegularSurface], dict[str, Surface])
+ assert _is_type_compatible(list[RegularSurface], list[Surface | None])
+ assert _is_type_compatible(
+ dict[str, list[RegularSurface]],
+ dict[str, list[Surface | None]],
+ )
+ assert _is_type_compatible(RegularSurface | TextOutput, Surface | TextOutput)
+ assert not _is_type_compatible(list[int], list[str])
+ assert not _is_type_compatible(dict[int, RegularSurface], dict[str, Surface])
+ assert not _is_type_compatible(RegularSurface | int, Surface)
+
+ def test_type_compatibility_handles_generic_edge_cases(self) -> None:
+ assert _is_type_compatible(list[RegularSurface], list)
+ assert not _is_type_compatible(list, list[Surface])
+ assert not _is_type_compatible(list[int], dict[int, int])
+ assert not _is_type_compatible(Literal["produced"], Literal["expected"])
+ assert not _is_type_compatible("Produced", "Expected")
+ assert not _is_type_compatible(tuple[int, str], tuple[int])
+
+ def test_type_compatibility_handles_variadic_tuples(self) -> None:
+ assert _is_type_compatible(tuple[RegularSurface, ...], tuple[Surface, ...])
+ assert _is_type_compatible(
+ tuple[RegularSurface, TextOutput],
+ tuple[Surface | TextOutput, ...],
+ )
+ assert not _is_type_compatible(tuple[RegularSurface, int], tuple[Surface, ...])
+
+ def test_type_name_handles_special_annotations(self) -> None:
+ assert _type_name(Any) == "Any"
+ assert _type_name(None) == "None"
+ assert _type_name(type(None)) == "None"
+ assert _type_name(Ellipsis) == "..."
+
+ def test_parameterized_fan_in_types_are_compared_recursively(self) -> None:
+ workflow = (
+ Workflow.builder("generic_fan_in")
+ .add_task(ProduceGenericSurfaces)
+ .add_task(
+ ConsumeGenericSurfaces,
+ depends_on={"surfaces": (ProduceGenericSurfaces, "surfaces")},
+ )
+ .build()
+ )
+
+ assert workflow.result_task is ConsumeGenericSurfaces
+
+ def test_parameterized_fan_in_error_shows_complete_annotations(self) -> None:
+ with pytest.raises(
+ WorkflowDefinitionError,
+ match=(
+ r"outputs dict\[int, RegularSurface\].*"
+ r"expects dict\[str, Surface \| None\]"
+ ),
+ ):
+ (
+ Workflow.builder("bad_generic_fan_in")
+ .add_task(ProduceBadGenericSurfaces)
+ .add_task(
+ ConsumeGenericSurfaces,
+ depends_on={"surfaces": (ProduceBadGenericSurfaces, "surfaces")},
+ )
+ .build()
+ )
+
+ def test_collection_and_config_cannot_supply_same_field(self) -> None:
+ with pytest.raises(WorkflowDefinitionError, match="supplied by both"):
+ (
+ Workflow.builder("conflicting_sources")
+ .add_task(ProduceSurface)
+ .add_task(
+ CollectSurfaceList,
+ depends_on={"surfaces": collect(ProduceSurface)},
+ config_fields=["surfaces"],
+ )
+ .build()
+ )
+
+
+class TestCollectionVisualization:
+ def test_collection_uses_explicit_junction_node(self) -> None:
+ workflow = (
+ Workflow.builder("collection_viz")
+ .add_task(ProduceSurface, name="top")
+ .add_task(ProduceEnvelope, name="base")
+ .add_task(
+ CollectSurfaceList,
+ depends_on={
+ "surfaces": collect("top", ("base", "surface")),
+ },
+ )
+ .build()
+ )
+
+ diagram = workflow.to_mermaid()
+
+ assert '_collect_collect_surface_list_surfaces_{{"collect surfaces"}}' in diagram
+ assert "top -->|0: RegularSurface| _collect_collect_surface_list_surfaces_" in diagram
+ assert "base -->|1: .surface: RegularSurface|" in diagram
+ assert "-->|surfaces: list‹Surface›|" in diagram
+
+ def test_keyed_collection_edges_use_aliases(self) -> None:
+ workflow = (
+ Workflow.builder("keyed_collection_viz")
+ .add_task(ProduceSurface, name="top")
+ .add_task(
+ CollectSurfaceDict,
+ depends_on={"surfaces": collect({"top_alias": "top"})},
+ )
+ .build()
+ )
+
+ diagram = workflow.to_mermaid()
+
+ assert "top -->|top_alias: RegularSurface|" in diagram
+ assert "-->|surfaces: dict‹str, Surface›|" in diagram
+
+
+class TestCollectionYaml:
+ def test_yaml_collection_end_to_end(self, tmp_path: Path) -> None:
+ module = "tests.test_collections"
+ workflow_path = tmp_path / "workflow.yaml"
+ workflow_path.write_text(
+ f"""\
+workflow:
+ name: yaml_collection
+ tasks:
+ - task: {module}.ProduceSurface
+ name: first
+ - task: {module}.ProduceEnvelope
+ name: second
+ - task: {module}.CollectSurfaceList
+ depends_on:
+ surfaces:
+ collect:
+ - first
+ - [second, surface]
+"""
+ )
+ input_path = tmp_path / "input.yaml"
+ input_path.write_text("value: 9\n")
+
+ result = load_workflow_from_yaml(workflow_path, input_path).run()
+
+ assert result.status == JobStatus.COMPLETED
+ assert result.result == SurfaceNames(names=["first-9", "second-9"])
+
+ def test_yaml_keyed_collection_end_to_end(self, tmp_path: Path) -> None:
+ module = "tests.test_collections"
+ workflow_path = tmp_path / "workflow.yaml"
+ workflow_path.write_text(
+ f"""\
+workflow:
+ name: yaml_keyed_collection
+ tasks:
+ - task: {module}.ProduceSurface
+ name: top_task
+ - task: {module}.ProduceEnvelope
+ name: base_task
+ - task: {module}.CollectSurfaceDict
+ depends_on:
+ surfaces:
+ collect:
+ top: top_task
+ base: [base_task, surface]
+"""
+ )
+ input_path = tmp_path / "input.yaml"
+ input_path.write_text("value: 5\n")
+
+ result = load_workflow_from_yaml(workflow_path, input_path).run()
+
+ assert result.status == JobStatus.COMPLETED
+ assert result.result == SurfaceNames(names=["top:top_task-5", "base:base_task-5"])
+
+ def test_yaml_empty_collection(self, tmp_path: Path) -> None:
+ module = "tests.test_collections"
+ workflow_path = tmp_path / "workflow.yaml"
+ workflow_path.write_text(
+ f"""\
+workflow:
+ name: yaml_empty_collection
+ tasks:
+ - task: {module}.CollectSurfaceList
+ depends_on:
+ surfaces:
+ collect: []
+"""
+ )
+ input_path = tmp_path / "input.yaml"
+ input_path.write_text("{}\n")
+
+ result = load_workflow_from_yaml(workflow_path, input_path).run()
+
+ assert result.status == JobStatus.COMPLETED
+ assert result.result == SurfaceNames(names=[])
+
+ @pytest.mark.parametrize(
+ ("dependency_yaml", "message"),
+ [
+ ("collect: [[producer]]", "Collection member must be"),
+ ("collect: [123]", "Collection member must be"),
+ ("collect: producer", "must contain a list or mapping"),
+ ("collect: {1: producer}", "Collection keys must be strings"),
+ ("unexpected: producer", "Invalid dependency"),
+ ],
+ )
+ def test_invalid_yaml_collection_forms(
+ self,
+ tmp_path: Path,
+ dependency_yaml: str,
+ message: str,
+ ) -> None:
+ module = "tests.test_collections"
+ workflow_path = tmp_path / "workflow.yaml"
+ workflow_path.write_text(
+ f"""\
+workflow:
+ name: invalid_collection
+ tasks:
+ - task: {module}.ProduceSurface
+ name: producer
+ - task: {module}.CollectSurfaceList
+ depends_on:
+ surfaces:
+ {dependency_yaml}
+"""
+ )
+ input_path = tmp_path / "input.yaml"
+ input_path.write_text("value: 1\n")
+
+ with pytest.raises(ConfigLoadError, match=message):
+ load_workflow_from_yaml(workflow_path, input_path)
+
+ def test_configured_root_without_per_task_input_is_rejected(self, tmp_path: Path) -> None:
+ workflow_path = tmp_path / "workflow.yaml"
+ workflow_path.write_text(
+ """\
+workflow:
+ name: missing_task_configuration
+ tasks:
+ - task: tests.conftest.ConfigOnlyTask
+ config_fields: [path, count]
+"""
+ )
+ input_path = tmp_path / "input.yaml"
+ input_path.write_text("{}\n")
+
+ with pytest.raises(ConfigLoadError, match="no per-task input configuration"):
+ load_workflow_from_yaml(workflow_path, input_path)
+
+ def test_duplicate_yaml_collection_key_is_rejected(self, tmp_path: Path) -> None:
+ module = "tests.test_collections"
+ workflow_path = tmp_path / "workflow.yaml"
+ workflow_path.write_text(
+ f"""\
+workflow:
+ name: duplicate_key
+ tasks:
+ - task: {module}.ProduceSurface
+ name: producer
+ - task: {module}.CollectSurfaceDict
+ depends_on:
+ surfaces:
+ collect:
+ top: producer
+ top: producer
+"""
+ )
+ input_path = tmp_path / "input.yaml"
+ input_path.write_text("value: 1\n")
+
+ with pytest.raises(ConfigLoadError, match="duplicate key"):
+ load_workflow_from_yaml(workflow_path, input_path)
diff --git a/tests/test_exceptions.py b/tests/test_exceptions.py
index e3eeb18..6c37656 100644
--- a/tests/test_exceptions.py
+++ b/tests/test_exceptions.py
@@ -37,6 +37,18 @@ def test_task_output_type_error(self) -> None:
def test_task_timeout_error(self) -> None:
assert issubclass(TaskTimeoutError, TaskExecutionError)
+ def test_workflow_task_error(self) -> None:
+ from types import SimpleNamespace
+
+ from taskmaestro.exceptions import WorkflowTaskError
+
+ assert issubclass(WorkflowTaskError, TaskExecutionError)
+ fake_job = SimpleNamespace(failed_task="step", error="kaboom")
+ exc = WorkflowTaskError("inner", fake_job)
+ assert exc.workflow_name == "inner"
+ assert exc.inner_job is fake_job
+ assert str(exc) == "Inner workflow 'inner' failed at task 'step': kaboom"
+
def test_exception_messages(self) -> None:
exc = CycleDetectedError("cycle found")
assert str(exc) == "cycle found"
diff --git a/tests/test_hooks.py b/tests/test_hooks.py
index 8254514..eaf5064 100644
--- a/tests/test_hooks.py
+++ b/tests/test_hooks.py
@@ -26,6 +26,7 @@
FailingTask,
FanInTask,
NumberInput,
+ NumberOutput,
)
@@ -143,6 +144,35 @@ def test_writes_json_files(self, ctx: ExecutionContext, tmp_path: Path) -> None:
double_data = json.loads(double_path.read_text())
assert double_data["value"] == 12
+ def test_task_name_cannot_escape_output_dir(
+ self, ctx: ExecutionContext, tmp_path: Path
+ ) -> None:
+ """Path separators and '..' in a task name are escaped, not interpreted."""
+
+ class Traversal(Task[NumberInput, NumberOutput]):
+ name = "../escaped"
+
+ def run(self, input: NumberInput, ctx: ExecutionContext) -> NumberOutput:
+ return NumberOutput(value=input.value)
+
+ out_dir = tmp_path / "sandbox" / "results"
+ wf = Workflow(name="test", tasks=[Traversal])
+ job = Job(workflow=wf, config=NumberInput(value=1))
+ Runner(hooks=[ResultPersistenceHook(output_dir=out_dir)]).run(job, ctx=ctx)
+
+ written = list(tmp_path.rglob("*.json"))
+ assert len(written) == 1
+ assert written[0].parent == out_dir
+ assert written[0].name == "..%2Fescaped.json"
+ assert not (tmp_path / "sandbox" / "escaped.json").exists()
+
+ def test_distinct_names_do_not_collide(self, tmp_path: Path) -> None:
+ """Escaping '%' keeps 'a%2Fb' and 'a/b' on different filenames."""
+ from taskmaestro.hooks.persistence import _safe
+
+ assert _safe("a/b") != _safe("a%2Fb")
+ assert "/" not in _safe("a/b")
+
class TestHookErrorHandling:
def test_hook_error_swallowed(self, ctx: ExecutionContext) -> None:
@@ -156,6 +186,50 @@ def on_task_start(self, job: Job[Any], task: Task[Any, Any]) -> None:
result = Runner(hooks=[BrokenHook()]).run(job, ctx=ctx)
assert result.status == JobStatus.COMPLETED
+ def test_hook_warning_carries_exception_and_category(self, ctx: ExecutionContext) -> None:
+ """The warning names the exception, uses HookError, and attaches it as source."""
+ import warnings
+
+ from taskmaestro import HookError
+
+ class BrokenHook(BaseHook):
+ def on_task_complete(
+ self, job: Job[Any], task: Task[Any, Any], output: BaseModel
+ ) -> None:
+ raise KeyError("missing-service")
+
+ wf = Workflow(name="test", tasks=[AddOne])
+ job = Job(workflow=wf, config=NumberInput(value=1))
+ with warnings.catch_warnings(record=True) as caught:
+ warnings.simplefilter("always")
+ Runner(hooks=[BrokenHook()]).run(job, ctx=ctx)
+
+ hook_warnings = [w for w in caught if issubclass(w.category, HookError)]
+ assert len(hook_warnings) == 1
+ message = str(hook_warnings[0].message)
+ assert "BrokenHook raised during task_complete" in message
+ assert "KeyError('missing-service')" in message
+ assert isinstance(hook_warnings[0].source, KeyError)
+ # HookError is a UserWarning so existing filters still apply.
+ assert issubclass(HookError, UserWarning)
+
+ def test_hook_warning_can_be_escalated(self, ctx: ExecutionContext) -> None:
+ """Users may opt into strictness with a warnings filter on HookError."""
+ import warnings
+
+ from taskmaestro import HookError
+
+ class BrokenHook(BaseHook):
+ def on_job_start(self, job: Job[Any]) -> None:
+ raise RuntimeError("nope")
+
+ wf = Workflow(name="test", tasks=[AddOne])
+ job = Job(workflow=wf, config=NumberInput(value=1))
+ with warnings.catch_warnings():
+ warnings.simplefilter("error", HookError)
+ with pytest.raises(HookError, match="RuntimeError\\('nope'\\)"):
+ Runner(hooks=[BrokenHook()]).run(job, ctx=ctx)
+
def test_multiple_hooks(self, ctx: ExecutionContext) -> None:
wf = Workflow(name="test", tasks=[AddOne])
job = Job(workflow=wf, config=NumberInput(value=1))
diff --git a/tests/test_job.py b/tests/test_job.py
index 578d915..0348038 100644
--- a/tests/test_job.py
+++ b/tests/test_job.py
@@ -35,6 +35,7 @@ def test_initial_state(self) -> None:
job = Job(workflow=wf, config=NumberInput(value=1))
assert job.result is None
assert job.error is None
+ assert job.exception is None
assert job.failed_task is None
assert job.started_at is None
assert job.completed_at is None
diff --git a/tests/test_mapping.py b/tests/test_mapping.py
new file mode 100644
index 0000000..4611d02
--- /dev/null
+++ b/tests/test_mapping.py
@@ -0,0 +1,880 @@
+"""Tests for sequential mapped task expansion."""
+
+from __future__ import annotations
+
+import signal
+from pathlib import Path
+from typing import Any, ClassVar
+
+import pytest
+from pydantic import BaseModel, ConfigDict, Field
+
+from taskmaestro import (
+ EmptyConfig,
+ ExecutionContext,
+ Job,
+ JobConfiguration,
+ JobStatus,
+ MappedOutput,
+ MappedTaskExecutionError,
+ Runner,
+ Task,
+ TaskMap,
+ Workflow,
+ WorkflowDefinitionError,
+ collect,
+ workflow_task,
+)
+from taskmaestro.hooks import LoggingHook, ResultPersistenceHook, TimingHook
+from taskmaestro.hooks.base import BaseHook
+from taskmaestro.yaml_config import ConfigLoadError, load_workflow_from_yaml
+from tests.conftest import AddOne, MergeTask, NumberInput, NumberOutput, StringOutput
+
+
+class MappedInput(BaseModel):
+ base: NumberOutput
+ item_name: str
+ amount: int
+ multiplier: int
+
+
+class MappedNumber(Task[MappedInput, NumberOutput]):
+ name = "mapped_number"
+ seen: ClassVar[list[tuple[int, str, str]]] = []
+
+ def run(self, input: MappedInput, ctx: ExecutionContext) -> NumberOutput:
+ self.seen.append((id(self), input.item_name, ctx.correlation_id))
+ if input.amount < 0:
+ raise ValueError(f"negative amount for {input.item_name}")
+ return NumberOutput(value=input.base.value + input.amount * input.multiplier)
+
+
+class EnvelopeOutput(BaseModel):
+ number: NumberOutput
+
+
+class ProduceEnvelope(Task[NumberInput, EnvelopeOutput]):
+ name = "produce_envelope_for_map"
+
+ def run(self, input: NumberInput, ctx: ExecutionContext) -> EnvelopeOutput:
+ return EnvelopeOutput(number=NumberOutput(value=input.value))
+
+
+class MappedCollectionInput(BaseModel):
+ bases: list[NumberOutput]
+ item_name: str
+ amount: int
+
+
+class MappedCollection(Task[MappedCollectionInput, NumberOutput]):
+ name = "mapped_collection"
+
+ def run(self, input: MappedCollectionInput, ctx: ExecutionContext) -> NumberOutput:
+ return NumberOutput(value=sum(item.value for item in input.bases) + input.amount)
+
+
+class MappedOnlyInput(BaseModel):
+ item_name: str
+ amount: int
+
+
+class MappedOnly(Task[MappedOnlyInput, NumberOutput]):
+ name = "mapped_only"
+
+ def run(self, input: MappedOnlyInput, ctx: ExecutionContext) -> NumberOutput:
+ return NumberOutput(value=input.amount)
+
+
+class MappedWrongOutput(Task[MappedOnlyInput, NumberOutput]):
+ name = "mapped_wrong_output"
+
+ def run(self, input: MappedOnlyInput, ctx: ExecutionContext) -> NumberOutput:
+ return StringOutput(text="wrong") # type: ignore[return-value]
+
+
+class MappedSlow(Task[MappedOnlyInput, NumberOutput]):
+ name = "mapped_slow"
+ timeout_seconds = 10
+
+ def run(self, input: MappedOnlyInput, ctx: ExecutionContext) -> NumberOutput:
+ return NumberOutput(value=input.amount)
+
+
+class AggregateInput(BaseModel):
+ values: dict[str, NumberOutput]
+
+
+class SumAggregate(Task[AggregateInput, NumberOutput]):
+ name = "sum_aggregate"
+
+ def run(self, input: AggregateInput, ctx: ExecutionContext) -> NumberOutput:
+ return NumberOutput(value=sum(value.value for value in input.values.values()))
+
+
+def _mapped_workflow(*, error_mode: str = "fail_fast") -> Workflow:
+ return (
+ Workflow.builder("mapped", result_task=SumAggregate)
+ .add_task(AddOne)
+ .add_task(
+ MappedNumber,
+ depends_on={"base": AddOne},
+ config_fields=["multiplier"],
+ mapped_over=TaskMap(
+ over="items",
+ key_as="item_name",
+ value_as="amount",
+ error_mode=error_mode, # type: ignore[arg-type]
+ ),
+ )
+ .add_task(SumAggregate, depends_on={"values": (MappedNumber, "root")})
+ .build()
+ )
+
+
+def _mapped_job(workflow: Workflow, items: dict[Any, Any]) -> Job[NumberInput]:
+ return Job(
+ workflow,
+ NumberInput(value=10),
+ job_configuration=JobConfiguration({"mapped_number": {"items": items, "multiplier": 2}}),
+ )
+
+
+class RecordingMapHook(BaseHook):
+ def __init__(self) -> None:
+ self.events: list[str] = []
+
+ def on_map_item_start(self, job: Job[Any], task: Task[Any, Any], key: str) -> None:
+ self.events.append(f"start:{task.name}[{key}]")
+
+ def on_map_item_complete(
+ self, job: Job[Any], task: Task[Any, Any], key: str, output: object
+ ) -> None:
+ self.events.append(f"complete:{task.name}[{key}]")
+
+ def on_map_item_fail(
+ self, job: Job[Any], task: Task[Any, Any], key: str, error: Exception
+ ) -> None:
+ self.events.append(f"fail:{task.name}[{key}]")
+
+
+class TestTaskMap:
+ @pytest.mark.parametrize("field", ["over", "key_as", "value_as"])
+ def test_fields_must_not_be_empty(self, field: str) -> None:
+ values = {"over": "items", "key_as": "item_name", "value_as": "amount"}
+ values[field] = ""
+ with pytest.raises(ValueError, match=f"TaskMap.{field}"):
+ TaskMap(**values) # type: ignore[arg-type]
+
+ def test_injected_fields_must_differ(self) -> None:
+ with pytest.raises(ValueError, match="must be different"):
+ TaskMap(over="items", key_as="item", value_as="item")
+
+ def test_error_mode_is_validated_at_runtime(self) -> None:
+ with pytest.raises(ValueError, match="error_mode"):
+ TaskMap(
+ over="items",
+ key_as="key",
+ value_as="value",
+ error_mode="invalid", # type: ignore[arg-type]
+ )
+
+
+class TestMappedWorkflowValidation:
+ def test_mapping_metadata_and_effective_output(self) -> None:
+ workflow = _mapped_workflow()
+
+ assert workflow.is_mapped_task("mapped_number")
+ assert workflow.get_task_map("mapped_number") is not None
+ assert workflow.get_task_map("add_one") is None
+ assert workflow.get_output_annotation("mapped_number") == MappedOutput[NumberOutput]
+ assert workflow.get_output_annotation("add_one") is NumberOutput
+
+ @pytest.mark.parametrize("key_as,value_as", [("missing", "amount"), ("item_name", "missing")])
+ def test_map_fields_must_exist(self, key_as: str, value_as: str) -> None:
+ with pytest.raises(WorkflowDefinitionError, match="Map field 'missing'"):
+ (
+ Workflow.builder("bad")
+ .add_task(
+ MappedOnly,
+ mapped_over=TaskMap(over="items", key_as=key_as, value_as=value_as),
+ )
+ .build()
+ )
+
+ def test_key_field_must_accept_strings(self) -> None:
+ class NumericKeyInput(BaseModel):
+ key: int
+ amount: int
+
+ class NumericKeyTask(Task[NumericKeyInput, NumberOutput]):
+ def run(self, input: NumericKeyInput, ctx: ExecutionContext) -> NumberOutput:
+ return NumberOutput(value=input.amount)
+
+ with pytest.raises(WorkflowDefinitionError, match="must accept strings"):
+ (
+ Workflow.builder("bad")
+ .add_task(
+ NumericKeyTask,
+ mapped_over=TaskMap(over="items", key_as="key", value_as="amount"),
+ )
+ .build()
+ )
+
+ def test_map_fields_cannot_be_config_fields(self) -> None:
+ with pytest.raises(WorkflowDefinitionError, match="cannot also be config_fields"):
+ (
+ Workflow.builder("bad")
+ .add_task(
+ MappedOnly,
+ config_fields=["amount"],
+ mapped_over=TaskMap(over="items", key_as="item_name", value_as="amount"),
+ )
+ .build()
+ )
+
+ def test_map_fields_cannot_be_dependencies(self) -> None:
+ with pytest.raises(WorkflowDefinitionError, match="cannot also be dependencies"):
+ (
+ Workflow.builder("bad")
+ .add_task(AddOne)
+ .add_task(
+ MappedNumber,
+ depends_on={"amount": AddOne, "base": AddOne},
+ config_fields=["multiplier"],
+ mapped_over=TaskMap(over="items", key_as="item_name", value_as="amount"),
+ )
+ .build()
+ )
+
+ def test_all_required_fields_must_be_covered(self) -> None:
+ with pytest.raises(WorkflowDefinitionError, match=r"multiplier.*not covered"):
+ (
+ Workflow.builder("bad")
+ .add_task(AddOne)
+ .add_task(
+ MappedNumber,
+ depends_on={"base": AddOne},
+ mapped_over=TaskMap(over="items", key_as="item_name", value_as="amount"),
+ )
+ .build()
+ )
+
+ def test_mapped_task_requires_named_dependencies(self) -> None:
+ with pytest.raises(WorkflowDefinitionError, match="requires named field dependencies"):
+ (
+ Workflow.builder("bad")
+ .add_task(AddOne)
+ .add_task(
+ MappedNumber,
+ depends_on=AddOne,
+ config_fields=["multiplier"],
+ mapped_over=TaskMap(over="items", key_as="item_name", value_as="amount"),
+ )
+ .build()
+ )
+
+ def test_mapped_output_type_is_checked_downstream(self) -> None:
+ class BadAggregateInput(BaseModel):
+ values: dict[str, StringOutput]
+
+ class BadAggregate(Task[BadAggregateInput, NumberOutput]):
+ def run(self, input: BadAggregateInput, ctx: ExecutionContext) -> NumberOutput:
+ return NumberOutput(value=0)
+
+ with pytest.raises(WorkflowDefinitionError, match="Fan-in type mismatch"):
+ (
+ Workflow.builder("bad")
+ .add_task(
+ MappedOnly,
+ mapped_over=TaskMap(over="items", key_as="item_name", value_as="amount"),
+ )
+ .add_task(BadAggregate, depends_on={"values": (MappedOnly, "root")})
+ .build()
+ )
+
+ def test_can_route_root_dictionary_from_mapped_output(self) -> None:
+ workflow = (
+ Workflow.builder("mapped_root_route")
+ .add_task(
+ MappedOnly,
+ mapped_over=TaskMap(over="items", key_as="item_name", value_as="amount"),
+ )
+ .add_task(SumAggregate, depends_on={"values": (MappedOnly, "root")})
+ .build()
+ )
+ assert workflow.result_task is SumAggregate
+
+ def test_mapped_upstream_with_single_dependency_and_config_is_rejected(self) -> None:
+ with pytest.raises(WorkflowDefinitionError, match="must be connected through"):
+ (
+ Workflow.builder("bad", result_task=MergeTask)
+ .add_task(
+ MappedOnly,
+ mapped_over=TaskMap(over="items", key_as="item_name", value_as="amount"),
+ )
+ .add_task(MergeTask, depends_on=MappedOnly, config_fields=["label"])
+ .build()
+ )
+
+ def test_unknown_mapped_output_field_is_rejected(self) -> None:
+ with pytest.raises(WorkflowDefinitionError, match="Field 'value' not found"):
+ (
+ Workflow.builder("bad")
+ .add_task(
+ MappedOnly,
+ mapped_over=TaskMap(over="items", key_as="item_name", value_as="amount"),
+ )
+ .add_task(
+ SumAggregate,
+ depends_on={"values": (MappedOnly, "value")},
+ )
+ .build()
+ )
+
+ def test_collection_can_route_root_from_mapped_output(self) -> None:
+ class NestedAggregateInput(BaseModel):
+ values: list[dict[str, NumberOutput]]
+
+ class NestedAggregate(Task[NestedAggregateInput, NumberOutput]):
+ def run(self, input: NestedAggregateInput, ctx: ExecutionContext) -> NumberOutput:
+ return NumberOutput(value=0)
+
+ workflow = (
+ Workflow.builder("mapped_collection_route")
+ .add_task(
+ MappedOnly,
+ mapped_over=TaskMap(over="items", key_as="item_name", value_as="amount"),
+ )
+ .add_task(
+ NestedAggregate,
+ depends_on={"values": collect((MappedOnly, "root"))},
+ )
+ .build()
+ )
+ assert workflow.result_task is NestedAggregate
+
+ def test_mapped_result_workflow_can_be_wrapped(self) -> None:
+ workflow = (
+ Workflow.builder("mapped_result")
+ .add_task(
+ MappedOnly,
+ mapped_over=TaskMap(over="items", key_as="item_name", value_as="amount"),
+ )
+ .build()
+ )
+
+ Wrapped = workflow_task(
+ workflow,
+ job_configuration=JobConfiguration({"mapped_only": {"items": {"one": 1}}}),
+ )
+ outer = Workflow("outer", [Wrapped])
+
+ result = Runner().run(Job(outer, EmptyConfig()))
+
+ assert result.status == JobStatus.COMPLETED
+ assert result.result == MappedOutput[NumberOutput](root={"one": NumberOutput(value=1)})
+
+
+class TestMappedJobValidation:
+ def test_job_configuration_is_required(self) -> None:
+ workflow = _mapped_workflow()
+ with pytest.raises(WorkflowDefinitionError, match="requires JobConfiguration"):
+ Job(workflow, NumberInput(value=1))
+
+ def test_map_source_is_required(self) -> None:
+ workflow = _mapped_workflow()
+ config = JobConfiguration({"mapped_number": {"multiplier": 2}})
+ with pytest.raises(WorkflowDefinitionError, match="requires configuration field 'items'"):
+ Job(workflow, NumberInput(value=1), job_configuration=config)
+
+ def test_map_source_must_be_mapping(self) -> None:
+ workflow = _mapped_workflow()
+ config = JobConfiguration({"mapped_number": {"items": [1, 2], "multiplier": 2}})
+ with pytest.raises(WorkflowDefinitionError, match="must be a mapping"):
+ Job(workflow, NumberInput(value=1), job_configuration=config)
+
+ def test_map_keys_must_be_strings(self) -> None:
+ workflow = _mapped_workflow()
+ with pytest.raises(WorkflowDefinitionError, match="must be strings"):
+ _mapped_job(workflow, {1: 2})
+
+ def test_map_values_are_validated(self) -> None:
+ workflow = _mapped_workflow()
+ with pytest.raises(WorkflowDefinitionError, match="Invalid mapping item 'bad'"):
+ _mapped_job(workflow, {"bad": "not-an-int"})
+
+ def test_map_values_preserve_model_config_and_nested_models(self) -> None:
+ class Resource:
+ pass
+
+ class ResourceInput(BaseModel):
+ model_config = ConfigDict(arbitrary_types_allowed=True)
+ key: str
+ value: Resource | NumberOutput
+
+ seen: list[Resource | NumberOutput] = []
+
+ class ResourceTask(Task[ResourceInput, NumberOutput]):
+ def run(self, input: ResourceInput, ctx: ExecutionContext) -> NumberOutput:
+ seen.append(input.value)
+ return NumberOutput(value=1)
+
+ resource = Resource()
+ workflow = (
+ Workflow.builder("resources")
+ .add_task(ResourceTask, mapped_over=TaskMap("items", "key", "value"))
+ .build()
+ )
+ job = Job(
+ workflow,
+ EmptyConfig(),
+ job_configuration=JobConfiguration(
+ {"ResourceTask": {"items": {"object": resource, "model": {"value": 3}}}}
+ ),
+ )
+
+ assert Runner().run(job).status == JobStatus.COMPLETED
+ assert seen == [resource, NumberOutput(value=3)]
+
+ def test_map_values_preserve_field_constraints(self) -> None:
+ class PositiveInput(BaseModel):
+ key: str
+ value: int = Field(gt=0)
+
+ class PositiveTask(Task[PositiveInput, NumberOutput]):
+ def run(self, input: PositiveInput, ctx: ExecutionContext) -> NumberOutput:
+ return NumberOutput(value=input.value)
+
+ workflow = (
+ Workflow.builder("positive")
+ .add_task(PositiveTask, mapped_over=TaskMap("items", "key", "value"))
+ .build()
+ )
+ with pytest.raises(WorkflowDefinitionError, match="Invalid mapping item 'bad'"):
+ Job(
+ workflow,
+ EmptyConfig(),
+ job_configuration=JobConfiguration({"PositiveTask": {"items": {"bad": -1}}}),
+ )
+
+
+class TestMappedExecution:
+ def setup_method(self) -> None:
+ MappedNumber.seen = []
+
+ def test_only_map_source_is_stripped_from_config_values(self) -> None:
+ """Mapped items receive every configured value except the map source."""
+ seen: list[dict[str, Any]] = []
+
+ class OpenInput(BaseModel):
+ model_config = ConfigDict(extra="allow")
+ item_name: str
+ amount: int
+
+ class OpenMapped(Task[OpenInput, NumberOutput]):
+ name = "open_mapped"
+
+ def run(self, input: OpenInput, ctx: ExecutionContext) -> NumberOutput:
+ seen.append(input.model_extra or {})
+ return NumberOutput(value=input.amount)
+
+ workflow = (
+ Workflow.builder("open")
+ .add_task(
+ OpenMapped,
+ mapped_over=TaskMap(over="items", key_as="item_name", value_as="amount"),
+ )
+ .build()
+ )
+ job = Job(
+ workflow,
+ EmptyConfig(),
+ job_configuration=JobConfiguration(
+ {"open_mapped": {"items": {"only": 1}, "passthrough": "yes"}}
+ ),
+ )
+
+ result = Runner().run(job)
+
+ assert result.status == JobStatus.COMPLETED
+ assert seen == [{"passthrough": "yes"}]
+
+ def test_executes_sequentially_and_aggregates_output(self) -> None:
+ workflow = _mapped_workflow()
+ job = _mapped_job(workflow, {"first": 1, "second": 2, "third": 3})
+ hook = RecordingMapHook()
+
+ result = Runner(hooks=[hook]).run(job)
+
+ assert result.status == JobStatus.COMPLETED
+ assert result.result == NumberOutput(value=45)
+ mapped_result = next(r for r in result.task_results if r.task_name == "mapped_number")
+ assert isinstance(mapped_result.output, MappedOutput)
+ assert list(mapped_result.output.root) == ["first", "second", "third"]
+ assert [name for _instance, name, _ctx in MappedNumber.seen] == [
+ "first",
+ "second",
+ "third",
+ ]
+ assert len({instance for instance, _name, _ctx in MappedNumber.seen}) == 3
+ assert hook.events == [
+ "start:mapped_number[first]",
+ "complete:mapped_number[first]",
+ "start:mapped_number[second]",
+ "complete:mapped_number[second]",
+ "start:mapped_number[third]",
+ "complete:mapped_number[third]",
+ ]
+ assert [r.task_name for r in result.mapped_item_results["mapped_number"]] == [
+ "mapped_number[first]",
+ "mapped_number[second]",
+ "mapped_number[third]",
+ ]
+ correlation_ids = [ctx_id for _instance, _name, ctx_id in MappedNumber.seen]
+ assert len(set(correlation_ids)) == 3
+ assert all(
+ ctx_id.startswith(job.task_results[0].task_name) is False for ctx_id in correlation_ids
+ )
+
+ def test_routed_and_collection_dependencies_are_shared_by_items(self) -> None:
+ routed_workflow = (
+ Workflow.builder("routed_map")
+ .add_task(ProduceEnvelope)
+ .add_task(
+ MappedNumber,
+ depends_on={"base": (ProduceEnvelope, "number")},
+ config_fields=["multiplier"],
+ mapped_over=TaskMap(over="items", key_as="item_name", value_as="amount"),
+ )
+ .build()
+ )
+ routed_job = Job(
+ routed_workflow,
+ NumberInput(value=5),
+ job_configuration=JobConfiguration(
+ {"mapped_number": {"items": {"one": 2}, "multiplier": 3}}
+ ),
+ )
+ assert Runner().run(routed_job).result == MappedOutput[NumberOutput](
+ root={"one": NumberOutput(value=11)}
+ )
+
+ collection_workflow = (
+ Workflow.builder("collection_map")
+ .add_task(AddOne, name="first")
+ .add_task(AddOne, name="second")
+ .add_task(
+ MappedCollection,
+ depends_on={"bases": collect("first", "second")},
+ mapped_over=TaskMap(over="items", key_as="item_name", value_as="amount"),
+ )
+ .build()
+ )
+ collection_job = Job(
+ collection_workflow,
+ NumberInput(value=4),
+ job_configuration=JobConfiguration({"mapped_collection": {"items": {"one": 1}}}),
+ )
+ assert Runner().run(collection_job).result == MappedOutput[NumberOutput](
+ root={"one": NumberOutput(value=11)}
+ )
+
+ def test_empty_mapping_produces_empty_dictionary(self) -> None:
+ workflow = (
+ Workflow.builder("empty_map")
+ .add_task(
+ MappedOnly,
+ mapped_over=TaskMap(over="items", key_as="item_name", value_as="amount"),
+ )
+ .build()
+ )
+ job = Job(
+ workflow,
+ EmptyConfig(),
+ job_configuration=JobConfiguration({"mapped_only": {"items": {}}}),
+ )
+
+ result = Runner().run(job)
+
+ assert result.status == JobStatus.COMPLETED
+ assert result.result == MappedOutput[NumberOutput](root={})
+ assert result.mapped_item_results["mapped_only"] == []
+
+ def test_fail_fast_stops_after_first_failure(self) -> None:
+ workflow = _mapped_workflow()
+ job = _mapped_job(workflow, {"good": 1, "bad": -1, "later": 3})
+ hook = RecordingMapHook()
+
+ result = Runner(hooks=[hook]).run(job)
+
+ assert result.status == JobStatus.FAILED
+ assert result.failed_task == "mapped_number"
+ assert "bad" in (result.error or "")
+ assert [r.task_name for r in result.mapped_item_results["mapped_number"]] == [
+ "mapped_number[good]",
+ "mapped_number[bad]",
+ ]
+ assert "start:mapped_number[later]" not in hook.events
+
+ def test_collect_all_records_every_failure(self) -> None:
+ workflow = _mapped_workflow(error_mode="collect_all")
+ job = _mapped_job(workflow, {"bad_one": -1, "good": 2, "bad_two": -2})
+
+ result = Runner().run(job)
+
+ assert result.status == JobStatus.FAILED
+ assert "bad_one" in (result.error or "")
+ assert "bad_two" in (result.error or "")
+ assert len(result.mapped_item_results["mapped_number"]) == 3
+
+ @pytest.mark.skipif(not hasattr(signal, "SIGALRM"), reason="SIGALRM unavailable")
+ @pytest.mark.parametrize("job_timeout", [True, False])
+ def test_collect_all_stops_only_for_job_timeout(self, job_timeout: bool) -> None:
+ seen: list[str] = []
+
+ class AlarmTask(Task[MappedOnlyInput, NumberOutput]):
+ timeout_seconds = None if job_timeout else 60
+
+ def run(self, input: MappedOnlyInput, ctx: ExecutionContext) -> NumberOutput:
+ seen.append(input.item_name)
+ if input.item_name == "first":
+ # Exercise the installed handler without waiting for a real deadline.
+ signal.raise_signal(signal.SIGALRM)
+ return NumberOutput(value=input.amount)
+
+ workflow = (
+ Workflow.builder("timeout")
+ .add_task(
+ AlarmTask,
+ mapped_over=TaskMap("items", "item_name", "amount", "collect_all"),
+ )
+ .build()
+ )
+ job = Job(
+ workflow,
+ EmptyConfig(),
+ job_configuration=JobConfiguration(
+ {"AlarmTask": {"items": {"first": 1, "second": 2}}}
+ ),
+ )
+ previous_handler = signal.getsignal(signal.SIGALRM)
+ try:
+ result = Runner().run(job, timeout_seconds=60 if job_timeout else None)
+ finally:
+ signal.alarm(0)
+ signal.signal(signal.SIGALRM, previous_handler)
+
+ assert result.status == JobStatus.FAILED
+ assert "timed out" in (result.error or "")
+ assert seen == (["first"] if job_timeout else ["first", "second"])
+ assert len(result.mapped_item_results["AlarmTask"]) == len(seen)
+
+ def test_item_input_validation_is_recorded_as_item_failure(self) -> None:
+ workflow = _mapped_workflow()
+ job = Job(
+ workflow,
+ NumberInput(value=1),
+ job_configuration=JobConfiguration({"mapped_number": {"items": {"one": 1}}}),
+ )
+ hook = RecordingMapHook()
+
+ result = Runner(hooks=[hook]).run(job)
+
+ assert result.status == JobStatus.FAILED
+ assert result.mapped_item_results["mapped_number"][0].status.value == "failed"
+ assert hook.events == [
+ "start:mapped_number[one]",
+ "fail:mapped_number[one]",
+ ]
+
+ def test_wrong_item_output_fails_mapped_task(self) -> None:
+ workflow = (
+ Workflow.builder("wrong_output")
+ .add_task(
+ MappedWrongOutput,
+ mapped_over=TaskMap(over="items", key_as="item_name", value_as="amount"),
+ )
+ .build()
+ )
+ job = Job(
+ workflow,
+ EmptyConfig(),
+ job_configuration=JobConfiguration({"mapped_wrong_output": {"items": {"one": 1}}}),
+ )
+
+ result = Runner().run(job)
+
+ assert result.status == JobStatus.FAILED
+ assert "expected NumberOutput" in (result.error or "")
+
+ def test_item_timeout_setup_and_cleanup(self, monkeypatch: pytest.MonkeyPatch) -> None:
+ workflow = (
+ Workflow.builder("mapped_timeout")
+ .add_task(
+ MappedSlow,
+ mapped_over=TaskMap(over="items", key_as="item_name", value_as="amount"),
+ )
+ .build()
+ )
+ job = Job(
+ workflow,
+ EmptyConfig(),
+ job_configuration=JobConfiguration({"mapped_slow": {"items": {"one": 1}}}),
+ )
+ calls: list[tuple[float, str]] = []
+
+ def fake_alarm(seconds: float, label: str, **_kwargs: object) -> bool:
+ calls.append((seconds, label))
+ return True
+
+ monkeypatch.setattr(Runner, "_set_alarm", staticmethod(fake_alarm))
+ result = Runner().run(job)
+
+ assert result.status == JobStatus.COMPLETED
+ assert calls == [(10, "mapped_slow[one]")]
+
+ def test_child_context_shares_services_and_has_safe_unique_paths(self, tmp_path: Path) -> None:
+ parent = ExecutionContext(correlation_id="parent", scratch_dir=tmp_path)
+ service = object()
+ parent.register("service", service)
+
+ first = parent.child(task_name="load surfaces", item_key="a/b")
+ second = parent.child(task_name="load surfaces", item_key="a_b")
+
+ assert first.parent_correlation_id == "parent"
+ assert first.resolve("service") is service
+ assert first.logger is parent.logger
+ assert first.scratch_dir != second.scratch_dir
+ assert first.correlation_id.startswith("parent:load_surfaces_a_b:")
+
+ def test_mapped_exception_retains_errors(self) -> None:
+ error = ValueError("bad")
+ exc = MappedTaskExecutionError("mapped", {"item": error})
+ assert exc.errors == {"item": error}
+ assert str(exc) == "Mapped task 'mapped' failed: item: bad"
+
+
+class TestMappedHooks:
+ def test_base_hook_handles_mapped_completion_and_failure(self) -> None:
+ success = _mapped_job(_mapped_workflow(), {"good": 1})
+ failure = _mapped_job(_mapped_workflow(), {"bad": -1})
+
+ assert Runner(hooks=[BaseHook()]).run(success).status == JobStatus.COMPLETED
+ assert Runner(hooks=[BaseHook()]).run(failure).status == JobStatus.FAILED
+
+ def test_builtin_hooks_record_and_persist_items(
+ self, tmp_path: Path, caplog: pytest.LogCaptureFixture
+ ) -> None:
+ workflow = _mapped_workflow()
+ job = _mapped_job(workflow, {"one/unsafe": 1})
+ timing = TimingHook()
+ persistence = ResultPersistenceHook(tmp_path)
+
+ with caplog.at_level("INFO", logger="taskmaestro.hooks.logging"):
+ result = Runner(hooks=[LoggingHook(), timing, persistence]).run(job)
+
+ assert result.status == JobStatus.COMPLETED
+ assert "one/unsafe" in timing.mapped_item_timings["mapped_number"]
+ assert (tmp_path / "mapped_number[one%2Funsafe].json").exists()
+ assert (tmp_path / "mapped_number.json").exists()
+ assert any("Map item started: mapped_number[one/unsafe]" in m for m in caplog.messages)
+ assert any("Map item completed: mapped_number[one/unsafe]" in m for m in caplog.messages)
+
+ def test_persisted_items_do_not_collide_even_when_parent_fails(self, tmp_path: Path) -> None:
+ job = _mapped_job(
+ _mapped_workflow(), {"a/b": 1, "a\\b": 2, "a_b": 3, "a%2Fb": 4, "bad": -1}
+ )
+
+ result = Runner(hooks=[ResultPersistenceHook(tmp_path)]).run(job)
+
+ assert result.status == JobStatus.FAILED
+ assert not (tmp_path / "mapped_number.json").exists()
+ for filename_key, value in [("a%2Fb", 13), ("a%5Cb", 15), ("a_b", 17), ("a%252Fb", 19)]:
+ path = tmp_path / f"mapped_number[{filename_key}].json"
+ assert NumberOutput.model_validate_json(path.read_text()) == NumberOutput(value=value)
+
+ def test_builtin_hooks_record_item_failure(self, caplog: pytest.LogCaptureFixture) -> None:
+ workflow = _mapped_workflow()
+ job = _mapped_job(workflow, {"bad": -1})
+ timing = TimingHook()
+
+ with caplog.at_level("INFO", logger="taskmaestro.hooks.logging"):
+ Runner(hooks=[LoggingHook(), timing]).run(job)
+
+ assert "bad" in timing.mapped_item_timings["mapped_number"]
+ assert any("Map item failed: mapped_number[bad]" in m for m in caplog.messages)
+
+
+class TestMappedYamlAndVisualization:
+ def test_yaml_mapped_task_end_to_end(self, tmp_path: Path) -> None:
+ module = "tests.test_mapping"
+ workflow_path = tmp_path / "workflow.yaml"
+ workflow_path.write_text(
+ f"""\
+workflow:
+ name: yaml_map
+ tasks:
+ - task: {module}.MappedOnly
+ map:
+ over: items
+ key_as: item_name
+ value_as: amount
+ error_mode: fail_fast
+"""
+ )
+ input_path = tmp_path / "input.yaml"
+ input_path.write_text(
+ """\
+mapped_only:
+ items:
+ first: 1
+ second: 2
+"""
+ )
+
+ loaded = load_workflow_from_yaml(workflow_path, input_path)
+ result = loaded.run()
+
+ assert loaded.workflow.get_config_fields("mapped_only") == set()
+ assert result.status == JobStatus.COMPLETED
+ assert result.result == MappedOutput[NumberOutput](
+ root={
+ "first": NumberOutput(value=1),
+ "second": NumberOutput(value=2),
+ }
+ )
+
+ def test_yaml_rejects_invalid_error_mode(self, tmp_path: Path) -> None:
+ module = "tests.test_mapping"
+ workflow_path = tmp_path / "workflow.yaml"
+ workflow_path.write_text(
+ f"""\
+workflow:
+ name: bad_yaml_map
+ tasks:
+ - task: {module}.MappedOnly
+ map:
+ over: items
+ key_as: item_name
+ value_as: amount
+ error_mode: invalid
+"""
+ )
+ input_path = tmp_path / "input.yaml"
+ input_path.write_text("mapped_only: {items: {}}\n")
+
+ with pytest.raises(ConfigLoadError, match="YAML schema validation error"):
+ load_workflow_from_yaml(workflow_path, input_path)
+
+ def test_mermaid_marks_mapped_node_and_output_type(self) -> None:
+ workflow = _mapped_workflow()
+
+ diagram = workflow.to_mermaid(
+ job_configuration=JobConfiguration(
+ {"mapped_number": {"items": {"one": 1}, "multiplier": 2}}
+ )
+ )
+
+ assert 'mapped_number["mapped_number
map over: items"]' in diagram
+ assert "items, multiplier" in diagram
+ assert ".root: dict‹str, NumberOutput›" in diagram
diff --git a/tests/test_runner.py b/tests/test_runner.py
index 10838c1..01e1bb3 100644
--- a/tests/test_runner.py
+++ b/tests/test_runner.py
@@ -3,7 +3,7 @@
from __future__ import annotations
import pytest
-from pydantic import BaseModel
+from pydantic import BaseModel, ConfigDict, ValidationError
from taskmaestro import (
EmptyConfig,
@@ -83,6 +83,9 @@ def test_task_failure(self, ctx: ExecutionContext) -> None:
assert result.failed_task == "failing_task"
assert result.error is not None
assert "intentionally" in result.error
+ # The original exception object is retained alongside its string form.
+ assert isinstance(result.exception, ValueError)
+ assert str(result.exception) == result.error
def test_output_type_mismatch(self, ctx: ExecutionContext) -> None:
wf = Workflow(name="test", tasks=[WrongOutputTask])
@@ -261,6 +264,194 @@ def run(self, input: NumberInput, ctx: ExecutionContext) -> NumberOutput:
assert result.status == JobStatus.FAILED
assert "timed out" in (result.error or "")
+ @pytest.mark.skipif(
+ not hasattr(__import__("signal"), "SIGALRM"),
+ reason="signal.SIGALRM not available on this platform",
+ )
+ def test_job_timeout_survives_task_with_own_timeout(self, ctx: ExecutionContext) -> None:
+ """A task's own alarm must not cancel the job deadline for later tasks."""
+ import time
+
+ class QuickWithTimeout(Task[NumberInput, NumberOutput]):
+ name = "quick"
+ timeout_seconds = 30
+
+ def run(self, input: NumberInput, ctx: ExecutionContext) -> NumberOutput:
+ return NumberOutput(value=input.value)
+
+ class SlowNoTimeout(Task[NumberOutput, NumberOutput]):
+ name = "slow_no_timeout"
+
+ def run(self, input: NumberOutput, ctx: ExecutionContext) -> NumberOutput:
+ time.sleep(5)
+ return input
+
+ wf = Workflow(name="test", tasks=[QuickWithTimeout, SlowNoTimeout])
+ job = Job(workflow=wf, config=NumberInput(value=1))
+ start = time.monotonic()
+ result = Runner().run(job, ctx=ctx, timeout_seconds=0.5)
+ assert time.monotonic() - start < 3
+ assert result.status == JobStatus.FAILED
+ assert result.failed_task == "slow_no_timeout"
+ assert "Job timed out after 0.5s" in (result.error or "")
+
+ @pytest.mark.skipif(
+ not hasattr(__import__("signal"), "SIGALRM"),
+ reason="signal.SIGALRM not available on this platform",
+ )
+ def test_job_timeout_survives_nested_workflow_task(self, ctx: ExecutionContext) -> None:
+ """An inner workflow's runner must not cancel the outer job deadline."""
+ import time
+
+ class InnerWithTimeout(Task[NumberInput, NumberOutput]):
+ name = "inner_with_timeout"
+ timeout_seconds = 30
+
+ def run(self, input: NumberInput, ctx: ExecutionContext) -> NumberOutput:
+ return NumberOutput(value=input.value)
+
+ class SlowNoTimeout(Task[NumberOutput, NumberOutput]):
+ name = "slow_no_timeout"
+
+ def run(self, input: NumberOutput, ctx: ExecutionContext) -> NumberOutput:
+ time.sleep(5)
+ return input
+
+ inner = Workflow(name="inner", tasks=[InnerWithTimeout]).as_task(name="inner")
+ wf = Workflow(name="outer", tasks=[inner, SlowNoTimeout])
+ job = Job(workflow=wf, config=NumberInput(value=1))
+ start = time.monotonic()
+ result = Runner().run(job, ctx=ctx, timeout_seconds=0.5)
+ assert time.monotonic() - start < 3
+ assert result.status == JobStatus.FAILED
+ assert result.failed_task == "slow_no_timeout"
+
+ @pytest.mark.skipif(
+ not hasattr(__import__("signal"), "SIGALRM"),
+ reason="signal.SIGALRM not available on this platform",
+ )
+ def test_expired_job_deadline_fails_next_task_immediately(self, ctx: ExecutionContext) -> None:
+ """If the deadline passes during a task, the following task is not started."""
+ import time
+
+ ran: list[str] = []
+
+ class Sleeper(Task[NumberInput, NumberOutput]):
+ name = "sleeper"
+
+ def run(self, input: NumberInput, ctx: ExecutionContext) -> NumberOutput:
+ ran.append(self.name)
+ time.sleep(0.3)
+ return NumberOutput(value=input.value)
+
+ class Never(Task[NumberOutput, NumberOutput]):
+ name = "never"
+
+ def run(self, input: NumberOutput, ctx: ExecutionContext) -> NumberOutput:
+ ran.append(self.name)
+ return input
+
+ wf = Workflow(name="test", tasks=[Sleeper, Never])
+ job = Job(workflow=wf, config=NumberInput(value=1))
+ # Deadline expires while Sleeper is running; Sleeper itself is only
+ # interrupted by the alarm, but Never must not run at all.
+ result = Runner().run(job, ctx=ctx, timeout_seconds=0.2)
+ assert result.status == JobStatus.FAILED
+ assert ran == ["sleeper"]
+
+ @pytest.mark.skipif(
+ not hasattr(__import__("signal"), "SIGALRM"),
+ reason="signal.SIGALRM not available on this platform",
+ )
+ def test_deadline_already_expired_before_next_task(self, ctx: ExecutionContext) -> None:
+ """A task that swallows the alarm and overruns the deadline still stops the job."""
+ import time
+ from contextlib import suppress
+
+ from taskmaestro.exceptions import TaskTimeoutError
+
+ ran: list[str] = []
+
+ class SwallowsAlarm(Task[NumberInput, NumberOutput]):
+ name = "swallows_alarm"
+
+ def run(self, input: NumberInput, ctx: ExecutionContext) -> NumberOutput:
+ ran.append(self.name)
+ # A misbehaving task that ignores the job deadline.
+ with suppress(TaskTimeoutError):
+ time.sleep(0.6)
+ return NumberOutput(value=input.value)
+
+ class Never(Task[NumberOutput, NumberOutput]):
+ name = "never"
+
+ def run(self, input: NumberOutput, ctx: ExecutionContext) -> NumberOutput:
+ ran.append(self.name)
+ return input
+
+ wf = Workflow(name="test", tasks=[SwallowsAlarm, Never])
+ job = Job(workflow=wf, config=NumberInput(value=1))
+ result = Runner().run(job, ctx=ctx, timeout_seconds=0.2)
+
+ assert ran == ["swallows_alarm"]
+ assert result.status == JobStatus.FAILED
+ assert result.failed_task == "never"
+ assert result.error == "Job timed out after 0.2s"
+ assert [r.status for r in result.task_results] == [
+ TaskStatus.COMPLETED,
+ TaskStatus.FAILED,
+ ]
+
+ def test_arm_raises_when_deadline_already_passed(self) -> None:
+ import time
+
+ from taskmaestro.exceptions import TaskTimeoutError
+ from taskmaestro.runner import _Deadline
+
+ deadline = _Deadline(job_timeout=1.0, job_deadline=time.monotonic() - 1)
+ with pytest.raises(TaskTimeoutError, match=r"Job timed out after 1\.0s"):
+ Runner()._arm(None, "task", deadline)
+
+ @pytest.mark.skipif(
+ not hasattr(__import__("signal"), "SIGALRM"),
+ reason="signal.SIGALRM not available on this platform",
+ )
+ def test_sub_second_timeout_is_not_truncated(self, ctx: ExecutionContext) -> None:
+ """timeout_seconds=1.9 must allow a 1.4s task to finish (was truncated to 1s)."""
+ import time
+
+ class MidTask(Task[NumberInput, NumberOutput]):
+ name = "mid"
+ timeout_seconds = 1.9
+
+ def run(self, input: NumberInput, ctx: ExecutionContext) -> NumberOutput:
+ time.sleep(1.4)
+ return NumberOutput(value=input.value)
+
+ wf = Workflow(name="test", tasks=[MidTask])
+ job = Job(workflow=wf, config=NumberInput(value=1))
+ result = Runner().run(job, ctx=ctx)
+ assert result.status == JobStatus.COMPLETED
+
+ @pytest.mark.skipif(
+ not hasattr(__import__("signal"), "SIGALRM"),
+ reason="signal.SIGALRM not available on this platform",
+ )
+ def test_previous_sigalrm_handler_restored(self, ctx: ExecutionContext) -> None:
+ import signal
+
+ def sentinel(signum: int, frame: object) -> None: # pragma: no cover
+ pass
+
+ previous = signal.signal(signal.SIGALRM, sentinel)
+ try:
+ wf = Workflow(name="test", tasks=[AddOne])
+ job = Job(workflow=wf, config=NumberInput(value=1))
+ Runner().run(job, ctx=ctx, timeout_seconds=60)
+ assert signal.getsignal(signal.SIGALRM) is sentinel
+ finally:
+ signal.signal(signal.SIGALRM, previous)
+
class TestAlarmUnavailable:
def test_alarm_unavailable_warns(self, ctx: ExecutionContext) -> None:
@@ -276,6 +467,53 @@ def test_alarm_unavailable_warns(self, ctx: ExecutionContext) -> None:
result = Runner().run(job, ctx=ctx, timeout_seconds=60)
assert result.status == JobStatus.COMPLETED
+ @pytest.mark.skipif(
+ not hasattr(__import__("signal"), "SIGALRM"),
+ reason="signal.SIGALRM not available on this platform",
+ )
+ def test_timeouts_in_non_main_thread_warn_and_continue(self) -> None:
+ """signal.signal() raises ValueError off the main thread; the job must still finish."""
+ import threading
+ import warnings
+
+ outcome: dict[str, object] = {}
+
+ def worker() -> None:
+ wf = Workflow(name="test", tasks=[AddOne, Double])
+ job = Job(workflow=wf, config=NumberInput(value=1))
+ with warnings.catch_warnings(record=True) as caught:
+ warnings.simplefilter("always")
+ try:
+ result = Runner().run(job, timeout_seconds=60)
+ except Exception as exc: # pragma: no cover - the bug under test
+ outcome["exc"] = exc
+ return
+ outcome["status"] = result.status
+ outcome["warnings"] = [str(w.message) for w in caught]
+
+ thread = threading.Thread(target=worker)
+ thread.start()
+ thread.join()
+
+ assert "exc" not in outcome, outcome.get("exc")
+ assert outcome["status"] == JobStatus.COMPLETED
+ messages = outcome["warnings"]
+ assert isinstance(messages, list)
+ assert len(messages) == 1 # warned once per run, not per task
+ assert "signal.alarm not available" in messages[0]
+
+ def test_arming_failure_marks_task_failed_not_running(self, ctx: ExecutionContext) -> None:
+ """An unexpected error while arming the timer is recorded as a task failure."""
+ from unittest.mock import patch
+
+ wf = Workflow(name="test", tasks=[SlowTask])
+ job = Job(workflow=wf, config=NumberInput(value=1))
+ with patch.object(Runner, "_arm", side_effect=RuntimeError("boom")):
+ result = Runner().run(job, ctx=ctx)
+ assert result.status == JobStatus.FAILED
+ assert result.failed_task == "slow_task"
+ assert result.error == "boom"
+
class TestContextIntegration:
def test_context_auto_created(self) -> None:
@@ -410,6 +648,26 @@ def test_fan_in_merge(self, ctx: ExecutionContext) -> None:
# AddOne: 3+1=4, FanInWithConfig: "hello:4"
assert result.result.combined == "hello:4" # type: ignore[union-attr]
+ def test_extra_config_values_reach_input_model(self, ctx: ExecutionContext) -> None:
+ """Configured values are passed through to the input model even when they
+ are not listed in config_fields, so the model decides how to treat them."""
+
+ class StrictInput(BaseModel):
+ model_config = ConfigDict(extra="forbid")
+ path: str
+
+ class StrictTask(Task[StrictInput, NumberOutput]):
+ name = "strict_task"
+
+ def run(self, input: StrictInput, ctx: ExecutionContext) -> NumberOutput:
+ return NumberOutput(value=len(input.path))
+
+ wf = Workflow.builder("strict").add_task(StrictTask, config_fields=["path"]).build()
+ jc = JobConfiguration({"strict_task": {"path": "/data", "unexpected": 1}})
+ job = Job(wf, EmptyConfig(), job_configuration=jc)
+ with pytest.raises(ValidationError, match="unexpected"):
+ Runner().run(job, ctx=ctx)
+
def test_backward_compat_no_config(self, ctx: ExecutionContext) -> None:
"""Workflow without config_fields runs normally."""
wf = Workflow(name="compat", tasks=[AddOne, Double])
diff --git a/tests/test_workflow.py b/tests/test_workflow.py
index 58f2c69..01d2d38 100644
--- a/tests/test_workflow.py
+++ b/tests/test_workflow.py
@@ -49,10 +49,73 @@ def test_single_task_workflow(self) -> None:
assert wf.result_task is AddOne
assert wf.topological_order() == [("add_one", AddOne)]
+ def test_explicit_result_task_overrides_last(self) -> None:
+ wf = Workflow(name="test", tasks=[AddOne, Double], result_task=AddOne)
+ assert wf.result_task is AddOne
+ assert wf.result_task_name == "add_one"
+
def test_empty_workflow(self) -> None:
wf = Workflow(name="empty")
assert wf._tasks == {}
+ def test_result_task_without_tasks_is_recorded(self) -> None:
+ """tasks=None with result_task keeps the name; validation happens on use."""
+ wf = Workflow(name="empty", result_task=AddOne)
+ assert wf.result_task_name == "add_one"
+
+ def test_empty_task_list_is_rejected(self) -> None:
+ """An explicit empty list is a definition error, unlike tasks=None."""
+ with pytest.raises(WorkflowDefinitionError, match="empty task list"):
+ Workflow(name="empty", tasks=[])
+
+ def test_duplicate_names_raise_not_cycle(self) -> None:
+ """Linear shorthand rejects duplicates instead of reporting a bogus cycle."""
+ with pytest.raises(WorkflowDefinitionError, match="Duplicate task name 'add_one'"):
+ Workflow(name="dup", tasks=[AddOne, AddOne])
+
+ def test_unregistered_result_task_raises(self) -> None:
+ with pytest.raises(
+ WorkflowDefinitionError, match=r"result_task 'double' is not registered"
+ ):
+ Workflow(name="test", tasks=[AddOne], result_task=Double)
+
+ def test_subclass_output_is_accepted_on_single_edge(self) -> None:
+ """Single-dep edges use type compatibility, not identity."""
+
+ class BaseOut(BaseModel):
+ value: int
+
+ class RichOut(BaseOut):
+ extra: str = ""
+
+ class Producer(Task[NumberInput, RichOut]):
+ name = "producer"
+
+ def run(self, input: NumberInput, ctx: ExecutionContext) -> RichOut:
+ return RichOut(value=input.value)
+
+ class Consumer(Task[BaseOut, NumberOutput]):
+ name = "consumer"
+
+ def run(self, input: BaseOut, ctx: ExecutionContext) -> NumberOutput:
+ return NumberOutput(value=input.value)
+
+ wf = Workflow(name="sub", tasks=[Producer, Consumer])
+ assert wf.result_task is Consumer
+
+ def test_incompatible_single_edge_message_is_complete(self) -> None:
+ class Consumer(Task[NumberInput, NumberOutput]):
+ name = "consumer"
+
+ def run(self, input: NumberInput, ctx: ExecutionContext) -> NumberOutput:
+ return NumberOutput(value=1)
+
+ with pytest.raises(
+ WorkflowDefinitionError,
+ match=r"add_one outputs NumberOutput but consumer expects NumberInput",
+ ):
+ Workflow(name="bad", tasks=[AddOne, Consumer])
+
class TestDAGWorkflow:
def test_fan_in_workflow(self) -> None:
@@ -219,6 +282,52 @@ def test_ambiguous_sinks_raises(self) -> None:
class TestOutputFieldRouting:
"""Tests for Feature 2: output field routing via tuple deps."""
+ def test_generic_field_mismatch_message_keeps_type_args(self) -> None:
+ """The message says list[int], not just 'list'."""
+
+ class ListOut(BaseModel):
+ items: list[int]
+
+ class Producer(Task[NumberInput, ListOut]):
+ name = "producer"
+
+ def run(self, input: NumberInput, ctx: ExecutionContext) -> ListOut:
+ return ListOut(items=[])
+
+ with pytest.raises(
+ WorkflowDefinitionError,
+ match=r"producer\.items is list\[int\] but double expects NumberOutput",
+ ):
+ (
+ Workflow.builder("bad")
+ .add_task(Producer)
+ .add_task(Double, depends_on=(Producer, "items"))
+ .build()
+ )
+
+ def test_field_ref_accepts_subclass(self) -> None:
+ """Field-ref edges use type compatibility, not identity."""
+
+ class RichNumber(NumberOutput):
+ note: str = ""
+
+ class Wrapped(BaseModel):
+ inner: RichNumber
+
+ class Producer(Task[NumberInput, Wrapped]):
+ name = "producer"
+
+ def run(self, input: NumberInput, ctx: ExecutionContext) -> Wrapped:
+ return Wrapped(inner=RichNumber(value=input.value))
+
+ wf = (
+ Workflow.builder("ok")
+ .add_task(Producer)
+ .add_task(Double, depends_on=(Producer, "inner"))
+ .build()
+ )
+ assert wf.result_task is Double
+
def test_valid_field_ref(self) -> None:
"""Single field ref validates and builds."""
@@ -538,6 +647,19 @@ def test_result_task_as_string(self) -> None:
assert wf.result_task is Double
assert wf.result_task_name == "my_double"
+ def test_unknown_result_task_string_raises_at_build(self) -> None:
+ """A result_task name that was never added fails in build(), not later."""
+ with pytest.raises(
+ WorkflowDefinitionError,
+ match=r"result_task 'nope' is not registered.*known tasks: \['add_one', 'double'\]",
+ ):
+ (
+ Workflow.builder(name="bad", result_task="nope")
+ .add_task(AddOne)
+ .add_task(Double, depends_on=AddOne)
+ .build()
+ )
+
def test_dep_not_found_string_raises(self) -> None:
"""String dependency that doesn't exist raises."""
with pytest.raises(WorkflowDefinitionError, match="not registered"):
diff --git a/tests/test_workflow_task.py b/tests/test_workflow_task.py
index a61d1e5..2e5e1bc 100644
--- a/tests/test_workflow_task.py
+++ b/tests/test_workflow_task.py
@@ -331,6 +331,57 @@ def test_inner_failure_surfaces(self) -> None:
assert "inner_failing" in result.error # type: ignore[operator]
assert "inner task broke" in result.error # type: ignore[operator]
+ def test_inner_failure_is_workflow_task_error_with_chain(self) -> None:
+ """The outer job keeps a WorkflowTaskError carrying the inner Job and cause."""
+ from taskmaestro import TaskExecutionError, WorkflowTaskError
+
+ inner_wf = Workflow("failing_inner", tasks=[InnerFailing])
+ SubTask = workflow_task(inner_wf, name="fail_sub")
+ outer_wf = Workflow.builder("outer").add_task(SubTask).build()
+ job = Job(outer_wf, InnerInput(value=1))
+ result = Runner().run(job, ctx=ExecutionContext())
+
+ exc = result.exception
+ assert isinstance(exc, WorkflowTaskError)
+ assert isinstance(exc, TaskExecutionError)
+ assert exc.workflow_name == "failing_inner"
+ assert str(exc) == (
+ "Inner workflow 'failing_inner' failed at task 'inner_failing': inner task broke"
+ )
+
+ # Original exception is chained, not flattened to a string.
+ assert isinstance(exc.__cause__, ValueError)
+ assert str(exc.__cause__) == "inner task broke"
+
+ # The inner Job is preserved for post-mortem inspection.
+ inner = exc.inner_job
+ assert inner.status == JobStatus.FAILED
+ assert inner.failed_task == "inner_failing"
+ assert inner.exception is exc.__cause__
+ assert [(r.task_name, r.status.value) for r in inner.task_results] == [
+ ("inner_failing", "failed")
+ ]
+
+ def test_nested_failure_chains_through_two_levels(self) -> None:
+ """Errors from a doubly-nested workflow remain walkable via __cause__."""
+ from taskmaestro import WorkflowTaskError
+
+ leaf_wf = Workflow("leaf", tasks=[InnerFailing])
+ LeafTask = workflow_task(leaf_wf, name="leaf_task")
+ mid_wf = Workflow.builder("mid").add_task(LeafTask).build()
+ MidTask = workflow_task(mid_wf, name="mid_task")
+ outer_wf = Workflow.builder("outer").add_task(MidTask).build()
+
+ result = Runner().run(Job(outer_wf, InnerInput(value=1)), ctx=ExecutionContext())
+
+ outer_exc = result.exception
+ assert isinstance(outer_exc, WorkflowTaskError)
+ assert outer_exc.workflow_name == "mid"
+ mid_exc = outer_exc.__cause__
+ assert isinstance(mid_exc, WorkflowTaskError)
+ assert mid_exc.workflow_name == "leaf"
+ assert isinstance(mid_exc.__cause__, ValueError)
+
# ============================================================
# TestContextSharing
diff --git a/tests/test_yaml_config.py b/tests/test_yaml_config.py
index 64dc547..b412929 100644
--- a/tests/test_yaml_config.py
+++ b/tests/test_yaml_config.py
@@ -5,6 +5,7 @@
from pathlib import Path
import pytest
+import yaml
from pydantic import BaseModel, ValidationError
from taskmaestro import (
@@ -17,6 +18,7 @@
TaskConfig,
YamlWorkflowConfig,
_coerce_hook_params,
+ _yaml_load,
import_class,
load_workflow_from_yaml,
run_workflow_from_yaml,
@@ -124,6 +126,56 @@ def _write_input_yaml(tmp_path: Path, content: str) -> Path:
# ============================================================
+class TestYamlMergeKeys:
+ def test_nested_merges_allow_overrides_and_reused_anchors(self) -> None:
+ text = """\
+defaults: &defaults {value: 1, other: 2}
+override: &override {value: 3}
+merged: &merged
+ <<: [*override, *defaults]
+ other: 4
+first: {<<: *merged}
+second: {<<: *merged, value: 5}
+"""
+ result = _yaml_load(text)
+
+ assert result == yaml.safe_load(text)
+ assert result["first"] == {"value": 3, "other": 4}
+ assert result["second"] == {"value": 5, "other": 4}
+
+ @pytest.mark.parametrize(
+ "text",
+ [
+ "value: 1\nvalue: 2\n",
+ "<<: {value: 1}\nvalue: 2\nvalue: 3\n",
+ "<<: {value: 1, value: 2}\n",
+ "<<: {<<: {value: 1, value: 2}}\n",
+ ],
+ )
+ def test_explicit_duplicates_are_still_rejected(self, text: str) -> None:
+ with pytest.raises(yaml.constructor.ConstructorError, match="duplicate key"):
+ _yaml_load(text)
+
+ def test_workflow_and_input_yaml_support_merges(self, tmp_path: Path) -> None:
+ workflow_path = _write_workflow_yaml(
+ tmp_path,
+ f"""\
+defaults: &defaults
+ task: {THIS_MODULE}.UpperText
+workflow:
+ name: merged
+ tasks:
+ - <<: *defaults
+""",
+ )
+ input_path = _write_input_yaml(tmp_path, "<<: {text: default}\ntext: override\n")
+
+ result = load_workflow_from_yaml(workflow_path, input_path).run()
+
+ assert result.status == JobStatus.COMPLETED
+ assert result.result == TextOutput(text="OVERRIDE")
+
+
class TestImportClass:
def test_valid_import(self) -> None:
cls = import_class(f"{THIS_MODULE}.UpperText")
@@ -887,6 +939,153 @@ def test_frozen(self, tmp_path: Path) -> None:
class TestYamlNamedInstances:
"""Tests for YAML configs with name: field on tasks."""
+ def test_ambiguous_class_path_dependency_raises(self, tmp_path: Path) -> None:
+ """The same class under two names cannot be referenced by class path."""
+ wf_path = _write_workflow_yaml(
+ tmp_path,
+ f"""\
+workflow:
+ name: ambiguous
+ tasks:
+ - task: {THIS_MODULE}.UpperText
+ - task: {THIS_MODULE}.ReverseText
+ name: rev_a
+ depends_on: {THIS_MODULE}.UpperText
+ - task: {THIS_MODULE}.ReverseText
+ name: rev_b
+ depends_on: {THIS_MODULE}.UpperText
+ - task: {THIS_MODULE}.TextLength
+ depends_on: {THIS_MODULE}.ReverseText
+""",
+ )
+ in_path = _write_input_yaml(tmp_path, "text: hello\n")
+ with pytest.raises(
+ ConfigLoadError,
+ match=(
+ rf"Dependency '{THIS_MODULE}\.ReverseText' for task "
+ rf"'{THIS_MODULE}\.TextLength' is ambiguous; it matches "
+ r"\['rev_a', 'rev_b'\]\. Use the instance name\."
+ ),
+ ):
+ load_workflow_from_yaml(wf_path, in_path)
+
+ def test_ambiguous_class_path_resolved_by_instance_name(self, tmp_path: Path) -> None:
+ """Using the instance name disambiguates; both instances run."""
+ wf_path = _write_workflow_yaml(
+ tmp_path,
+ f"""\
+workflow:
+ name: disambiguated
+ result_task: length_b
+ tasks:
+ - task: {THIS_MODULE}.UpperText
+ - task: {THIS_MODULE}.ReverseText
+ name: rev_a
+ depends_on: {THIS_MODULE}.UpperText
+ - task: {THIS_MODULE}.ReverseText
+ name: rev_b
+ depends_on: rev_a
+ - task: {THIS_MODULE}.TextLength
+ name: length_b
+ depends_on: rev_b
+""",
+ )
+ in_path = _write_input_yaml(tmp_path, "text: hello\n")
+ loaded = load_workflow_from_yaml(wf_path, in_path)
+ assert loaded.workflow.get_dependencies("rev_b") == "rev_a"
+ result = loaded.run()
+ assert result.status == JobStatus.COMPLETED
+ assert result.result.length == 5 # type: ignore[union-attr]
+
+ def test_ambiguous_result_task_raises(self, tmp_path: Path) -> None:
+ wf_path = _write_workflow_yaml(
+ tmp_path,
+ f"""\
+workflow:
+ name: ambiguous_result
+ result_task: {THIS_MODULE}.ReverseText
+ tasks:
+ - task: {THIS_MODULE}.UpperText
+ - task: {THIS_MODULE}.ReverseText
+ name: rev_a
+ depends_on: {THIS_MODULE}.UpperText
+ - task: {THIS_MODULE}.ReverseText
+ name: rev_b
+ depends_on: {THIS_MODULE}.UpperText
+""",
+ )
+ in_path = _write_input_yaml(tmp_path, "text: hello\n")
+ with pytest.raises(
+ ConfigLoadError,
+ match=rf"result_task '{THIS_MODULE}\.ReverseText' is ambiguous",
+ ):
+ load_workflow_from_yaml(wf_path, in_path)
+
+ def test_same_inner_workflow_file_twice(self, tmp_path: Path) -> None:
+ """Two workflow: entries for one file get distinct tasks and wiring."""
+ (tmp_path / "inner.yaml").write_text(
+ f"""\
+workflow:
+ name: inner
+ tasks:
+ - task: {THIS_MODULE}.ReverseText
+"""
+ )
+ wf_path = _write_workflow_yaml(
+ tmp_path,
+ f"""\
+workflow:
+ name: outer
+ tasks:
+ - task: {THIS_MODULE}.UpperText
+ - workflow: inner.yaml
+ name: first_reverse
+ depends_on: {THIS_MODULE}.UpperText
+ - workflow: inner.yaml
+ name: second_reverse
+ depends_on: first_reverse
+""",
+ )
+ in_path = _write_input_yaml(tmp_path, "text: hello\n")
+ loaded = load_workflow_from_yaml(wf_path, in_path)
+ assert list(loaded.workflow._tasks) == ["upper_text", "first_reverse", "second_reverse"]
+ assert loaded.workflow.get_dependencies("second_reverse") == "first_reverse"
+ result = loaded.run()
+ assert result.status == JobStatus.COMPLETED
+ assert result.result.text == "HELLO" # type: ignore[union-attr]
+
+ def test_same_inner_workflow_file_referenced_by_path_is_ambiguous(
+ self, tmp_path: Path
+ ) -> None:
+ (tmp_path / "inner.yaml").write_text(
+ f"""\
+workflow:
+ name: inner
+ tasks:
+ - task: {THIS_MODULE}.ReverseText
+"""
+ )
+ wf_path = _write_workflow_yaml(
+ tmp_path,
+ f"""\
+workflow:
+ name: outer
+ tasks:
+ - task: {THIS_MODULE}.UpperText
+ - workflow: inner.yaml
+ name: a
+ depends_on: {THIS_MODULE}.UpperText
+ - workflow: inner.yaml
+ name: b
+ depends_on: {THIS_MODULE}.UpperText
+ - task: {THIS_MODULE}.TextLength
+ depends_on: inner.yaml
+""",
+ )
+ in_path = _write_input_yaml(tmp_path, "text: hello\n")
+ with pytest.raises(ConfigLoadError, match=r"'inner\.yaml'.*is ambiguous.*\['a', 'b'\]"):
+ load_workflow_from_yaml(wf_path, in_path)
+
def test_named_instances_yaml(self, tmp_path: Path) -> None:
"""YAML with name: field on tasks loads and resolves dependencies correctly."""
wf_path = _write_workflow_yaml(
@@ -954,7 +1153,7 @@ def test_result_task_not_found_raises(self, tmp_path: Path) -> None:
""",
)
in_path = _write_input_yaml(tmp_path, "text: hello\n")
- with pytest.raises(ConfigLoadError, match=r"result_task.*not found"):
+ with pytest.raises(ConfigLoadError, match=r"^result_task 'nonexistent_task' not found$"):
load_workflow_from_yaml(wf_path, in_path)
def test_named_result_task(self, tmp_path: Path) -> None:
@@ -1061,6 +1260,176 @@ def run(self, input: DownstreamInput, ctx: ExecutionContext) -> DownstreamOutput
return DownstreamOutput(result=f"{input.label}:{input.path}:{input.flag}")
+class AmbiguousPayload(BaseModel):
+ a: int
+
+
+class AmbiguousInput(BaseModel):
+ """Root input whose sole field shares its name with the task below."""
+
+ payload: AmbiguousPayload
+
+
+class AmbiguousRoot(Task[AmbiguousInput, TextOutput]):
+ name = "payload"
+
+ def run(self, input: AmbiguousInput, ctx: ExecutionContext) -> TextOutput:
+ return TextOutput(text=str(input.payload.a))
+
+
+class TestInputMode:
+ """workflow.input_mode controls flat vs per-task interpretation of input.yaml."""
+
+ def _ambiguous_workflow(self, tmp_path: Path, input_mode: str | None) -> Path:
+ mode_line = f" input_mode: {input_mode}\n" if input_mode else ""
+ return _write_workflow_yaml(
+ tmp_path,
+ f"""\
+workflow:
+ name: ambiguous
+{mode_line} tasks:
+ - task: {THIS_MODULE}.AmbiguousRoot
+""",
+ )
+
+ def test_auto_refuses_to_guess_when_both_readings_valid(self, tmp_path: Path) -> None:
+ """A flat input whose only key equals a task name is rejected under auto."""
+ wf_path = self._ambiguous_workflow(tmp_path, None)
+ in_path = _write_input_yaml(tmp_path, "payload:\n a: 1\n")
+ with pytest.raises(
+ ConfigLoadError,
+ match=(
+ r"Input file is ambiguous: its top-level keys \['payload'\] are task names, "
+ r"but the mapping is also a valid AmbiguousInput for root task 'payload'\. "
+ r"Set workflow\.input_mode to 'flat' or 'per_task'\."
+ ),
+ ):
+ load_workflow_from_yaml(wf_path, in_path)
+
+ def test_flat_forces_root_input_reading(self, tmp_path: Path) -> None:
+ wf_path = self._ambiguous_workflow(tmp_path, "flat")
+ in_path = _write_input_yaml(tmp_path, "payload:\n a: 7\n")
+ loaded = load_workflow_from_yaml(wf_path, in_path)
+ assert loaded.job.job_configuration is None
+ assert isinstance(loaded.job.config, AmbiguousInput)
+ assert loaded.run().result.text == "7" # type: ignore[union-attr]
+
+ def test_per_task_forces_config_reading(self, tmp_path: Path) -> None:
+ wf_path = self._ambiguous_workflow(tmp_path, "per_task")
+ in_path = _write_input_yaml(tmp_path, "payload:\n payload:\n a: 3\n")
+ loaded = load_workflow_from_yaml(wf_path, in_path)
+ assert loaded.job.job_configuration is not None
+ assert loaded.workflow.get_config_fields("payload") == {"payload"}
+ assert loaded.run().result.text == "3" # type: ignore[union-attr]
+
+ def test_auto_still_picks_per_task_when_flat_reading_is_invalid(self, tmp_path: Path) -> None:
+ """The heuristic is kept for the unambiguous common case."""
+ wf_path = _write_workflow_yaml(
+ tmp_path,
+ f"""\
+workflow:
+ name: per_task_ok
+ tasks:
+ - task: {THIS_MODULE}.PerTaskRoot
+""",
+ )
+ in_path = _write_input_yaml(tmp_path, 'per_task_root:\n egrid_path: "/x"\n')
+ loaded = load_workflow_from_yaml(wf_path, in_path)
+ assert loaded.job.job_configuration is not None
+
+ def test_per_task_rejects_unknown_task_key(self, tmp_path: Path) -> None:
+ wf_path = _write_workflow_yaml(
+ tmp_path,
+ f"""\
+workflow:
+ name: strict
+ input_mode: per_task
+ tasks:
+ - task: {THIS_MODULE}.PerTaskRoot
+""",
+ )
+ in_path = _write_input_yaml(tmp_path, 'per_task_rooot:\n egrid_path: "/x"\n')
+ with pytest.raises(
+ ConfigLoadError,
+ match=(
+ r"input_mode is 'per_task' but top-level key 'per_task_rooot' is not a task "
+ r"name \(known tasks: \['per_task_root'\]\)"
+ ),
+ ):
+ load_workflow_from_yaml(wf_path, in_path)
+
+ def test_per_task_rejects_non_mapping_value(self, tmp_path: Path) -> None:
+ wf_path = _write_workflow_yaml(
+ tmp_path,
+ f"""\
+workflow:
+ name: strict
+ input_mode: per_task
+ tasks:
+ - task: {THIS_MODULE}.PerTaskRoot
+""",
+ )
+ in_path = _write_input_yaml(tmp_path, "per_task_root: 42\n")
+ with pytest.raises(
+ ConfigLoadError,
+ match=r"value for task 'per_task_root' is not a mapping \(got int\)",
+ ):
+ load_workflow_from_yaml(wf_path, in_path)
+
+ def test_per_task_allows_null_value(self, tmp_path: Path) -> None:
+ """``task_name:`` with no value means 'configured, no fields'."""
+ wf_path = _write_workflow_yaml(
+ tmp_path,
+ f"""\
+workflow:
+ name: strict
+ input_mode: per_task
+ tasks:
+ - task: {THIS_MODULE}.PerTaskRoot
+""",
+ )
+ in_path = _write_input_yaml(tmp_path, "per_task_root:\n")
+ # Accepted mode-wise; the task then has no config_fields and no job
+ # input, which surfaces as a wrapped Job validation error.
+ with pytest.raises(
+ ConfigLoadError,
+ match=r"Job validation failed: Root task 'per_task_root' expects input type",
+ ):
+ load_workflow_from_yaml(wf_path, in_path)
+
+ def test_auto_with_dag_root_is_ambiguity_checked(self, tmp_path: Path) -> None:
+ """In DAG mode the root is the entry without depends_on, not entry 0."""
+ wf_path = _write_workflow_yaml(
+ tmp_path,
+ f"""\
+workflow:
+ name: dag_ambiguous
+ tasks:
+ - task: {THIS_MODULE}.TextLength
+ depends_on: payload
+ - task: {THIS_MODULE}.AmbiguousRoot
+""",
+ )
+ in_path = _write_input_yaml(tmp_path, "payload:\n a: 1\n")
+ with pytest.raises(ConfigLoadError, match="Input file is ambiguous"):
+ load_workflow_from_yaml(wf_path, in_path)
+
+ def test_invalid_input_mode_is_schema_error(self, tmp_path: Path) -> None:
+ wf_path = _write_workflow_yaml(
+ tmp_path,
+ f"""\
+workflow:
+ name: bad
+ input_mode: sideways
+ tasks:
+ - task: {THIS_MODULE}.UpperText
+""",
+ )
+ in_path = _write_input_yaml(tmp_path, "text: hi\n")
+ with pytest.raises(ConfigLoadError, match="YAML schema validation error"):
+ load_workflow_from_yaml(wf_path, in_path)
+
+
class TestPerTaskConfig:
"""Tests for per-task YAML config format."""
@@ -1161,6 +1530,154 @@ def test_per_task_empty_config(self, tmp_path: Path) -> None:
# ============================================================
+class TestLinearModeViaBuilder:
+ """Linear-mode YAML (no depends_on) must honour the same rules as DAG mode."""
+
+ def test_name_override_is_honoured(self, tmp_path: Path) -> None:
+ wf_path = _write_workflow_yaml(
+ tmp_path,
+ f"""\
+workflow:
+ name: lin
+ tasks:
+ - task: {THIS_MODULE}.UpperText
+ name: shout
+ - task: {THIS_MODULE}.ReverseText
+ name: flip
+""",
+ )
+ in_path = _write_input_yaml(tmp_path, "text: hello\n")
+ loaded = load_workflow_from_yaml(wf_path, in_path)
+
+ assert list(loaded.workflow._tasks) == ["shout", "flip"]
+ assert loaded.workflow.get_dependencies("flip") == "shout"
+ assert loaded.workflow.result_task_name == "flip"
+ result = loaded.run()
+ assert result.status == JobStatus.COMPLETED
+ assert [r.task_name for r in result.task_results] == ["shout", "flip"]
+
+ def test_per_task_config_with_name_override(self, tmp_path: Path) -> None:
+ """Per-task input keyed by the overridden name is wired to the right task."""
+ wf_path = _write_workflow_yaml(
+ tmp_path,
+ f"""\
+workflow:
+ name: lin
+ tasks:
+ - task: {THIS_MODULE}.PerTaskRoot
+ name: my_root
+""",
+ )
+ in_path = _write_input_yaml(
+ tmp_path,
+ """\
+my_root:
+ egrid_path: "/data/x.egrid"
+""",
+ )
+ loaded = load_workflow_from_yaml(wf_path, in_path)
+ assert loaded.workflow.get_config_fields("my_root") == {"egrid_path"}
+ result = loaded.run()
+ assert result.status == JobStatus.COMPLETED
+ assert result.result.path == "/data/x.egrid" # type: ignore[union-attr]
+
+ def test_unknown_config_field_is_rejected_at_load(self, tmp_path: Path) -> None:
+ """Config fields are validated (previously bypassed in linear mode)."""
+ wf_path = _write_workflow_yaml(
+ tmp_path,
+ f"""\
+workflow:
+ name: lin
+ tasks:
+ - task: {THIS_MODULE}.PerTaskRoot
+""",
+ )
+ in_path = _write_input_yaml(
+ tmp_path,
+ """\
+per_task_root:
+ egrid_path: "/data/x.egrid"
+ bogus: 1
+""",
+ )
+ with pytest.raises(ConfigLoadError, match="Config field 'bogus' not found"):
+ load_workflow_from_yaml(wf_path, in_path)
+
+ def test_type_mismatch_is_wrapped(self, tmp_path: Path) -> None:
+ """A linear chain with incompatible types raises ConfigLoadError, not a raw error."""
+ wf_path = _write_workflow_yaml(
+ tmp_path,
+ f"""\
+workflow:
+ name: lin
+ tasks:
+ - task: {THIS_MODULE}.UpperText
+ - task: {THIS_MODULE}.TextLength
+ - task: {THIS_MODULE}.ReverseText
+""",
+ )
+ in_path = _write_input_yaml(tmp_path, "text: hello\n")
+ with pytest.raises(ConfigLoadError, match=r"Workflow validation failed.*Type mismatch"):
+ load_workflow_from_yaml(wf_path, in_path)
+
+ def test_duplicate_names_are_wrapped(self, tmp_path: Path) -> None:
+ wf_path = _write_workflow_yaml(
+ tmp_path,
+ f"""\
+workflow:
+ name: lin
+ tasks:
+ - task: {THIS_MODULE}.UpperText
+ - task: {THIS_MODULE}.UpperText
+""",
+ )
+ in_path = _write_input_yaml(tmp_path, "text: hello\n")
+ with pytest.raises(ConfigLoadError, match="Duplicate task name 'upper_text'"):
+ load_workflow_from_yaml(wf_path, in_path)
+
+ def test_result_task_by_instance_name(self, tmp_path: Path) -> None:
+ wf_path = _write_workflow_yaml(
+ tmp_path,
+ f"""\
+workflow:
+ name: lin
+ result_task: shout
+ tasks:
+ - task: {THIS_MODULE}.UpperText
+ name: shout
+ - task: {THIS_MODULE}.ReverseText
+""",
+ )
+ in_path = _write_input_yaml(tmp_path, "text: hello\n")
+ loaded = load_workflow_from_yaml(wf_path, in_path)
+ assert loaded.workflow.result_task_name == "shout"
+
+ def test_job_validation_error_is_wrapped(self, tmp_path: Path) -> None:
+ """Errors raised while constructing the Job surface as ConfigLoadError."""
+ from unittest.mock import patch
+
+ from taskmaestro.exceptions import WorkflowDefinitionError
+
+ wf_path = _write_workflow_yaml(
+ tmp_path,
+ f"""\
+workflow:
+ name: lin
+ tasks:
+ - task: {THIS_MODULE}.UpperText
+""",
+ )
+ in_path = _write_input_yaml(tmp_path, "text: hello\n")
+ with (
+ patch(
+ "taskmaestro.yaml_config.Job.__init__",
+ side_effect=WorkflowDefinitionError("boom"),
+ ),
+ pytest.raises(ConfigLoadError, match="Job validation failed: boom"),
+ ):
+ load_workflow_from_yaml(wf_path, in_path)
+
+
class TestWorkflowTaskYaml:
"""Tests for YAML workflow: references (workflow_task via YAML)."""
@@ -1168,6 +1685,83 @@ def _write_yaml(self, path: Path, content: str) -> Path:
path.write_text(content)
return path
+ def test_self_referencing_workflow_is_rejected(self, tmp_path: Path) -> None:
+ outer_path = self._write_yaml(
+ tmp_path / "outer.yaml",
+ """\
+workflow:
+ name: loop
+ tasks:
+ - workflow: outer.yaml
+""",
+ )
+ in_path = self._write_yaml(tmp_path / "input.yaml", "text: hello\n")
+ with pytest.raises(ConfigLoadError, match="Recursive workflow reference"):
+ load_workflow_from_yaml(outer_path, in_path)
+
+ def test_mutually_referencing_workflows_are_rejected(self, tmp_path: Path) -> None:
+ self._write_yaml(
+ tmp_path / "a.yaml",
+ """\
+workflow:
+ name: a
+ tasks:
+ - workflow: b.yaml
+""",
+ )
+ self._write_yaml(
+ tmp_path / "b.yaml",
+ f"""\
+workflow:
+ name: b
+ tasks:
+ - workflow: ../{tmp_path.name}/a.yaml
+""",
+ )
+ in_path = self._write_yaml(tmp_path / "input.yaml", "text: hello\n")
+ with pytest.raises(ConfigLoadError, match="Recursive workflow reference") as excinfo:
+ load_workflow_from_yaml(tmp_path / "a.yaml", in_path)
+ # Both files appear in the reported chain.
+ assert "a.yaml" in str(excinfo.value)
+ assert "b.yaml" in str(excinfo.value)
+
+ def test_reuse_of_inner_workflow_is_not_a_cycle(self, tmp_path: Path) -> None:
+ """Only files on the *current* nesting chain count as recursion."""
+ self._write_yaml(
+ tmp_path / "leaf.yaml",
+ f"""\
+workflow:
+ name: leaf
+ tasks:
+ - task: {THIS_MODULE}.ReverseText
+""",
+ )
+ self._write_yaml(
+ tmp_path / "mid.yaml",
+ """\
+workflow:
+ name: mid
+ tasks:
+ - workflow: leaf.yaml
+ name: inner_leaf
+""",
+ )
+ outer_path = self._write_yaml(
+ tmp_path / "outer.yaml",
+ f"""\
+workflow:
+ name: outer
+ tasks:
+ - task: {THIS_MODULE}.UpperText
+ - workflow: mid.yaml
+ name: via_mid
+ depends_on: {THIS_MODULE}.UpperText
+""",
+ )
+ in_path = self._write_yaml(tmp_path / "input.yaml", "text: hello\n")
+ loaded = load_workflow_from_yaml(outer_path, in_path)
+ assert loaded.run().status == JobStatus.COMPLETED
+
def test_workflow_ref_basic(self, tmp_path: Path) -> None:
"""Outer YAML references inner YAML via workflow:, end-to-end."""
self._write_yaml(