Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
11 changes: 6 additions & 5 deletions src/_pytask/execute.py
Original file line number Diff line number Diff line change
Expand Up @@ -47,8 +47,9 @@
from _pytask.state import get_node_change_info
from _pytask.state import has_node_changed
from _pytask.state import update_states
from _pytask.task_io import concrete_product_nodes
from _pytask.task_io import task_node_leaves
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

Expand Down Expand Up @@ -285,7 +286,7 @@ def pytask_execute_task(session: Session, task: PTask) -> bool:
)
raise ValueError(msg)

nodes = tree_leaves(task.produces["return"])
nodes = task_node_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):
Expand All @@ -301,9 +302,9 @@ def pytask_execute_task_teardown(session: Session, task: PTask) -> None:
return

collect_provisional_products(session, task)
missing_nodes: list[Any] = [
node for node in tree_leaves(task.produces) if not node.state()
]
# Collecting provisional products must leave only concrete nodes.
product_nodes = concrete_product_nodes(task.produces)
missing_nodes = [node for node in product_nodes if not node.state()]
if missing_nodes:
paths = session.config["paths"]
files = [format_node_name(i, paths).plain for i in missing_nodes]
Expand Down
50 changes: 50 additions & 0 deletions src/_pytask/task_io.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,50 @@
"""Typed operations on the recursive trees used for task inputs and outputs."""

from __future__ import annotations

from typing import TYPE_CHECKING
from typing import Any
from typing import TypeVar
from typing import cast

from _pytask.node_protocols import NodeTree
from _pytask.node_protocols import PNode
from _pytask.node_protocols import PProvisionalNode
from _pytask.node_protocols import TaskIO
from _pytask.node_protocols import TaskNode
from _pytask.tree_util import PyTree
from _pytask.tree_util import _pytree

if TYPE_CHECKING:
from collections.abc import Callable

_T = TypeVar("_T")


def task_node_leaves(tree: NodeTree | TaskIO) -> list[TaskNode]:
"""Flatten a task tree while retaining its node leaf type."""
# optree's dynamic PyTree alias cannot express pytask's static recursive alias.
return cast("list[TaskNode]", _pytree.leaves(cast("Any", tree), none_is_leaf=True))


def concrete_product_nodes(products: TaskIO) -> list[PNode]:
"""Return products after provisional products have been collected."""
nodes: list[PNode] = []
for node in task_node_leaves(products):
if isinstance(node, PProvisionalNode) or not isinstance(node, PNode):
msg = f"Uncollected provisional product: {node!r}"
raise TypeError(msg)
nodes.append(node)
return nodes


def map_task_io(
func: Callable[[tuple[Any, ...], TaskNode], _T],
tree: TaskIO,
) -> dict[str, PyTree[_T]]:
"""Map task nodes while preserving the top-level argument mapping."""
# optree preserves the dict root and recursively replaces each leaf.
return cast(
"dict[str, PyTree[_T]]",
_pytree.map_with_path(func, cast("Any", tree), none_is_leaf=True),
)
9 changes: 9 additions & 0 deletions src/pytask/task_io.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,9 @@
"""Typed operations on task input and output trees."""

from __future__ import annotations

from _pytask.task_io import concrete_product_nodes
from _pytask.task_io import map_task_io
from _pytask.task_io import task_node_leaves

__all__ = ["concrete_product_nodes", "map_task_io", "task_node_leaves"]
39 changes: 39 additions & 0 deletions tests/test_task_io.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,39 @@
"""Tests for task-specific operations on recursive node trees."""

from __future__ import annotations

from typing import TYPE_CHECKING

import pytest

from _pytask.nodes import DirectoryNode
from _pytask.nodes import PythonNode
from pytask.task_io import concrete_product_nodes
from pytask.task_io import map_task_io
from pytask.task_io import task_node_leaves

if TYPE_CHECKING:
from _pytask.node_protocols import TaskIO


def test_task_io_preserves_nested_nodes_and_paths() -> None:
first = PythonNode(name="first")
second = PythonNode(name="second")
products: TaskIO = {"return": [first, {"nested": (second,)}]}

assert task_node_leaves(products) == [first, second]
assert task_node_leaves(products["return"]) == [first, second]
assert concrete_product_nodes(products) == [first, second]
assert map_task_io(lambda path, node: (path, node.name), products) == {
"return": [
(("return", 0), "first"),
{"nested": ((("return", 1, "nested", 0), "second"),)},
]
}


def test_concrete_products_reject_uncollected_provisional_node() -> None:
products: TaskIO = {"return": [DirectoryNode()]}

with pytest.raises(TypeError, match="Uncollected provisional product"):
concrete_product_nodes(products)
42 changes: 21 additions & 21 deletions uv.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

Loading