diff --git a/CLAUDE.md b/AGENTS.md similarity index 70% rename from CLAUDE.md rename to AGENTS.md index 0439165..f8b0969 100644 --- a/CLAUDE.md +++ b/AGENTS.md @@ -1,7 +1,9 @@ -# CLAUDE.md +# AGENTS.md Typed DAG task workflow library with Pydantic models, lifecycle hooks, and fail-fast semantics. +For a runnable CLI/YAML example and the validate → graph → run loop, see [docs/agent-quickstart.md](docs/agent-quickstart.md). Run its commands from the repository root. + ## Commands ```bash @@ -21,6 +23,10 @@ mypy taskmaestro # type check (strict mode) | `taskmaestro/context.py` | `ExecutionContext` with correlation ID, logger, scratch dir, service registry | | `taskmaestro/task.py` | `Task[I, O]` ABC, type introspection (`get_input_type`, `get_output_type`) | | `taskmaestro/workflow.py` | `Workflow` (linear + DAG), `WorkflowBuilder`, validation (cycles, types, fan-in) | +| `taskmaestro/dependencies.py`, `taskmaestro/mapping.py` | Task handles, output references, `collect()`, mapped-task configuration | +| `taskmaestro/workflow_task.py` | Nested workflows wrapped as tasks | +| `taskmaestro/yaml_config.py` | YAML parsing, task imports, workflow and input validation | +| `taskmaestro/cli.py`, `taskmaestro/discovery.py` | CLI (`run`, `validate`, `graph`, `tasks list/describe`, `workflow describe`), plugin entry-point discovery | | `taskmaestro/job.py` | `Job[C]`, `JobStatus`, `TaskStatus`, `TaskResult` dataclass | | `taskmaestro/runner.py` | `Runner` — topological execution, timeout via `signal.alarm`, hook dispatch | | `taskmaestro/hooks/base.py` | `Event` StrEnum, `Hook` protocol, `BaseHook` no-op base | @@ -32,7 +38,7 @@ 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 +- **Timeouts**: `signal.setitimer`/`SIGALRM` (Unix only, main thread); gracefully warns if unavailable. All armed timeouts go through the process-wide `_AlarmScheduler`, which always arms the nearest expiry, so nested runs (`workflow_task`) cannot cancel or extend an enclosing run's timeouts - **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 @@ -40,6 +46,6 @@ mypy taskmaestro # type check (strict mode) ## Testing Conventions - Shared fixtures and reusable tasks/models in `tests/conftest.py` -- Tests organized by module: `test_exceptions`, `test_context`, `test_task`, `test_workflow`, `test_job`, `test_runner`, `test_hooks` +- Tests organized by module in `tests/test_*.py`, including CLI, YAML, mapping, nested workflows, and discovery - Timeout tests skip on non-Unix (no `signal.SIGALRM`) - Use `RecordingHook` pattern to assert event sequences diff --git a/README.md b/README.md index f96ff4f..c6df63e 100644 --- a/README.md +++ b/README.md @@ -6,6 +6,8 @@ A Python 3.12+ library for defining and executing typed DAG task workflows with Pydantic models, lifecycle hooks, and fail-fast semantics. +For an agent-friendly, runnable YAML example with validation and troubleshooting steps, see the [agent quickstart](docs/agent-quickstart.md). + ## Installation ```bash @@ -199,6 +201,8 @@ job = Job(workflow=workflow, config=EmptyConfig(), job_configuration=job_config) result = Runner().run(job) ``` +Every declared config field, on root, dependent and mapped tasks alike, must have a value in the `JobConfiguration`; otherwise `Job(...)` raises `WorkflowDefinitionError`. + ## Nested Workflows A workflow can be wrapped as a typed task and used inside a larger workflow. Its input @@ -470,6 +474,25 @@ without importing plugin modules, or `get_registered_task(name)` and `get_registered_workflow(name)` to load one plugin. Duplicate names and invalid plugin types raise `PluginLoadError`. +Agents can inspect installed task plugins from the CLI without loading every plugin: + +```bash +taskmaestro tasks list --json +taskmaestro tasks describe acme.prepare --json +``` + +`list --json` emits `{"tasks": ["acme.prepare", ...]}` in sorted order; an empty list +means no task plugins are installed. `describe --json` loads only the named plugin +and emits its registered `identifier`, task `name`, `timeout_seconds`, and Pydantic +`input_schema` / `output_schema` (JSON Schema objects). These commands inspect +**installed entry points**, not task classes local to a workflow YAML file. Without +`--json`, `list` prints one identifier per line and `describe` prints indented JSON. +For runtime-only Python objects (such as `ObjectModel[rips.EclipseCase]`), schema +fields include `"not": {}`, `"x-taskmaestro-opaque": true`, and +`"x-taskmaestro-python-type"`. They cannot be supplied as JSON; wire them from +upstream tasks or a Python context instead. Other unsupported schema constructs, +unknown identifiers, and invalid plugins report an error on stderr and exit with status 2. + ## YAML Configuration Workflows can be defined entirely in YAML instead of Python. A `task:` value may be @@ -618,9 +641,62 @@ Installed packages provide a `taskmaestro` command for YAML workflows: taskmaestro validate workflow.yaml --input input.yaml taskmaestro graph workflow.yaml --input input.yaml taskmaestro run workflow.yaml --input input.yaml --log-level INFO +taskmaestro workflow describe workflow.yaml --json ``` -`run` prints the final output as JSON and returns a nonzero exit code when the workflow fails. `graph` prints Mermaid markup. +By default, `run` prints the final output as JSON and reports errors on stderr; +`validate` prints a human-readable confirmation. `graph` prints Mermaid markup. + +For automation, use `--json` with `validate` or `run`: + +```bash +taskmaestro validate workflow.yaml --input input.yaml --json +taskmaestro run workflow.yaml --input input.yaml --json +``` + +Successful validation emits `{"status":"valid","workflow":"..."}`. A successful +run emits `{"status":"completed","workflow":"...","result":{...}}`. A failure +emits one JSON object on stdout, for example: + +```json +{"status":"failed","workflow":"example","failed_task":"prepare","error":{"code":"task_failed","type":"ValueError","message":"Task failed","task":"prepare","field":null,"issues":[]}} +``` + +Errors contain `code`, exception `type`, a safe `message`, nullable `task` and +`field`, and `issues` (field paths and error codes, without input values). +Loading failures have status `invalid` and `code: "configuration_error"`; task +failures have status `failed` and `code: "task_failed"`. If a completed result +cannot be encoded as JSON (e.g. a Python-only object), `run --json` returns +`code: "serialization_error"`. Missing configuration fields produce `issues` +with code `missing`; when no structured field metadata is available, `field` +is `null`. Use text mode when you need the original exception message. + +Exit codes: `0` success, `1` task or result-serialization failure, `2` workflow +configuration failure. In JSON mode logs and ordinary Python `print()` output +from imports/tasks go to stderr, reserving stdout for the result document. +Application code can still write directly to file descriptor 1; JSON mode is +not a sandbox. Error objects intentionally omit raw exception messages, but +application-generated stderr may contain sensitive data. + +Inspect a workflow **before its input file is complete** with +`taskmaestro workflow describe workflow.yaml --json`. The JSON contains the +workflow name and result task plus topologically ordered task instances. Each +instance includes its `name`, `python_type`, `depends_on` references (`task` / +`field`), `config_fields`, `required_input_fields`, optional `map` configuration, +and Pydantic input/output schemas. Collection dependencies include their kind +(`positional` or `keyed`) and member references. `python_type` identifies the +loaded class; it may differ from the plugin entry-point identifier used in YAML. + +Pass `--input input.yaml` to check that required configuration fields and map +sources are supplied. Then `provided_config_fields` lists **field names only**; +without `--input`, it is `null`. This check does not run tasks or hooks and does +not validate every configured value's runtime type. Loading YAML still imports +Python modules (and nested workflows); **do not inspect untrusted YAML or plugins** +under a privileged account. In JSON mode inspection failures have status +`invalid`, an `error` object, and exit code 2. Without `--json`, inspection +prints indented JSON. + +`python -m taskmaestro ...` is equivalent, which is useful when the scripts directory is not on `PATH`. ## Examples diff --git a/docs/agent-quickstart.md b/docs/agent-quickstart.md new file mode 100644 index 0000000..26d25ad --- /dev/null +++ b/docs/agent-quickstart.md @@ -0,0 +1,102 @@ +# Agent quickstart + +Use this as a short, reproducible path from a typed task to a validated workflow. Commands below run from the **repository root**. The example uses only Taskmaestro's declared dependencies; it needs no API keys or external services. + +## Set up + +Python 3.12+ is required. In a fresh checkout: + +```bash +python3 -m venv .venv +.venv/bin/python -m pip install -e . +``` + +If you already have Taskmaestro installed in another environment, use that environment's Python instead of `.venv/bin/python` below. + +## Inspect the example + +- [`examples/agent_quickstart/tasks.py`](../examples/agent_quickstart/tasks.py) defines `AddOne` and `Double`. Each accepts and returns a Pydantic `Number` model. Tasks implement `run(input, ctx)`. +- [`examples/agent_quickstart/workflow.yaml`](../examples/agent_quickstart/workflow.yaml) gives the tasks stable instance names, configures `add_one.value` from the input file, and routes `add_one`'s output to `double`. +- [`examples/agent_quickstart/input.yaml`](../examples/agent_quickstart/input.yaml) supplies `value: 5` to `add_one`. + +The YAML loader imports `tasks.AddOne` and `tasks.Double` from the workflow file's directory when invoked through the CLI. No package installation or Python path changes are needed for these local tasks. + +## Validate, inspect, run + +```bash +.venv/bin/python -m taskmaestro validate examples/agent_quickstart/workflow.yaml --input examples/agent_quickstart/input.yaml +.venv/bin/python -m taskmaestro graph examples/agent_quickstart/workflow.yaml --input examples/agent_quickstart/input.yaml +.venv/bin/python -m taskmaestro run examples/agent_quickstart/workflow.yaml --input examples/agent_quickstart/input.yaml +``` + +`validate` prints `Workflow 'agent_quickstart' is valid`; `graph` prints Mermaid text with `add_one -->|Number| double`; `run` prints JSON: + +```json +{ + "value": 12 +} +``` + +The calculation is `(5 + 1) * 2`. The CLI writes errors to stderr and returns a nonzero exit code on failure. `validate` loads and checks the workflow and job configuration **without executing tasks**; `run` executes them. + +## Fix a validation error + +To reproduce a missing configuration field without changing the checked-in input, create a temporary empty input file: + +```bash +printf '{}\n' > /tmp/taskmaestro-agent-bad-input.yaml +.venv/bin/python -m taskmaestro validate examples/agent_quickstart/workflow.yaml --input /tmp/taskmaestro-agent-bad-input.yaml +``` + +This exits with code 2 and reports: + +```text +Configuration error: Job validation failed: Task 'add_one' is missing configuration fields ['value'] +``` + +The workflow declares `value` in `config_fields`, so the input must include `add_one: {value: 5}`. Validate again with the checked-in input, then run. Remove the temporary file when done: + +```bash +rm /tmp/taskmaestro-agent-bad-input.yaml +``` + +## When generating your own workflow + +1. Define each task's input and output as Pydantic models in a Python module. Subclass `Task[InputModel, OutputModel]` and implement `run(input, ctx)`. +2. List tasks in a workflow YAML file. Prefer explicit `name:` values; dependency references and top-level input YAML keys refer to **instance names**. Use `depends_on` to connect tasks and `config_fields` for fields supplied from input YAML. +3. Put input data under the task instance name in a separate YAML file. Run `validate`, then `graph`, then `run`. When validation fails, correct the named task/field before retrying. +4. For runtime failures, check stderr for the failing task. A successful validation does not test the task's `run()` logic or guarantee external services are available. + +## Discover installed task plugins + +If tasks are published by an installed package through the `taskmaestro.tasks` entry-point group, inspect their identifiers and model schemas before generating a workflow: + +```bash +.venv/bin/python -m taskmaestro tasks list --json +.venv/bin/python -m taskmaestro tasks describe acme.prepare --json +``` + +The list command returns a sorted JSON object such as `{"tasks": ["acme.prepare"]}`. The describe command returns the chosen task's `identifier`, `name`, `timeout_seconds`, and Pydantic `input_schema` / `output_schema` JSON Schema objects. Fields containing runtime-only Python objects are marked `x-taskmaestro-opaque` and `x-taskmaestro-python-type`, with `"not": {}` because no JSON value can satisfy them; route these values from upstream tasks rather than inventing JSON input. Replace `acme.prepare` with an identifier from your list; if you have no installed task plugins, the list is empty. The example tasks above are **local Python classes**, not installed plugins, so they will not appear in `tasks list`. + +## Inspect a workflow before writing input + +Inspection needs only the workflow YAML, so use it to find required fields and routing before creating an input file: + +```bash +.venv/bin/python -m taskmaestro workflow describe examples/agent_quickstart/workflow.yaml --json +``` + +This reports the result task `double`, a root `add_one` with `config_fields: ["value"]`, the dependency on `add_one`, and schemas for each task. Add `--input examples/agent_quickstart/input.yaml` to check that required configuration fields are present; the output includes `provided_config_fields` **names**, not their values. Inspection does not execute tasks or instantiate hooks, but it **does import Python modules** named in YAML (including nested workflows). Only inspect trusted workflow files and plugins. Runtime input values are checked when the tasks run, not completely by inspection. + +## Parse results and errors as JSON + +Use `--json` with `validate` and `run` when you need a stable response instead of parsing text from stderr: + +```bash +.venv/bin/python -m taskmaestro validate examples/agent_quickstart/workflow.yaml --input examples/agent_quickstart/input.yaml --json +.venv/bin/python -m taskmaestro run examples/agent_quickstart/workflow.yaml --input examples/agent_quickstart/input.yaml --json +``` + +Validation returns `{"status": "valid", "workflow": "agent_quickstart"}`; a successful run returns `{"status": "completed", "workflow": "agent_quickstart", "result": {"value": 12}}`. On failure, stdout contains a single JSON object with status `invalid` (configuration error) or `failed` (execution/serialization error), and an `error` containing `code`, `type`, `message`, `task`, `field`, and `issues`. Missing metadata is `null` or an empty list. Exit codes are 0 for success, 1 for a task/serialization failure, and 2 for an invalid configuration. Error messages omit raw exception details and input values; stderr may still contain application logs or prints. See the [CLI section of the README](../README.md#command-line-interface) for the full contract. + +For fan-in, collections, mapping, nested workflows, and the Python builder API, see the [README](../README.md) and the other examples under `examples/`. diff --git a/examples/agent_quickstart/input.yaml b/examples/agent_quickstart/input.yaml new file mode 100644 index 0000000..b1b8bd5 --- /dev/null +++ b/examples/agent_quickstart/input.yaml @@ -0,0 +1,2 @@ +add_one: + value: 5 diff --git a/examples/agent_quickstart/tasks.py b/examples/agent_quickstart/tasks.py new file mode 100644 index 0000000..4cf59da --- /dev/null +++ b/examples/agent_quickstart/tasks.py @@ -0,0 +1,19 @@ +"""Small, dependency-free tasks for the agent quickstart.""" + +from pydantic import BaseModel + +from taskmaestro import ExecutionContext, Task + + +class Number(BaseModel): + value: int + + +class AddOne(Task[Number, Number]): + def run(self, input: Number, ctx: ExecutionContext) -> Number: + return Number(value=input.value + 1) + + +class Double(Task[Number, Number]): + def run(self, input: Number, ctx: ExecutionContext) -> Number: + return Number(value=input.value * 2) diff --git a/examples/agent_quickstart/workflow.yaml b/examples/agent_quickstart/workflow.yaml new file mode 100644 index 0000000..93051b6 --- /dev/null +++ b/examples/agent_quickstart/workflow.yaml @@ -0,0 +1,9 @@ +workflow: + name: agent_quickstart + tasks: + - task: tasks.AddOne + name: add_one + config_fields: [value] + - task: tasks.Double + name: double + depends_on: add_one diff --git a/examples/image_processing/input.yaml b/examples/image_processing/input.yaml index 350c9e1..3240e4c 100644 --- a/examples/image_processing/input.yaml +++ b/examples/image_processing/input.yaml @@ -3,5 +3,6 @@ # Run: # python examples/image_processing/pipeline.py --yaml --input input.yaml +# Relative paths are resolved against this example's directory. load_image: - image_path: "taskmaestro.png" + image_path: "../../taskmaestro.png" diff --git a/examples/image_processing/pipeline.py b/examples/image_processing/pipeline.py index 6c999d4..81d184c 100644 --- a/examples/image_processing/pipeline.py +++ b/examples/image_processing/pipeline.py @@ -30,12 +30,14 @@ import hashlib import struct from pathlib import Path +from typing import cast from pydantic import BaseModel from taskmaestro import ( ExecutionContext, Job, + JobStatus, Runner, Task, Workflow, @@ -209,14 +211,21 @@ def run(self, input: AnalysisInput, ctx: ExecutionContext) -> ImageAnalysis: # --------------------------------------------------------------------------- +EXAMPLE_DIR = Path(__file__).resolve().parent + + class LoadImage(Task[ImageInput, ImagePath]): - """Resolve the image path to an absolute path.""" + """Resolve the image path to an absolute path. + + Relative paths are resolved against this example's directory rather than + the current working directory, so the example runs from anywhere. + """ name = "load_image" def run(self, input: ImageInput, ctx: ExecutionContext) -> ImagePath: ctx.logger.info("Loading image path: %s", input.image_path) - resolved = str(Path(input.image_path).resolve()) + resolved = str((EXAMPLE_DIR / input.image_path).resolve()) return ImagePath(path=resolved) @@ -280,7 +289,12 @@ def print_report( outer_workflow: Workflow, ) -> None: """Print the analysis report, timings, and Mermaid diagrams.""" - report: ReportOutput = result.result # type: ignore[assignment] + if result.status != JobStatus.COMPLETED: + print(f" Job status: {result.status}") + print(f" Failed task: {result.failed_task}") + print(f" Error: {result.error}") + return + report = cast(ReportOutput, result.result) print("=" * 60) print(f" {report.title}") @@ -303,8 +317,7 @@ def print_report( def run_python_mode() -> None: """Run the pipeline using the Python API.""" - _dir = Path(__file__).resolve().parent - image_path = str(_dir / ".." / ".." / "taskmaestro.png") + image_path = "../../taskmaestro.png" # relative to EXAMPLE_DIR # Build the outer workflow outer_workflow = ( @@ -345,20 +358,18 @@ def run_yaml_mode(workflow_path: str, input_path: str) -> None: def main() -> None: import argparse - _dir = Path(__file__).resolve().parent - parser = argparse.ArgumentParser(description="Image Processing Pipeline example") parser.add_argument( "--yaml", metavar="FILE", nargs="?", - const=str(_dir / "workflow.yaml"), + const=str(EXAMPLE_DIR / "workflow.yaml"), help="Load workflow from a YAML config file (default: workflow.yaml)", ) parser.add_argument( "--input", metavar="FILE", - default=str(_dir / "input.yaml"), + default=str(EXAMPLE_DIR / "input.yaml"), help="Input YAML file (default: input.yaml)", ) args = parser.parse_args() diff --git a/examples/text_analysis/pipeline.py b/examples/text_analysis/pipeline.py index 0785c3b..67ef1a3 100644 --- a/examples/text_analysis/pipeline.py +++ b/examples/text_analysis/pipeline.py @@ -25,6 +25,7 @@ import re from collections import Counter +from typing import cast from pydantic import BaseModel @@ -33,6 +34,7 @@ ExecutionContext, Job, JobConfiguration, + JobStatus, Runner, Task, Workflow, @@ -395,7 +397,12 @@ def run(self, input: ReportInput, ctx: ExecutionContext) -> AnalysisReport: def print_report(result: Job[EmptyConfig], timing: TimingHook, workflow: Workflow) -> None: """Print the analysis report, timings, and Mermaid diagram.""" - report: AnalysisReport = result.result # type: ignore[assignment] + if result.status != JobStatus.COMPLETED: + print(f" Job status: {result.status}") + print(f" Failed task: {result.failed_task}") + print(f" Error: {result.error}") + return + report = cast(AnalysisReport, result.result) print("=" * 60) print(f" {report.title} — Analysis Report") diff --git a/examples/text_analysis/workflow.yaml b/examples/text_analysis/workflow.yaml index 63a1f94..80bfeb0 100644 --- a/examples/text_analysis/workflow.yaml +++ b/examples/text_analysis/workflow.yaml @@ -21,7 +21,6 @@ workflow: stop_words: pipeline.GenerateStopWords - task: pipeline.ScoreReadability depends_on: pipeline.PrepareText - config_fields: [text, title] - task: pipeline.BuildReport depends_on: stats: pipeline.ComputeWordStats diff --git a/pyproject.toml b/pyproject.toml index 5d73e31..69f7e90 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -38,7 +38,7 @@ Issues = "https://github.com/OPM/taskmaestro/issues" [tool.hatch.build.targets.sdist] exclude = [ - "CLAUDE.md", + "AGENTS.md", ] [tool.hatch.build.targets.wheel] diff --git a/taskmaestro/__main__.py b/taskmaestro/__main__.py new file mode 100644 index 0000000..19900a4 --- /dev/null +++ b/taskmaestro/__main__.py @@ -0,0 +1,8 @@ +"""Allow ``python -m taskmaestro`` as an alias for the ``taskmaestro`` command.""" + +import sys + +from taskmaestro.cli import main + +if __name__ == "__main__": + sys.exit(main()) diff --git a/taskmaestro/cli.py b/taskmaestro/cli.py index fb9139a..a0fa09b 100644 --- a/taskmaestro/cli.py +++ b/taskmaestro/cli.py @@ -3,15 +3,27 @@ from __future__ import annotations import argparse +import json import logging import sys -from collections.abc import Sequence +from collections.abc import Iterator, Sequence +from contextlib import contextmanager, redirect_stdout +from dataclasses import asdict 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 +from pydantic import ValidationError +from pydantic.errors import PydanticInvalidForJsonSchema, PydanticSchemaGenerationError +from pydantic.json_schema import GenerateJsonSchema, JsonSchemaValue +from pydantic_core import core_schema + +from taskmaestro.dependencies import CollectionRef, OutputRef +from taskmaestro.discovery import get_registered_task, registered_task_names +from taskmaestro.exceptions import ConfigLoadError, PluginLoadError, WorkflowDefinitionError +from taskmaestro.job import EmptyConfig, Job, JobStatus +from taskmaestro.task import get_input_type, get_output_type +from taskmaestro.workflow import Workflow +from taskmaestro.yaml_config import LoadedWorkflow, _load_workflow_only, load_workflow_from_yaml def _add_workflow_arguments(parser: argparse.ArgumentParser) -> None: @@ -19,24 +31,100 @@ def _add_workflow_arguments(parser: argparse.ArgumentParser) -> None: 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) +@contextmanager +def _workflow_imports(path: str) -> Iterator[None]: + """Make task modules next to a YAML file importable for the duration of loading.""" original_path = sys.path.copy() - sys.path.insert(0, workflow_dir) + sys.path.insert(0, str(Path(path).resolve().parent)) try: - return load_workflow_from_yaml(args.workflow, args.input) + yield finally: sys.path[:] = original_path +def _load(args: argparse.Namespace) -> LoadedWorkflow: + with _workflow_imports(args.workflow): + return load_workflow_from_yaml(args.workflow, args.input) + + +def _error(code: str, exc: Exception, message: str, *, task: str | None = None) -> dict[str, Any]: + """Return safe diagnostics from an exception and its causes.""" + issues: list[dict[str, str]] = [] + cause: BaseException | None = exc + while cause is not None: + if isinstance(cause, ValidationError): + issues = [ + {"field": ".".join(map(str, issue["loc"])), "code": issue["type"]} + for issue in cause.errors(include_input=False, include_context=False) + ] + break + if isinstance(cause, WorkflowDefinitionError): + task = task or cause.task_name + issues = [{"field": field, "code": "missing"} for field in cause.fields] + if issues: + break + cause = cause.__cause__ + return { + "code": code, + "type": type(exc).__name__, + "message": message, + "task": task, + "field": issues[0]["field"] if issues else None, + "issues": issues, + } + + 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 args.json: + # Plugin imports and user tasks may print; reserve stdout for one JSON document. + with redirect_stdout(sys.stderr): + result = _load(args).run() + else: + result = _load(args).run() + if args.json: + if result.status == JobStatus.FAILED: + assert result.exception is not None + print( + json.dumps( + { + "status": "failed", + "workflow": result.workflow.name, + "failed_task": result.failed_task, + "error": _error( + "task_failed", result.exception, "Task failed", task=result.failed_task + ), + } + ) + ) + return 1 + assert result.result is not None + try: + with redirect_stdout(sys.stderr): + output = json.loads(result.result.model_dump_json()) + except Exception as exc: + # Tasks may return Python-only objects; keep stdout valid JSON even then. + print( + json.dumps( + { + "status": "failed", + "workflow": result.workflow.name, + "failed_task": None, + "error": _error( + "serialization_error", exc, "Result is not JSON serializable" + ), + } + ) + ) + return 1 + print( + json.dumps({"status": "completed", "workflow": result.workflow.name, "result": output}) + ) + return 0 if result.status == JobStatus.FAILED: print( f"Workflow failed at {result.failed_task}: {result.error}", @@ -49,8 +137,15 @@ def _run(args: argparse.Namespace) -> int: def _validate(args: argparse.Namespace) -> int: - loaded = _load(args) - print(f"Workflow '{loaded.workflow.name}' is valid") + if args.json: + with redirect_stdout(sys.stderr): + loaded = _load(args) + else: + loaded = _load(args) + if args.json: + print(json.dumps({"status": "valid", "workflow": loaded.workflow.name})) + else: + print(f"Workflow '{loaded.workflow.name}' is valid") return 0 @@ -65,6 +160,137 @@ def _graph(args: argparse.Namespace) -> int: return 0 +def _tasks_list(args: argparse.Namespace) -> int: + names = sorted(registered_task_names()) + if args.json: + print(json.dumps({"tasks": names})) + else: + print("\n".join(names)) + return 0 + + +class _TaskSchemaGenerator(GenerateJsonSchema): + """Describe runtime-only Python objects without pretending they accept JSON.""" + + def is_instance_schema(self, schema: core_schema.IsInstanceSchema) -> JsonSchemaValue: + cls = schema["cls"] + return { + "not": {}, # No JSON value can satisfy an isinstance check for this object. + "x-taskmaestro-opaque": True, + "x-taskmaestro-python-type": f"{cls.__module__}.{cls.__qualname__}", + } + + +def _output_ref(ref: OutputRef) -> dict[str, str | None]: + return {"task": ref.task_name, "field": ref.output_field} + + +def _dependency_spec(deps: Any) -> Any: + """Render validated dependencies, including field routing and collections.""" + if deps is None: + return None + if isinstance(deps, str): + return {"task": deps, "field": None} + if isinstance(deps, tuple): + return {"task": deps[0], "field": deps[1]} + result: dict[str, Any] = {} + for field, ref in deps.items(): + if isinstance(ref, CollectionRef): + members: list[Any] | dict[str, Any] + if ref.kind == "keyed": + members = {key: _output_ref(item) for key, item in ref.keyed_members} + else: + members = [_output_ref(item) for item in ref.positional_members] + result[field] = {"collect": {"kind": ref.kind, "members": members}} + elif isinstance(ref, tuple): + result[field] = {"task": ref[0], "field": ref[1]} + else: + result[field] = {"task": ref, "field": None} + return result + + +def _workflow_description( + workflow: Workflow, *, configured: dict[str, list[str]] | None +) -> dict[str, Any]: + tasks: list[dict[str, Any]] = [] + for name, task in workflow.topological_order(): + input_type = get_input_type(task) + output_type = workflow.get_output_annotation(name) + task_map = workflow.get_task_map(name) + tasks.append( + { + "name": name, + "python_type": f"{task.__module__}.{task.__qualname__}", + "depends_on": _dependency_spec(workflow.get_dependencies(name)), + "config_fields": sorted(workflow.get_config_fields(name)), + "provided_config_fields": configured[name] if configured is not None else None, + "required_input_fields": sorted( + field for field, info in input_type.model_fields.items() if info.is_required() + ), + "map": asdict(task_map) if task_map is not None else None, + "input_schema": input_type.model_json_schema( + schema_generator=_TaskSchemaGenerator + ), + "output_schema": output_type.model_json_schema( + schema_generator=_TaskSchemaGenerator + ), + } + ) + return {"workflow": workflow.name, "result_task": workflow.result_task_name, "tasks": tasks} + + +def _workflow_describe(args: argparse.Namespace) -> int: + try: + with redirect_stdout(sys.stderr), _workflow_imports(args.workflow): + workflow, config = _load_workflow_only( + Path(args.workflow), Path(args.input) if args.input else None + ) + configured = None + if args.input is not None: + assert config is not None + # Check supplied config without constructing hooks or executing any tasks. + Job(workflow, EmptyConfig(), job_configuration=config) + configured = { + name: sorted(config.config_fields_for_task(name)) + for name, _task in workflow.topological_order() + } + description = _workflow_description(workflow, configured=configured) + except WorkflowDefinitionError as exc: + raise ConfigLoadError(f"Job validation failed: {exc}") from exc + except ( + TypeError, + ValueError, + PydanticInvalidForJsonSchema, + PydanticSchemaGenerationError, + ) as exc: + raise ConfigLoadError(f"Cannot describe workflow: {exc}") from exc + print(json.dumps(description, indent=None if args.json else 2)) + return 0 + + +def _tasks_describe(args: argparse.Namespace) -> int: + task = get_registered_task(args.name) + try: + input_type = get_input_type(task) + output_type = get_output_type(task) + description = { + "identifier": args.name, + "name": task.name, + "timeout_seconds": task.timeout_seconds, + "input_schema": input_type.model_json_schema(schema_generator=_TaskSchemaGenerator), + "output_schema": output_type.model_json_schema(schema_generator=_TaskSchemaGenerator), + } + except ( + TypeError, + ValueError, + PydanticInvalidForJsonSchema, + PydanticSchemaGenerationError, + ) as exc: + raise PluginLoadError(f"Cannot describe task '{args.name}': {exc}") from exc + print(json.dumps(description, indent=None if args.json else 2)) + return 0 + + def build_parser() -> argparse.ArgumentParser: """Build the public command-line parser.""" parser = argparse.ArgumentParser(prog="taskmaestro") @@ -77,16 +303,46 @@ def build_parser() -> argparse.ArgumentParser: choices=["DEBUG", "INFO", "WARNING", "ERROR", "CRITICAL"], default="INFO", ) + run_parser.add_argument( + "--json", action="store_true", help="Print structured JSON result or error" + ) run_parser.set_defaults(handler=_run) validate_parser = subparsers.add_parser("validate", help="Validate a YAML workflow") _add_workflow_arguments(validate_parser) + validate_parser.add_argument( + "--json", action="store_true", help="Print structured JSON result or error" + ) 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) + tasks_parser = subparsers.add_parser("tasks", help="Discover installed task plugins") + tasks_subparsers = tasks_parser.add_subparsers(dest="tasks_command", required=True) + list_parser = tasks_subparsers.add_parser("list", help="List registered task identifiers") + list_parser.add_argument("--json", action="store_true", help="Print machine-readable JSON") + list_parser.set_defaults(handler=_tasks_list) + describe_parser = tasks_subparsers.add_parser("describe", help="Describe a registered task") + describe_parser.add_argument("name", help="Registered task identifier (not a class path)") + describe_parser.add_argument("--json", action="store_true", help="Print single-line JSON") + describe_parser.set_defaults(handler=_tasks_describe) + + workflow_parser = subparsers.add_parser("workflow", help="Inspect YAML workflows") + workflow_subparsers = workflow_parser.add_subparsers(dest="workflow_command", required=True) + workflow_describe_parser = workflow_subparsers.add_parser( + "describe", help="Describe a workflow without running tasks" + ) + workflow_describe_parser.add_argument("workflow", help="Path to the workflow YAML file") + workflow_describe_parser.add_argument( + "--input", help="Optional input YAML to check required configuration fields" + ) + workflow_describe_parser.add_argument( + "--json", action="store_true", help="Print single-line JSON" + ) + workflow_describe_parser.set_defaults(handler=_workflow_describe) + return parser @@ -97,5 +353,27 @@ def main(argv: Sequence[str] | None = None) -> int: handler: Any = args.handler return int(handler(args)) except ConfigLoadError as exc: - print(f"Configuration error: {exc}", file=sys.stderr) + if getattr(args, "json", False) and ( + args.command in ("run", "validate") + or (args.command == "workflow" and args.workflow_command == "describe") + ): + print( + json.dumps( + { + "status": "invalid", + "error": _error( + "configuration_error", exc, "Workflow configuration is invalid" + ), + } + ) + ) + else: + print(f"Configuration error: {exc}", file=sys.stderr) + return 2 + except PluginLoadError as exc: + print(f"Plugin error: {exc}", file=sys.stderr) return 2 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/taskmaestro/exceptions.py b/taskmaestro/exceptions.py index bf6006e..d60bf15 100644 --- a/taskmaestro/exceptions.py +++ b/taskmaestro/exceptions.py @@ -10,6 +10,17 @@ class WorkflowRunnerError(Exception): class WorkflowDefinitionError(WorkflowRunnerError): """Raised at workflow construction time for invalid definitions.""" + def __init__( + self, + message: str, + *, + task_name: str | None = None, + fields: list[str] | None = None, + ) -> None: + super().__init__(message) + self.task_name = task_name + self.fields = fields or [] + class CycleDetectedError(WorkflowDefinitionError): """Dependency graph contains a cycle.""" diff --git a/taskmaestro/job.py b/taskmaestro/job.py index 568df2a..db35e95 100644 --- a/taskmaestro/job.py +++ b/taskmaestro/job.py @@ -100,19 +100,22 @@ def __init__( self.task_results: list[TaskResult] = [] self.mapped_item_results: dict[str, list[TaskResult]] = {} + # Map sources first: their errors are more specific than a generic + # missing-configuration-field error for the same mapped task. + self._validate_task_maps() 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.""" + """Ensure every declared configuration field has a supplied value. + + This applies to every task, not only roots: a dependent or mapped task + whose configuration is missing would otherwise run with an incomplete + (or, for a single whole-output dependency, the wrong) input. + """ 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) - ): + if not expected: continue supplied = ( self.job_configuration.config_fields_for_task(task_name) @@ -122,24 +125,21 @@ def _validate_task_configuration(self) -> None: missing = expected - supplied if missing: raise WorkflowDefinitionError( - f"Task '{task_name}' is missing configuration fields {sorted(missing)}" + f"Task '{task_name}' is missing configuration fields {sorted(missing)}", + task_name=task_name, + 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: - # Configured and mapped roots do not consume job.config directly. - config_fields = self.workflow.get_config_fields(task_name) - 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) - if not isinstance(config, expected_input): - raise WorkflowDefinitionError( - f"Root task '{task_name}' expects input type " - f"{expected_input.__name__} but got {type(config).__name__}" - ) + # Configured and mapped roots do not consume job.config directly. + for task_name in self.workflow.input_root_names(): + expected_input = get_input_type(self.workflow._tasks[task_name]) + if not isinstance(config, expected_input): + raise WorkflowDefinitionError( + 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.""" diff --git a/taskmaestro/runner.py b/taskmaestro/runner.py index 34ea08d..55cd725 100644 --- a/taskmaestro/runner.py +++ b/taskmaestro/runner.py @@ -3,6 +3,7 @@ from __future__ import annotations import signal +import threading import time import warnings from collections.abc import Mapping @@ -40,21 +41,106 @@ class HookError(UserWarning): """ +@dataclass(eq=False) +class _Timer: + """One armed timeout: an absolute expiry and the error to raise when it fires.""" + + expiry: float + error_type: type[TaskTimeoutError] + message: str + fired: bool = False + + +class _AlarmScheduler: + """Multiplex the process's single ``ITIMER_REAL``/``SIGALRM`` across timers. + + Nested runs (e.g. a ``workflow_task`` whose inner workflow has its own + timeouts) share one process timer. Every armed timeout is registered here + and the OS timer always tracks the nearest pending expiry, so an inner run + arming or cancelling its own timers can neither extend nor cancel the + timeouts of an enclosing run. The SIGALRM handler that was installed before + the first timer was added is restored when the last timer is removed. + """ + + def __init__(self) -> None: + self._timers: list[_Timer] = [] + self._previous_handler: Any = None + self._installed = False + + def add(self, timer: _Timer) -> None: + """Register *timer* and re-arm the OS timer. + + Raises ``ValueError`` off the main thread (signals are delivered to the + main thread only) and ``AttributeError``/``OSError`` where ``SIGALRM`` + is unavailable. + """ + if threading.current_thread() is not threading.main_thread(): + raise ValueError("signal timers only work in the main thread") + if not self._installed: + self._previous_handler = signal.signal(signal.SIGALRM, self._handle) + self._installed = True + self._timers.append(timer) + self._reschedule() + + def remove(self, timer: _Timer) -> None: + """Unregister *timer*; restore the previous handler once none remain.""" + self._timers.remove(timer) + if self._timers: + self._reschedule() + return + _start_timer(0) + with suppress(AttributeError, OSError, ValueError, TypeError): # pragma: no cover + signal.signal(signal.SIGALRM, self._previous_handler) + self._previous_handler = None + self._installed = False + + def _nearest_pending(self) -> _Timer | None: + pending = [timer for timer in self._timers if not timer.fired] + return min(pending, key=lambda timer: timer.expiry) if pending else None + + def _reschedule(self) -> None: + nearest = self._nearest_pending() + if nearest is None: + _start_timer(0) + else: + _start_timer(max(nearest.expiry - time.monotonic(), 1e-6)) + + def _handle(self, signum: int, frame: Any) -> None: + nearest = self._nearest_pending() + if nearest is None: + return # stale signal: every registered timer has already fired + nearest.fired = True + # The OS timer is one-shot; keep the remaining timers armed. + self._reschedule() + raise nearest.error_type(nearest.message) + + +def _start_timer(seconds: float) -> None: + """Start (or with ``0`` cancel) the process's one-shot real-time timer.""" + setitimer = getattr(signal, "setitimer", None) + if setitimer is not None: + setitimer(signal.ITIMER_REAL, seconds) + else: # pragma: no cover - every SIGALRM platform has setitimer + signal.alarm(max(1, int(seconds + 0.999999)) if seconds else 0) + + +_ALARMS = _AlarmScheduler() + + @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. + 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. ``timer`` is + the unit of work's currently registered :class:`_Timer`, if any. """ 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 + timer: _Timer | None = field(default=None, repr=False) def remaining(self) -> float | None: """Seconds left until the job deadline, or ``None`` if there is none.""" @@ -107,8 +193,6 @@ def run( try: for task_name, task_cls in workflow.topological_order(): - task = task_cls() - task.name = task_name # instance-level override for named instances deps = workflow.get_dependencies(task_name) config_fields = workflow.get_config_fields(task_name) task_map = workflow.get_task_map(task_name) @@ -125,59 +209,15 @@ def run( if task_map is None or key != task_map.over } - # 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) - task_input = input_type.model_validate(config_values) - else: - task_input = job.config - elif isinstance(deps, str): - if config_values: - # Single dep with config: decompose upstream, merge with config - input_type = get_input_type(task_cls) - 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 - } - merged.update(config_values) - task_input = input_type.model_validate(merged) - else: - task_input = outputs[deps] - elif isinstance(deps, tuple): - upstream_name, field_name = deps - task_input = getattr(outputs[upstream_name], field_name) - elif isinstance(deps, dict): - input_type = get_input_type(task_cls) - field_values: dict[str, object] = {} - for fname, upstream_ref in deps.items(): - 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: - field_values[fname] = outputs[upstream_ref] - if config_values: - field_values.update(config_values) - task_input = input_type.model_validate(field_values) - else: - task_input = job.config # pragma: no cover - task_started = datetime.now() - self._emit(Event.TASK_START, job, task) - + task: Task[Any, Any] | None = None try: - # Arming happens inside the guarded block so that an expired - # job deadline or an unusable timer is recorded as a task + # Instantiation, input assembly and arming all happen inside + # the guarded block so that any failure is recorded as a task # failure rather than escaping with the job left RUNNING. + task = task_cls() + task.name = task_name # instance-level override for named instances + self._emit(Event.TASK_START, job, task) deadline.check() if task_map is not None: output = self._run_mapped_task( @@ -193,8 +233,17 @@ def run( deadline, ) else: + task_input = self._build_task_input( + job, task_cls, deps, config_values, outputs + ) self._arm(task.timeout_seconds, task.name, deadline) - output = task.run(task_input, ctx) + try: + output = task.run(task_input, ctx) + finally: + # Stop the clock as soon as the task returns: output + # checks and hooks must not be interrupted by (or + # swallow) this task's timeout. + self._disarm(deadline) # Validate output matches declared type expected_output_type = get_output_type(task_cls) @@ -205,10 +254,10 @@ def run( ) duration = (datetime.now() - task_started).total_seconds() - outputs[task.name] = output + outputs[task_name] = output job.task_results.append( TaskResult( - task_name=task.name, + task_name=task_name, status=TaskStatus.COMPLETED, output=output, started_at=task_started, @@ -221,11 +270,11 @@ def run( job.status = JobStatus.FAILED job.error = str(exc) job.exception = exc - job.failed_task = task.name + job.failed_task = task_name job.completed_at = datetime.now() job.task_results.append( TaskResult( - task_name=task.name, + task_name=task_name, status=TaskStatus.FAILED, output=None, started_at=task_started, @@ -233,14 +282,15 @@ def run( error=str(exc), ) ) - self._emit(Event.TASK_FAIL, job, task, exc) + # Without an instance (the constructor raised) there is no + # task to report; the job-level failure is still emitted. + if task is not None: + self._emit(Event.TASK_FAIL, job, task, exc) self._emit(Event.JOB_FAIL, job) return job - finally: - self._disarm(deadline) finally: + # Safety net; each unit of work already disarms right after running. self._disarm(deadline) - self._restore_handler(deadline) job.status = JobStatus.COMPLETED job.result = outputs[workflow.result_task_name] @@ -248,6 +298,57 @@ def run( self._emit(Event.JOB_COMPLETE, job) return job + @staticmethod + def _build_task_input( + job: Job[Any], + task_cls: type[Task[Any, Any]], + deps: Any, + config_values: dict[str, Any], + outputs: dict[str, BaseModel], + ) -> Any: + """Assemble one (unmapped) task's input from upstream outputs and config.""" + if deps is None: + if config_values: + # Root task with config: build input from config values + return get_input_type(task_cls).model_validate(config_values) + return job.config + if isinstance(deps, str): + if not config_values: + return outputs[deps] + # Single dep with config: decompose upstream, merge with config + input_type = get_input_type(task_cls) + upstream_data = outputs[deps].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 + } + merged.update(config_values) + return input_type.model_validate(merged) + if isinstance(deps, tuple): + upstream_name, field_name = deps + return getattr(outputs[upstream_name], field_name) + # Fan-in: one named input field per upstream reference. + field_values = Runner._resolve_named_dependencies(deps, outputs) + field_values.update(config_values) + return get_input_type(task_cls).model_validate(field_values) + + @staticmethod + def _resolve_named_dependencies( + deps: Mapping[str, Any], + outputs: dict[str, BaseModel], + ) -> dict[str, object]: + """Resolve a fan-in dependency mapping to input field values.""" + values: dict[str, object] = {} + for field_name, ref in deps.items(): + if isinstance(ref, CollectionRef): + values[field_name] = Runner._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] + return values + def _run_mapped_task( self, job: Job[Any], @@ -282,10 +383,13 @@ def _run_mapped_task( 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) + self._arm(item_task.timeout_seconds, f"{parent_task.name}[{key}]", deadline) + try: + output = item_task.run(item_input, item_ctx) + finally: + self._disarm(deadline) if not isinstance(output, expected_output_type): raise TaskOutputTypeError( f"Task '{parent_task.name}[{key}]' returned " @@ -319,8 +423,6 @@ def _run_mapped_task( 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) @@ -334,16 +436,9 @@ def _mapped_shared_values( 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: dict[str, object] = ( + self._resolve_named_dependencies(deps, outputs) if isinstance(deps, dict) else {} + ) values.update(config_values) return values @@ -358,18 +453,18 @@ def _resolve_output_ref( return output return getattr(output, ref.output_field) + @staticmethod 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 + Runner._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 + key: Runner._resolve_output_ref(ref, outputs) for key, ref in collection.keyed_members } def _arm(self, task_timeout: float | None, label: str, deadline: _Deadline) -> None: @@ -396,35 +491,23 @@ def _set_alarm( deadline: _Deadline | None = None, job_timeout: bool = False, ) -> bool: - """Install a SIGALRM handler and start a one-shot timer. + """Register a one-shot timeout with the process-wide alarm scheduler. - 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. + Uses ``signal.setitimer`` for sub-second precision. 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) + timer = _Timer(time.monotonic() + seconds, 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 + _ALARMS.add(timer) except (AttributeError, OSError, ValueError): - # ValueError: signal.signal() called outside the main thread. + # ValueError: signals can only be used from the main thread. if deadline is None or not deadline.warned: if deadline is not None: deadline.warned = True @@ -434,26 +517,17 @@ def _handler(signum: int, frame: Any) -> None: stacklevel=2, ) return False + if deadline is not None: + deadline.timer = timer + return True @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: + """Cancel this run's pending timer, leaving enclosing runs' timers armed.""" + if deadline.timer is None: return - with suppress(AttributeError, OSError, ValueError, TypeError): # pragma: no cover - signal.signal(signal.SIGALRM, deadline.previous_handler) - deadline.handler_installed = False + _ALARMS.remove(deadline.timer) + deadline.timer = None def _emit(self, event: Event, *args: object) -> None: """Dispatch event to all hooks, swallowing any hook errors. diff --git a/taskmaestro/visualization.py b/taskmaestro/visualization.py index 8cc81a1..954c569 100644 --- a/taskmaestro/visualization.py +++ b/taskmaestro/visualization.py @@ -186,13 +186,12 @@ def to_mermaid( inner_wf = getattr(task_cls, "_inner_workflow", None) if inner_wf is not None: workflow_task_nodes[task_name] = inner_wf - # Find inner root(s): tasks with no dependencies and no config_fields - inner_tasks = inner_wf.topological_order() - for iname, _icls in inner_tasks: - ideps = inner_wf.get_dependencies(iname) - if ideps is None and not inner_wf.get_config_fields(iname): - target_redirect[task_name] = f"{task_name}__{iname}" - break + # Edges into the subgraph enter at the inner task consuming the + # wrapper's input (the same root workflow_task derives its input + # type from); configured and mapped roots are fed by config. + inner_roots = inner_wf.input_root_names() + if inner_roots: + target_redirect[task_name] = f"{task_name}__{inner_roots[0]}" # Result task → source redirect source_redirect[task_name] = f"{task_name}__{inner_wf.result_task_name}" diff --git a/taskmaestro/workflow.py b/taskmaestro/workflow.py index cc0556d..9a66add 100644 --- a/taskmaestro/workflow.py +++ b/taskmaestro/workflow.py @@ -273,6 +273,19 @@ 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 input_root_names(self) -> list[str]: + """Return the root tasks that consume the job input, in declaration order. + + These are tasks without dependencies that are neither configured + (``config_fields``) nor mapped; configured and mapped roots take their + input from the ``JobConfiguration`` instead. + """ + return [ + name + for name, deps in self._dependencies.items() + if deps is None and not self.get_config_fields(name) and not self.is_mapped_task(name) + ] + 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]) @@ -294,7 +307,7 @@ def _validate_unique_names(self) -> None: 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. + validation order documented in AGENTS.md remains visible here. """ def _validate_references(self) -> None: @@ -469,6 +482,15 @@ def _validate_types(self) -> None: elif isinstance(deps, tuple): # Single dependency, specific output field upstream_name, field_name = deps + if cf: + # The routed field value becomes the whole task input, so + # there is nothing to merge configuration values into. + raise WorkflowDefinitionError( + f"Task '{name}' depends on the output field " + f"'{upstream_name}.{field_name}' and cannot also declare " + f"config_fields; use a named dependency such as " + f"depends_on={{'': ({upstream_name!r}, {field_name!r})}}" + ) upstream_output = self.get_output_annotation(upstream_name) upstream_fields = upstream_output.model_fields if field_name not in upstream_fields: @@ -901,10 +923,26 @@ def set_result_task(self, task: TaskReference) -> WorkflowBuilder: return self def build(self) -> Workflow: - """Finalize and validate the workflow. Returns an immutable Workflow.""" - # Resolve result_task ref + """Finalize and validate the workflow. + + Returns an independent snapshot of the builder's current state: later + calls on this builder never modify an already built workflow, and + calling ``build()`` again yields a new, separately validated one. + """ + source = self._workflow + workflow = Workflow.__new__(Workflow) + workflow.name = source.name + workflow._tasks = dict(source._tasks) + # Fan-in dicts are copied because validation may rewrite their entries. + workflow._dependencies = { + name: dict(deps) if isinstance(deps, dict) else deps + for name, deps in source._dependencies.items() + } + workflow._config_fields = { + name: set(fields) for name, fields in source._config_fields.items() + } + workflow._task_maps = dict(source._task_maps) ref = self._result_task_ref - if ref is not None: - self._workflow._result_task_name = self._resolve_dep_ref(ref) - self._workflow._validate() - return self._workflow + workflow._result_task_name = self._resolve_dep_ref(ref) if ref is not None else None + workflow._validate() + return workflow diff --git a/taskmaestro/workflow_task.py b/taskmaestro/workflow_task.py index 09fac92..c923921 100644 --- a/taskmaestro/workflow_task.py +++ b/taskmaestro/workflow_task.py @@ -45,13 +45,8 @@ def workflow_task( 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 and not workflow.is_mapped_task(task_name): - roots.append((task_name, workflow._tasks[task_name])) + # Root tasks consuming the job input (not configured, not mapped). + roots = [(name, workflow._tasks[name]) for name in workflow.input_root_names()] all_roots_configured = False diff --git a/taskmaestro/yaml_config.py b/taskmaestro/yaml_config.py index 5594440..c18d631 100644 --- a/taskmaestro/yaml_config.py +++ b/taskmaestro/yaml_config.py @@ -119,7 +119,17 @@ def flatten_mapping(self, node: yaml.nodes.MappingNode) -> None: if key_node.tag == "tag:yaml.org,2002:merge" else self.construct_object(key_node) ) - if key in keys: + try: + is_duplicate = key in keys + except TypeError as exc: + # Match PyYAML's own error for unhashable keys (e.g. ``? [a, b]``). + raise yaml.constructor.ConstructorError( + "while constructing a mapping", + node.start_mark, + f"found unhashable key ({exc})", + key_node.start_mark, + ) from exc + if is_duplicate: raise yaml.constructor.ConstructorError( "while constructing a mapping", node.start_mark, @@ -147,7 +157,13 @@ def import_class(dotted_path: str) -> type[Any]: try: module = importlib.import_module(module_path) except ModuleNotFoundError as exc: - raise ConfigLoadError(f"Cannot import module '{module_path}': {exc}") from exc + if exc.name is not None and _is_module_or_parent(exc.name, module_path): + raise ConfigLoadError(f"Cannot import module '{module_path}': {exc}") from exc + # The module exists but one of its own imports is missing. + raise ConfigLoadError(f"Error while importing module '{module_path}': {exc!r}") from exc + except Exception as exc: + # SyntaxError, ImportError, or any error raised by module-level code. + raise ConfigLoadError(f"Error while importing module '{module_path}': {exc!r}") from exc try: cls = getattr(module, class_name) except AttributeError: @@ -155,6 +171,11 @@ def import_class(dotted_path: str) -> type[Any]: return cls # type: ignore[no-any-return] +def _is_module_or_parent(missing: str, module_path: str) -> bool: + """Return whether *missing* is *module_path* itself or one of its parent packages.""" + return module_path == missing or module_path.startswith(f"{missing}.") + + def _coerce_hook_params(hook_cls: type[Any], params: dict[str, Any]) -> dict[str, Any]: """Inspect hook __init__ signature and coerce string values to Path where annotated.""" coerced = dict(params) @@ -221,13 +242,14 @@ def _load_workflow_only( workflow_path: Path, input_path: Path | None = None, *, - _ancestors: frozenset[Path] = frozenset(), + _ancestors: tuple[Path, ...] = (), ) -> 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. ``_ancestors`` holds - the resolved paths of every enclosing workflow file so that a self- or + the resolved paths of every enclosing workflow file, outermost first, so + that the reported reference chain is in nesting order and a self- or mutually-referencing ``workflow:`` entry is rejected instead of recursing without bound. @@ -237,9 +259,9 @@ def _load_workflow_only( resolved_path = workflow_path.resolve() if resolved_path in _ancestors: - chain = " -> ".join(str(p) for p in (*sorted(_ancestors), resolved_path)) + chain = " -> ".join(str(p) for p in (*_ancestors, resolved_path)) raise ConfigLoadError(f"Recursive workflow reference: {chain}") - ancestors = _ancestors | {resolved_path} + ancestors = (*_ancestors, resolved_path) # 1. Parse workflow YAML try: @@ -567,11 +589,14 @@ def load_workflow_from_yaml(workflow_path: str | Path, input_path: str | Path) - hook_cls = import_class(hook_config.hook) if not (isinstance(hook_cls, type) and issubclass(hook_cls, BaseHook)): raise ConfigLoadError(f"'{hook_config.hook}' is not a BaseHook subclass") - coerced_params = _coerce_hook_params(hook_cls, hook_config.params) try: + coerced_params = _coerce_hook_params(hook_cls, hook_config.params) hooks.append(hook_cls(**coerced_params)) - except TypeError as exc: - raise ConfigLoadError(f"Cannot instantiate hook '{hook_config.hook}': {exc}") from exc + except Exception as exc: + # TypeError for bad parameters, but a hook's __init__ may raise anything. + raise ConfigLoadError( + f"Cannot instantiate hook '{hook_config.hook}': {type(exc).__name__}: {exc}" + ) from exc runner = Runner(hooks=hooks) diff --git a/tests/test_agent_quickstart.py b/tests/test_agent_quickstart.py new file mode 100644 index 0000000..dde832b --- /dev/null +++ b/tests/test_agent_quickstart.py @@ -0,0 +1,39 @@ +"""Keep the agent quickstart's CLI commands and expected output working.""" + +import json +import subprocess +import sys +from pathlib import Path + +EXAMPLE = Path(__file__).resolve().parents[1] / "examples" / "agent_quickstart" +WORKFLOW = EXAMPLE / "workflow.yaml" +INPUT = EXAMPLE / "input.yaml" + + +def _cli(command: str, input_path: Path) -> subprocess.CompletedProcess[str]: + return subprocess.run( + [sys.executable, "-m", "taskmaestro", command, str(WORKFLOW), "--input", str(input_path)], + capture_output=True, + text=True, + check=False, + ) + + +def test_agent_quickstart(tmp_path: Path) -> None: + validated = _cli("validate", INPUT) + assert validated.returncode == 0, validated.stderr + assert validated.stdout.strip() == "Workflow 'agent_quickstart' is valid" + + graph = _cli("graph", INPUT) + assert graph.returncode == 0, graph.stderr + assert "add_one -->|Number| double" in graph.stdout + + run = _cli("run", INPUT) + assert run.returncode == 0, run.stderr + assert json.loads(run.stdout) == {"value": 12} + + bad_input = tmp_path / "bad-input.yaml" + bad_input.write_text("{}\n", encoding="utf-8") + invalid = _cli("validate", bad_input) + assert invalid.returncode == 2 + assert "Task 'add_one' is missing configuration fields ['value']" in invalid.stderr diff --git a/tests/test_cli.py b/tests/test_cli.py index c25e74a..90bcb1d 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -2,10 +2,51 @@ from __future__ import annotations +import json import sys +from importlib.metadata import EntryPoint from pathlib import Path -from taskmaestro.cli import main +import pytest +from pydantic import BaseModel +from pydantic_core import core_schema + +from taskmaestro import ExecutionContext, ObjectModel, Task +from taskmaestro.cli import _dependency_spec, main +from taskmaestro.dependencies import CollectionRef, OutputRef + + +class ExternalClient: + """Represents a Python-only value such as a live gRPC client.""" + + +class ClientHandle(ObjectModel[ExternalClient]): + pass + + +class OpaqueInput(BaseModel): + handle: ClientHandle + amount: int + + +class OpaqueTask(Task[OpaqueInput, ClientHandle]): + def run(self, input: OpaqueInput, ctx: ExecutionContext) -> ClientHandle: + return input.handle + + +class UnsupportedValue: + @classmethod + def __get_pydantic_core_schema__(cls, source: object, handler: object) -> object: + return core_schema.no_info_plain_validator_function(lambda value: value) + + +class UnsupportedInput(BaseModel): + value: UnsupportedValue + + +class UnsupportedTask(Task[UnsupportedInput, UnsupportedInput]): + def run(self, input: UnsupportedInput, ctx: ExecutionContext) -> UnsupportedInput: + return input def _files(tmp_path: Path, task: str = "Increment") -> tuple[Path, Path]: @@ -31,6 +72,13 @@ class Fail(Task[NumberInput, NumberOutput]): def run(self, input: NumberInput, ctx: ExecutionContext) -> NumberOutput: raise ValueError("intentional failure") + +class Noisy(Task[NumberInput, NumberOutput]): + name = "noisy" + + def run(self, input: NumberInput, ctx: ExecutionContext) -> NumberOutput: + print("task diagnostic") + return NumberOutput(value=input.value) """, encoding="utf-8", ) @@ -91,6 +139,185 @@ def test_validate_reports_success(tmp_path: Path, capsys: object) -> None: assert "Workflow 'cli_test' is valid" in captured.out +def test_validate_json_success(tmp_path: Path, capsys: pytest.CaptureFixture[str]) -> None: + workflow, input_path = _files(tmp_path) + + assert main(["validate", str(workflow), "--input", str(input_path), "--json"]) == 0 + captured = capsys.readouterr() + assert json.loads(captured.out) == {"status": "valid", "workflow": "cli_test"} + assert captured.err == "" + + +def test_run_json_success(tmp_path: Path, capsys: pytest.CaptureFixture[str]) -> None: + workflow, input_path = _files(tmp_path) + + assert main(["run", str(workflow), "--input", str(input_path), "--json"]) == 0 + captured = capsys.readouterr() + assert json.loads(captured.out) == { + "status": "completed", + "workflow": "cli_test", + "result": {"value": 5}, + } + assert captured.err == "" + + +def test_run_json_redirects_task_stdout( + tmp_path: Path, capsys: pytest.CaptureFixture[str] +) -> None: + workflow, input_path = _files(tmp_path, "Noisy") + + assert main(["run", str(workflow), "--input", str(input_path), "--json"]) == 0 + captured = capsys.readouterr() + assert json.loads(captured.out)["result"] == {"value": 4} + assert "task diagnostic" in captured.err + + +def test_run_json_failure(tmp_path: Path, capsys: pytest.CaptureFixture[str]) -> None: + workflow, input_path = _files(tmp_path, "Fail") + + assert main(["run", str(workflow), "--input", str(input_path), "--json"]) == 1 + captured = capsys.readouterr() + assert json.loads(captured.out) == { + "status": "failed", + "workflow": "cli_test", + "failed_task": "fail", + "error": { + "code": "task_failed", + "type": "ValueError", + "message": "Task failed", + "task": "fail", + "field": None, + "issues": [], + }, + } + assert captured.err == "" + + +def test_run_json_validation_failure_omits_input_values( + tmp_path: Path, capsys: pytest.CaptureFixture[str] +) -> None: + workflow, input_path = _files(tmp_path) + input_path.write_text("increment:\n value: secret-value\n", encoding="utf-8") + + assert main(["run", str(workflow), "--input", str(input_path), "--json"]) == 1 + captured = capsys.readouterr() + assert "secret-value" not in captured.out + captured.err + result = json.loads(captured.out) + assert result["failed_task"] == "increment" + assert result["error"]["type"] == "ValidationError" + assert result["error"]["field"] == "value" + assert result["error"]["issues"] == [{"field": "value", "code": "int_parsing"}] + + +def test_validate_json_redirects_import_stdout( + tmp_path: Path, capsys: pytest.CaptureFixture[str] +) -> None: + workflow, input_path = _files(tmp_path) + (tmp_path / "pipeline.py").write_text( + (tmp_path / "pipeline.py").read_text(encoding="utf-8") + "\nprint('import diagnostic')\n", + encoding="utf-8", + ) + sys.modules.pop("pipeline", None) + + assert main(["validate", str(workflow), "--input", str(input_path), "--json"]) == 0 + captured = capsys.readouterr() + assert json.loads(captured.out)["status"] == "valid" + assert "import diagnostic" in captured.err + + +def test_validate_json_configuration_failure( + tmp_path: Path, capsys: pytest.CaptureFixture[str] +) -> None: + workflow, input_path = _files(tmp_path) + input_path.write_text("{}\n", encoding="utf-8") + + assert main(["validate", str(workflow), "--input", str(input_path), "--json"]) == 2 + captured = capsys.readouterr() + assert json.loads(captured.out) == { + "status": "invalid", + "error": { + "code": "configuration_error", + "type": "ConfigLoadError", + "message": "Workflow configuration is invalid", + "task": None, + "field": None, + "issues": [], + }, + } + assert captured.err == "" + + +def test_validate_json_missing_config_identifies_task_and_field( + tmp_path: Path, capsys: pytest.CaptureFixture[str] +) -> None: + workflow, input_path = _files(tmp_path) + workflow.write_text( + "workflow:\n name: cli_test\n tasks:\n" + " - task: pipeline.Increment\n config_fields: [value]\n", + encoding="utf-8", + ) + input_path.write_text("{}\n", encoding="utf-8") + + assert main(["validate", str(workflow), "--input", str(input_path), "--json"]) == 2 + result = json.loads(capsys.readouterr().out) + assert result["error"]["task"] == "increment" + assert result["error"]["field"] == "value" + assert result["error"]["issues"] == [{"field": "value", "code": "missing"}] + + +def test_validate_json_schema_failure_omits_input_values( + tmp_path: Path, capsys: pytest.CaptureFixture[str] +) -> None: + workflow, input_path = _files(tmp_path) + workflow.write_text( + workflow.read_text(encoding="utf-8") + "runner:\n timeout_seconds: secret-value\n", + encoding="utf-8", + ) + + assert main(["validate", str(workflow), "--input", str(input_path), "--json"]) == 2 + captured = capsys.readouterr() + assert "secret-value" not in captured.out + captured.err + result = json.loads(captured.out) + assert result["error"]["field"] == "runner.timeout_seconds" + assert result["error"]["issues"][0]["code"] == "float_parsing" + + +def test_run_json_unserializable_output( + tmp_path: Path, capsys: pytest.CaptureFixture[str] +) -> None: + (tmp_path / "opaque_pipeline.py").write_text( + """\ +from taskmaestro import EmptyConfig, ExecutionContext, ObjectModel, Task + +class Resource: + pass + +class Handle(ObjectModel[Resource]): + pass + +class GetHandle(Task[EmptyConfig, Handle]): + def run(self, input: EmptyConfig, ctx: ExecutionContext) -> Handle: + return Handle(value=Resource()) +""", + encoding="utf-8", + ) + workflow = tmp_path / "workflow.yaml" + workflow.write_text( + "workflow:\n name: opaque\n tasks:\n - task: opaque_pipeline.GetHandle\n", + encoding="utf-8", + ) + input_path = tmp_path / "input.yaml" + input_path.write_text("{}\n", encoding="utf-8") + + assert main(["run", str(workflow), "--input", str(input_path), "--json"]) == 1 + captured = capsys.readouterr() + result = json.loads(captured.out) + assert result["status"] == "failed" + assert result["error"]["code"] == "serialization_error" + assert result["failed_task"] is None + assert captured.err == "" + + def test_graph_prints_mermaid(tmp_path: Path, capsys: object) -> None: workflow, input_path = _files(tmp_path) @@ -110,3 +337,343 @@ def test_configuration_errors_return_two(tmp_path: Path, capsys: object) -> None assert status == 2 captured = capsys.readouterr() # type: ignore[attr-defined] assert "Configuration error: Cannot read file" in captured.err + + +def test_error_inside_task_module_is_configuration_error( + tmp_path: Path, capsys: pytest.CaptureFixture[str] +) -> None: + workflow, input_path = _files(tmp_path) + (tmp_path / "pipeline.py").write_text("import missing_dependency_xyz\n", encoding="utf-8") + sys.modules.pop("pipeline", None) + + status = main(["validate", str(workflow), "--input", str(input_path)]) + + assert status == 2 + err = capsys.readouterr().err + assert "Configuration error: Error while importing module 'pipeline'" in err + assert "missing_dependency_xyz" in err + + +@pytest.mark.parametrize("module", ["taskmaestro", "taskmaestro.cli"]) +def test_module_entry_points_run_the_cli( + module: str, + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, + capsys: pytest.CaptureFixture[str], +) -> None: + """``python -m taskmaestro`` and ``python -m taskmaestro.cli`` invoke main().""" + import runpy + import warnings + + workflow, input_path = _files(tmp_path) + monkeypatch.setattr( + sys, "argv", [module, "validate", str(workflow), "--input", str(input_path)] + ) + with warnings.catch_warnings(): + # Re-executing an already imported module as __main__ warns; that is expected here. + warnings.simplefilter("ignore", RuntimeWarning) + with pytest.raises(SystemExit) as excinfo: + runpy.run_module(module, run_name="__main__", alter_sys=True) + assert excinfo.value.code == 0 + assert "Workflow 'cli_test' is valid" in capsys.readouterr().out + + +def test_workflow_describe_without_input( + tmp_path: Path, capsys: pytest.CaptureFixture[str] +) -> None: + workflow, _input_path = _files(tmp_path) + original_path = sys.path.copy() + + assert main(["workflow", "describe", str(workflow), "--json"]) == 0 + assert sys.path == original_path + captured = capsys.readouterr() + assert captured.err == "" + description = json.loads(captured.out) + assert description["workflow"] == "cli_test" + assert description["result_task"] == "increment" + assert len(description["tasks"]) == 1 + task = description["tasks"][0] + assert task["name"] == "increment" + assert task["python_type"] == "pipeline.Increment" + assert task["depends_on"] is None + assert task["required_input_fields"] == ["value"] + assert task["config_fields"] == [] + assert task["provided_config_fields"] is None + assert task["input_schema"]["properties"]["value"]["type"] == "integer" + assert task["output_schema"]["properties"]["value"]["type"] == "integer" + assert main(["workflow", "describe", str(workflow)]) == 0 + assert json.loads(capsys.readouterr().out) == description + + +def test_workflow_describe_redirects_import_stdout( + tmp_path: Path, capsys: pytest.CaptureFixture[str] +) -> None: + workflow, _input_path = _files(tmp_path) + (tmp_path / "pipeline.py").write_text( + (tmp_path / "pipeline.py").read_text(encoding="utf-8") + "\nprint('import diagnostic')\n", + encoding="utf-8", + ) + sys.modules.pop("pipeline", None) + + assert main(["workflow", "describe", str(workflow), "--json"]) == 0 + captured = capsys.readouterr() + assert json.loads(captured.out)["workflow"] == "cli_test" + assert "import diagnostic" in captured.err + + +def test_workflow_describe_with_input_does_not_execute( + tmp_path: Path, capsys: pytest.CaptureFixture[str] +) -> None: + workflow, input_path = _files(tmp_path, "Fail") + # Inspections should not instantiate hooks or execute the task's failing run(). + input_path.write_text("fail:\n value: private-token\n", encoding="utf-8") + workflow.write_text( + workflow.read_text(encoding="utf-8") + "runner:\n hooks:\n - hook: nonexistent.hook\n", + encoding="utf-8", + ) + + assert main(["workflow", "describe", str(workflow), "--input", str(input_path), "--json"]) == 0 + description = json.loads(capsys.readouterr().out) + assert description["tasks"][0]["provided_config_fields"] == ["value"] + assert "private-token" not in json.dumps(description) + + +def test_workflow_describe_missing_input_reports_json_error( + tmp_path: Path, capsys: pytest.CaptureFixture[str] +) -> None: + workflow, input_path = _files(tmp_path) + input_path.write_text("{}\n", encoding="utf-8") + + assert main(["workflow", "describe", str(workflow), "--input", str(input_path), "--json"]) == 2 + result = json.loads(capsys.readouterr().out) + assert result["status"] == "invalid" + assert result["error"]["code"] == "configuration_error" + + +def test_workflow_describe_reports_missing_config_fields( + tmp_path: Path, capsys: pytest.CaptureFixture[str] +) -> None: + workflow, input_path = _files(tmp_path) + workflow.write_text( + "workflow:\n name: cli_test\n tasks:\n" + " - task: pipeline.Increment\n config_fields: [value]\n", + encoding="utf-8", + ) + input_path.write_text("{}\n", encoding="utf-8") + + assert main(["workflow", "describe", str(workflow), "--input", str(input_path), "--json"]) == 2 + result = json.loads(capsys.readouterr().out) + assert result["error"]["task"] == "increment" + assert result["error"]["issues"] == [{"field": "value", "code": "missing"}] + + +def test_workflow_dependency_routing_and_positional_collection() -> None: + """The inspector renders validated tuple and ordered collection references.""" + assert _dependency_spec(("producer", "content")) == { + "task": "producer", + "field": "content", + } + assert _dependency_spec( + { + "items": CollectionRef( + "positional", + positional_members=(OutputRef("first"), OutputRef("second", "content")), + ) + } + ) == { + "items": { + "collect": { + "kind": "positional", + "members": [ + {"task": "first", "field": None}, + {"task": "second", "field": "content"}, + ], + } + } + } + + +def test_workflow_describe_fan_in_collections_and_mapping( + capsys: pytest.CaptureFixture[str], monkeypatch: pytest.MonkeyPatch +) -> None: + workflow = Path(__file__).resolve().parents[1] / "examples/release_pipeline/workflow.yaml" + + # Earlier CLI tests import a different local module also named "pipeline". + with monkeypatch.context() as patch: + patch.delitem(sys.modules, "pipeline", raising=False) + assert main(["workflow", "describe", str(workflow), "--json"]) == 0 + patch.delitem(sys.modules, "pipeline", raising=False) + tasks = {task["name"]: task for task in json.loads(capsys.readouterr().out)["tasks"]} + checks = tasks["validate_release"]["depends_on"]["checks"]["collect"] + assert checks == { + "kind": "keyed", + "members": { + "tests": {"task": "run_tests", "field": None}, + "lint": {"task": "run_lint", "field": None}, + "types": {"task": "check_types", "field": None}, + }, + } + assert tasks["build_targets"]["map"] == { + "over": "targets", + "key_as": "target_name", + "value_as": "settings", + "error_mode": "collect_all", + } + assert tasks["create_release_manifest"]["depends_on"]["artifacts"] == { + "task": "build_targets", + "field": "root", + } + assert tasks["build_targets"]["output_schema"]["additionalProperties"] == { + "$ref": "#/$defs/Artifact" + } + + +def test_tasks_list_does_not_import_plugins( + monkeypatch: pytest.MonkeyPatch, capsys: pytest.CaptureFixture[str] +) -> None: + entries = [ + EntryPoint(name="z.broken", value="missing_module:Task", group="taskmaestro.tasks"), + EntryPoint( + name="a.valid", value="tests.test_discovery:ExampleTask", group="taskmaestro.tasks" + ), + ] + monkeypatch.setattr("taskmaestro.discovery.entry_points", lambda *, group: entries) + + assert main(["tasks", "list", "--json"]) == 0 + assert json.loads(capsys.readouterr().out) == {"tasks": ["a.valid", "z.broken"]} + assert main(["tasks", "list"]) == 0 + assert capsys.readouterr().out == "a.valid\nz.broken\n" + + +def test_tasks_list_rejects_duplicate_identifiers( + monkeypatch: pytest.MonkeyPatch, capsys: pytest.CaptureFixture[str] +) -> None: + entry = EntryPoint( + name="duplicate", value="tests.test_discovery:ExampleTask", group="taskmaestro.tasks" + ) + monkeypatch.setattr("taskmaestro.discovery.entry_points", lambda *, group: [entry, entry]) + + assert main(["tasks", "list", "--json"]) == 2 + captured = capsys.readouterr() + assert captured.out == "" + assert "Plugin error: Multiple entry points named 'duplicate'" in captured.err + + +def test_tasks_describe_json( + monkeypatch: pytest.MonkeyPatch, capsys: pytest.CaptureFixture[str] +) -> None: + entries = [ + EntryPoint( + name="example.increment", + value="tests.test_discovery:ExampleTask", + group="taskmaestro.tasks", + ) + ] + monkeypatch.setattr("taskmaestro.discovery.entry_points", lambda *, group: entries) + + assert main(["tasks", "describe", "example.increment", "--json"]) == 0 + description = json.loads(capsys.readouterr().out) + assert description["identifier"] == "example.increment" + assert description["name"] == "ExampleTask" + assert description["timeout_seconds"] is None + assert description["input_schema"]["properties"]["value"]["type"] == "integer" + assert description["input_schema"]["required"] == ["value"] + assert description["output_schema"]["properties"]["value"]["type"] == "integer" + assert main(["tasks", "describe", "example.increment"]) == 0 + assert json.loads(capsys.readouterr().out) == description + + +def test_tasks_describe_runtime_only_objects( + monkeypatch: pytest.MonkeyPatch, capsys: pytest.CaptureFixture[str] +) -> None: + entries = [ + EntryPoint( + name="example.opaque", value="tests.test_cli:OpaqueTask", group="taskmaestro.tasks" + ) + ] + monkeypatch.setattr("taskmaestro.discovery.entry_points", lambda *, group: entries) + + assert main(["tasks", "describe", "example.opaque", "--json"]) == 0 + description = json.loads(capsys.readouterr().out) + input_schema = description["input_schema"] + handle_schema = input_schema["$defs"]["ClientHandle"]["properties"]["value"] + assert handle_schema["not"] == {} + assert handle_schema["x-taskmaestro-opaque"] is True + assert handle_schema["x-taskmaestro-python-type"] == "tests.test_cli.ExternalClient" + assert input_schema["properties"]["amount"]["type"] == "integer" + output_schema = description["output_schema"] + assert output_schema["properties"]["value"]["x-taskmaestro-opaque"] is True + + +def test_tasks_describe_unsupported_schema_reports_plugin_error( + monkeypatch: pytest.MonkeyPatch, capsys: pytest.CaptureFixture[str] +) -> None: + entry = EntryPoint( + name="unsupported", value="tests.test_cli:UnsupportedTask", group="taskmaestro.tasks" + ) + monkeypatch.setattr("taskmaestro.discovery.entry_points", lambda *, group: [entry]) + + assert main(["tasks", "describe", "unsupported", "--json"]) == 2 + captured = capsys.readouterr() + assert captured.out == "" + assert "Plugin error: Cannot describe task 'unsupported'" in captured.err + assert "PlainValidatorFunctionSchema" in captured.err + + +def test_workflow_describe_unsupported_schema_reports_configuration_error( + tmp_path: Path, capsys: pytest.CaptureFixture[str] +) -> None: + workflow = tmp_path / "workflow.yaml" + workflow.write_text( + "workflow:\n name: unsupported\n tasks:\n - task: tests.test_cli.UnsupportedTask\n", + encoding="utf-8", + ) + + assert main(["workflow", "describe", str(workflow), "--json"]) == 2 + captured = capsys.readouterr() + assert json.loads(captured.out)["error"]["code"] == "configuration_error" + assert captured.err == "" + + +def test_tasks_describe_unknown_plugin( + monkeypatch: pytest.MonkeyPatch, capsys: pytest.CaptureFixture[str] +) -> None: + monkeypatch.setattr("taskmaestro.discovery.entry_points", lambda *, group: []) + + assert main(["tasks", "describe", "missing", "--json"]) == 2 + captured = capsys.readouterr() + assert captured.out == "" + assert "Plugin error: No task entry point named 'missing'" in captured.err + + +def test_tasks_describe_broken_plugin( + monkeypatch: pytest.MonkeyPatch, capsys: pytest.CaptureFixture[str] +) -> None: + entries = [ + EntryPoint( + name="broken", value="missing_taskmaestro_plugin:Task", group="taskmaestro.tasks" + ) + ] + monkeypatch.setattr("taskmaestro.discovery.entry_points", lambda *, group: entries) + + assert main(["tasks", "list", "--json"]) == 0 + assert json.loads(capsys.readouterr().out) == {"tasks": ["broken"]} + assert main(["tasks", "describe", "broken", "--json"]) == 2 + captured = capsys.readouterr() + assert captured.out == "" + assert "Plugin error: Cannot load task entry point 'broken'" in captured.err + + +def test_python_dash_m_exit_code(tmp_path: Path) -> None: + """The real interpreter invocation propagates main()'s exit status.""" + import subprocess + + workflow, input_path = _files(tmp_path, "Fail") + completed = subprocess.run( + [sys.executable, "-m", "taskmaestro", "run", str(workflow), "--input", str(input_path)], + capture_output=True, + text=True, + check=False, + ) + assert completed.returncode == 1 + assert "Workflow failed at fail: intentional failure" in completed.stderr diff --git a/tests/test_job.py b/tests/test_job.py index 54ac73f..b998324 100644 --- a/tests/test_job.py +++ b/tests/test_job.py @@ -18,6 +18,8 @@ ConfigOnlyTask, Double, FanInTask, + FanInWithConfigTask, + MergeTask, NumberInput, NumberOutput, ) @@ -128,6 +130,40 @@ def test_missing_declared_configuration_is_rejected(self) -> None: with pytest.raises(WorkflowDefinitionError, match="missing configuration fields"): Job(workflow=workflow, config=EmptyConfig()) + def test_missing_configuration_on_dependent_task_is_rejected(self) -> None: + """Dependent tasks' config fields are checked too, not only roots'.""" + workflow = ( + Workflow.builder(name="cfg") + .add_task(AddOne) + .add_task(MergeTask, depends_on=AddOne, config_fields=["label"]) + .build() + ) + + with pytest.raises( + WorkflowDefinitionError, + match=r"Task 'merge_task' is missing configuration fields \['label'\]", + ): + Job(workflow=workflow, config=NumberInput(value=1)) + with pytest.raises(WorkflowDefinitionError, match="missing configuration fields"): + Job( + workflow=workflow, + config=NumberInput(value=1), + job_configuration=JobConfiguration({"merge_task": {}}), + ) + + def test_missing_configuration_on_fan_in_task_is_rejected(self) -> None: + workflow = ( + Workflow.builder(name="cfg") + .add_task(AddOne) + .add_task(FanInWithConfigTask, depends_on={"a": AddOne}, config_fields=["extra"]) + .build() + ) + + with pytest.raises( + WorkflowDefinitionError, match=r"missing configuration fields \['extra'\]" + ): + Job(workflow=workflow, config=NumberInput(value=1)) + def test_job_configuration_stored(self) -> None: wf = ( Workflow.builder(name="cfg") diff --git a/tests/test_mapping.py b/tests/test_mapping.py index 101961e..64c0821 100644 --- a/tests/test_mapping.py +++ b/tests/test_mapping.py @@ -405,6 +405,14 @@ def test_job_configuration_is_required(self) -> None: with pytest.raises(WorkflowDefinitionError, match="requires JobConfiguration"): Job(workflow, NumberInput(value=1)) + def test_mapped_task_config_fields_are_required(self) -> None: + workflow = _mapped_workflow() + config = JobConfiguration({"mapped_number": {"items": {"one": 1}}}) + with pytest.raises( + WorkflowDefinitionError, match=r"missing configuration fields \['multiplier'\]" + ): + Job(workflow, NumberInput(value=1), job_configuration=config) + def test_map_source_is_required(self) -> None: workflow = _mapped_workflow() config = JobConfiguration({"mapped_number": {"multiplier": 2}}) @@ -698,7 +706,9 @@ def test_item_input_validation_is_recorded_as_item_failure(self) -> None: job = Job( workflow, NumberInput(value=1), - job_configuration=JobConfiguration({"mapped_number": {"items": {"one": 1}}}), + job_configuration=JobConfiguration( + {"mapped_number": {"multiplier": "not-an-int", "items": {"one": 1}}} + ), ) hook = RecordingMapHook() diff --git a/tests/test_runner.py b/tests/test_runner.py index 01e1bb3..1d93a68 100644 --- a/tests/test_runner.py +++ b/tests/test_runner.py @@ -2,6 +2,8 @@ from __future__ import annotations +from typing import Any + import pytest from pydantic import BaseModel, ConfigDict, ValidationError @@ -16,6 +18,7 @@ Workflow, ) from taskmaestro.exceptions import JobStateError +from taskmaestro.hooks.base import BaseHook from taskmaestro.job import TaskStatus from tests.conftest import ( AddOne, @@ -103,6 +106,70 @@ def test_rerun_guard(self, ctx: ExecutionContext) -> None: with pytest.raises(JobStateError, match="Cannot run job"): runner.run(job, ctx=ctx) + def test_invalid_config_value_fails_job(self, ctx: ExecutionContext) -> None: + """Input validation errors are task failures, not escaping exceptions.""" + wf = ( + Workflow.builder("merge") + .add_task(AddOne) + .add_task(MergeTask, depends_on=AddOne, config_fields=["label"]) + .build() + ) + jc = JobConfiguration({"merge_task": {"label": ["not", "a", "string"]}}) + events: list[str] = [] + + class Recorder(BaseHook): + def on_task_start(self, job: Job[Any], task: Task[Any, Any]) -> None: + events.append(f"task_start:{task.name}") + + def on_task_fail(self, job: Job[Any], task: Task[Any, Any], error: Exception) -> None: + events.append(f"task_fail:{task.name}") + + def on_job_fail(self, job: Job[Any]) -> None: + events.append("job_fail") + + job = Job(wf, NumberInput(value=1), job_configuration=jc) + result = Runner(hooks=[Recorder()]).run(job, ctx=ctx) + assert result.status == JobStatus.FAILED + assert result.failed_task == "merge_task" + assert isinstance(result.exception, ValidationError) + assert result.completed_at is not None + assert [r.status for r in result.task_results] == [ + TaskStatus.COMPLETED, + TaskStatus.FAILED, + ] + assert events[-3:] == ["task_start:merge_task", "task_fail:merge_task", "job_fail"] + + def test_task_constructor_failure_fails_job(self, ctx: ExecutionContext) -> None: + """A task whose constructor raises fails the job without a TASK_FAIL event.""" + + class BrokenInit(Task[NumberInput, NumberOutput]): + name = "broken_init" + + def __init__(self) -> None: + raise RuntimeError("cannot construct") + + def run( + self, input: NumberInput, ctx: ExecutionContext + ) -> NumberOutput: # pragma: no cover - never constructed + return NumberOutput(value=input.value) + + events: list[str] = [] + + class Recorder(BaseHook): + def on_task_fail(self, job: Job[Any], task: Task[Any, Any], error: Exception) -> None: + events.append("task_fail") # pragma: no cover - must not be called + + def on_job_fail(self, job: Job[Any]) -> None: + events.append("job_fail") + + job = Job(Workflow(name="test", tasks=[BrokenInit]), NumberInput(value=1)) + result = Runner(hooks=[Recorder()]).run(job, ctx=ctx) + assert result.status == JobStatus.FAILED + assert result.failed_task == "broken_init" + assert result.error == "cannot construct" + assert result.task_results[0].task_name == "broken_init" + assert events == ["job_fail"] + def test_failed_task_results(self, ctx: ExecutionContext) -> None: wf = Workflow(name="test", tasks=[FailingTask]) job = Job(workflow=wf, config=NumberInput(value=1)) @@ -326,6 +393,99 @@ def run(self, input: NumberOutput, ctx: ExecutionContext) -> NumberOutput: 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_nested_run_does_not_cancel_outer_task_timeout(self, ctx: ExecutionContext) -> None: + """An inner workflow cancelling its own alarm must not cancel the wrapper's.""" + import signal + 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 InnerSlow(Task[NumberOutput, NumberOutput]): + name = "inner_slow" + + def run(self, input: NumberOutput, ctx: ExecutionContext) -> NumberOutput: + time.sleep(5) + return input # pragma: no cover - interrupted by the outer alarm + + wrapper = Workflow(name="inner", tasks=[InnerWithTimeout, InnerSlow]).as_task( + name="wrapper" + ) + wrapper.timeout_seconds = 0.3 + handler_before = signal.getsignal(signal.SIGALRM) + job = Job(Workflow(name="outer", tasks=[wrapper]), NumberInput(value=1)) + start = time.monotonic() + result = Runner().run(job, ctx=ctx) + assert time.monotonic() - start < 3 + assert result.status == JobStatus.FAILED + assert result.failed_task == "wrapper" + assert "wrapper timed out after 0.3s" in (result.error or "") + assert signal.getsignal(signal.SIGALRM) is handler_before + assert signal.getitimer(signal.ITIMER_REAL) == (0.0, 0.0) + + @pytest.mark.skipif( + not hasattr(__import__("signal"), "SIGALRM"), + reason="signal.SIGALRM not available on this platform", + ) + def test_nested_run_does_not_cancel_outer_job_timeout(self, ctx: ExecutionContext) -> None: + """The job deadline also interrupts a nested workflow that is the last task.""" + 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 InnerSlow(Task[NumberOutput, NumberOutput]): + name = "inner_slow" + + def run(self, input: NumberOutput, ctx: ExecutionContext) -> NumberOutput: + time.sleep(5) + return input # pragma: no cover - interrupted by the job alarm + + wrapper = Workflow(name="inner", tasks=[InnerWithTimeout, InnerSlow]).as_task( + name="wrapper" + ) + job = Job(Workflow(name="outer", tasks=[wrapper]), NumberInput(value=1)) + start = time.monotonic() + result = Runner().run(job, ctx=ctx, timeout_seconds=0.3) + assert time.monotonic() - start < 3 + assert result.status == JobStatus.FAILED + assert "Job timed out after 0.3s" in (result.error or "") + + @pytest.mark.skipif( + not hasattr(__import__("signal"), "SIGALRM"), + reason="signal.SIGALRM not available on this platform", + ) + def test_stale_alarm_is_ignored(self) -> None: + """A SIGALRM arriving after every registered timer fired raises nothing.""" + import signal + import time + + from taskmaestro.exceptions import TaskTimeoutError + from taskmaestro.runner import _AlarmScheduler, _Timer + + scheduler = _AlarmScheduler() + timer = _Timer(time.monotonic() + 60, TaskTimeoutError, "t timed out") + scheduler.add(timer) + try: + with pytest.raises(TaskTimeoutError, match="t timed out"): + signal.raise_signal(signal.SIGALRM) + assert timer.fired + signal.raise_signal(signal.SIGALRM) # stale: must not raise + finally: + scheduler.remove(timer) + @pytest.mark.skipif( not hasattr(__import__("signal"), "SIGALRM"), reason="signal.SIGALRM not available on this platform", @@ -453,6 +613,116 @@ def sentinel(signum: int, frame: object) -> None: # pragma: no cover signal.signal(signal.SIGALRM, previous) +class TestTimerNotArmedDuringHooks: + """The task timer is stopped as soon as the task returns or raises.""" + + @pytest.mark.skipif( + not hasattr(__import__("signal"), "SIGALRM"), + reason="signal.SIGALRM not available on this platform", + ) + def test_timer_idle_in_task_hooks(self, ctx: ExecutionContext) -> None: + import signal + + class Timed(Task[NumberInput, NumberOutput]): + name = "timed" + timeout_seconds = 30 + + def run(self, input: NumberInput, ctx: ExecutionContext) -> NumberOutput: + if input.value < 0: + raise ValueError("negative") + return NumberOutput(value=input.value) + + timers: dict[str, tuple[float, float]] = {} + + class Probe(BaseHook): + def on_task_complete( + self, job: Job[Any], task: Task[Any, Any], output: BaseModel + ) -> None: + timers["complete"] = signal.getitimer(signal.ITIMER_REAL) + + def on_task_fail(self, job: Job[Any], task: Task[Any, Any], error: Exception) -> None: + timers["fail"] = signal.getitimer(signal.ITIMER_REAL) + + wf = Workflow(name="test", tasks=[Timed]) + Runner(hooks=[Probe()]).run(Job(wf, NumberInput(value=1)), ctx=ctx, timeout_seconds=60) + Runner(hooks=[Probe()]).run(Job(wf, NumberInput(value=-1)), ctx=ctx, timeout_seconds=60) + assert timers == {"complete": (0.0, 0.0), "fail": (0.0, 0.0)} + + @pytest.mark.skipif( + not hasattr(__import__("signal"), "SIGALRM"), + reason="signal.SIGALRM not available on this platform", + ) + def test_timer_idle_in_map_item_hooks(self) -> None: + import signal + + from taskmaestro import TaskMap + + class ItemInput(BaseModel): + key: str + value: int + + class TimedItem(Task[ItemInput, NumberOutput]): + name = "timed_item" + timeout_seconds = 30 + + def run(self, input: ItemInput, ctx: ExecutionContext) -> NumberOutput: + if input.value < 0: + raise ValueError("negative") + return NumberOutput(value=input.value) + + timers: dict[str, tuple[float, float]] = {} + + class Probe(BaseHook): + def on_map_item_complete( + self, job: Job[Any], task: Task[Any, Any], key: str, output: BaseModel + ) -> None: + timers[f"complete:{key}"] = signal.getitimer(signal.ITIMER_REAL) + + def on_map_item_fail( + self, job: Job[Any], task: Task[Any, Any], key: str, error: Exception + ) -> None: + timers[f"fail:{key}"] = signal.getitimer(signal.ITIMER_REAL) + + wf = ( + Workflow.builder("mapped") + .add_task(TimedItem, mapped_over=TaskMap("items", "key", "value", "collect_all")) + .build() + ) + jc = JobConfiguration({"timed_item": {"items": {"ok": 1, "bad": -1}}}) + result = Runner(hooks=[Probe()]).run(Job(wf, EmptyConfig(), job_configuration=jc)) + assert result.status == JobStatus.FAILED + assert timers == {"complete:ok": (0.0, 0.0), "fail:bad": (0.0, 0.0)} + + @pytest.mark.skipif( + not hasattr(__import__("signal"), "SIGALRM"), + reason="signal.SIGALRM not available on this platform", + ) + def test_slow_completion_hook_does_not_time_out_task(self, ctx: ExecutionContext) -> None: + """A hook outlasting the task's timeout must not turn success into a lost timeout.""" + import time + import warnings + + class Quick(Task[NumberInput, NumberOutput]): + name = "quick" + timeout_seconds = 0.1 + + def run(self, input: NumberInput, ctx: ExecutionContext) -> NumberOutput: + return NumberOutput(value=input.value) + + class SlowHook(BaseHook): + def on_task_complete( + self, job: Job[Any], task: Task[Any, Any], output: BaseModel + ) -> None: + time.sleep(0.3) + + with warnings.catch_warnings(): + warnings.simplefilter("error") # a swallowed timeout would surface as HookError + result = Workflow(name="test", tasks=[Quick]).run( + NumberInput(value=1), hooks=[SlowHook()] + ) + assert result.status == JobStatus.COMPLETED + + class TestAlarmUnavailable: def test_alarm_unavailable_warns(self, ctx: ExecutionContext) -> None: """When signal.alarm is unavailable, a warning is issued and execution proceeds.""" @@ -665,8 +935,11 @@ def run(self, input: StrictInput, ctx: ExecutionContext) -> NumberOutput: 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) + result = Runner().run(job, ctx=ctx) + assert result.status == JobStatus.FAILED + assert result.failed_task == "strict_task" + assert isinstance(result.exception, ValidationError) + assert "unexpected" in (result.error or "") def test_backward_compat_no_config(self, ctx: ExecutionContext) -> None: """Workflow without config_fields runs normally.""" diff --git a/tests/test_visualization.py b/tests/test_visualization.py index 5f2a6ea..e078256 100644 --- a/tests/test_visualization.py +++ b/tests/test_visualization.py @@ -432,6 +432,55 @@ def test_non_workflow_task_nodes_normal(self) -> None: assert 'stringify["stringify"]' in result assert "stringify -->|StringOutput| _end_" in result + def test_outer_edges_skip_mapped_inner_root(self) -> None: + """A mapped inner root is fed by config, so edges enter at the input root.""" + from taskmaestro import JobConfiguration, TaskMap + from tests.conftest import AddOne, NumberInput, NumberOutput + + class ItemInput(BaseModel): + key: str + value: int + + class AMapped(Task[ItemInput, NumberOutput]): + name = "a_mapped" # sorts before the input root + + def run(self, input: ItemInput, ctx: ExecutionContext) -> NumberOutput: + return NumberOutput(value=input.value) # pragma: no cover - never run + + class CombineInput(BaseModel): + mapped: dict[str, NumberOutput] + entry: NumberOutput + + class Combine(Task[CombineInput, NumberOutput]): + name = "combine" + + def run(self, input: CombineInput, ctx: ExecutionContext) -> NumberOutput: + return input.entry # pragma: no cover - never run + + inner = ( + Workflow.builder("inner") + .add_task(AMapped, mapped_over=TaskMap("items", "key", "value")) + .add_task(AddOne) + .add_task(Combine, depends_on={"mapped": AMapped, "entry": AddOne}) + .build() + ) + wrapped = inner.as_task( + name="wrapped", + job_configuration=JobConfiguration({"a_mapped": {"items": {"x": 1}}}), + ) + + class Source(Task[NumberInput, NumberInput]): + name = "source" + + def run(self, input: NumberInput, ctx: ExecutionContext) -> NumberInput: + return input # pragma: no cover - never run + + outer = Workflow("outer", [Source, wrapped]) + result = to_mermaid(outer) + + assert "source -->|NumberInput| wrapped__add_one" in result + assert "source -->|NumberInput| wrapped__a_mapped" not in result + def test_three_task_inner_workflow_subgraph(self) -> None: """workflow_task wrapping a 3-task chain renders all inner nodes/edges.""" from tests.conftest import AddOne, Double, Stringify diff --git a/tests/test_workflow.py b/tests/test_workflow.py index a9cae58..a1dced7 100644 --- a/tests/test_workflow.py +++ b/tests/test_workflow.py @@ -275,6 +275,67 @@ def test_handle_can_select_result_task(self) -> None: assert workflow.result_task_name == "add_one" +class TestBuilderSnapshot: + """build() returns an independent workflow, not the builder's live state.""" + + def test_builder_changes_after_build_do_not_leak(self) -> None: + builder = Workflow.builder("snapshot") + builder.add_task(AddOne) + workflow = builder.build() + + builder.add_task(Double, depends_on=AddOne) + + assert [name for name, _ in workflow.topological_order()] == ["add_one"] + assert workflow.result_task_name == "add_one" + result = workflow.run(NumberInput(value=1)) + assert result.status == JobStatus.COMPLETED + assert result.result == NumberOutput(value=2) + + def test_builder_can_be_extended_and_rebuilt(self) -> None: + builder = Workflow.builder("rebuild") + builder.add_task(AddOne) + first = builder.build() + builder.add_task(Double, depends_on=AddOne) + second = builder.build() + + assert first is not second + assert first.result_task_name == "add_one" + assert second.result_task_name == "double" + assert second.run(NumberInput(value=1)).result == NumberOutput(value=4) + + def test_validation_rewrites_do_not_touch_builder_state(self) -> None: + """Mapped fan-in unwrapping rewrites the built workflow's own dependency copy.""" + from taskmaestro import MappedOutput, TaskMap + + class ItemInput(BaseModel): + key: str + value: int + + class Item(Task[ItemInput, NumberOutput]): + name = "item" + + def run(self, input: ItemInput, ctx: ExecutionContext) -> NumberOutput: + return NumberOutput(value=input.value) # pragma: no cover - never run + + class SumInput(BaseModel): + values: dict[str, NumberOutput] + + class Sum(Task[SumInput, NumberOutput]): + name = "sum" + + def run(self, input: SumInput, ctx: ExecutionContext) -> NumberOutput: + return NumberOutput(value=0) # pragma: no cover - never run + + builder = Workflow.builder("mapped") + builder.add_task(Item, mapped_over=TaskMap("items", "key", "value")) + builder.add_task(Sum, depends_on={"values": Item}) + workflow = builder.build() + + assert workflow.get_dependencies("sum") == {"values": ("item", "root")} + assert builder._workflow.get_dependencies("sum") == {"values": "item"} + assert workflow.get_output_annotation("item") == MappedOutput[NumberOutput] + + class TestDAGWorkflow: def test_fan_in_workflow(self) -> None: wf = ( @@ -463,6 +524,56 @@ def run(self, input: NumberInput, ctx: ExecutionContext) -> ListOut: .build() ) + def test_field_ref_with_config_fields_raises(self) -> None: + """A field-routed input has no fields to merge config into, so reject it.""" + + class Wrapped(BaseModel): + inner: NumberInput + note: str + + class Producer(Task[NumberInput, Wrapped]): + name = "producer" + + def run(self, input: NumberInput, ctx: ExecutionContext) -> Wrapped: + return Wrapped(inner=input, note="") # pragma: no cover - never run + + with pytest.raises( + WorkflowDefinitionError, + match=r"output field 'producer\.inner' and cannot also declare config_fields", + ): + ( + Workflow.builder("bad") + .add_task(Producer) + .add_task(AddOne, depends_on=(Producer, "inner"), config_fields=["value"]) + .build() + ) + + def test_named_field_ref_with_config_fields_merges(self) -> None: + """The suggested named-dependency form combines a routed field with config.""" + + class Wrapped(BaseModel): + number: NumberOutput + + class Producer(Task[NumberInput, Wrapped]): + name = "producer" + + def run(self, input: NumberInput, ctx: ExecutionContext) -> Wrapped: + return Wrapped(number=NumberOutput(value=input.value)) + + wf = ( + Workflow.builder("ok") + .add_task(Producer) + .add_task( + FanInWithConfigTask, + depends_on={"a": (Producer, "number")}, + config_fields=["extra"], + ) + .build() + ) + result = wf.run(NumberInput(value=4), task_config={"fan_in_with_config": {"extra": "x"}}) + assert result.status == JobStatus.COMPLETED + assert result.result.combined == "x:4" # type: ignore[union-attr] + def test_field_ref_accepts_subclass(self) -> None: """Field-ref edges use type compatibility, not identity.""" @@ -701,6 +812,32 @@ def test_fan_in_deps_dict(self) -> None: assert deps == {"a": "add_one", "b": "add_one_b"} +class TestInputRootNames: + def test_excludes_configured_mapped_and_dependent_tasks(self) -> None: + from taskmaestro import TaskMap + + class ItemInput(BaseModel): + key: str + value: int + + class Mapped(Task[ItemInput, NumberOutput]): + name = "mapped" + + def run(self, input: ItemInput, ctx: ExecutionContext) -> NumberOutput: + return NumberOutput(value=input.value) # pragma: no cover - never run + + wf = ( + Workflow.builder("roots", result_task="double") + .add_task(AddOneB) + .add_task(ConfigOnlyTask, config_fields=["path", "count"]) + .add_task(Mapped, mapped_over=TaskMap("items", "key", "value")) + .add_task(AddOne) + .add_task(Double, depends_on=AddOne) + .build() + ) + assert wf.input_root_names() == ["add_one_b", "add_one"] + + class TestNamedTaskInstances: """Tests for using the same Task class with different names.""" diff --git a/tests/test_yaml_config.py b/tests/test_yaml_config.py index 30bde67..c339e61 100644 --- a/tests/test_yaml_config.py +++ b/tests/test_yaml_config.py @@ -2,6 +2,7 @@ from __future__ import annotations +import re from pathlib import Path import pytest @@ -14,6 +15,7 @@ JobStatus, Task, ) +from taskmaestro.hooks.base import BaseHook from taskmaestro.yaml_config import ( TaskConfig, YamlWorkflowConfig, @@ -42,6 +44,14 @@ def run(self, input: TextInput, ctx: ExecutionContext) -> TextOutput: return TextOutput(text=input.text.upper()) +class RejectingHook(BaseHook): + """Hook whose constructor validates its parameters.""" + + def __init__(self, limit: int) -> None: + if limit < 0: + raise ValueError("limit must be >= 0") + + class ReverseText(Task[TextOutput, TextOutput]): name = "reverse_text" @@ -156,6 +166,19 @@ def test_explicit_duplicates_are_still_rejected(self, text: str) -> None: with pytest.raises(yaml.constructor.ConstructorError, match="duplicate key"): _yaml_load(text) + @pytest.mark.parametrize("text", ["? [a, b]\n: 1\n", "? {a: 1}\n: 1\n"]) + def test_unhashable_keys_are_yaml_errors(self, text: str) -> None: + with pytest.raises(yaml.constructor.ConstructorError, match="found unhashable key"): + _yaml_load(text) + + def test_unhashable_key_in_input_is_config_error(self, tmp_path: Path) -> None: + wf_path = _write_workflow_yaml( + tmp_path, f"workflow:\n name: w\n tasks:\n - task: {THIS_MODULE}.UpperText\n" + ) + in_path = _write_input_yaml(tmp_path, "? [a, b]\n: 1\n") + with pytest.raises(ConfigLoadError, match=r"(?s)Input YAML parse error.*unhashable key"): + load_workflow_from_yaml(wf_path, in_path) + def test_workflow_and_input_yaml_support_merges(self, tmp_path: Path) -> None: workflow_path = _write_workflow_yaml( tmp_path, @@ -188,6 +211,42 @@ def test_nonexistent_module(self) -> None: with pytest.raises(ConfigLoadError, match="Cannot import module"): import_class("nonexistent.module.ClassName") + @pytest.mark.parametrize( + ("source", "expected"), + [ + ("import missing_dependency_xyz\n", "No module named 'missing_dependency_xyz'"), + ("def broken(:\n", "SyntaxError"), + ("raise RuntimeError('module init failed')\n", "module init failed"), + ("from os import no_such_name\n", "ImportError"), + ], + ) + def test_errors_inside_the_module_are_config_errors( + self, + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, + source: str, + expected: str, + ) -> None: + """Failures raised while importing an existing module are wrapped, with context.""" + import sys + import zlib + + module_name = f"broken_module_{zlib.crc32(source.encode())}" + (tmp_path / f"{module_name}.py").write_text(source, encoding="utf-8") + monkeypatch.syspath_prepend(str(tmp_path)) + monkeypatch.delitem(sys.modules, module_name, raising=False) + + with pytest.raises(ConfigLoadError) as excinfo: + import_class(f"{module_name}.Anything") + message = str(excinfo.value) + assert message.startswith(f"Error while importing module '{module_name}'") + assert expected in message + assert excinfo.value.__cause__ is not None + + def test_missing_subpackage_is_reported_as_missing_module(self) -> None: + with pytest.raises(ConfigLoadError, match=r"Cannot import module 'taskmaestro\.nope'"): + import_class("taskmaestro.nope.ClassName") + def test_nonexistent_class(self) -> None: with pytest.raises(ConfigLoadError, match="has no attribute"): import_class(f"{THIS_MODULE}.NonexistentClass") @@ -607,6 +666,42 @@ def test_hook_bad_params(self, tmp_path: Path) -> None: with pytest.raises(ConfigLoadError, match="Cannot instantiate hook"): load_workflow_from_yaml(wf_path, in_path) + @pytest.mark.parametrize( + ("hook", "params", "expected"), + [ + (f"{THIS_MODULE}.RejectingHook", "limit: -1", "ValueError: limit must be >= 0"), + ( + "taskmaestro.hooks.persistence.ResultPersistenceHook", + "output_dir: 5", + "TypeError: ", + ), + ], + ) + def test_hook_constructor_errors_are_config_errors( + self, tmp_path: Path, hook: str, params: str, expected: str + ) -> None: + wf_path = _write_workflow_yaml( + tmp_path, + f"""\ +workflow: + name: bad_hook + tasks: + - task: {THIS_MODULE}.UpperText +runner: + hooks: + - hook: {hook} + params: + {params} +""", + ) + in_path = _write_input_yaml(tmp_path, "upper_text:\n text: hello\n") + with pytest.raises( + ConfigLoadError, match=re.escape(f"Cannot instantiate hook '{hook}'") + ) as excinfo: + load_workflow_from_yaml(wf_path, in_path) + assert expected in str(excinfo.value) + assert excinfo.value.__cause__ is not None + 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, "upper_text:\n text: hello\n") @@ -640,6 +735,27 @@ def test_list_depends_on_field_routing(self, tmp_path: Path) -> None: assert result.status == JobStatus.COMPLETED assert result.result.text == "OLLEH" # type: ignore[union-attr] + def test_list_depends_on_with_input_values_raises(self, tmp_path: Path) -> None: + """Input values for a field-routed task are rejected, not silently dropped.""" + wf_path = _write_workflow_yaml( + tmp_path, + f"""\ +workflow: + name: list_field_ref + tasks: + - task: {THIS_MODULE}.WrapText + - task: {THIS_MODULE}.ReverseText + depends_on: + - {THIS_MODULE}.WrapText + - inner +""", + ) + in_path = _write_input_yaml( + tmp_path, "wrap_text:\n text: hello\nreverse_text:\n text: ignored\n" + ) + with pytest.raises(ConfigLoadError, match="cannot also declare config_fields"): + load_workflow_from_yaml(wf_path, in_path) + def test_dict_fan_in_with_list_field_ref(self, tmp_path: Path) -> None: """Dict depends_on with list-form field refs (fan-in + field routing).""" wf_path = _write_workflow_yaml( @@ -1688,6 +1804,20 @@ def test_mutually_referencing_workflows_are_rejected(self, tmp_path: Path) -> No assert "a.yaml" in str(excinfo.value) assert "b.yaml" in str(excinfo.value) + def test_recursive_reference_chain_is_in_nesting_order(self, tmp_path: Path) -> None: + """The message follows the actual reference path, not alphabetical order.""" + for name, target in (("c", "a"), ("a", "b"), ("b", "c")): + self._write_yaml( + tmp_path / f"{name}.yaml", + f"workflow:\n name: {name}\n tasks:\n - workflow: {target}.yaml\n", + ) + in_path = self._write_yaml(tmp_path / "input.yaml", "{}\n") + with pytest.raises(ConfigLoadError) as excinfo: + load_workflow_from_yaml(tmp_path / "c.yaml", in_path) + root = tmp_path.resolve() + expected = " -> ".join(str(root / f"{n}.yaml") for n in ("c", "a", "b", "c")) + assert str(excinfo.value) == f"Recursive workflow reference: {expected}" + 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(