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..f96ff4f 100644
--- a/README.md
+++ b/README.md
@@ -110,16 +110,55 @@ 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()`. Prefer `builder.task()` and task handles for unambiguous dependencies; the fluent `add_task()` API remains supported. The builder accepts `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. |
| **Hooks** | Subclass `BaseHook` and override methods like `on_job_start`, `on_task_complete`, etc. Hook errors are swallowed and reported via `warnings.warn()`, so they never crash the job. Built-ins: `LoggingHook`, `TimingHook`, `ResultPersistenceHook`. |
| **ObjectModel** | Generic `ObjectModel[T]` base model for wrapping arbitrary (non-Pydantic) objects. Enables `arbitrary_types_allowed` so fields can hold native library objects like database connections or API clients. |
+For common cases, run a workflow directly without constructing `Job` and `Runner`:
+
+```python
+result = workflow.run(
+ Input(value=5),
+ task_config={"configured_task": {"option": "value"}},
+ hooks=[LoggingHook()],
+)
+```
+
+The explicit `Job` and `Runner` API remains available for advanced lifecycle control.
+
+## Task Handles
+
+`builder.task()` adds a task and returns a handle to that specific instance. Handles avoid ambiguous class and string references, especially when the same task class is registered more than once:
+
+```python
+builder = Workflow.builder(name="parallel_wells")
+model = builder.task(LoadModel)
+well_1 = builder.task(LoadWellPath, name="well_1", depends_on=model)
+well_2 = builder.task(LoadWellPath, name="well_2", depends_on=model)
+builder.task(Process, name="proc_1", depends_on=well_1)
+builder.task(Process, name="proc_2", depends_on=well_2)
+workflow = builder.build()
+```
+
+Use `handle.field("field_name")` to route one output field. Named input dependencies can be passed directly to `task()`, while keyword arguments to `collect()` provide a concise keyed collection:
+
+```python
+merged = builder.task(
+ MergeResults,
+ primary=producer.field("result"),
+ checks=collect(tests=tests, lint=lint, types=types),
+)
+builder.set_result_task(merged)
+```
+
+Handles are accepted anywhere dependency references are accepted. A handle from a different builder is rejected. Use the `depends_on` dictionary form when an input field conflicts with a reserved builder argument such as `name` or `config_fields`. The existing fluent `add_task()` API remains fully supported for backward compatibility.
+
## Named Task Instances
-The same Task class can appear multiple times in a workflow with different names. Use the `name=` parameter in `add_task()`:
+The same Task class can appear multiple times in a workflow with different names. With the fluent API, use the `name=` parameter in `add_task()`:
```python
workflow = (
@@ -231,6 +270,155 @@ 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
+builder = Workflow.builder("create_grid")
+connection = builder.task(ConnectToResInsight)
+surfaces = builder.map_task(
+ LoadRegularSurface,
+ name="load_surfaces",
+ depends_on={"resinsight": connection},
+ config_fields=["unit"],
+ over="surfaces",
+ key_as="surface_name",
+ value_as="path",
+ error_mode="fail_fast",
+)
+builder.task(CreateGrid, depends_on={"surfaces": surfaces})
+workflow = builder.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.
+When a mapped task is connected to a named `dict[str, O]` input field, its
+`root` value is unwrapped automatically:
+
+```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:
@@ -324,8 +512,9 @@ context:
```yaml
# input.yaml
-text: "Python is a high-level programming language..."
-title: "Python Overview"
+prepare_text:
+ text: "Python is a high-level programming language..."
+ title: "Python Overview"
```
Load and run:
@@ -341,7 +530,18 @@ result = loaded.run()
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.
+YAML input always uses per-task configuration: every top-level key in `input.yaml` must be a registered task instance name, and its value must be a mapping or `null`. Unknown task names and scalar task values are rejected. Fields are validated against the task's input model and can configure root tasks, downstream tasks, and mapped tasks.
+
+YAML also supports named task instances (`name:`), fan-in dictionaries, and output field routing via `[task, field]` lists. Named instances use their instance name as the input key:
+
+```yaml
+load_well_path_1:
+ path: first.dev
+load_well_path_2:
+ path: second.dev
+```
+
+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
@@ -405,19 +605,33 @@ 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
```
+## Command-Line Interface
+
+Installed packages provide a `taskmaestro` command for YAML workflows:
+
+```bash
+taskmaestro validate workflow.yaml --input input.yaml
+taskmaestro graph workflow.yaml --input input.yaml
+taskmaestro run workflow.yaml --input input.yaml --log-level INFO
+```
+
+`run` prints the final output as JSON and returns a nonzero exit code when the workflow fails. `graph` prints Mermaid markup.
+
## Examples
-Three full example pipelines are included in the `examples/` directory:
+Four full example pipelines are included in the `examples/` directory:
| Example | Features |
|---|---|
| `examples/text_analysis/` | DAG with fan-out/fan-in, output field routing, inline `Inputs`/`Outputs` classes, YAML config, Mermaid visualization |
| `examples/resinsight/` | `ObjectModel[T]` for gRPC objects, `JobConfiguration` with per-task config, named task instances, `config_fields`, YAML config |
| `examples/image_processing/` | Nested workflows through `Workflow.as_task()` and YAML `workflow:`, typed boundaries, expanded Mermaid subgraph |
+| `examples/release_pipeline/` | Keyed `collect()` dependencies, mapped tasks, mapped output routing, per-task YAML config |
Run an example:
diff --git a/examples/image_processing/input.yaml b/examples/image_processing/input.yaml
index 2d4553d..350c9e1 100644
--- a/examples/image_processing/input.yaml
+++ b/examples/image_processing/input.yaml
@@ -3,4 +3,5 @@
# Run:
# python examples/image_processing/pipeline.py --yaml --input input.yaml
-image_path: "taskmaestro.png"
+load_image:
+ image_path: "taskmaestro.png"
diff --git a/examples/release_pipeline/input.yaml b/examples/release_pipeline/input.yaml
new file mode 100644
index 0000000..c309dfa
--- /dev/null
+++ b/examples/release_pipeline/input.yaml
@@ -0,0 +1,27 @@
+# Per-task configuration for the release pipeline.
+
+load_package:
+ name: taskmaestro-demo
+ version: 1.0.0
+ files:
+ src/demo.py: |
+ def greet(name: str) -> str:
+ return f"Hello {name}"
+ tests/test_demo.py: |
+ def test_greet():
+ assert True
+
+build_targets:
+ targets:
+ linux-x64:
+ operating_system: linux
+ architecture: x86_64
+ extension: tar.gz
+ windows-x64:
+ operating_system: windows
+ architecture: x86_64
+ extension: zip
+ macos-arm64:
+ operating_system: macos
+ architecture: arm64
+ extension: tar.gz
diff --git a/examples/release_pipeline/pipeline.py b/examples/release_pipeline/pipeline.py
new file mode 100644
index 0000000..acd695d
--- /dev/null
+++ b/examples/release_pipeline/pipeline.py
@@ -0,0 +1,379 @@
+"""Example: release validation, collected dependencies, and mapped builds.
+
+The fixed validation tasks fan out from one package snapshot. Their results are
+collected into a keyed dictionary before a build task is expanded over the
+configured release targets.
+
+ LoadPackage ─┬─ RunTests ──┐
+ ├─ RunLint ───┼─ ValidateRelease ── BuildTarget[*] ── Manifest
+ └─ CheckTypes ┘
+
+Run:
+ python examples/release_pipeline/pipeline.py
+ python examples/release_pipeline/pipeline.py --yaml
+"""
+
+from __future__ import annotations
+
+import argparse
+import hashlib
+import json
+import logging
+from pathlib import Path
+from typing import Any, cast
+
+from pydantic import BaseModel
+
+from taskmaestro import (
+ EmptyConfig,
+ ExecutionContext,
+ Job,
+ JobConfiguration,
+ Runner,
+ Task,
+ Workflow,
+ collect,
+)
+from taskmaestro.hooks import LoggingHook, TimingHook
+
+# ---------------------------------------------------------------------------
+# Models
+# ---------------------------------------------------------------------------
+
+
+class PackageInput(BaseModel):
+ """Source package supplied as job configuration."""
+
+ name: str
+ version: str
+ files: dict[str, str]
+
+
+class PackageSnapshot(BaseModel):
+ """Immutable package information passed to each validation task."""
+
+ name: str
+ version: str
+ files: dict[str, str]
+ source_digest: str
+
+
+class CheckResult(BaseModel):
+ """Common output type that allows validation results to be collected."""
+
+ passed: bool
+ message: str
+
+
+class ValidateReleaseInput(BaseModel):
+ package: PackageSnapshot
+ checks: dict[str, CheckResult]
+
+
+class ValidatedRelease(BaseModel):
+ package: PackageSnapshot
+ checks: dict[str, CheckResult]
+
+
+class TargetSettings(BaseModel):
+ operating_system: str
+ architecture: str
+ extension: str
+
+
+class BuildTargetInput(BaseModel):
+ release: ValidatedRelease
+ target_name: str
+ settings: TargetSettings
+
+
+class Artifact(BaseModel):
+ target: str
+ path: str
+ sha256: str
+
+
+class ManifestInput(BaseModel):
+ artifacts: dict[str, Artifact]
+
+
+class ReleaseManifest(BaseModel):
+ manifest_path: str
+ artifacts: dict[str, Artifact]
+
+
+# ---------------------------------------------------------------------------
+# Tasks
+# ---------------------------------------------------------------------------
+
+
+class LoadPackage(Task[PackageInput, PackageSnapshot]):
+ name = "load_package"
+
+ def run(self, input: PackageInput, ctx: ExecutionContext) -> PackageSnapshot:
+ serialized_files = json.dumps(input.files, sort_keys=True).encode()
+ digest = hashlib.sha256(serialized_files).hexdigest()
+ ctx.logger.info(
+ "Loaded %s %s with %d files",
+ input.name,
+ input.version,
+ len(input.files),
+ )
+ return PackageSnapshot(
+ name=input.name,
+ version=input.version,
+ files=input.files,
+ source_digest=digest,
+ )
+
+
+class RunTests(Task[PackageSnapshot, CheckResult]):
+ name = "run_tests"
+
+ def run(self, input: PackageSnapshot, ctx: ExecutionContext) -> CheckResult:
+ failing_files = [name for name, content in input.files.items() if "FAIL_TEST" in content]
+ return CheckResult(
+ passed=not failing_files,
+ message=(
+ "All tests passed"
+ if not failing_files
+ else f"Test failures in: {', '.join(failing_files)}"
+ ),
+ )
+
+
+class RunLint(Task[PackageSnapshot, CheckResult]):
+ name = "run_lint"
+
+ def run(self, input: PackageSnapshot, ctx: ExecutionContext) -> CheckResult:
+ files_with_tabs = [name for name, content in input.files.items() if "\t" in content]
+ return CheckResult(
+ passed=not files_with_tabs,
+ message=(
+ "No lint errors"
+ if not files_with_tabs
+ else f"Tabs found in: {', '.join(files_with_tabs)}"
+ ),
+ )
+
+
+class CheckTypes(Task[PackageSnapshot, CheckResult]):
+ name = "check_types"
+
+ def run(self, input: PackageSnapshot, ctx: ExecutionContext) -> CheckResult:
+ ignored_files = [
+ name for name, content in input.files.items() if "# type: ignore" in content
+ ]
+ return CheckResult(
+ passed=not ignored_files,
+ message=(
+ "Type checks passed"
+ if not ignored_files
+ else f"Unchecked types in: {', '.join(ignored_files)}"
+ ),
+ )
+
+
+class ValidateRelease(Task[ValidateReleaseInput, ValidatedRelease]):
+ name = "validate_release"
+
+ def run(self, input: ValidateReleaseInput, ctx: ExecutionContext) -> ValidatedRelease:
+ failed = [name for name, result in input.checks.items() if not result.passed]
+ if failed:
+ details = "; ".join(f"{name}: {input.checks[name].message}" for name in failed)
+ raise ValueError(f"Release validation failed: {details}")
+
+ ctx.logger.info("All %d release checks passed", len(input.checks))
+ return ValidatedRelease(package=input.package, checks=input.checks)
+
+
+class BuildTarget(Task[BuildTargetInput, Artifact]):
+ name = "build_target"
+
+ def run(self, input: BuildTargetInput, ctx: ExecutionContext) -> Artifact:
+ package = input.release.package
+ settings = input.settings
+ ctx.scratch_dir.mkdir(parents=True, exist_ok=True)
+
+ filename = (
+ f"{package.name}-{package.version}-{settings.operating_system}-"
+ f"{settings.architecture}.{settings.extension}"
+ )
+ artifact_path = ctx.scratch_dir / filename
+ payload = {
+ "package": package.name,
+ "version": package.version,
+ "source_digest": package.source_digest,
+ "target": input.target_name,
+ "operating_system": settings.operating_system,
+ "architecture": settings.architecture,
+ }
+ artifact_path.write_text(json.dumps(payload, indent=2), encoding="utf-8")
+ digest = hashlib.sha256(artifact_path.read_bytes()).hexdigest()
+ ctx.logger.info("Built %s at %s", input.target_name, artifact_path)
+ return Artifact(
+ target=input.target_name,
+ path=str(artifact_path),
+ sha256=digest,
+ )
+
+
+class CreateReleaseManifest(Task[ManifestInput, ReleaseManifest]):
+ name = "create_release_manifest"
+
+ def run(self, input: ManifestInput, ctx: ExecutionContext) -> ReleaseManifest:
+ ctx.scratch_dir.mkdir(parents=True, exist_ok=True)
+ manifest_path = ctx.scratch_dir / "release-manifest.json"
+ manifest_path.write_text(
+ json.dumps(
+ {name: artifact.model_dump() for name, artifact in input.artifacts.items()},
+ indent=2,
+ ),
+ encoding="utf-8",
+ )
+ return ReleaseManifest(
+ manifest_path=str(manifest_path),
+ artifacts=input.artifacts,
+ )
+
+
+# ---------------------------------------------------------------------------
+# Workflow and execution
+# ---------------------------------------------------------------------------
+
+
+def build_workflow() -> Workflow:
+ """Build the release DAG using unambiguous task handles."""
+ builder = Workflow.builder("release_pipeline")
+ package = builder.task(
+ LoadPackage,
+ config_fields=["name", "version", "files"],
+ )
+ tests = builder.task(RunTests, depends_on=package)
+ lint = builder.task(RunLint, depends_on=package)
+ types = builder.task(CheckTypes, depends_on=package)
+ validated = builder.task(
+ ValidateRelease,
+ package=package,
+ checks=collect(tests=tests, lint=lint, types=types),
+ )
+ builds = builder.map_task(
+ BuildTarget,
+ name="build_targets",
+ release=validated,
+ over="targets",
+ key_as="target_name",
+ value_as="settings",
+ error_mode="collect_all",
+ )
+ builder.task(CreateReleaseManifest, artifacts=builds)
+ return builder.build()
+
+
+def sample_job_configuration() -> JobConfiguration:
+ return JobConfiguration(
+ {
+ "load_package": {
+ "name": "taskmaestro-demo",
+ "version": "1.0.0",
+ "files": {
+ "src/demo.py": "def greet(name: str) -> str:\n return f'Hello {name}'\n",
+ "tests/test_demo.py": "def test_greet():\n assert True\n",
+ },
+ },
+ "build_targets": {
+ "targets": {
+ "linux-x64": {
+ "operating_system": "linux",
+ "architecture": "x86_64",
+ "extension": "tar.gz",
+ },
+ "windows-x64": {
+ "operating_system": "windows",
+ "architecture": "x86_64",
+ "extension": "zip",
+ },
+ "macos-arm64": {
+ "operating_system": "macos",
+ "architecture": "arm64",
+ "extension": "tar.gz",
+ },
+ }
+ },
+ }
+ )
+
+
+def print_result(result: Job[Any], timing: TimingHook, workflow: Workflow) -> None:
+ print("=" * 60)
+ print("Release pipeline")
+ print("=" * 60)
+ print(f"Status: {result.status}")
+ if result.error:
+ print(f"Error: {result.error}")
+ return
+
+ assert result.result is not None
+ manifest = cast(ReleaseManifest, result.result)
+ print(f"Manifest: {manifest.manifest_path}")
+ for name, artifact in manifest.artifacts.items():
+ print(f" {name:16s} {artifact.path}")
+ print(f"Duration: {timing.job_duration:.4f}s")
+ print("\nMermaid diagram:")
+ print("```mermaid")
+ print(workflow.to_mermaid(), end="")
+ print("```")
+
+
+def run_python_mode() -> None:
+ workflow = build_workflow()
+ timing = TimingHook()
+ result = Runner(hooks=[LoggingHook(), timing]).run(
+ Job(
+ workflow,
+ EmptyConfig(),
+ job_configuration=sample_job_configuration(),
+ )
+ )
+ print_result(result, timing, workflow)
+
+
+def run_yaml_mode(workflow_path: str, input_path: str) -> None:
+ from taskmaestro.yaml_config import load_workflow_from_yaml
+
+ loaded = load_workflow_from_yaml(workflow_path, input_path)
+ result = loaded.run()
+ timing = next(hook for hook in loaded.runner.hooks if isinstance(hook, TimingHook))
+ print_result(result, timing, loaded.workflow)
+
+
+def main() -> None:
+ logging.basicConfig(
+ level=logging.INFO,
+ format="%(levelname)s %(name)s — %(message)s",
+ )
+ example_dir = Path(__file__).resolve().parent
+ parser = argparse.ArgumentParser(description="Mapped release pipeline example")
+ parser.add_argument(
+ "--yaml",
+ metavar="FILE",
+ nargs="?",
+ const=str(example_dir / "workflow.yaml"),
+ help="Load YAML workflow (default: workflow.yaml)",
+ )
+ parser.add_argument(
+ "--input",
+ metavar="FILE",
+ default=str(example_dir / "input.yaml"),
+ help="Input YAML file (default: input.yaml)",
+ )
+ args = parser.parse_args()
+
+ if args.yaml:
+ run_yaml_mode(args.yaml, args.input)
+ else:
+ run_python_mode()
+
+
+if __name__ == "__main__":
+ main()
diff --git a/examples/release_pipeline/workflow.yaml b/examples/release_pipeline/workflow.yaml
new file mode 100644
index 0000000..b02d1d4
--- /dev/null
+++ b/examples/release_pipeline/workflow.yaml
@@ -0,0 +1,44 @@
+# Release pipeline using keyed collection dependencies and mapped execution.
+# Run with: python examples/release_pipeline/pipeline.py --yaml
+
+workflow:
+ name: release_pipeline
+ tasks:
+ - task: pipeline.LoadPackage
+
+ - task: pipeline.RunTests
+ depends_on: pipeline.LoadPackage
+
+ - task: pipeline.RunLint
+ depends_on: pipeline.LoadPackage
+
+ - task: pipeline.CheckTypes
+ depends_on: pipeline.LoadPackage
+
+ - task: pipeline.ValidateRelease
+ depends_on:
+ package: pipeline.LoadPackage
+ checks:
+ collect:
+ tests: pipeline.RunTests
+ lint: pipeline.RunLint
+ types: pipeline.CheckTypes
+
+ - task: pipeline.BuildTarget
+ name: build_targets
+ map:
+ over: targets
+ key_as: target_name
+ value_as: settings
+ error_mode: collect_all
+ depends_on:
+ release: pipeline.ValidateRelease
+
+ - task: pipeline.CreateReleaseManifest
+ depends_on:
+ artifacts: [build_targets, root]
+
+runner:
+ hooks:
+ - hook: taskmaestro.hooks.logging.LoggingHook
+ - hook: taskmaestro.hooks.timing.TimingHook
diff --git a/pyproject.toml b/pyproject.toml
index 2e4a407..5d73e31 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",
]
@@ -27,6 +28,9 @@ dependencies = [
"pyyaml>=6.0,<7.0",
]
+[project.scripts]
+taskmaestro = "taskmaestro.cli:main"
+
[project.urls]
Homepage = "https://github.com/OPM/taskmaestro"
Repository = "https://github.com/OPM/taskmaestro"
diff --git a/taskmaestro/__init__.py b/taskmaestro/__init__.py
index f613347..83d85c2 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 OutputHandle, TaskHandle, 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,23 @@
"CycleDetectedError",
"EmptyConfig",
"ExecutionContext",
+ "HookError",
"IncompleteInputError",
"Job",
"JobConfiguration",
"JobStateError",
"JobStatus",
"LoadedWorkflow",
+ "MappedOutput",
+ "MappedTaskExecutionError",
"ObjectModel",
+ "OutputHandle",
"PluginLoadError",
"Runner",
"Task",
"TaskExecutionError",
+ "TaskHandle",
+ "TaskMap",
"TaskOutputTypeError",
"TaskResult",
"TaskStatus",
@@ -64,6 +74,8 @@
"WorkflowBuilder",
"WorkflowDefinitionError",
"WorkflowRunnerError",
+ "WorkflowTaskError",
+ "collect",
"get_registered_task",
"get_registered_workflow",
"load_workflow_from_yaml",
diff --git a/taskmaestro/cli.py b/taskmaestro/cli.py
new file mode 100644
index 0000000..fb9139a
--- /dev/null
+++ b/taskmaestro/cli.py
@@ -0,0 +1,101 @@
+"""Command-line interface for validating, visualizing, and running workflows."""
+
+from __future__ import annotations
+
+import argparse
+import logging
+import sys
+from collections.abc import Sequence
+from pathlib import Path
+from typing import Any
+
+from taskmaestro.exceptions import ConfigLoadError
+from taskmaestro.job import JobStatus
+from taskmaestro.yaml_config import LoadedWorkflow, load_workflow_from_yaml
+
+
+def _add_workflow_arguments(parser: argparse.ArgumentParser) -> None:
+ parser.add_argument("workflow", help="Path to the workflow YAML file")
+ parser.add_argument("--input", required=True, help="Path to the input YAML file")
+
+
+def _load(args: argparse.Namespace) -> LoadedWorkflow:
+ """Load YAML with its directory available for local task imports."""
+ workflow_dir = str(Path(args.workflow).resolve().parent)
+ original_path = sys.path.copy()
+ sys.path.insert(0, workflow_dir)
+ try:
+ return load_workflow_from_yaml(args.workflow, args.input)
+ finally:
+ sys.path[:] = original_path
+
+
+def _run(args: argparse.Namespace) -> int:
+ logging.basicConfig(
+ level=getattr(logging, args.log_level),
+ format="%(levelname)s %(name)s — %(message)s",
+ force=True,
+ )
+ result = _load(args).run()
+ if result.status == JobStatus.FAILED:
+ print(
+ f"Workflow failed at {result.failed_task}: {result.error}",
+ file=sys.stderr,
+ )
+ return 1
+ assert result.result is not None
+ print(result.result.model_dump_json(indent=2))
+ return 0
+
+
+def _validate(args: argparse.Namespace) -> int:
+ loaded = _load(args)
+ print(f"Workflow '{loaded.workflow.name}' is valid")
+ return 0
+
+
+def _graph(args: argparse.Namespace) -> int:
+ loaded = _load(args)
+ print(
+ loaded.workflow.to_mermaid(
+ job_configuration=loaded.job.job_configuration,
+ ),
+ end="",
+ )
+ return 0
+
+
+def build_parser() -> argparse.ArgumentParser:
+ """Build the public command-line parser."""
+ parser = argparse.ArgumentParser(prog="taskmaestro")
+ subparsers = parser.add_subparsers(dest="command", required=True)
+
+ run_parser = subparsers.add_parser("run", help="Run a YAML workflow")
+ _add_workflow_arguments(run_parser)
+ run_parser.add_argument(
+ "--log-level",
+ choices=["DEBUG", "INFO", "WARNING", "ERROR", "CRITICAL"],
+ default="INFO",
+ )
+ run_parser.set_defaults(handler=_run)
+
+ validate_parser = subparsers.add_parser("validate", help="Validate a YAML workflow")
+ _add_workflow_arguments(validate_parser)
+ validate_parser.set_defaults(handler=_validate)
+
+ graph_parser = subparsers.add_parser("graph", help="Print a Mermaid workflow graph")
+ _add_workflow_arguments(graph_parser)
+ graph_parser.set_defaults(handler=_graph)
+
+ return parser
+
+
+def main(argv: Sequence[str] | None = None) -> int:
+ """Run the Taskmaestro command-line interface."""
+ args = build_parser().parse_args(argv)
+ try:
+ handler: Any = args.handler
+ return int(handler(args))
+ except ConfigLoadError as exc:
+ print(f"Configuration error: {exc}", file=sys.stderr)
+ return 2
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..5ec7800
--- /dev/null
+++ b/taskmaestro/dependencies.py
@@ -0,0 +1,131 @@
+"""Dependency references used to collect multiple task outputs."""
+
+from __future__ import annotations
+
+from collections.abc import Mapping
+from dataclasses import dataclass, field
+from typing import Any, Generic, Literal, TypeVar, overload
+
+from pydantic import BaseModel
+
+from taskmaestro.exceptions import WorkflowDefinitionError
+from taskmaestro.task import Task
+
+O = TypeVar("O", bound=BaseModel)
+T = TypeVar("T")
+
+
+@dataclass(frozen=True)
+class OutputHandle(Generic[T]):
+ """Reference to one field of a task instance's output."""
+
+ task_name: str
+ field_name: str
+ annotation: Any
+ _owner: object = field(repr=False)
+
+
+@dataclass(frozen=True)
+class TaskHandle(Generic[O]):
+ """Unambiguous reference to one registered task instance."""
+
+ name: str
+ output_type: type[BaseModel]
+ _owner: object = field(repr=False)
+
+ def field(self, name: str) -> OutputHandle[Any]:
+ """Return a validated reference to a named output field."""
+ fields = self.output_type.model_fields
+ if name not in fields:
+ raise WorkflowDefinitionError(
+ f"Field '{name}' not found on {self.output_type.__name__} (output of {self.name})"
+ )
+ return OutputHandle(
+ task_name=self.name,
+ field_name=name,
+ annotation=fields[name].annotation,
+ _owner=self._owner,
+ )
+
+
+type TaskReference = type[Task[Any, Any]] | str | TaskHandle[Any]
+type OutputReference = TaskReference | tuple[TaskReference, str] | OutputHandle[Any]
+
+
+@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: ...
+
+
+@overload
+def collect(**members: OutputReference) -> CollectionDependency: ...
+
+
+def collect(
+ *members: OutputReference | Mapping[str, OutputReference],
+ **keyed_members: OutputReference,
+) -> CollectionDependency:
+ """Collect several upstream outputs into one list or dictionary input field.
+
+ Positional members target ``list[T]`` fields. A single mapping argument or
+ keyword arguments target ``dict[str, T]`` fields. Members may be task
+ classes, task handles, registered names, output handles, or
+ ``(task, output_field)`` references.
+ """
+ if keyed_members:
+ if members:
+ raise TypeError("collect() accepts either positional members or keyword members")
+ return CollectionDependency("keyed", keyed_members=tuple(keyed_members.items()))
+
+ 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..568df2a 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,45 @@ 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_task_configuration()
self._validate_root_task_inputs(config)
+ self._validate_task_maps()
+
+ def _validate_task_configuration(self) -> None:
+ """Ensure every declared configuration field has a supplied value."""
+ for task_name in self.workflow._tasks:
+ expected = self.workflow.get_config_fields(task_name)
+ if (
+ not expected
+ or self.workflow.get_dependencies(task_name) is not None
+ or self.workflow.is_mapped_task(task_name)
+ ):
+ continue
+ supplied = (
+ self.job_configuration.config_fields_for_task(task_name)
+ if self.job_configuration is not None
+ else set()
+ )
+ missing = expected - supplied
+ if missing:
+ raise WorkflowDefinitionError(
+ f"Task '{task_name}' is missing configuration fields {sorted(missing)}"
+ )
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 +140,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..8cc81a1 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] = {}
@@ -147,8 +196,14 @@ def to_mermaid(
# Result task → source redirect
source_redirect[task_name] = f"{task_name}__{inner_wf.result_task_name}"
- # Start and end nodes
- lines.append(' _start_(("start"))')
+ # A start node represents external job input. Fully configured and
+ # self-contained workflows have no such input, so avoid an orphan node.
+ has_start_node = any(
+ workflow.get_dependencies(name) is None and name not in configured_tasks
+ for name, _cls in tasks
+ )
+ if has_start_node:
+ lines.append(' _start_(("start"))')
lines.append(' _end_(("end"))')
# JobConfiguration node (if there are configured tasks)
@@ -193,7 +248,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 +269,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..cc0556d 100644
--- a/taskmaestro/workflow.py
+++ b/taskmaestro/workflow.py
@@ -2,27 +2,45 @@
from __future__ import annotations
-from typing import TYPE_CHECKING, Any, Union
+import types
+import typing
+from collections.abc import Mapping
+from typing import TYPE_CHECKING, Any, Literal, TypeVar, cast, get_args, get_origin
from pydantic import BaseModel
if TYPE_CHECKING:
- from taskmaestro.job import JobConfiguration
-
+ from taskmaestro.context import ExecutionContext
+ from taskmaestro.hooks.base import BaseHook
+ from taskmaestro.job import Job, JobConfiguration
+
+from taskmaestro.dependencies import (
+ CollectionDependency,
+ CollectionRef,
+ OutputHandle,
+ OutputRef,
+ OutputReference,
+ TaskHandle,
+ TaskReference,
+)
from taskmaestro.exceptions import (
CycleDetectedError,
IncompleteInputError,
WorkflowDefinitionError,
)
+from taskmaestro.mapping import MappedOutput, TaskMap
from taskmaestro.task import Task, get_input_type, get_output_type
+O = TypeVar("O", bound=BaseModel)
+
# Stored dependency types after name resolution:
# None — root task
# 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 +51,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 +157,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
@@ -87,6 +190,35 @@ def builder(
"""Return a builder for DAG construction."""
return WorkflowBuilder(name, result_task=result_task)
+ def run(
+ self,
+ input: BaseModel,
+ *,
+ task_config: JobConfiguration | dict[str, dict[str, Any]] | None = None,
+ hooks: list[BaseHook] | None = None,
+ ctx: ExecutionContext | None = None,
+ timeout_seconds: float | None = None,
+ ) -> Job[Any]:
+ """Create and run a job with sensible defaults.
+
+ ``task_config`` accepts either an existing :class:`JobConfiguration`
+ or the nested dictionary used to construct one. Use :class:`Runner`
+ and :class:`Job` directly when more control over their lifecycle is
+ required.
+ """
+ from taskmaestro.job import Job, JobConfiguration
+ from taskmaestro.runner import Runner
+
+ job_configuration = (
+ task_config
+ if isinstance(task_config, JobConfiguration)
+ else JobConfiguration(task_config)
+ if task_config is not None
+ else None
+ )
+ job = Job(self, input, job_configuration=job_configuration)
+ return Runner(hooks=hooks).run(job, ctx=ctx, timeout_seconds=timeout_seconds)
+
def topological_order(self) -> list[tuple[str, type[Task[Any, Any]]]]:
"""Return (name, task_class) pairs in a valid execution order (Kahn's algorithm)."""
in_degree: dict[str, int] = {name: 0 for name in self._tasks}
@@ -133,18 +265,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 +334,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 +394,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 +415,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 +443,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 +460,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 +478,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 +497,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 +522,26 @@ 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 self.is_mapped_task(upstream_ref) and not _is_type_compatible(
+ resolved_type, field_annotation
+ ):
+ root_type = resolved_type.model_fields["root"].annotation
+ if root_type is not None and _is_type_compatible(
+ root_type, field_annotation
+ ):
+ deps[field_name] = (upstream_ref, "root")
+ resolved_type = root_type
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 +551,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 +564,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,9 +680,11 @@ 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
+ self._result_task_ref: type[Task[Any, Any]] | str | TaskHandle[Any] | None = result_task
+ self._handle_owner = object()
def _resolve_dep_name(self, dep_cls: type[Task[Any, Any]]) -> str:
"""Resolve a class reference to its registered name.
@@ -412,32 +707,63 @@ def _resolve_dep_name(self, dep_cls: type[Task[Any, Any]]) -> str:
f"Use a string name to disambiguate."
)
- def _resolve_dep_ref(
- self,
- dep: type[Task[Any, Any]] | str,
- ) -> str:
- """Resolve a dependency reference (class or string) to a registered name.
+ def _resolve_dep_ref(self, dep: TaskReference) -> str:
+ """Resolve a class, name, or task handle to a registered name.
String references are accepted as-is (validated at build time),
allowing forward references to tasks not yet added.
"""
if isinstance(dep, str):
return dep
+ if isinstance(dep, TaskHandle):
+ self._validate_handle_owner(dep._owner)
+ if dep.name not in self._workflow._tasks:
+ raise WorkflowDefinitionError(
+ f"Task handle '{dep.name}' is not registered in this workflow"
+ )
+ return dep.name
return self._resolve_dep_name(dep)
+ def _validate_handle_owner(self, owner: object) -> None:
+ if owner is not self._handle_owner:
+ raise WorkflowDefinitionError("Task handle belongs to a different workflow builder")
+
+ def _resolve_output_reference(self, ref: OutputReference) -> OutputRef:
+ """Resolve a public task/output-field reference."""
+ if isinstance(ref, OutputHandle):
+ self._validate_handle_owner(ref._owner)
+ return OutputRef(ref.task_name, ref.field_name)
+ 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]],
*,
name: str | None = None,
depends_on: (
- 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]]
- | None
+ OutputReference | Mapping[str, OutputReference | 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.
@@ -446,11 +772,14 @@ def add_task(
``depends_on`` accepts:
- ``None`` — root task (no upstream)
- - ``TaskClass`` — single upstream, whole output
+ - ``TaskClass`` or ``TaskHandle`` — single upstream, whole output
- ``"task_name"`` — single upstream by registered name
- - ``(TaskClass | "name", "field")`` — single upstream, specific output field
- - ``{"field": TaskClass | "name", ...}`` — fan-in, whole outputs
- - ``{"field": (TaskClass | "name", "f"), ...}`` — fan-in with field routing
+ - ``(task_reference, "field")`` or ``handle.field("field")`` — output field
+ - ``{"field": task_reference, ...}`` — fan-in, whole outputs
+ - ``{"field": output_reference, ...}`` — 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
@@ -460,31 +789,115 @@ def add_task(
if depends_on is None:
wf._dependencies[task_name] = None
+ elif isinstance(depends_on, OutputHandle):
+ resolved = self._resolve_output_reference(depends_on)
+ wf._dependencies[task_name] = (
+ resolved.task_name,
+ cast(str, resolved.output_field),
+ )
elif isinstance(depends_on, tuple):
dep_ref, field = depends_on
resolved_name = self._resolve_dep_ref(dep_ref)
wf._dependencies[task_name] = (resolved_name, field)
- elif isinstance(depends_on, dict):
- resolved: dict[str, DepValue] = {}
+ elif isinstance(depends_on, Mapping):
+ resolved_dependencies: dict[str, FanInValue] = {}
for field, dep in depends_on.items():
- if isinstance(dep, tuple):
+ if isinstance(dep, CollectionDependency):
+ resolved_dependencies[field] = self._resolve_collection(dep)
+ elif isinstance(dep, OutputHandle):
+ output_ref = self._resolve_output_reference(dep)
+ resolved_dependencies[field] = (
+ output_ref.task_name,
+ cast(str, output_ref.output_field),
+ )
+ elif isinstance(dep, tuple):
dep_ref, dep_field = dep
resolved_name = self._resolve_dep_ref(dep_ref)
- resolved[field] = (resolved_name, dep_field)
+ resolved_dependencies[field] = (resolved_name, dep_field)
else:
- resolved[field] = self._resolve_dep_ref(dep)
- wf._dependencies[task_name] = resolved
- elif isinstance(depends_on, str):
- resolved_name = self._resolve_dep_ref(depends_on)
- wf._dependencies[task_name] = resolved_name
+ resolved_dependencies[field] = self._resolve_dep_ref(dep)
+ wf._dependencies[task_name] = resolved_dependencies
else:
- # Class reference
- resolved_name = self._resolve_dep_name(depends_on)
+ resolved_name = self._resolve_dep_ref(depends_on)
wf._dependencies[task_name] = resolved_name
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
+
+ def task(
+ self,
+ task_cls: type[Task[Any, O]],
+ *,
+ name: str | None = None,
+ depends_on: (
+ OutputReference | Mapping[str, OutputReference | CollectionDependency] | None
+ ) = None,
+ config_fields: list[str] | None = None,
+ mapped_over: TaskMap | None = None,
+ **input_dependencies: OutputReference | CollectionDependency,
+ ) -> TaskHandle[O]:
+ """Add a task and return an unambiguous handle to that instance.
+
+ Unlike :meth:`add_task`, this method does not return the builder and is
+ intended for local-variable-based graph construction. Named input
+ dependencies may be passed directly as keyword arguments instead of
+ through ``depends_on``.
+ """
+ if input_dependencies:
+ if depends_on is not None:
+ raise WorkflowDefinitionError(
+ "Use either depends_on or keyword input dependencies, not both"
+ )
+ depends_on = input_dependencies
+ self.add_task(
+ task_cls,
+ name=name,
+ depends_on=depends_on,
+ config_fields=config_fields,
+ mapped_over=mapped_over,
+ )
+ task_name = name if name is not None else task_cls.name
+ output_type = cast(type[BaseModel], self._workflow.get_output_annotation(task_name))
+ return TaskHandle(task_name, output_type, self._handle_owner)
+ def map_task(
+ self,
+ task_cls: type[Task[Any, O]],
+ *,
+ over: str,
+ key_as: str,
+ value_as: str,
+ error_mode: Literal["fail_fast", "collect_all"] = "fail_fast",
+ name: str | None = None,
+ depends_on: (
+ OutputReference | Mapping[str, OutputReference | CollectionDependency] | None
+ ) = None,
+ config_fields: list[str] | None = None,
+ **input_dependencies: OutputReference | CollectionDependency,
+ ) -> TaskHandle[MappedOutput[O]]:
+ """Add a task mapped over configured items and return its handle."""
+ handle = self.task(
+ task_cls,
+ name=name,
+ depends_on=depends_on,
+ config_fields=config_fields,
+ mapped_over=TaskMap(
+ over=over,
+ key_as=key_as,
+ value_as=value_as,
+ error_mode=error_mode,
+ ),
+ **input_dependencies,
+ )
+ return cast(TaskHandle[MappedOutput[O]], handle)
+
+ def set_result_task(self, task: TaskReference) -> WorkflowBuilder:
+ """Select the result task, accepting a class, name, or task handle."""
+ self._result_task_ref = task
return self
def build(self) -> Workflow:
@@ -492,9 +905,6 @@ def build(self) -> Workflow:
# Resolve result_task ref
ref = self._result_task_ref
if ref is not None:
- if isinstance(ref, str):
- self._workflow._result_task_name = ref
- else:
- self._workflow._result_task_name = self._resolve_dep_name(ref)
+ self._workflow._result_task_name = self._resolve_dep_ref(ref)
self._workflow._validate()
return self._workflow
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..5594440 100644
--- a/taskmaestro/yaml_config.py
+++ b/taskmaestro/yaml_config.py
@@ -9,20 +9,31 @@
from typing import Any
import yaml
-from pydantic import BaseModel, Field, ValidationError, model_validator
+from pydantic import BaseModel, ConfigDict, 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.task import Task
from taskmaestro.workflow import Workflow, WorkflowBuilder
# --- 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:
@@ -69,6 +81,8 @@ class ContextConfig(BaseModel):
class WorkflowSectionConfig(BaseModel):
"""Workflow section of the YAML config."""
+ model_config = ConfigDict(extra="forbid")
+
name: str
result_task: str | None = None
tasks: list[TaskConfig] = Field(min_length=1)
@@ -85,6 +99,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 +173,23 @@ 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
+
+
@dataclass(frozen=True)
class LoadedWorkflow:
"""A fully resolved workflow ready to execute."""
@@ -153,19 +220,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 +256,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 +271,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 +284,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,93 +303,149 @@ 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")
-
- # 6. Detect linear vs DAG mode
- has_depends_on = any(tc.depends_on is not None for tc in config.workflow.tasks)
+ 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}'"
+ )
- # 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)
+ 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}'"
+ )
- 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
+ # 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
)
+ # 6b. Input YAML always uses per-task configuration. Top-level keys are
+ # registered task instance names and each value is a mapping or null.
+ entry_by_name = {entry.registered_name: entry for entry in entries}
+ for key, value in raw_input.items():
+ if key not in entry_by_name:
+ raise ConfigLoadError(
+ f"Input top-level key '{key}' is not a task name "
+ f"(known tasks: {sorted(entry_by_name)})"
+ )
+ if value is not None and not isinstance(value, dict):
+ raise ConfigLoadError(
+ f"Input value for task '{key}' must be a mapping or null "
+ f"(got {type(value).__name__})"
+ )
+
per_task_data: dict[str, dict[str, Any]] = {}
per_task_cfg_fields: dict[str, list[str]] = {}
- if is_per_task_config:
- 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())
+ for task_name, task_values in raw_input.items():
+ per_task_data[task_name] = dict(task_values) if task_values else {}
+ if task_values:
+ 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,
+ 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)
- 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,
- )
- 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 +460,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,23 +485,30 @@ 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
- if is_per_task_config:
- job_configuration = JobConfiguration(per_task_data)
-
- return workflow, job_configuration
+ # 9. Input YAML always becomes per-task job configuration.
+ return workflow, JobConfiguration(per_task_data)
def load_workflow_from_yaml(workflow_path: str | Path, input_path: str | Path) -> LoadedWorkflow:
@@ -371,7 +524,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 +535,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:
@@ -400,28 +553,13 @@ def load_workflow_from_yaml(workflow_path: str | Path, input_path: str | Path) -
# 4. Build workflow and job_configuration via shared helper
workflow, job_configuration = _load_workflow_only(workflow_path, input_path)
- # 5. Validate input and build Job
- job: Job[Any]
- if job_configuration is not None:
+ # 5. Validate per-task input and build Job. The root input is EmptyConfig
+ # because all external YAML values are delivered through JobConfiguration.
+ assert job_configuration is not None
+ try:
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)
+ except WorkflowDefinitionError as exc:
+ raise ConfigLoadError(f"Job validation failed: {exc}") from exc
# 6. Instantiate hooks
hooks: list[BaseHook] = []
diff --git a/tests/test_cli.py b/tests/test_cli.py
new file mode 100644
index 0000000..c25e74a
--- /dev/null
+++ b/tests/test_cli.py
@@ -0,0 +1,112 @@
+"""Tests for the Taskmaestro command-line interface."""
+
+from __future__ import annotations
+
+import sys
+from pathlib import Path
+
+from taskmaestro.cli import main
+
+
+def _files(tmp_path: Path, task: str = "Increment") -> tuple[Path, Path]:
+ (tmp_path / "pipeline.py").write_text(
+ """\
+from pydantic import BaseModel
+from taskmaestro import ExecutionContext, Task
+
+class NumberInput(BaseModel):
+ value: int
+
+class NumberOutput(BaseModel):
+ value: int
+
+class Increment(Task[NumberInput, NumberOutput]):
+ name = "increment"
+
+ def run(self, input: NumberInput, ctx: ExecutionContext) -> NumberOutput:
+ return NumberOutput(value=input.value + 1)
+
+class Fail(Task[NumberInput, NumberOutput]):
+ name = "fail"
+
+ def run(self, input: NumberInput, ctx: ExecutionContext) -> NumberOutput:
+ raise ValueError("intentional failure")
+""",
+ encoding="utf-8",
+ )
+ sys.modules.pop("pipeline", None)
+ workflow = tmp_path / "workflow.yaml"
+ workflow.write_text(
+ f"""\
+workflow:
+ name: cli_test
+ tasks:
+ - task: pipeline.{task}
+""",
+ encoding="utf-8",
+ )
+ input_path = tmp_path / "input.yaml"
+ input_path.write_text(f"{task.lower()}:\n value: 4\n", encoding="utf-8")
+ return workflow, input_path
+
+
+def test_run_prints_json_result(tmp_path: Path, capsys: object) -> None:
+ workflow, input_path = _files(tmp_path)
+
+ status = main(
+ [
+ "run",
+ str(workflow),
+ "--input",
+ str(input_path),
+ "--log-level",
+ "DEBUG",
+ ]
+ )
+
+ assert status == 0
+ captured = capsys.readouterr() # type: ignore[attr-defined]
+ assert '"value": 5' in captured.out
+
+
+def test_run_reports_failed_job(tmp_path: Path, capsys: object) -> None:
+ workflow, input_path = _files(tmp_path, "Fail")
+
+ status = main(["run", str(workflow), "--input", str(input_path)])
+
+ assert status == 1
+ captured = capsys.readouterr() # type: ignore[attr-defined]
+ assert "Workflow failed at fail: intentional failure" in captured.err
+
+
+def test_validate_reports_success(tmp_path: Path, capsys: object) -> None:
+ workflow, input_path = _files(tmp_path)
+ original_path = sys.path.copy()
+
+ status = main(["validate", str(workflow), "--input", str(input_path)])
+
+ assert status == 0
+ assert sys.path == original_path
+ captured = capsys.readouterr() # type: ignore[attr-defined]
+ assert "Workflow 'cli_test' is valid" in captured.out
+
+
+def test_graph_prints_mermaid(tmp_path: Path, capsys: object) -> None:
+ workflow, input_path = _files(tmp_path)
+
+ status = main(["graph", str(workflow), "--input", str(input_path)])
+
+ assert status == 0
+ captured = capsys.readouterr() # type: ignore[attr-defined]
+ assert "graph TD" in captured.out
+ assert 'increment["increment"]' in captured.out
+
+
+def test_configuration_errors_return_two(tmp_path: Path, capsys: object) -> None:
+ missing = tmp_path / "missing.yaml"
+
+ status = main(["validate", str(missing), "--input", str(missing)])
+
+ assert status == 2
+ captured = capsys.readouterr() # type: ignore[attr-defined]
+ assert "Configuration error: Cannot read file" in captured.err
diff --git a/tests/test_collections.py b/tests/test_collections.py
new file mode 100644
index 0000000..f82e590
--- /dev/null
+++ b/tests/test_collections.py
@@ -0,0 +1,593 @@
+"""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_keyword_collection_accepts_task_and_output_handles(self) -> None:
+ builder = Workflow.builder("handle_collection")
+ top = builder.task(ProduceSurface, name="top_task")
+ base = builder.task(ProduceEnvelope, name="base_task")
+ builder.task(
+ CollectSurfaceDict,
+ depends_on={
+ "surfaces": collect(top=top, base=base.field("surface")),
+ },
+ )
+
+ result = Runner().run(Job(builder.build(), NumberInput(value=4)))
+
+ assert result.status == JobStatus.COMPLETED
+ assert result.result == SurfaceNames(names=["top:top_task-4", "base:base_task-4"])
+
+ def test_collect_rejects_mixed_positional_and_keyword_members(self) -> None:
+ with pytest.raises(TypeError, match="either positional members or keyword members"):
+ collect(ProduceSurface, top=ProduceSurface)
+
+ 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("first:\n value: 9\nsecond:\n 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("top_task:\n value: 5\nbase_task:\n 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("producer:\n 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="missing configuration fields"):
+ 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_discovery.py b/tests/test_discovery.py
index 71804d7..a0f4f9f 100644
--- a/tests/test_discovery.py
+++ b/tests/test_discovery.py
@@ -81,7 +81,7 @@ def test_registered_task_can_be_used_in_yaml(plugin_entry_points: None, tmp_path
"workflow:\n name: entry_point_workflow\n tasks:\n - task: example.increment\n"
)
input_path = tmp_path / "input.yaml"
- input_path.write_text("value: 4\n")
+ input_path.write_text("ExampleTask:\n value: 4\n")
result = load_workflow_from_yaml(workflow_path, input_path).run()
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..54ac73f 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
@@ -117,6 +118,16 @@ def test_root_task_with_config_fields_skips_validation(self) -> None:
job = Job(workflow=wf, config=EmptyConfig(), job_configuration=jc)
assert job.status == JobStatus.PENDING
+ def test_missing_declared_configuration_is_rejected(self) -> None:
+ workflow = (
+ Workflow.builder(name="cfg")
+ .add_task(ConfigOnlyTask, config_fields=["path", "count"])
+ .build()
+ )
+
+ with pytest.raises(WorkflowDefinitionError, match="missing configuration fields"):
+ Job(workflow=workflow, config=EmptyConfig())
+
def test_job_configuration_stored(self) -> None:
wf = (
Workflow.builder(name="cfg")
diff --git a/tests/test_mapping.py b/tests/test_mapping.py
new file mode 100644
index 0000000..101961e
--- /dev/null
+++ b/tests/test_mapping.py
@@ -0,0 +1,904 @@
+"""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
+
+ def test_mapped_task_output_is_automatically_unwrapped(self) -> None:
+ builder = Workflow.builder("mapped_handles")
+ base = builder.task(AddOne)
+ mapped = builder.map_task(
+ MappedNumber,
+ over="items",
+ key_as="item_name",
+ value_as="amount",
+ config_fields=["multiplier"],
+ base=base,
+ )
+ builder.task(SumAggregate, values=mapped)
+ workflow = builder.build()
+ job = Job(
+ workflow,
+ NumberInput(value=10),
+ job_configuration=JobConfiguration(
+ {"mapped_number": {"items": {"one": 1, "two": 2}, "multiplier": 2}}
+ ),
+ )
+
+ assert workflow.get_dependencies("sum_aggregate") == {"values": ("mapped_number", "root")}
+ assert Runner().run(job).result == NumberOutput(value=28)
+
+ @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_visualization.py b/tests/test_visualization.py
index b9e1204..5f2a6ea 100644
--- a/tests/test_visualization.py
+++ b/tests/test_visualization.py
@@ -318,9 +318,9 @@ def test_config_fields_root_task_no_start_edge(self) -> None:
)
result = to_mermaid(wf)
- # No start edge for configured root task
- assert "_start_ -->|" not in result
- # But job config node and dashed edge are present
+ # No orphan start node: JobConfiguration is the workflow's source.
+ assert "_start_" not in result
+ # The job config node and dashed edge are present.
assert '_job_config_[("JobConfiguration")]' in result
assert "_job_config_ -.->|" in result
diff --git a/tests/test_workflow.py b/tests/test_workflow.py
index 58f2c69..a9cae58 100644
--- a/tests/test_workflow.py
+++ b/tests/test_workflow.py
@@ -2,16 +2,25 @@
from __future__ import annotations
+from typing import Any
+
import pytest
from pydantic import BaseModel
from taskmaestro import (
+ EmptyConfig,
ExecutionContext,
+ Job,
+ JobConfiguration,
+ JobStatus,
+ Runner,
Task,
+ TaskHandle,
Workflow,
WorkflowDefinitionError,
)
from taskmaestro.exceptions import CycleDetectedError, IncompleteInputError
+from taskmaestro.hooks.base import BaseHook
from tests.conftest import (
AddOne,
AddOneB,
@@ -28,6 +37,53 @@
)
+class TestWorkflowRun:
+ def test_run_executes_workflow_with_context_hooks_and_timeout(self) -> None:
+ class ContextTask(Task[NumberInput, NumberOutput]):
+ name = "context_task"
+
+ def run(self, input: NumberInput, ctx: ExecutionContext) -> NumberOutput:
+ return NumberOutput(value=input.value + ctx.resolve("offset"))
+
+ class CompleteHook(BaseHook):
+ completed = False
+
+ def on_job_complete(self, job: Job[Any]) -> None:
+ self.completed = True
+
+ workflow = Workflow("convenience", tasks=[ContextTask])
+ ctx = ExecutionContext()
+ ctx.register("offset", 2)
+ hook = CompleteHook()
+
+ result = workflow.run(
+ NumberInput(value=3),
+ hooks=[hook],
+ ctx=ctx,
+ timeout_seconds=10,
+ )
+
+ assert result.status == JobStatus.COMPLETED
+ assert result.result == NumberOutput(value=5)
+ assert hook.completed
+
+ @pytest.mark.parametrize("as_object", [False, True])
+ def test_run_accepts_task_configuration_dictionary_or_object(self, as_object: bool) -> None:
+ workflow = (
+ Workflow.builder("configured_run")
+ .add_task(ConfigOnlyTask, config_fields=["path", "count"])
+ .build()
+ )
+ values = {"config_only_task": {"path": "item", "count": 2}}
+ task_config = JobConfiguration(values) if as_object else values
+
+ result = workflow.run(EmptyConfig(), task_config=task_config)
+
+ assert result.status == JobStatus.COMPLETED
+ assert result.result is not None
+ assert result.result.model_dump() == {"summary": "itemx2"}
+
+
class TestLinearWorkflow:
def test_valid_two_task_chain(self) -> None:
wf = Workflow(name="test", tasks=[AddOne, Double])
@@ -49,10 +105,175 @@ 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 TestTaskHandles:
+ def test_task_adds_task_and_returns_typed_handle(self) -> None:
+ builder = Workflow.builder("handles")
+ add_one = builder.task(AddOne)
+ builder.add_task(Double, depends_on=add_one)
+ workflow = builder.build()
+
+ assert isinstance(add_one, TaskHandle)
+ assert add_one.name == "add_one"
+ assert add_one.output_type is NumberOutput
+ result = Runner().run(Job(workflow, NumberInput(value=3)))
+ assert result.status == JobStatus.COMPLETED
+ assert result.result == NumberOutput(value=8)
+
+ def test_handles_disambiguate_repeated_task_class(self) -> None:
+ builder = Workflow.builder("named_handles")
+ first = builder.task(AddOne, name="first")
+ second = builder.task(AddOne, name="second")
+ builder.task(FanInTask, a=first, b=second)
+
+ result = Runner().run(Job(builder.build(), NumberInput(value=2)))
+
+ assert result.result == FanInOutput(total=6)
+
+ def test_keyword_dependencies_cannot_be_mixed_with_depends_on(self) -> None:
+ builder = Workflow.builder("mixed_dependencies")
+ first = builder.task(AddOne)
+ second = builder.task(AddOneB)
+
+ with pytest.raises(WorkflowDefinitionError, match="either depends_on or keyword"):
+ builder.task(FanInTask, depends_on={"a": first}, b=second)
+
+ def test_output_field_handle_routes_field(self) -> None:
+ class Envelope(BaseModel):
+ number: NumberOutput
+
+ class ProduceEnvelope(Task[NumberInput, Envelope]):
+ name = "produce_handle_envelope"
+
+ def run(self, input: NumberInput, ctx: ExecutionContext) -> Envelope:
+ return Envelope(number=NumberOutput(value=input.value + 1))
+
+ builder = Workflow.builder("field_handle")
+ envelope = builder.task(ProduceEnvelope)
+ builder.task(Double, depends_on=envelope.field("number"))
+
+ result = Runner().run(Job(builder.build(), NumberInput(value=3)))
+ assert result.result == NumberOutput(value=8)
+
+ def test_output_field_handle_can_be_used_in_fan_in(self) -> None:
+ class Envelope(BaseModel):
+ number: NumberOutput
+
+ class ProduceEnvelope(Task[NumberInput, Envelope]):
+ name = "produce_fan_in_envelope"
+
+ def run(self, input: NumberInput, ctx: ExecutionContext) -> Envelope:
+ return Envelope(number=NumberOutput(value=input.value + 1))
+
+ builder = Workflow.builder("field_handle_fan_in")
+ envelope = builder.task(ProduceEnvelope)
+ other = builder.task(AddOneB)
+ builder.task(
+ FanInTask,
+ depends_on={"a": envelope.field("number"), "b": other},
+ )
+
+ result = builder.build().run(NumberInput(value=3))
+ assert result.result == FanInOutput(total=8)
+
+ def test_output_field_handle_rejects_unknown_field_immediately(self) -> None:
+ builder = Workflow.builder("bad_field_handle")
+ handle = builder.task(AddOne)
+
+ with pytest.raises(WorkflowDefinitionError, match="Field 'missing' not found"):
+ handle.field("missing")
+
+ def test_handle_from_another_builder_is_rejected(self) -> None:
+ first_builder = Workflow.builder("first")
+ foreign = first_builder.task(AddOne)
+ second_builder = Workflow.builder("second")
+
+ with pytest.raises(WorkflowDefinitionError, match="different workflow builder"):
+ second_builder.task(Double, depends_on=foreign)
+
+ def test_stale_handle_is_rejected(self) -> None:
+ builder = Workflow.builder("stale")
+ stale = builder.task(AddOne)
+ del builder._workflow._tasks[stale.name]
+
+ with pytest.raises(WorkflowDefinitionError, match="is not registered"):
+ builder.task(Double, depends_on=stale)
+
+ def test_handle_can_select_result_task(self) -> None:
+ builder = Workflow.builder("explicit_result")
+ selected = builder.task(AddOne)
+ builder.task(AddOneB)
+ workflow = builder.set_result_task(selected).build()
+
+ assert workflow.result_task_name == "add_one"
+
class TestDAGWorkflow:
def test_fan_in_workflow(self) -> None:
@@ -219,6 +440,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 +805,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..30bde67 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,59 @@ 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,
+ "upper_text: {<<: &defaults {text: default}, text: 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")
@@ -245,7 +300,7 @@ def test_linear_workflow_end_to_end(self, tmp_path: Path) -> None:
- task: {THIS_MODULE}.ReverseText
""",
)
- in_path = _write_input_yaml(tmp_path, "text: hello\n")
+ in_path = _write_input_yaml(tmp_path, "upper_text:\n text: hello\n")
loaded = load_workflow_from_yaml(wf_path, in_path)
assert loaded.workflow.name == "linear_test"
@@ -274,7 +329,7 @@ def test_dag_workflow_end_to_end(self, tmp_path: Path) -> None:
length: {THIS_MODULE}.TextLength
""",
)
- in_path = _write_input_yaml(tmp_path, "text: hello\n")
+ in_path = _write_input_yaml(tmp_path, "upper_text:\n text: hello\n")
loaded = load_workflow_from_yaml(wf_path, in_path)
result = loaded.run()
@@ -293,7 +348,7 @@ def test_minimal_config(self, tmp_path: Path) -> None:
- task: {THIS_MODULE}.UpperText
""",
)
- in_path = _write_input_yaml(tmp_path, "text: hello\n")
+ in_path = _write_input_yaml(tmp_path, "upper_text:\n text: hello\n")
loaded = load_workflow_from_yaml(wf_path, in_path)
result = loaded.run()
assert result.status == JobStatus.COMPLETED
@@ -315,7 +370,7 @@ def test_hooks_with_params(self, tmp_path: Path) -> None:
- hook: taskmaestro.hooks.timing.TimingHook
""",
)
- in_path = _write_input_yaml(tmp_path, "text: hello\n")
+ in_path = _write_input_yaml(tmp_path, "upper_text:\n text: hello\n")
loaded = load_workflow_from_yaml(wf_path, in_path)
assert len(loaded.runner.hooks) == 2
@@ -335,7 +390,7 @@ def test_persistence_hook_path_coercion(self, tmp_path: Path) -> None:
output_dir: "{output_dir}"
""",
)
- in_path = _write_input_yaml(tmp_path, "text: hello\n")
+ in_path = _write_input_yaml(tmp_path, "upper_text:\n text: hello\n")
loaded = load_workflow_from_yaml(wf_path, in_path)
result = loaded.run()
assert result.status == JobStatus.COMPLETED
@@ -356,7 +411,7 @@ def test_context_services(self, tmp_path: Path) -> None:
name: test
""",
)
- in_path = _write_input_yaml(tmp_path, "text: hello\n")
+ in_path = _write_input_yaml(tmp_path, "upper_text:\n text: hello\n")
loaded = load_workflow_from_yaml(wf_path, in_path)
assert loaded.context.correlation_id == "test-run-42"
assert loaded.context.resolve("multiplier") == 3
@@ -375,7 +430,7 @@ def test_context_scratch_dir(self, tmp_path: Path) -> None:
scratch_dir: "{scratch}"
""",
)
- in_path = _write_input_yaml(tmp_path, "text: hello\n")
+ in_path = _write_input_yaml(tmp_path, "upper_text:\n text: hello\n")
loaded = load_workflow_from_yaml(wf_path, in_path)
assert loaded.context.scratch_dir == scratch
@@ -389,7 +444,7 @@ def test_bad_import_path(self, tmp_path: Path) -> None:
- task: nonexistent.module.BadTask
""",
)
- in_path = _write_input_yaml(tmp_path, "text: hello\n")
+ in_path = _write_input_yaml(tmp_path, "upper_text:\n text: hello\n")
with pytest.raises(ConfigLoadError, match="Cannot import module"):
load_workflow_from_yaml(wf_path, in_path)
@@ -403,7 +458,7 @@ def test_not_a_task_class(self, tmp_path: Path) -> None:
- task: {THIS_MODULE}.TextInput
""",
)
- in_path = _write_input_yaml(tmp_path, "text: hello\n")
+ in_path = _write_input_yaml(tmp_path, "upper_text:\n text: hello\n")
with pytest.raises(ConfigLoadError, match="not a Task subclass"):
load_workflow_from_yaml(wf_path, in_path)
@@ -418,7 +473,7 @@ def test_bad_yaml_syntax(self, tmp_path: Path) -> None:
input:
""",
)
- in_path = _write_input_yaml(tmp_path, "text: hello\n")
+ in_path = _write_input_yaml(tmp_path, "upper_text:\n text: hello\n")
with pytest.raises(ConfigLoadError, match="YAML parse error"):
load_workflow_from_yaml(wf_path, in_path)
@@ -432,13 +487,13 @@ def test_input_validation_error(self, tmp_path: Path) -> None:
- task: {THIS_MODULE}.UpperText
""",
)
- in_path = _write_input_yaml(tmp_path, "wrong_field: 123\n")
- with pytest.raises(ConfigLoadError, match="Input validation error"):
+ in_path = _write_input_yaml(tmp_path, "upper_text:\n wrong_field: 123\n")
+ with pytest.raises(ConfigLoadError, match="Workflow validation failed"):
load_workflow_from_yaml(wf_path, in_path)
def test_file_not_found(self, tmp_path: Path) -> None:
wf_path = tmp_path / "nonexistent.yaml"
- in_path = _write_input_yaml(tmp_path, "text: hello\n")
+ in_path = _write_input_yaml(tmp_path, "upper_text:\n text: hello\n")
with pytest.raises(ConfigLoadError, match="Cannot read file"):
load_workflow_from_yaml(wf_path, in_path)
@@ -496,7 +551,7 @@ def test_explicit_result_task(self, tmp_path: Path) -> None:
- task: {THIS_MODULE}.ReverseText
""",
)
- in_path = _write_input_yaml(tmp_path, "text: hello\n")
+ in_path = _write_input_yaml(tmp_path, "upper_text:\n text: hello\n")
loaded = load_workflow_from_yaml(wf_path, in_path)
assert loaded.workflow.result_task.name == "upper_text"
@@ -512,7 +567,7 @@ def test_dependency_not_found(self, tmp_path: Path) -> None:
depends_on: nonexistent.module.Missing
""",
)
- in_path = _write_input_yaml(tmp_path, "text: hello\n")
+ in_path = _write_input_yaml(tmp_path, "upper_text:\n text: hello\n")
with pytest.raises(ConfigLoadError, match="not found"):
load_workflow_from_yaml(wf_path, in_path)
@@ -529,7 +584,7 @@ def test_not_a_hook_class(self, tmp_path: Path) -> None:
- hook: {THIS_MODULE}.TextInput
""",
)
- in_path = _write_input_yaml(tmp_path, "text: hello\n")
+ in_path = _write_input_yaml(tmp_path, "upper_text:\n text: hello\n")
with pytest.raises(ConfigLoadError, match="not a BaseHook subclass"):
load_workflow_from_yaml(wf_path, in_path)
@@ -548,19 +603,19 @@ def test_hook_bad_params(self, tmp_path: Path) -> None:
nonexistent_param: true
""",
)
- in_path = _write_input_yaml(tmp_path, "text: hello\n")
+ in_path = _write_input_yaml(tmp_path, "upper_text:\n text: hello\n")
with pytest.raises(ConfigLoadError, match="Cannot instantiate hook"):
load_workflow_from_yaml(wf_path, in_path)
def test_workflow_yaml_not_a_mapping(self, tmp_path: Path) -> None:
wf_path = _write_workflow_yaml(tmp_path, "- item1\n- item2\n")
- in_path = _write_input_yaml(tmp_path, "text: hello\n")
+ in_path = _write_input_yaml(tmp_path, "upper_text:\n text: hello\n")
with pytest.raises(ConfigLoadError, match="must contain a mapping"):
load_workflow_from_yaml(wf_path, in_path)
def test_schema_validation_error(self, tmp_path: Path) -> None:
wf_path = _write_workflow_yaml(tmp_path, "workflow:\n name: test\n")
- in_path = _write_input_yaml(tmp_path, "text: hello\n")
+ in_path = _write_input_yaml(tmp_path, "upper_text:\n text: hello\n")
with pytest.raises(ConfigLoadError, match="YAML schema validation error"):
load_workflow_from_yaml(wf_path, in_path)
@@ -579,7 +634,7 @@ def test_list_depends_on_field_routing(self, tmp_path: Path) -> None:
- inner
""",
)
- in_path = _write_input_yaml(tmp_path, "text: hello\n")
+ in_path = _write_input_yaml(tmp_path, "wrap_text:\n text: hello\n")
loaded = load_workflow_from_yaml(wf_path, in_path)
result = loaded.run()
assert result.status == JobStatus.COMPLETED
@@ -606,7 +661,7 @@ def test_dict_fan_in_with_list_field_ref(self, tmp_path: Path) -> None:
length: {THIS_MODULE}.TextLength
""",
)
- in_path = _write_input_yaml(tmp_path, "text: hello\n")
+ in_path = _write_input_yaml(tmp_path, "wrap_text:\n text: hello\n")
loaded = load_workflow_from_yaml(wf_path, in_path)
result = loaded.run()
assert result.status == JobStatus.COMPLETED
@@ -632,7 +687,7 @@ def test_dict_fan_in_list_ref_invalid_length(self, tmp_path: Path) -> None:
length: {THIS_MODULE}.TextLength
""",
)
- in_path = _write_input_yaml(tmp_path, "text: hello\n")
+ in_path = _write_input_yaml(tmp_path, "wrap_text:\n text: hello\n")
with pytest.raises(ConfigLoadError, match="List dep must be"):
load_workflow_from_yaml(wf_path, in_path)
@@ -655,7 +710,7 @@ def test_dict_fan_in_list_ref_not_found(self, tmp_path: Path) -> None:
length: {THIS_MODULE}.TextLength
""",
)
- in_path = _write_input_yaml(tmp_path, "text: hello\n")
+ in_path = _write_input_yaml(tmp_path, "wrap_text:\n text: hello\n")
with pytest.raises(ConfigLoadError, match="not found"):
load_workflow_from_yaml(wf_path, in_path)
@@ -674,7 +729,7 @@ def test_fan_in_string_dep_not_found(self, tmp_path: Path) -> None:
length: {THIS_MODULE}.UpperText
""",
)
- in_path = _write_input_yaml(tmp_path, "text: hello\n")
+ in_path = _write_input_yaml(tmp_path, "upper_text:\n text: hello\n")
with pytest.raises(ConfigLoadError, match="not found"):
load_workflow_from_yaml(wf_path, in_path)
@@ -693,7 +748,7 @@ def test_workflow_validation_failed(self, tmp_path: Path) -> None:
depends_on: {THIS_MODULE}.UpperText
""",
)
- in_path = _write_input_yaml(tmp_path, "text: hello\n")
+ in_path = _write_input_yaml(tmp_path, "upper_text:\n text: hello\n")
with pytest.raises(ConfigLoadError, match="Workflow validation failed"):
load_workflow_from_yaml(wf_path, in_path)
@@ -709,7 +764,7 @@ def test_timeout_seconds_passed(self, tmp_path: Path) -> None:
timeout_seconds: 60
""",
)
- in_path = _write_input_yaml(tmp_path, "text: hello\n")
+ in_path = _write_input_yaml(tmp_path, "upper_text:\n text: hello\n")
loaded = load_workflow_from_yaml(wf_path, in_path)
assert loaded._timeout_seconds == 60
@@ -735,7 +790,7 @@ def test_list_depends_on_end_to_end(self, tmp_path: Path) -> None:
depends_on: {THIS_MODULE}.UpperText
""",
)
- in_path = _write_input_yaml(tmp_path, "text: hello\n")
+ in_path = _write_input_yaml(tmp_path, "upper_text:\n text: hello\n")
loaded = load_workflow_from_yaml(wf_path, in_path)
result = loaded.run()
assert result.status == JobStatus.COMPLETED
@@ -760,7 +815,7 @@ def test_dict_with_list_field_ref(self, tmp_path: Path) -> None:
length: {THIS_MODULE}.TextLength
""",
)
- in_path = _write_input_yaml(tmp_path, "text: hello\n")
+ in_path = _write_input_yaml(tmp_path, "upper_text:\n text: hello\n")
loaded = load_workflow_from_yaml(wf_path, in_path)
result = loaded.run()
assert result.status == JobStatus.COMPLETED
@@ -782,7 +837,7 @@ def test_invalid_list_length_raises(self, tmp_path: Path) -> None:
- extra
""",
)
- in_path = _write_input_yaml(tmp_path, "text: hello\n")
+ in_path = _write_input_yaml(tmp_path, "upper_text:\n text: hello\n")
with pytest.raises(ConfigLoadError, match="List depends_on must be"):
load_workflow_from_yaml(wf_path, in_path)
@@ -801,7 +856,7 @@ def test_list_dep_not_found_raises(self, tmp_path: Path) -> None:
- text
""",
)
- in_path = _write_input_yaml(tmp_path, "text: hello\n")
+ in_path = _write_input_yaml(tmp_path, "upper_text:\n text: hello\n")
with pytest.raises(ConfigLoadError, match="not found"):
load_workflow_from_yaml(wf_path, in_path)
@@ -818,7 +873,7 @@ def test_convenience_function(self, tmp_path: Path) -> None:
- task: {THIS_MODULE}.ReverseText
""",
)
- in_path = _write_input_yaml(tmp_path, "text: hello\n")
+ in_path = _write_input_yaml(tmp_path, "upper_text:\n text: hello\n")
job = run_workflow_from_yaml(wf_path, in_path)
assert job.status == JobStatus.COMPLETED
assert job.result is not None
@@ -841,7 +896,7 @@ def test_run_method(self, tmp_path: Path) -> None:
- task: {THIS_MODULE}.UpperText
""",
)
- in_path = _write_input_yaml(tmp_path, "text: hello\n")
+ in_path = _write_input_yaml(tmp_path, "upper_text:\n text: hello\n")
loaded = load_workflow_from_yaml(wf_path, in_path)
result = loaded.run()
assert result.status == JobStatus.COMPLETED
@@ -856,7 +911,7 @@ def test_components_accessible(self, tmp_path: Path) -> None:
- task: {THIS_MODULE}.UpperText
""",
)
- in_path = _write_input_yaml(tmp_path, "text: hello\n")
+ in_path = _write_input_yaml(tmp_path, "upper_text:\n text: hello\n")
loaded = load_workflow_from_yaml(wf_path, in_path)
assert loaded.workflow is not None
assert loaded.runner is not None
@@ -873,7 +928,7 @@ def test_frozen(self, tmp_path: Path) -> None:
- task: {THIS_MODULE}.UpperText
""",
)
- in_path = _write_input_yaml(tmp_path, "text: hello\n")
+ in_path = _write_input_yaml(tmp_path, "upper_text:\n text: hello\n")
loaded = load_workflow_from_yaml(wf_path, in_path)
with pytest.raises(AttributeError):
loaded.workflow = None # type: ignore[misc]
@@ -887,6 +942,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, "upper_text:\n 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, "upper_text:\n 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, "upper_text:\n 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, "upper_text:\n 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, "upper_text:\n 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(
@@ -905,7 +1107,10 @@ def test_named_instances_yaml(self, tmp_path: Path) -> None:
depends_on: upper_1
""",
)
- in_path = _write_input_yaml(tmp_path, "text: hello\n")
+ in_path = _write_input_yaml(
+ tmp_path,
+ "upper_1:\n text: hello\nupper_2:\n text: hello\n",
+ )
loaded = load_workflow_from_yaml(wf_path, in_path)
result = loaded.run()
assert result.status == JobStatus.COMPLETED
@@ -933,7 +1138,7 @@ def test_named_instances_fan_in(self, tmp_path: Path) -> None:
length: length_1
""",
)
- in_path = _write_input_yaml(tmp_path, "text: hello\n")
+ in_path = _write_input_yaml(tmp_path, "upper_text:\n text: hello\n")
loaded = load_workflow_from_yaml(wf_path, in_path)
result = loaded.run()
assert result.status == JobStatus.COMPLETED
@@ -953,8 +1158,8 @@ def test_result_task_not_found_raises(self, tmp_path: Path) -> None:
depends_on: {THIS_MODULE}.UpperText
""",
)
- in_path = _write_input_yaml(tmp_path, "text: hello\n")
- with pytest.raises(ConfigLoadError, match=r"result_task.*not found"):
+ in_path = _write_input_yaml(tmp_path, "upper_text:\n text: hello\n")
+ 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:
@@ -972,7 +1177,7 @@ def test_named_result_task(self, tmp_path: Path) -> None:
depends_on: my_upper
""",
)
- in_path = _write_input_yaml(tmp_path, "text: hello\n")
+ in_path = _write_input_yaml(tmp_path, "my_upper:\n text: hello\n")
loaded = load_workflow_from_yaml(wf_path, in_path)
assert loaded.workflow.result_task_name == "my_upper"
@@ -1061,6 +1266,133 @@ 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 TestPerTaskInputFormat:
+ """Input YAML is always keyed by registered task instance name."""
+
+ def test_task_keyed_root_input(self, tmp_path: Path) -> None:
+ wf_path = _write_workflow_yaml(
+ tmp_path,
+ f"""\
+workflow:
+ name: task_keyed
+ tasks:
+ - task: {THIS_MODULE}.AmbiguousRoot
+""",
+ )
+ 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_flat_input_is_rejected(self, tmp_path: Path) -> None:
+ wf_path = _write_workflow_yaml(
+ tmp_path,
+ f"""\
+workflow:
+ name: task_keyed
+ tasks:
+ - task: {THIS_MODULE}.UpperText
+""",
+ )
+ in_path = _write_input_yaml(tmp_path, "text: hello\n")
+
+ with pytest.raises(ConfigLoadError, match="top-level key 'text' is not a task name"):
+ load_workflow_from_yaml(wf_path, in_path)
+
+ def test_unknown_task_key_is_rejected(self, tmp_path: Path) -> None:
+ wf_path = _write_workflow_yaml(
+ tmp_path,
+ f"""\
+workflow:
+ name: strict
+ 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 top-level key 'per_task_rooot' is not a task name "
+ r"\(known tasks: \['per_task_root'\]\)"
+ ),
+ ):
+ load_workflow_from_yaml(wf_path, in_path)
+
+ def test_task_value_must_be_mapping_or_null(self, tmp_path: Path) -> None:
+ wf_path = _write_workflow_yaml(
+ tmp_path,
+ f"""\
+workflow:
+ name: strict
+ tasks:
+ - task: {THIS_MODULE}.PerTaskRoot
+""",
+ )
+ in_path = _write_input_yaml(tmp_path, "per_task_root: 42\n")
+
+ with pytest.raises(
+ ConfigLoadError,
+ match=r"Input value for task 'per_task_root' must be a mapping or null \(got int\)",
+ ):
+ load_workflow_from_yaml(wf_path, in_path)
+
+ def test_null_task_value_is_empty_configuration(self, tmp_path: Path) -> None:
+ wf_path = _write_workflow_yaml(
+ tmp_path,
+ f"""\
+workflow:
+ name: strict
+ tasks:
+ - task: {THIS_MODULE}.PerTaskRoot
+""",
+ )
+ in_path = _write_input_yaml(tmp_path, "per_task_root:\n")
+
+ 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_input_mode_is_no_longer_supported(self, tmp_path: Path) -> None:
+ wf_path = _write_workflow_yaml(
+ tmp_path,
+ f"""\
+workflow:
+ name: old_mode
+ input_mode: per_task
+ tasks:
+ - task: {THIS_MODULE}.UpperText
+""",
+ )
+ in_path = _write_input_yaml(tmp_path, "upper_text:\n 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."""
@@ -1114,8 +1446,8 @@ def test_per_task_with_dag(self, tmp_path: Path) -> None:
assert result.status == JobStatus.COMPLETED
assert result.result.result == "my_label:/data/model.egrid:False" # type: ignore[union-attr]
- def test_flat_config_backward_compat(self, tmp_path: Path) -> None:
- """Flat config format still works when keys don't match task names."""
+ def test_linear_root_uses_task_keyed_input(self, tmp_path: Path) -> None:
+ """Linear workflows receive root input through the root task key."""
wf_path = _write_workflow_yaml(
tmp_path,
f"""\
@@ -1126,7 +1458,7 @@ def test_flat_config_backward_compat(self, tmp_path: Path) -> None:
- task: {THIS_MODULE}.ReverseText
""",
)
- in_path = _write_input_yaml(tmp_path, "text: hello\n")
+ in_path = _write_input_yaml(tmp_path, "upper_text:\n text: hello\n")
loaded = load_workflow_from_yaml(wf_path, in_path)
result = loaded.run()
assert result.status == JobStatus.COMPLETED
@@ -1161,6 +1493,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, "shout:\n 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, "upper_text:\n 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, "upper_text:\n 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, "shout:\n 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, "upper_text:\n 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 +1648,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", "upper_text:\n 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", "upper_text:\n 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", "upper_text:\n 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(
@@ -1192,7 +1749,7 @@ def test_workflow_ref_basic(self, tmp_path: Path) -> None:
depends_on: sub_pipeline
""",
)
- in_path = self._write_yaml(tmp_path / "input.yaml", "text: hello\n")
+ in_path = self._write_yaml(tmp_path / "input.yaml", "sub_pipeline:\n text: hello\n")
loaded = load_workflow_from_yaml(outer_path, in_path)
result = loaded.run()
@@ -1267,7 +1824,7 @@ def test_workflow_ref_path_resolution(self, tmp_path: Path) -> None:
name: sub
""",
)
- in_path = self._write_yaml(tmp_path / "input.yaml", "text: world\n")
+ in_path = self._write_yaml(tmp_path / "input.yaml", "sub:\n text: world\n")
loaded = load_workflow_from_yaml(outer_path, in_path)
result = loaded.run()
@@ -1304,7 +1861,7 @@ def test_inner_workflow_bad_yaml(self, tmp_path: Path) -> None:
depends_on: sub
""",
)
- in_path = self._write_yaml(tmp_path / "input.yaml", "text: hello\n")
+ in_path = self._write_yaml(tmp_path / "input.yaml", "upper_text:\n text: hello\n")
with pytest.raises(ConfigLoadError, match="YAML parse error"):
load_workflow_from_yaml(outer_path, in_path)
@@ -1322,7 +1879,7 @@ def test_inner_workflow_file_not_found(self, tmp_path: Path) -> None:
depends_on: sub
""",
)
- in_path = self._write_yaml(tmp_path / "input.yaml", "text: hello\n")
+ in_path = self._write_yaml(tmp_path / "input.yaml", "upper_text:\n text: hello\n")
with pytest.raises(ConfigLoadError, match="Cannot read file"):
load_workflow_from_yaml(outer_path, in_path)
@@ -1341,7 +1898,7 @@ def test_inner_workflow_not_a_mapping(self, tmp_path: Path) -> None:
depends_on: sub
""",
)
- in_path = self._write_yaml(tmp_path / "input.yaml", "text: hello\n")
+ in_path = self._write_yaml(tmp_path / "input.yaml", "upper_text:\n text: hello\n")
with pytest.raises(ConfigLoadError, match="must contain a mapping"):
load_workflow_from_yaml(outer_path, in_path)
@@ -1360,7 +1917,7 @@ def test_inner_workflow_schema_error(self, tmp_path: Path) -> None:
depends_on: sub
""",
)
- in_path = self._write_yaml(tmp_path / "input.yaml", "text: hello\n")
+ in_path = self._write_yaml(tmp_path / "input.yaml", "upper_text:\n text: hello\n")
with pytest.raises(ConfigLoadError, match="YAML schema validation error"):
load_workflow_from_yaml(outer_path, in_path)