diff --git a/src/_pytask/execute.py b/src/_pytask/execute.py index 51a25db7..81f530d5 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,8 @@ 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.explain import ChangeReason from _pytask.explain import NodeType from _pytask.explain import ReasonType @@ -49,7 +48,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 @@ -261,33 +259,13 @@ 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) + out = execute_task(task) if "return" in task.produces: structure_out = tree_structure(out) diff --git a/src/_pytask/execute_utils.py b/src/_pytask/execute_utils.py new file mode 100644 index 00000000..9b95b07b --- /dev/null +++ b/src/_pytask/execute_utils.py @@ -0,0 +1,40 @@ +"""Shared task invocation and return-product handling.""" + +from __future__ import annotations + +import inspect +from typing import TYPE_CHECKING +from typing import Any + +from _pytask.exceptions import NodeLoadError +from _pytask.tree_util import tree_map + +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: + 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) diff --git a/src/_pytask/provisional.py b/src/_pytask/provisional.py index e72ccc05..5b61fb9f 100644 --- a/src/_pytask/provisional.py +++ b/src/_pytask/provisional.py @@ -2,26 +2,23 @@ from __future__ import annotations -import inspect -import sys from typing import TYPE_CHECKING from typing import Any from _pytask.config import hookimpl -from _pytask.exceptions import NodeLoadError -from _pytask.node_protocols import PNode -from _pytask.node_protocols import PProvisionalNode +from _pytask.dag import create_dag_from_session +from _pytask.exceptions import CollectionError +from _pytask.execute_utils import execute_task 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.reports import ExecutionReport 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 @@ -30,6 +27,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 @@ -48,78 +47,96 @@ 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 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 + + execute_task(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 = 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) + # Generators must run to discover tasks even in simulation modes. + if session.config["dry_run"] or session.config["explain"]: + raise WouldBeExecuted + 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) + validate_unique_task_signatures(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 - ) - validate_unique_task_signatures(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) - session.should_stop = True - return None - - recreate_dag(session, task) - return True - - return None + except BaseException: + # 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 + + 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 15c6d75b..cc26fb91 100644 --- a/tests/test_provisional.py +++ b/tests/test_provisional.py @@ -5,6 +5,7 @@ import pytest from _pytask.pluginmanager import get_plugin_manager +from pytask import CollectionOutcome from pytask import ExitCode from pytask import TaskOutcome from pytask import build @@ -197,6 +198,32 @@ 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(): + counter = Path(__file__).parent / "counter.txt" + count = int(counter.read_text()) if counter.exists() else 0 + counter.write_text(str(count + 1)) + + @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_task_generator_return_annotation_is_rejected(runner, tmp_path): source = """ from pathlib import Path @@ -248,6 +275,54 @@ def task_child(produces=Path("child.txt")): assert "return" not in generator.produces +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(): + counter = Path(__file__).parent / "counter.txt" + count = int(counter.read_text()) if counter.exists() else 0 + counter.write_text(str(count + 1)) + + @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_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 tmp_path.joinpath("unrelated.txt").read_text() == "unrelated" + assert len(session.tasks) == 2 + assert [report.outcome for report in session.execution_reports].count( + TaskOutcome.FAIL + ) == 1 + 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_task_generator_cannot_define_products_with_argument(runner, tmp_path): source = """ from pathlib import Path @@ -400,5 +475,47 @@ 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.dag.nodes) == 2 + + +def test_generator_failed_dag_rebuild_restores_execution_state(tmp_path, monkeypatch): + source = """ + from pathlib import Path + from pytask import task + + def task_existing(produces=Path(__file__).with_name("same.txt")): + produces.write_text("existing") + + @task(is_generator=True) + def task_generator(): + @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) == 2 + assert any( + r.exc_info and "same output" in str(r.exc_info[1]) + for r in session.execution_reports + )