diff --git a/docs/source/conf.py b/docs/source/conf.py index a960390..128ccd0 100644 --- a/docs/source/conf.py +++ b/docs/source/conf.py @@ -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 ------------------------------------------------------------- diff --git a/pyproject.toml b/pyproject.toml index af84431..3394f02 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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"}, ] diff --git a/src/pytask_parallel/wrappers.py b/src/pytask_parallel/wrappers.py index ace745a..3d5446d 100644 --- a/src/pytask_parallel/wrappers.py +++ b/src/pytask_parallel/wrappers.py @@ -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 @@ -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] @@ -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( diff --git a/tests/test_execute.py b/tests/test_execute.py index f8b919e..4e6a316 100644 --- a/tests/test_execute.py +++ b/tests/test_execute.py @@ -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 @@ -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 = """