Skip to content
Merged
26 changes: 2 additions & 24 deletions src/_pytask/execute.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,6 @@

from __future__ import annotations

import inspect
import sys
import time
from typing import TYPE_CHECKING
Expand All @@ -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
Expand All @@ -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

Expand Down Expand Up @@ -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)
Expand Down
40 changes: 40 additions & 0 deletions src/_pytask/execute_utils.py
Original file line number Diff line number Diff line change
@@ -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)
163 changes: 90 additions & 73 deletions src/_pytask/provisional.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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


Expand All @@ -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 "<unknown>"
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
Expand Down
Loading
Loading