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
2 changes: 1 addition & 1 deletion docs/source/conf.py
Original file line number Diff line number Diff line change
Expand Up @@ -34,7 +34,7 @@
# The version, including alpha/beta/rc tags, but not commit hash and datestamps
release = version("pytask_parallel")
# The short X.Y version.
version = ".".join(release.split(".")[:2]) # ty: ignore[invalid-assignment]
version = ".".join(release.split(".")[:2])

# -- General configuration -------------------------------------------------------------

Expand Down
2 changes: 1 addition & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -50,7 +50,7 @@ test = [
]
typing = [
"pytask-parallel",
"ty>=0.0.8,<0.0.81",
"ty>=0.0.8,<0.0.85",
{include-group = "coiled"},
{include-group = "dask"},
]
Expand Down
46 changes: 30 additions & 16 deletions src/pytask_parallel/wrappers.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,7 @@
from attrs import define
from pytask import PNode
from pytask import PPathNode
from pytask import PProvisionalNode
from pytask import PTask
from pytask import PythonNode
from pytask import Traceback
Expand Down Expand Up @@ -255,7 +256,7 @@ def _handle_function_products(
raise ValueError(msg)

def _save_and_carry_over_product(
path: tuple[Any, ...], node: PNode
path: tuple[Any, ...], node: PNode | PProvisionalNode
) -> CarryOverPath | PythonNode | None:
argument = path[0]

Expand All @@ -275,25 +276,38 @@ def _save_and_carry_over_product(
for p in path[1:]:
value = value[p]

# If the node is a PythonNode, we need to carry it over to the main process.
if isinstance(node, PythonNode):
node.save(value=value)
return node
return _save_return_product(node, value, remote=remote)

# If the path is local and we are remote, we need to carry over the value to
# the main process as a PythonNode and save it later.
if isinstance(node, PPathNode) and is_local_path(node.path) and remote:
return PythonNode(value=value)
return cast(
"PyTree[CarryOverPath | PythonNode | None]",
tree_map_with_path(
_save_and_carry_over_product,
cast("Any", task.produces),
),
)

# If no condition applies, we save the value and do not carry it over. Like a
# remote path to S3.
node.save(value)

def _save_return_product(
node: PNode | PProvisionalNode,
value: Any, # noqa: ANN401
*,
remote: bool,
) -> CarryOverPath | PythonNode | None:
"""Save a concrete return product or defer a provisional one to collection."""
if isinstance(node, PProvisionalNode):
return None

return tree_map_with_path(
_save_and_carry_over_product,
cast("Any", task.produces),
)
# Python nodes must be carried back to the main process.
if isinstance(node, PythonNode):
node.save(value=value)
return node

# Local paths on remote workers are saved by the main process.
if isinstance(node, PPathNode) and is_local_path(node.path) and remote:
return PythonNode(value=value)

node.save(value)
return None


def _write_local_files_to_remote(
Expand Down
13 changes: 13 additions & 0 deletions tests/test_execute.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,12 +6,15 @@
from time import time

import pytest
from pytask import DirectoryNode
from pytask import ExitCode
from pytask import TaskWithoutPath
from pytask import build
from pytask import cli

from pytask_parallel import ParallelBackend
from pytask_parallel.execute import _Sleeper
from pytask_parallel.wrappers import _handle_function_products
from tests.conftest import restore_sys_path_and_module_after_test_execution
from tests.conftest import skip_if_deadlock

Expand All @@ -28,6 +31,16 @@
]


def test_provisional_return_product_is_not_saved() -> None:
task = TaskWithoutPath(
name="provisional",
function=lambda: None,
produces={"return": DirectoryNode()},
)

assert _handle_function_products(task, None) == {"return": None}


@pytest.mark.parametrize("parallel_backend", _IMPLEMENTED_BACKENDS)
def test_parallel_execution(tmp_path, parallel_backend):
source = """
Expand Down
Loading