diff --git a/CHANGELOG.md b/CHANGELOG.md index 18953991..52edc207 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -9,6 +9,8 @@ releases are available on [PyPI](https://pypi.org/project/pytask) and - [#995](https://github.com/pytask-dev/pytask/pull/995) documents that return values from task generators are ignored and adds a regression test for this behavior. +- [#996](https://github.com/pytask-dev/pytask/pull/996) raises collection errors when + task generators define products with return annotations or a `produces` argument. - [#977](https://github.com/pytask-dev/pytask/pull/977) allows marker expressions passed to `-m` to select tasks by marker keyword arguments. - [#976](https://github.com/pytask-dev/pytask/pull/976) makes invalid marker and diff --git a/docs/source/how_to_guides/provisional_nodes_and_task_generators.md b/docs/source/how_to_guides/provisional_nodes_and_task_generators.md index f0310a5f..f72a5cda 100644 --- a/docs/source/how_to_guides/provisional_nodes_and_task_generators.md +++ b/docs/source/how_to_guides/provisional_nodes_and_task_generators.md @@ -96,5 +96,6 @@ content to a `.txt` file. The generated tasks need to be decorated with `@task` to be collected. -Task generators are only used to define other tasks. Their return values are ignored -and cannot be used to define products. Define products on the generated tasks instead. +Task generators are only used to define other tasks. Their ordinary return values are +ignored. Task generators cannot define products with a `produces` function argument or +a return annotation. diff --git a/src/_pytask/collect_utils.py b/src/_pytask/collect_utils.py index cd9f4c7b..9965a280 100644 --- a/src/_pytask/collect_utils.py +++ b/src/_pytask/collect_utils.py @@ -156,8 +156,16 @@ def task_example() -> str: Read more about products in the documentation: http://tinyurl.com/pytask-return. """ +_ERROR_TASK_GENERATOR_PRODUCTS = """Task generators cannot define products with a +'produces' function argument. +""" + +_ERROR_TASK_GENERATOR_RETURN_ANNOTATION = """Task generators cannot define products +with a return annotation. +""" + -def parse_products_from_task_function( +def parse_products_from_task_function( # noqa: C901 session: Session, task_path: Path | None, task_name: str, node_path: Path, obj: Any ) -> dict[str, Any]: """Parse products from task function. @@ -171,6 +179,7 @@ def parse_products_from_task_function( """ has_return = False has_task_decorator = False + is_generator = isinstance(obj, TaskFunction) and obj.pytask_meta.is_generator out: dict[str, Any] = {} @@ -182,11 +191,23 @@ def parse_products_from_task_function( parameters_with_product_annot = _find_args_with_product_annotation(obj) parameters_with_node_annot = _find_args_with_node_annotation(obj) + if is_generator and "produces" in parameters: + raise NodeNotCollectedError(_ERROR_TASK_GENERATOR_PRODUCTS) + + if is_generator: + if "return" in parameters_with_node_annot: + raise NodeNotCollectedError(_ERROR_TASK_GENERATOR_RETURN_ANNOTATION) + + parameters_with_product_annot = [ + name for name in parameters_with_product_annot if name != "return" + ] + parameters_with_node_annot.pop("return", None) + # Allow to collect products from 'produces'. if "produces" in parameters and "produces" not in parameters_with_product_annot: parameters_with_product_annot.append("produces") - if "return" in parameters_with_node_annot: + if not is_generator and "return" in parameters_with_node_annot: parameters_with_product_annot.append("return") has_return = True @@ -227,7 +248,7 @@ def parse_products_from_task_function( out[parameter_name] = collected_products task_produces = obj.pytask_meta.produces if isinstance(obj, TaskFunction) else None - if task_produces: + if task_produces and not is_generator: has_task_decorator = True collected_products = _collect_nodes_and_provisional_nodes( _collect_product, diff --git a/tests/test_provisional.py b/tests/test_provisional.py index f181e548..15c6d75b 100644 --- a/tests/test_provisional.py +++ b/tests/test_provisional.py @@ -197,7 +197,7 @@ def task_copy( assert tmp_path.joinpath("b-copy.txt").exists() -def test_task_generator_return_value_is_ignored(runner, tmp_path): +def test_task_generator_return_annotation_is_rejected(runner, tmp_path): source = """ from pathlib import Path from typing import Annotated @@ -217,9 +217,54 @@ def task_child(produces=Path("child.txt")): result = runner.invoke(cli, [tmp_path.as_posix()]) - assert result.exit_code == ExitCode.OK + assert result.exit_code == ExitCode.COLLECTION_FAILED + assert "cannot define products" in result.output + assert "return annotation" in result.output + + +def test_task_generator_decorator_return_product_is_ignored(tmp_path): + source = """ + from pathlib import Path + from pytask import task + + @task(is_generator=True, produces=Path("returned.txt")) + def task_generator(): + @task + def task_child(produces=Path("child.txt")): + produces.write_text("child") + + return "ignored" + """ + tmp_path.joinpath("task_module.py").write_text(textwrap.dedent(source)) + + session = build(paths=tmp_path) + generator = next( + task for task in session.tasks if task.name.endswith("task_generator") + ) + + assert session.exit_code == ExitCode.OK assert tmp_path.joinpath("child.txt").read_text() == "child" assert not tmp_path.joinpath("returned.txt").exists() + assert "return" not in generator.produces + + +def test_task_generator_cannot_define_products_with_argument(runner, tmp_path): + source = """ + from pathlib import Path + from pytask import task + + @task(is_generator=True) + def task_generator(produces=Path("generator.txt")): + @task + def task_child(): + pass + """ + tmp_path.joinpath("task_module.py").write_text(textwrap.dedent(source)) + + result = runner.invoke(cli, [tmp_path.as_posix()]) + + assert result.exit_code == ExitCode.COLLECTION_FAILED + assert "cannot define products" in result.output def test_gracefully_fail_when_task_generator_raises_error(runner, tmp_path):