From eaa9051968721e0d6970c8c85ad4c174a96a39c1 Mon Sep 17 00:00:00 2001 From: Tobias Raabe Date: Sun, 2 Aug 2026 18:22:39 +0200 Subject: [PATCH 1/6] Fix task generator execution --- src/_pytask/provisional.py | 146 +++++++++++++++++++++++-------------- tests/test_provisional.py | 87 ++++++++++++++++++++++ 2 files changed, 178 insertions(+), 55 deletions(-) diff --git a/src/_pytask/provisional.py b/src/_pytask/provisional.py index 473b8637..29c7d0b8 100644 --- a/src/_pytask/provisional.py +++ b/src/_pytask/provisional.py @@ -3,11 +3,12 @@ from __future__ import annotations import inspect -import sys from typing import TYPE_CHECKING from typing import Any from _pytask.config import hookimpl +from _pytask.dag import create_dag_from_session +from _pytask.exceptions import CollectionError from _pytask.exceptions import NodeLoadError from _pytask.node_protocols import PNode from _pytask.node_protocols import PProvisionalNode @@ -17,7 +18,6 @@ from _pytask.provisional_utils import TASKS_WITH_PROVISIONAL_NODES from _pytask.provisional_utils import collect_provisional_nodes from _pytask.provisional_utils import recreate_dag -from _pytask.reports import ExecutionReport from _pytask.task_utils import COLLECTED_TASKS from _pytask.task_utils import parse_collected_tasks_with_task_marker from _pytask.tree_util import tree_map @@ -29,6 +29,8 @@ from collections.abc import Callable from collections.abc import Mapping + from _pytask.reports import CollectionReport + from _pytask.reports import ExecutionReport from _pytask.session import Session @@ -56,63 +58,97 @@ def _safe_load(node: PNode | PProvisionalNode, task: PTask, is_product: bool) -> @hookimpl -def pytask_execute_task(session: Session, task: PTask) -> None: +def pytask_execute_task(session: Session, task: PTask) -> bool | None: """Execute task generators and collect the tasks.""" - if is_task_generator(task): - kwargs = {} - for name, value in task.depends_on.items(): - kwargs[name] = tree_map(lambda x: _safe_load(x, task, False), value) - - parameters = inspect.signature(task.function).parameters - for name, value in task.produces.items(): - if name in parameters: - kwargs[name] = tree_map(lambda x: _safe_load(x, task, True), value) - - task.execute(**kwargs) - - # Parse tasks created with @task. - name_to_function: Mapping[str, Callable[..., Any] | PTask] - if isinstance(task, PTaskWithPath) and task.path in COLLECTED_TASKS: - tasks = COLLECTED_TASKS.pop(task.path) - name_to_function = parse_collected_tasks_with_task_marker(tasks) - elif None in COLLECTED_TASKS: - tasks = COLLECTED_TASKS.pop(None) - name_to_function = parse_collected_tasks_with_task_marker(tasks) - else: - msg = "The task generator {task.name!r} did not create any tasks." - raise RuntimeError(msg) - - new_reports = [] - for name, function in name_to_function.items(): - report = session.hook.pytask_collect_task_protocol( - session=session, - reports=session.collection_reports, - path=task.path if isinstance(task, PTaskWithPath) else None, - name=name, - obj=function, - ) + if not is_task_generator(task): + return None + + kwargs = {} + for name, value in task.depends_on.items(): + kwargs[name] = tree_map(lambda x: _safe_load(x, task, False), value) + + parameters = inspect.signature(task.function).parameters + for name, value in task.produces.items(): + if name in parameters: + kwargs[name] = tree_map(lambda x: _safe_load(x, task, True), value) + + task.execute(**kwargs) + + name_to_function: Mapping[str, Callable[..., Any] | PTask] + if isinstance(task, PTaskWithPath) and task.path in COLLECTED_TASKS: + tasks = COLLECTED_TASKS.pop(task.path) + name_to_function = parse_collected_tasks_with_task_marker(tasks) + elif None in COLLECTED_TASKS: + tasks = COLLECTED_TASKS.pop(None) + name_to_function = parse_collected_tasks_with_task_marker(tasks) + else: + msg = f"The task generator {task.name!r} did not create any tasks." + raise RuntimeError(msg) + + new_reports: list[CollectionReport] = [] + for name, function in name_to_function.items(): + report = session.hook.pytask_collect_task_protocol( + session=session, + reports=new_reports, + path=task.path if isinstance(task, PTaskWithPath) else None, + name=name, + obj=function, + ) + if report is not None: new_reports.append(report) - session.tasks.extend( - i.node - for i in new_reports - if i.outcome == CollectionOutcome.SUCCESS and isinstance(i.node, PTask) + session.collection_reports.extend(new_reports) + failed_reports = [ + report for report in new_reports if report.outcome == CollectionOutcome.FAIL + ] + if failed_reports: + _raise_error_for_failed_generated_tasks(task, failed_reports) + + generated_tasks = [ + report.node + for report in new_reports + if report.outcome == CollectionOutcome.SUCCESS + and isinstance(report.node, PTask) + ] + _commit_generated_tasks(session, generated_tasks) + return True + + +def _raise_error_for_failed_generated_tasks( + task: PTask, reports: list[CollectionReport] +) -> None: + """Raise one generator error containing every child collection failure.""" + lines = [f"The task generator {task.name!r} created tasks that failed collection:"] + for report in reports: + assert report.exc_info is not None + exception = report.exc_info[1] + node_name = report.node.name if report.node is not None else "" + lines.append(f"- {node_name!r}: {type(exception).__name__}: {exception}") + + error = CollectionError("\n".join(lines)) + first_exception = reports[0].exc_info + assert first_exception is not None + raise error from first_exception[1] + + +def _commit_generated_tasks(session: Session, generated_tasks: list[PTask]) -> None: + """Atomically commit generated tasks and the corresponding execution state.""" + previous_tasks = session.tasks + session.tasks = [*previous_tasks, *generated_tasks] + try: + session.hook.pytask_collect_modify_tasks(session=session, tasks=session.tasks) + dag = create_dag_from_session(session) + scheduler = ( + session.scheduler.rebuild(dag) if session.scheduler is not None else None ) - - try: - session.hook.pytask_collect_modify_tasks( - session=session, tasks=session.tasks - ) - # Append the last collection report after successful modification - if report: - session.collection_reports.append(report) - except Exception: # noqa: BLE001 # pragma: no cover - exec_report = ExecutionReport.from_task_and_exception( - task=task, exc_info=sys.exc_info() - ) - session.execution_reports.append(exec_report) - - recreate_dag(session, task) + except BaseException: + session.tasks = previous_tasks + raise + + previous_tasks[:] = session.tasks + session.tasks = previous_tasks + session.dag = dag + session.scheduler = scheduler @hookimpl diff --git a/tests/test_provisional.py b/tests/test_provisional.py index 203e784f..61bd8517 100644 --- a/tests/test_provisional.py +++ b/tests/test_provisional.py @@ -2,6 +2,7 @@ import textwrap +from pytask import CollectionOutcome from pytask import ExitCode from pytask import TaskOutcome from pytask import build @@ -193,6 +194,92 @@ def task_copy( assert tmp_path.joinpath("b-copy.txt").exists() +def test_task_generator_executes_once(tmp_path): + source = """ + from pathlib import Path + from pytask import task + + @task(is_generator=True) + def task_generator(produces=Path(__file__).parent / "generator.txt"): + counter = Path(__file__).parent / "counter.txt" + count = int(counter.read_text()) if counter.exists() else 0 + counter.write_text(str(count + 1)) + produces.write_text("generator") + + @task + def task_generated(produces=Path(__file__).parent / "generated.txt"): + produces.write_text("generated") + """ + tmp_path.joinpath("task_module.py").write_text(textwrap.dedent(source)) + + session = build(paths=tmp_path) + + assert session.exit_code == ExitCode.OK + assert tmp_path.joinpath("counter.txt").read_text() == "1" + assert tmp_path.joinpath("generated.txt").read_text() == "generated" + assert len(session.tasks) == 2 + assert len(session.execution_reports) == 2 + + +def test_failed_generated_task_collection_is_atomic(tmp_path): + source = """ + from pathlib import Path + from typing import Annotated + from pytask import task + + @task(is_generator=True) + def task_generator(produces=Path(__file__).parent / "generator.txt"): + counter = Path(__file__).parent / "counter.txt" + count = int(counter.read_text()) if counter.exists() else 0 + counter.write_text(str(count + 1)) + produces.write_text("generator") + + @task + def task_valid(produces=Path(__file__).parent / "valid.txt"): + produces.write_text("valid") + + @task + def task_invalid() -> Annotated[int, 1]: + return 1 + + def task_descendant( + path=Path(__file__).parent / "generator.txt", + produces=Path(__file__).parent / "descendant.txt", + ): + produces.write_text("descendant") + + def task_unrelated(produces=Path(__file__).parent / "unrelated.txt"): + produces.write_text("unrelated") + """ + tmp_path.joinpath("task_module.py").write_text(textwrap.dedent(source)) + + session = build(paths=tmp_path) + + assert session.exit_code == ExitCode.FAILED + assert tmp_path.joinpath("counter.txt").read_text() == "1" + assert not tmp_path.joinpath("valid.txt").exists() + assert not tmp_path.joinpath("descendant.txt").exists() + assert tmp_path.joinpath("unrelated.txt").read_text() == "unrelated" + assert len(session.tasks) == 3 + assert [report.outcome for report in session.execution_reports].count( + TaskOutcome.FAIL + ) == 1 + assert TaskOutcome.SKIP_PREVIOUS_FAILED in { + report.outcome for report in session.execution_reports + } + generated_reports = [ + report + for report in session.collection_reports + if report.node is not None + and report.node.name.endswith(("task_valid", "task_invalid")) + ] + assert len(generated_reports) == 2 + assert {report.outcome for report in generated_reports} == { + CollectionOutcome.SUCCESS, + CollectionOutcome.FAIL, + } + + def test_gracefully_fail_when_task_generator_raises_error(runner, tmp_path): source = """ from typing import Annotated From f02f0dc14f7e6c935cc83727a40ad2fe77cf1485 Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Sun, 2 Aug 2026 16:24:02 +0000 Subject: [PATCH 2/6] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- tests/test_mark.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/tests/test_mark.py b/tests/test_mark.py index 1fb7713e..62218508 100644 --- a/tests/test_mark.py +++ b/tests/test_mark.py @@ -191,8 +191,10 @@ def task_func(arg=arg): [ ( "foo or", - ("at column 7: expected not OR left parenthesis OR identifier; got end of " - "input"), + ( + "at column 7: expected not OR left parenthesis OR identifier; got end of " + "input" + ), ), ( "foo or or", From 0da432148639eae750ec29957cfc71d2ee3d7bde Mon Sep 17 00:00:00 2001 From: Tobias Raabe Date: Mon, 7 Sep 2026 01:03:22 +0200 Subject: [PATCH 3/6] fix --- src/_pytask/execute.py | 49 ++------------- src/_pytask/execute_utils.py | 61 ++++++++++++++++++ src/_pytask/provisional.py | 36 ++++------- tests/test_provisional.py | 117 ++++++++++++++++++++++++++++++++++- 4 files changed, 192 insertions(+), 71 deletions(-) create mode 100644 src/_pytask/execute_utils.py diff --git a/src/_pytask/execute.py b/src/_pytask/execute.py index 51a25db7..f644b03c 100644 --- a/src/_pytask/execute.py +++ b/src/_pytask/execute.py @@ -2,7 +2,6 @@ from __future__ import annotations -import inspect import sys import time from typing import TYPE_CHECKING @@ -21,8 +20,9 @@ from _pytask.dag_utils import descending_tasks from _pytask.dag_utils import node_and_neighbors from _pytask.exceptions import ExecutionError -from _pytask.exceptions import NodeLoadError from _pytask.exceptions import NodeNotFoundError +from _pytask.execute_utils import execute_task +from _pytask.execute_utils import save_return_products from _pytask.explain import ChangeReason from _pytask.explain import NodeType from _pytask.explain import ReasonType @@ -49,8 +49,6 @@ from _pytask.state import update_states from _pytask.traceback import remove_traceback_from_exc_info from _pytask.tree_util import tree_leaves -from _pytask.tree_util import tree_map -from _pytask.tree_util import tree_structure from _pytask.typing import is_task_generator if TYPE_CHECKING: @@ -261,53 +259,14 @@ def pytask_execute_task_setup(session: Session, task: PTask) -> None: # noqa: C node.root_dir.mkdir(parents=True, exist_ok=True) -def _safe_load(node: PNode | PProvisionalNode, task: PTask, *, is_product: bool) -> Any: - try: - return node.load(is_product=is_product) - except Exception as e: - msg = f"Exception while loading node {node.name!r} of task {task.name!r}" - raise NodeLoadError(msg) from e - - @hookimpl(trylast=True) def pytask_execute_task(session: Session, task: PTask) -> bool: """Execute task.""" if session.config["dry_run"] or session.config["explain"]: raise WouldBeExecuted - parameters = inspect.signature(task.function).parameters - - kwargs = {} - for name, value in task.depends_on.items(): - kwargs[name] = tree_map(lambda x: _safe_load(x, task, is_product=False), value) - - for name, value in task.produces.items(): - if name in parameters: - kwargs[name] = tree_map( - lambda x: _safe_load(x, task, is_product=True), value - ) - - out = task.execute(**kwargs) - - if "return" in task.produces: - structure_out = tree_structure(out) - structure_return = tree_structure(task.produces["return"]) - - # strict must be false when none is leaf. - if not structure_return.is_prefix(structure_out, strict=False): - msg = ( - f"The structure of the return annotation is not a subtree of the " - f"structure of the function return.\n\nFunction return: {structure_out}" - f"\n\nReturn annotation: {structure_return}" - ) - raise ValueError(msg) - - nodes = tree_leaves(task.produces["return"]) - values = structure_return.flatten_up_to(out) - for node, value in zip(nodes, values, strict=False): - if not isinstance(node, PProvisionalNode): - node.save(value) - + out = execute_task(task) + save_return_products(task, out) return True diff --git a/src/_pytask/execute_utils.py b/src/_pytask/execute_utils.py new file mode 100644 index 00000000..be575b82 --- /dev/null +++ b/src/_pytask/execute_utils.py @@ -0,0 +1,61 @@ +"""Shared task invocation and return-product handling.""" + +from __future__ import annotations + +import inspect +from typing import Any + +from _pytask.exceptions import NodeLoadError +from _pytask.node_protocols import PNode +from _pytask.node_protocols import PProvisionalNode +from _pytask.node_protocols import PTask +from _pytask.tree_util import tree_leaves +from _pytask.tree_util import tree_map +from _pytask.tree_util import tree_structure + + +def _safe_load(node: PNode | PProvisionalNode, task: PTask, *, is_product: bool) -> Any: + try: + return node.load(is_product=is_product) + except Exception as e: + msg = f"Exception while loading node {node.name!r} of task {task.name!r}" + raise NodeLoadError(msg) from e + + +def execute_task(task: PTask) -> Any: + """Load task arguments and invoke its body once.""" + parameters = inspect.signature(task.function).parameters + + kwargs = {} + for name, value in task.depends_on.items(): + kwargs[name] = tree_map(lambda x: _safe_load(x, task, is_product=False), value) + + for name, value in task.produces.items(): + if name in parameters: + kwargs[name] = tree_map( + lambda x: _safe_load(x, task, is_product=True), value + ) + + return task.execute(**kwargs) + + +def save_return_products(task: PTask, out: Any) -> None: + """Validate the return structure and save its declared products.""" + if "return" in task.produces: + structure_out = tree_structure(out) + structure_return = tree_structure(task.produces["return"]) + + # strict must be false when none is leaf. + if not structure_return.is_prefix(structure_out, strict=False): + msg = ( + f"The structure of the return annotation is not a subtree of the " + f"structure of the function return.\n\nFunction return: {structure_out}" + f"\n\nReturn annotation: {structure_return}" + ) + raise ValueError(msg) + + nodes = tree_leaves(task.produces["return"]) + values = structure_return.flatten_up_to(out) + for node, value in zip(nodes, values, strict=False): + if not isinstance(node, PProvisionalNode): + node.save(value) diff --git a/src/_pytask/provisional.py b/src/_pytask/provisional.py index 52c79a47..e3a0e1d9 100644 --- a/src/_pytask/provisional.py +++ b/src/_pytask/provisional.py @@ -2,26 +2,24 @@ from __future__ import annotations -import inspect from typing import TYPE_CHECKING from typing import Any from _pytask.config import hookimpl from _pytask.dag import create_dag_from_session from _pytask.exceptions import CollectionError -from _pytask.exceptions import NodeLoadError -from _pytask.node_protocols import PNode -from _pytask.node_protocols import PProvisionalNode +from _pytask.execute_utils import execute_task +from _pytask.execute_utils import save_return_products from _pytask.node_protocols import PTask from _pytask.node_protocols import PTaskWithPath from _pytask.outcomes import CollectionOutcome +from _pytask.outcomes import WouldBeExecuted from _pytask.provisional_utils import TASKS_WITH_PROVISIONAL_NODES from _pytask.provisional_utils import collect_provisional_nodes from _pytask.provisional_utils import recreate_dag from _pytask.task_utils import COLLECTED_TASKS from _pytask.task_utils import parse_collected_tasks_with_task_marker from _pytask.task_utils import validate_unique_task_signatures -from _pytask.tree_util import tree_map from _pytask.tree_util import tree_map_with_path from _pytask.typing import is_task_generator from pytask import TaskOutcome @@ -50,30 +48,13 @@ def pytask_execute_task_setup(session: Session, task: PTask) -> None: recreate_dag(session, task) -def _safe_load(node: PNode | PProvisionalNode, task: PTask, is_product: bool) -> Any: - try: - return node.load(is_product=is_product) - except Exception as e: - msg = f"Exception while loading node {node.name!r} of task {task.name!r}" - raise NodeLoadError(msg) from e - - @hookimpl def pytask_execute_task(session: Session, task: PTask) -> bool | None: """Execute task generators and collect the tasks.""" if not is_task_generator(task): return None - kwargs = {} - for name, value in task.depends_on.items(): - kwargs[name] = tree_map(lambda x: _safe_load(x, task, False), value) - - parameters = inspect.signature(task.function).parameters - for name, value in task.produces.items(): - if name in parameters: - kwargs[name] = tree_map(lambda x: _safe_load(x, task, True), value) - - task.execute(**kwargs) + out = execute_task(task) name_to_function: Mapping[str, Callable[..., Any] | PTask] if isinstance(task, PTaskWithPath) and task.path in COLLECTED_TASKS: @@ -112,6 +93,10 @@ def pytask_execute_task(session: Session, task: PTask) -> bool | None: and isinstance(report.node, PTask) ] _commit_generated_tasks(session, generated_tasks) + # Generators must run to discover tasks even in simulation modes. + if session.config["dry_run"] or session.config["explain"]: + raise WouldBeExecuted + save_return_products(task, out) return True @@ -144,8 +129,9 @@ def _commit_generated_tasks(session: Session, generated_tasks: list[PTask]) -> N session.scheduler.rebuild(dag) if session.scheduler is not None else None ) except BaseException: - # Keep the collected tasks available for diagnostics, but do not replace the - # existing DAG or scheduler after a failed validation or collection hook. + # Rejected children remain available in collection_reports for diagnostics. + # This restores list membership, not arbitrary plugin mutations to tasks. + session.tasks = previous_tasks session.should_stop = True raise diff --git a/tests/test_provisional.py b/tests/test_provisional.py index 823a65a6..e1e8ebca 100644 --- a/tests/test_provisional.py +++ b/tests/test_provisional.py @@ -417,5 +417,120 @@ def pytask_collect_modify_tasks(self, tasks): report.exc_info and "Conflicting task identities" in str(report.exc_info[1]) for report in session.execution_reports ) - assert len(session.tasks) == 3 + assert len(session.tasks) == 2 + assert len(session.collection_reports) == 3 assert len(session.dag.nodes) == 2 + + +@pytest.mark.parametrize("mode", ["build", "dry_run", "explain"]) +def test_generator_return_products(tmp_path, mode): + source = """ + from pathlib import Path + from typing import Annotated + from pytask import task + + @task(is_generator=True) + def task_generator() -> Annotated[str, Path(__file__).with_name("returned.txt")]: + counter = Path(__file__).with_name("counter.txt") + count = int(counter.read_text()) if counter.exists() else 0 + counter.write_text(str(count + 1)) + + @task + def task_child(): + pass + + return "value" + + def task_consumer( + path=Path(__file__).with_name("returned.txt"), + produces=Path(__file__).with_name("copied.txt"), + ): + produces.write_text(path.read_text()) + """ + tmp_path.joinpath("task_module.py").write_text(textwrap.dedent(source)) + + session = build(paths=tmp_path, **({mode: True} if mode != "build" else {})) + + assert session.exit_code == ExitCode.OK + assert tmp_path.joinpath("counter.txt").read_text() == "1" + assert len(session.execution_reports) == 3 + if mode == "build": + assert tmp_path.joinpath("returned.txt").read_text() == "value" + assert tmp_path.joinpath("copied.txt").read_text() == "value" + assert all(r.outcome == TaskOutcome.SUCCESS for r in session.execution_reports) + else: + assert not tmp_path.joinpath("returned.txt").exists() + assert not tmp_path.joinpath("copied.txt").exists() + assert all( + r.outcome == TaskOutcome.WOULD_BE_EXECUTED + for r in session.execution_reports + ) + + +def test_generator_invalid_return_structure(tmp_path): + source = """ + from pathlib import Path + from typing import Annotated + from pytask import task + + @task(is_generator=True) + def task_generator() -> Annotated[ + tuple[str, str], (Path("a.txt"), Path("b.txt")) + ]: + @task + def task_child(): + pass + + return "invalid" + """ + tmp_path.joinpath("task_module.py").write_text(textwrap.dedent(source)) + + session = build(paths=tmp_path) + + assert session.exit_code == ExitCode.FAILED + failures = [r for r in session.execution_reports if r.outcome == TaskOutcome.FAIL] + assert len(failures) == 1 + assert failures[0].exc_info is not None + assert "structure of the return annotation" in str(failures[0].exc_info[1]) + assert not tmp_path.joinpath("a.txt").exists() + assert not tmp_path.joinpath("b.txt").exists() + + +def test_generator_failed_dag_rebuild_restores_execution_state(tmp_path, monkeypatch): + source = """ + from pathlib import Path + from pytask import task + + @task(is_generator=True) + def task_generator(produces=Path(__file__).with_name("same.txt")): + @task + def task_child(produces=Path(__file__).with_name("same.txt")): + produces.write_text("child") + """ + tmp_path.joinpath("task_module.py").write_text(textwrap.dedent(source)) + previous = {} + + class Plugin: + @hookimpl + def pytask_execute_task(self, session): + previous.update( + tasks=session.tasks, dag=session.dag, scheduler=session.scheduler + ) + + pm = get_plugin_manager() + pm.register(Plugin()) + monkeypatch.setattr("_pytask.build.get_plugin_manager", lambda: pm) + + session = build(paths=tmp_path) + + assert session.exit_code == ExitCode.FAILED + assert session.should_stop + assert session.tasks is previous["tasks"] + assert session.dag is previous["dag"] + assert session.scheduler is previous["scheduler"] + assert len(session.tasks) == 1 + assert len(session.collection_reports) == 2 + assert any( + r.exc_info and "same output" in str(r.exc_info[1]) + for r in session.execution_reports + ) From 2ddf3f434b01ae4c0edd14f4102e5dade5970f58 Mon Sep 17 00:00:00 2001 From: Tobias Raabe Date: Sat, 19 Sep 2026 16:14:59 +0200 Subject: [PATCH 4/6] Ignore task generator return products --- src/_pytask/collect_utils.py | 11 +++++- src/_pytask/provisional.py | 4 +- tests/test_provisional.py | 73 ++++++++++++++++++++++++------------ 3 files changed, 59 insertions(+), 29 deletions(-) diff --git a/src/_pytask/collect_utils.py b/src/_pytask/collect_utils.py index cd9f4c7b..02b04d90 100644 --- a/src/_pytask/collect_utils.py +++ b/src/_pytask/collect_utils.py @@ -171,6 +171,7 @@ def parse_products_from_task_function( """ has_return = False has_task_decorator = False + is_generator = isinstance(obj, TaskFunction) and obj.pytask_meta.is_generator out: dict[str, Any] = {} @@ -182,11 +183,17 @@ def parse_products_from_task_function( parameters_with_product_annot = _find_args_with_product_annotation(obj) parameters_with_node_annot = _find_args_with_node_annotation(obj) + if is_generator: + parameters_with_product_annot = [ + name for name in parameters_with_product_annot if name != "return" + ] + parameters_with_node_annot.pop("return", None) + # Allow to collect products from 'produces'. if "produces" in parameters and "produces" not in parameters_with_product_annot: parameters_with_product_annot.append("produces") - if "return" in parameters_with_node_annot: + if not is_generator and "return" in parameters_with_node_annot: parameters_with_product_annot.append("return") has_return = True @@ -227,7 +234,7 @@ def parse_products_from_task_function( out[parameter_name] = collected_products task_produces = obj.pytask_meta.produces if isinstance(obj, TaskFunction) else None - if task_produces: + if task_produces and not is_generator: has_task_decorator = True collected_products = _collect_nodes_and_provisional_nodes( _collect_product, diff --git a/src/_pytask/provisional.py b/src/_pytask/provisional.py index e3a0e1d9..5b61fb9f 100644 --- a/src/_pytask/provisional.py +++ b/src/_pytask/provisional.py @@ -9,7 +9,6 @@ from _pytask.dag import create_dag_from_session from _pytask.exceptions import CollectionError from _pytask.execute_utils import execute_task -from _pytask.execute_utils import save_return_products from _pytask.node_protocols import PTask from _pytask.node_protocols import PTaskWithPath from _pytask.outcomes import CollectionOutcome @@ -54,7 +53,7 @@ def pytask_execute_task(session: Session, task: PTask) -> bool | None: if not is_task_generator(task): return None - out = execute_task(task) + execute_task(task) name_to_function: Mapping[str, Callable[..., Any] | PTask] if isinstance(task, PTaskWithPath) and task.path in COLLECTED_TASKS: @@ -96,7 +95,6 @@ def pytask_execute_task(session: Session, task: PTask) -> bool | None: # Generators must run to discover tasks even in simulation modes. if session.config["dry_run"] or session.config["explain"]: raise WouldBeExecuted - save_return_products(task, out) return True diff --git a/tests/test_provisional.py b/tests/test_provisional.py index 288377d8..39180d66 100644 --- a/tests/test_provisional.py +++ b/tests/test_provisional.py @@ -225,7 +225,7 @@ def task_generated(produces=Path(__file__).parent / "generated.txt"): assert len(session.execution_reports) == 2 -def test_task_generator_return_value_is_ignored(runner, tmp_path): +def test_task_generator_return_value_is_ignored(tmp_path): source = """ from pathlib import Path from typing import Annotated @@ -243,11 +243,15 @@ def task_child(produces=Path("child.txt")): """ tmp_path.joinpath("task_module.py").write_text(textwrap.dedent(source)) - result = runner.invoke(cli, [tmp_path.as_posix()]) + session = build(paths=tmp_path) + generator = next( + task for task in session.tasks if task.name.endswith("task_generator") + ) - assert result.exit_code == ExitCode.OK + assert session.exit_code == ExitCode.OK assert tmp_path.joinpath("child.txt").read_text() == "child" assert not tmp_path.joinpath("returned.txt").exists() + assert "return" not in generator.produces def test_failed_generated_task_collection_is_atomic(tmp_path): @@ -448,7 +452,7 @@ def pytask_collect_modify_tasks(self, tasks): @pytest.mark.parametrize("mode", ["build", "dry_run", "explain"]) -def test_generator_return_products(tmp_path, mode): +def test_generator_return_annotation_is_ignored(tmp_path, mode): source = """ from pathlib import Path from typing import Annotated @@ -461,38 +465,62 @@ def task_generator() -> Annotated[str, Path(__file__).with_name("returned.txt")] counter.write_text(str(count + 1)) @task - def task_child(): - pass + def task_child(produces=Path(__file__).with_name("child.txt")): + produces.write_text("child") return "value" - - def task_consumer( - path=Path(__file__).with_name("returned.txt"), - produces=Path(__file__).with_name("copied.txt"), - ): - produces.write_text(path.read_text()) """ tmp_path.joinpath("task_module.py").write_text(textwrap.dedent(source)) session = build(paths=tmp_path, **({mode: True} if mode != "build" else {})) + generator = next( + task for task in session.tasks if task.name.endswith("task_generator") + ) assert session.exit_code == ExitCode.OK assert tmp_path.joinpath("counter.txt").read_text() == "1" - assert len(session.execution_reports) == 3 + assert len(session.execution_reports) == 2 + assert "return" not in generator.produces if mode == "build": - assert tmp_path.joinpath("returned.txt").read_text() == "value" - assert tmp_path.joinpath("copied.txt").read_text() == "value" + assert tmp_path.joinpath("child.txt").read_text() == "child" + assert not tmp_path.joinpath("returned.txt").exists() assert all(r.outcome == TaskOutcome.SUCCESS for r in session.execution_reports) else: assert not tmp_path.joinpath("returned.txt").exists() - assert not tmp_path.joinpath("copied.txt").exists() + assert not tmp_path.joinpath("child.txt").exists() assert all( r.outcome == TaskOutcome.WOULD_BE_EXECUTED for r in session.execution_reports ) -def test_generator_invalid_return_structure(tmp_path): +def test_generator_task_decorator_return_product_is_ignored(tmp_path): + source = """ + from pathlib import Path + from pytask import task + + @task(is_generator=True, produces=Path("returned.txt")) + def task_generator(): + @task + def task_child(produces=Path("child.txt")): + produces.write_text("child") + + return "ignored" + """ + tmp_path.joinpath("task_module.py").write_text(textwrap.dedent(source)) + + session = build(paths=tmp_path) + generator = next( + task for task in session.tasks if task.name.endswith("task_generator") + ) + + assert session.exit_code == ExitCode.OK + assert tmp_path.joinpath("child.txt").read_text() == "child" + assert not tmp_path.joinpath("returned.txt").exists() + assert "return" not in generator.produces + + +def test_generator_invalid_return_structure_is_ignored(tmp_path): source = """ from pathlib import Path from typing import Annotated @@ -503,8 +531,8 @@ def task_generator() -> Annotated[ tuple[str, str], (Path("a.txt"), Path("b.txt")) ]: @task - def task_child(): - pass + def task_child(produces=Path("child.txt")): + produces.write_text("child") return "invalid" """ @@ -512,11 +540,8 @@ def task_child(): session = build(paths=tmp_path) - assert session.exit_code == ExitCode.FAILED - failures = [r for r in session.execution_reports if r.outcome == TaskOutcome.FAIL] - assert len(failures) == 1 - assert failures[0].exc_info is not None - assert "structure of the return annotation" in str(failures[0].exc_info[1]) + assert session.exit_code == ExitCode.OK + assert tmp_path.joinpath("child.txt").read_text() == "child" assert not tmp_path.joinpath("a.txt").exists() assert not tmp_path.joinpath("b.txt").exists() From 8835e51919153b16932a0475948bcb468a97ee34 Mon Sep 17 00:00:00 2001 From: Tobias Raabe Date: Sun, 20 Sep 2026 13:46:20 +0200 Subject: [PATCH 5/6] test: remove duplicate generator test --- tests/test_provisional.py | 26 -------------------------- 1 file changed, 26 deletions(-) diff --git a/tests/test_provisional.py b/tests/test_provisional.py index f9ab9a94..4d62b778 100644 --- a/tests/test_provisional.py +++ b/tests/test_provisional.py @@ -480,32 +480,6 @@ def pytask_collect_modify_tasks(self, tasks): assert len(session.dag.nodes) == 2 -def test_generator_task_decorator_return_product_is_ignored(tmp_path): - source = """ - from pathlib import Path - from pytask import task - - @task(is_generator=True, produces=Path("returned.txt")) - def task_generator(): - @task - def task_child(produces=Path("child.txt")): - produces.write_text("child") - - return "ignored" - """ - tmp_path.joinpath("task_module.py").write_text(textwrap.dedent(source)) - - session = build(paths=tmp_path) - generator = next( - task for task in session.tasks if task.name.endswith("task_generator") - ) - - assert session.exit_code == ExitCode.OK - assert tmp_path.joinpath("child.txt").read_text() == "child" - assert not tmp_path.joinpath("returned.txt").exists() - assert "return" not in generator.produces - - def test_generator_failed_dag_rebuild_restores_execution_state(tmp_path, monkeypatch): source = """ from pathlib import Path From c576918de9d47243ba364cce2ae4d331c5f24acd Mon Sep 17 00:00:00 2001 From: Tobias Raabe Date: Sun, 20 Sep 2026 14:42:18 +0200 Subject: [PATCH 6/6] refactor: keep return product saving in execute hook --- src/_pytask/execute.py | 23 +++++++++++++++++++++-- src/_pytask/execute_utils.py | 33 ++++++--------------------------- tests/test_provisional.py | 2 -- 3 files changed, 27 insertions(+), 31 deletions(-) diff --git a/src/_pytask/execute.py b/src/_pytask/execute.py index f644b03c..81f530d5 100644 --- a/src/_pytask/execute.py +++ b/src/_pytask/execute.py @@ -22,7 +22,6 @@ from _pytask.exceptions import ExecutionError from _pytask.exceptions import NodeNotFoundError from _pytask.execute_utils import execute_task -from _pytask.execute_utils import save_return_products from _pytask.explain import ChangeReason from _pytask.explain import NodeType from _pytask.explain import ReasonType @@ -49,6 +48,7 @@ from _pytask.state import update_states from _pytask.traceback import remove_traceback_from_exc_info from _pytask.tree_util import tree_leaves +from _pytask.tree_util import tree_structure from _pytask.typing import is_task_generator if TYPE_CHECKING: @@ -266,7 +266,26 @@ def pytask_execute_task(session: Session, task: PTask) -> bool: raise WouldBeExecuted out = execute_task(task) - save_return_products(task, out) + + if "return" in task.produces: + structure_out = tree_structure(out) + structure_return = tree_structure(task.produces["return"]) + + # strict must be false when none is leaf. + if not structure_return.is_prefix(structure_out, strict=False): + msg = ( + f"The structure of the return annotation is not a subtree of the " + f"structure of the function return.\n\nFunction return: {structure_out}" + f"\n\nReturn annotation: {structure_return}" + ) + raise ValueError(msg) + + nodes = tree_leaves(task.produces["return"]) + values = structure_return.flatten_up_to(out) + for node, value in zip(nodes, values, strict=False): + if not isinstance(node, PProvisionalNode): + node.save(value) + return True diff --git a/src/_pytask/execute_utils.py b/src/_pytask/execute_utils.py index be575b82..9b95b07b 100644 --- a/src/_pytask/execute_utils.py +++ b/src/_pytask/execute_utils.py @@ -3,15 +3,16 @@ from __future__ import annotations import inspect +from typing import TYPE_CHECKING from typing import Any from _pytask.exceptions import NodeLoadError -from _pytask.node_protocols import PNode -from _pytask.node_protocols import PProvisionalNode -from _pytask.node_protocols import PTask -from _pytask.tree_util import tree_leaves from _pytask.tree_util import tree_map -from _pytask.tree_util import tree_structure + +if TYPE_CHECKING: + from _pytask.node_protocols import PNode + from _pytask.node_protocols import PProvisionalNode + from _pytask.node_protocols import PTask def _safe_load(node: PNode | PProvisionalNode, task: PTask, *, is_product: bool) -> Any: @@ -37,25 +38,3 @@ def execute_task(task: PTask) -> Any: ) return task.execute(**kwargs) - - -def save_return_products(task: PTask, out: Any) -> None: - """Validate the return structure and save its declared products.""" - if "return" in task.produces: - structure_out = tree_structure(out) - structure_return = tree_structure(task.produces["return"]) - - # strict must be false when none is leaf. - if not structure_return.is_prefix(structure_out, strict=False): - msg = ( - f"The structure of the return annotation is not a subtree of the " - f"structure of the function return.\n\nFunction return: {structure_out}" - f"\n\nReturn annotation: {structure_return}" - ) - raise ValueError(msg) - - nodes = tree_leaves(task.produces["return"]) - values = structure_return.flatten_up_to(out) - for node, value in zip(nodes, values, strict=False): - if not isinstance(node, PProvisionalNode): - node.save(value) diff --git a/tests/test_provisional.py b/tests/test_provisional.py index 4d62b778..cc26fb91 100644 --- a/tests/test_provisional.py +++ b/tests/test_provisional.py @@ -476,7 +476,6 @@ def pytask_collect_modify_tasks(self, tasks): for report in session.execution_reports ) assert len(session.tasks) == 2 - assert len(session.collection_reports) == 3 assert len(session.dag.nodes) == 2 @@ -516,7 +515,6 @@ def pytask_execute_task(self, session): assert session.dag is previous["dag"] assert session.scheduler is previous["scheduler"] assert len(session.tasks) == 2 - assert len(session.collection_reports) == 3 assert any( r.exc_info and "same output" in str(r.exc_info[1]) for r in session.execution_reports