diff --git a/CHANGELOG.md b/CHANGELOG.md index e44b448..4ede98b 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -7,6 +7,15 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## [Unreleased] +### Added + +- Task inputs and outputs of a `"method"` task are validated and coerced with its type annotations. +- Task discovery reports the outputs of a `"method"` task. + +### Changed + +- **Breaking**: a `"method"` task has one output per field of its return type, instead of a single + `return_value`, when the return type is a pydantic model, dataclass, named tuple or typed dict. ## [5.1.0rc3] - 2026-09-08 diff --git a/doc/conf.py b/doc/conf.py index 8032b9e..12d2cc6 100644 --- a/doc/conf.py +++ b/doc/conf.py @@ -17,9 +17,11 @@ extensions = [ "sphinx.ext.autodoc", "sphinx.ext.autosummary", + "sphinx.ext.intersphinx", "sphinx.ext.viewcode", "sphinxcontrib.mermaid", "sphinx_autodoc_typehints", + "sphinx_design", "nbsphinx", "nbsphinx_link", "sphinx_copybutton", @@ -33,6 +35,11 @@ # https://github.com/spatialaudio/nbsphinx/issues/678 nbsphinx_requirejs_path = "" +intersphinx_mapping = { + "python": ("https://docs.python.org/3", None), + "pydantic": ("https://docs.pydantic.dev/latest", None), +} + autosummary_generate = True autodoc_default_flags = [ "members", diff --git a/doc/howtoguides.rst b/doc/howtoguides.rst index 991d381..14d9a51 100644 --- a/doc/howtoguides.rst +++ b/doc/howtoguides.rst @@ -9,6 +9,7 @@ How-to Guides howtoguides/workflow_conversion howtoguides/events howtoguides/task_discovery + howtoguides/function_task howtoguides/workflow_discovery howtoguides/notebook_task howtoguides/change_schema diff --git a/doc/howtoguides/function_task.rst b/doc/howtoguides/function_task.rst new file mode 100644 index 0000000..8f66b8d --- /dev/null +++ b/doc/howtoguides/function_task.rst @@ -0,0 +1,323 @@ +Python function as workflow task +================================ + +A Python function can be used as a task node in a workflow by using ``"method"`` as its ``task_type``. +By default, such task has a single output named ``return_value``. +It is possible to `Define multiple outputs`_. + +Use a function as a task +------------------------ + +Example with a ``range_info`` function returning 3 values in a dictionary: + +.. code:: python + + def range_info(a, b): + return { + "extent": abs(b - a), + "minimum": min(a, b), + "maximum": max(a, b), + } + +The corresponding workflow node must be declared with ``"method"`` as ``task_type``: + +.. code:: python + + range_info_node = { + "id": "task_range_info", + "task_type": "method", + "task_identifier": "__main__.range_info", + } + +Code to execute the ``range_info`` function as a task in a workflow with ``a=15`` and ``b=10`` as inputs: + +.. code:: python + + from ewokscore import execute_graph + + # Define a workflow which calls the range_info function as a task + workflow = { + "graph": {"id": "range_info_workflow"}, + "nodes": [ + { + "id": "task_range_info", + "task_type": "method", + "task_identifier": "__main__.range_info", + }, + ], + "links": [], + } + + # Define task inputs + inputs = [ + {"id": "task_range_info", "name": "a", "value": 15}, + {"id": "task_range_info", "name": "b", "value": 10}, + ] + + # Execute the workflow + result = execute_graph(workflow, inputs=inputs) + print(result) + +The task inputs are the arguments of the function, provided by name. An argument without a +default value is a required task input. For arguments that cannot be passed by name, see +`Provide inputs by position`_. + +The task output contains a single ``return_value`` field which is set to the function return value. + +In the ``range_info`` example, the result contains a single output field ``return_value``: + +.. code:: python + + {'return_value': {'extent': 5, 'minimum': 10, 'maximum': 15}} + +The docstring of the function describes the task. See :doc:`task_discovery` to make the functions +of a module discoverable as tasks. + + +Provide inputs by position +-------------------------- + +An argument that cannot be passed by name is provided as a positional task input, named by its +position: + +.. code:: python + + def range_info(a: float, b: float, /): + ... + +.. code:: python + + inputs = [ + {"id": "task_range_info", "name": 0, "value": 15}, + {"id": "task_range_info", "name": 1, "value": 10}, + ] + +Positional task inputs are also how a function receives ``*args``: + +.. code:: python + + def total(*args): + return sum(args) + +An argument that can be passed by name has to be provided by name. Providing it by position as +well makes the function receive two values for it, which fails the task. + +A function with ``**kwargs`` receives every task input that is not one of its arguments. + + +Define multiple outputs +----------------------- + +To declare a function that can be used as a :class:`~ewokscore.task.Task` with multiple output +fields, declare the function with one of the following kind of return type: + +.. tab-set:: + + .. tab-item:: pydantic model + + .. code:: python + + from ewokscore.model import BaseOutputModel + + class Result(BaseOutputModel): + extent: float + minimum: float + maximum: float + + .. tab-item:: dataclass + + .. code:: python + + from dataclasses import dataclass + + @dataclass + class Result: + extent: float + minimum: float + maximum: float + + .. tab-item:: typing.NamedTuple + + .. code:: python + + from typing import NamedTuple + + class Result(NamedTuple): + extent: float + minimum: float + maximum: float + + .. tab-item:: typing.TypedDict + + .. code:: python + + from typing import TypedDict + + class Result(TypedDict): + extent: float + minimum: float + maximum: float + + .. tab-item:: namedtuple + + .. code:: python + + from collections import namedtuple + + Result = namedtuple("Result", ["extent", "minimum", "maximum"]) + +The task output names are then defined by the fields of the return type: + +.. code:: python + + def range_info(a: float, b: float) -> Result: + return Result( + extent=abs(b - a), + minimum=min(a, b), + maximum=max(a, b), + ) + +In the example, the return value of the ``range_info`` function is available through 3 task outputs: + +.. code:: python + + {'maximum': 15, 'extent': 5, 'minimum': 10} + +The function does not have to construct the return type itself. Returning a mapping with the +field names as keys works as well: + +.. code:: python + + def range_info(a: float, b: float) -> Result: + return { + "extent": abs(b - a), + "minimum": min(a, b), + "maximum": max(a, b), + } + +This is also what a :class:`~typing.TypedDict` return type does, as a typed dict is a mapping. + + +Describe inputs and outputs +--------------------------- + +Use :func:`~pydantic.Field` to add a description to a task input or output. + +Describe a task input with :func:`~pydantic.Field` in the annotation of the corresponding +argument, and a task output in the return type: + +.. code:: python + + from typing import Annotated + + from pydantic import Field + from ewokscore.model import BaseOutputModel + + + class Result(BaseOutputModel): + extent: float = Field(..., description="Distance between the two numbers") + minimum: float = Field(..., description="Smallest of the two numbers") + maximum: float = Field(..., description="Largest of the two numbers") + + + def range_info( + a: Annotated[float, Field(description="First number")], + b: Annotated[float, Field(description="Second number", ge=0)] = 0, + ) -> Result: + return Result( + extent=abs(b - a), + minimum=min(a, b), + maximum=max(a, b), + ) + +Anything else :func:`~pydantic.Field` provides applies as well. In the example, ``b`` has a +default value which makes it an optional task input, and ``ge=0`` rejects a negative value for it. + +A task input can also be described with :func:`~pydantic.Field` as the default value of the +argument: + +.. code:: python + + def range_info( + a: float = Field(..., description="First number"), + b: float = Field(0, description="Second number", ge=0), + ) -> Result: + +Validate inputs and outputs +--------------------------- + +The type annotations of the function validate the task inputs and outputs. A task input that +cannot be coerced to the annotated argument type fails the task before the function is called, +and a return value that does not match the return type fails the task after it is called. + +Note that annotations coerce as well as validate: an ``int`` input is passed to an argument +annotated as ``float`` as a ``float``. + +The task inputs are only validated when every argument of the function can become a field of a +pydantic model. A positional-only, ``*args`` or ``**kwargs`` argument cannot, in which case the +task inputs are passed to the function unvalidated: + +.. code:: python + + def range_info(a: float, b: float, **kw) -> Result: # inputs are not validated + +The task outputs are validated in both cases, as the return type does not depend on the +arguments. + +Nothing of this is specific to *ewoks*: the function can still be imported and called directly. +Such a direct call is not validated though, because the validation is done by the task. + +Validate direct calls +--------------------- + +Add :func:`~pydantic.validate_call` to validate the arguments and the return value of a direct call +to the function with the same annotations: + +.. code:: python + + from pydantic import validate_call + + + @validate_call(validate_return=True) + def range_info(a: float, b: float) -> Result: + return Result( + extent=abs(b - a), + minimum=min(a, b), + maximum=max(a, b), + ) + +.. code:: python + + range_info("not a number", 10) # ValidationError + +The decorator does not change the task: the task inputs and outputs are derived from the +annotations of the undecorated function. + +Use its ``config`` argument when an argument is annotated with a type pydantic does not know: + +.. code:: python + + from pydantic import ConfigDict + + + @validate_call(config=ConfigDict(arbitrary_types_allowed=True), validate_return=True) + def spectrum_size(data: numpy.ndarray) -> int: + return data.size + + +Use a function as a task outside a workflow +------------------------------------------- + +The task class of a function is provided by :func:`~ewokscore.methodtask.get_method_task`: + +.. code:: python + + from ewokscore.methodtask import get_method_task + + task = get_method_task("__main__.range_info")(inputs={"a": 15, "b": 10}) + task.execute() + print(task.get_output_values()) + +.. code:: python + + {'extent': 5, 'minimum': 10, 'maximum': 15} diff --git a/pyproject.toml b/pyproject.toml index dd72324..68ff89a 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -55,6 +55,7 @@ doc = [ "sphinx >=4.5", "sphinxcontrib-mermaid >=0.7", "sphinx-autodoc-typehints >=1.16", + "sphinx-design", "nbsphinx", # v0.21 incompatible with nbsphinx_link: https://github.com/vidartf/nbsphinx-link/issues/22 "docutils < 0.21", diff --git a/src/ewokscore/inittask.py b/src/ewokscore/inittask.py index 3c8d193..9e42959 100644 --- a/src/ewokscore/inittask.py +++ b/src/ewokscore/inittask.py @@ -7,7 +7,7 @@ from ewoksutils.import_utils import import_qualname from .dynamictask import get_dynamically_task_class -from .methodtask import MethodExecutorTask +from .methodtask import get_method_task from .node import NodeIdType from .node import get_node_label from .notebooktask import NotebookExecutorTask @@ -155,8 +155,7 @@ def instantiate_task( return Task.instantiate(task_info["task_identifier"], **task_kwargs) if task_type == "method": - task_inputs[MethodExecutorTask.METHOD_ARGUMENT] = task_info["task_identifier"] - return MethodExecutorTask(**task_kwargs) + return get_method_task(task_info["task_identifier"])(**task_kwargs) if task_type == "ppfmethod": task_inputs[PpfMethodExecutorTask.METHOD_ARGUMENT] = task_info[ @@ -254,7 +253,7 @@ def get_task_class(node_id: NodeIdType, node_attrs: dict): if task_type == "class": return Task.get_subclass(task_info["task_identifier"]) if task_type == "method": - return MethodExecutorTask + return get_method_task(task_info["task_identifier"]) if task_type == "ppfmethod": return PpfMethodExecutorTask if task_type == "ppfport": diff --git a/src/ewokscore/methodtask.py b/src/ewokscore/methodtask.py index 239e31c..d547c5b 100644 --- a/src/ewokscore/methodtask.py +++ b/src/ewokscore/methodtask.py @@ -1,16 +1,47 @@ -from typing import Mapping +import dataclasses +import functools +import inspect +import sys +from collections.abc import Mapping +from typing import Any +from typing import Callable +from typing import Dict +from typing import List +from typing import Optional from typing import Set +from typing import Tuple +from typing import Type + +if sys.version_info < (3, 9): + from typing_extensions import get_type_hints +else: + from typing import get_type_hints from ewoksutils.import_utils import import_method +from pydantic import BaseModel +from pydantic import create_model +from .model import BaseInputModel +from .model import BaseOutputModel from .task import Task METHOD_ARGUMENT = "_method" +SINGLE_OUTPUT_NAME = "return_value" + +_UNMODELLABLE_KINDS = frozenset( + { + inspect.Parameter.POSITIONAL_ONLY, + inspect.Parameter.VAR_POSITIONAL, + inspect.Parameter.VAR_KEYWORD, + } +) class MethodExecutorTask( - Task, input_names=[METHOD_ARGUMENT], output_names=["return_value"] + Task, input_names=[METHOD_ARGUMENT], output_names=[SINGLE_OUTPUT_NAME] ): + """Executes the python function provided as the `_method` input.""" + METHOD_ARGUMENT = METHOD_ARGUMENT def _warn_unexpected_inputs(self, unexpected_names: Set[str]) -> None: @@ -19,7 +50,7 @@ def _warn_unexpected_inputs(self, unexpected_names: Set[str]) -> None: def _get_task_identifier(self, inputs: Mapping) -> str: return inputs.get(self.METHOD_ARGUMENT, self.class_registry_name()) - def run(self): + def run(self) -> None: kwargs = self.get_named_input_values() args = self.get_positional_input_values() fullname = kwargs.pop(self.METHOD_ARGUMENT) @@ -27,4 +58,179 @@ def run(self): result = method(*args, **kwargs) - self.outputs.return_value = result + self.outputs[SINGLE_OUTPUT_NAME] = result + + +class MethodTask(Task, register=False): + """Executes a python function. Use `get_method_task` to get the task class + of a specific function. + """ + + _METHOD: Optional[str] = None + + def __init_subclass__(subclass, task_identifier: Optional[str] = None, **kwargs): + if task_identifier is None: + super().__init_subclass__(**kwargs) + return + + method = import_method(task_identifier) + subclass._METHOD = task_identifier + super().__init_subclass__(**_task_arguments(method), **kwargs) + subclass.__doc__ = inspect.getdoc(method) + + def _warn_unexpected_inputs(self, unexpected_names: Set[str]) -> None: + pass + + def _get_task_identifier(self, inputs: Mapping) -> str: + return self._METHOD or self.class_registry_name() + + def run(self) -> None: + method = import_method(self._METHOD) + args = self.get_positional_input_values() + kwargs = self.get_named_input_values() + + result = method(*args, **kwargs) + + if self._OUTPUT_MODEL is None: + self.outputs[SINGLE_OUTPUT_NAME] = result + return + + for name in self._OUTPUT_MODEL.model_fields: + if isinstance(result, Mapping): + self.outputs[name] = result[name] + else: + self.outputs[name] = getattr(result, name) + + +@functools.lru_cache(maxsize=None) +def get_method_task(task_identifier: str) -> Type[MethodTask]: + """Task class that executes the function with the given qualified name.""" + + class GeneratedMethodTask( + MethodTask, task_identifier=task_identifier, register=False + ): + pass + + method_name = task_identifier.rsplit(".", 1)[-1] + name = "".join(part.title() for part in method_name.split("_")) + "Task" + GeneratedMethodTask.__name__ = name + GeneratedMethodTask.__qualname__ = name + return GeneratedMethodTask + + +def _task_arguments(method: Callable) -> Dict[str, Any]: + """Task input and output arguments derived from the signature of `method`.""" + arguments: Dict[str, Any] = dict() + + input_model = _input_model(method) + if input_model is None: + arguments.update(_input_arguments(method)) + else: + arguments["input_model"] = input_model + + output_model = _output_model(method) + if output_model is None: + arguments["output_names"] = [SINGLE_OUTPUT_NAME] + else: + arguments["output_model"] = output_model + + return arguments + + +def _type_hints(obj: Any) -> Dict[str, Any]: + """Type hints of `obj`, empty when an annotation cannot be resolved.""" + try: + return get_type_hints(obj, include_extras=True) + except NameError: + return dict() + + +def _input_model(method: Callable) -> Optional[Type[BaseInputModel]]: + """Input model with a field per argument of `method`, or `None` when the + arguments cannot be expressed as model fields. + """ + parameters = list(inspect.signature(method).parameters.values()) + if any(parameter.kind in _UNMODELLABLE_KINDS for parameter in parameters): + return None + + hints = _type_hints(method) + fields = { + parameter.name: ( + hints.get(parameter.name, Any), + ( + ... + if parameter.default is inspect.Parameter.empty + else parameter.default + ), + ) + for parameter in parameters + } + return create_model( + f"{method.__name__}InputModel", __base__=BaseInputModel, **fields + ) + + +def _output_model(method: Callable) -> Optional[Type[BaseOutputModel]]: + """Output model with a field per task output of `method`, or `None` when the + return type does not define task outputs. + """ + return_type = _type_hints(method).get("return") + if not isinstance(return_type, type): + return None + + if issubclass(return_type, BaseOutputModel): + return return_type + if issubclass(return_type, BaseModel): + return create_model( + f"{method.__name__}OutputModel", __base__=(BaseOutputModel, return_type) + ) + + names = _output_names(return_type) + if not names: + return None + + hints = _type_hints(return_type) + fields = {name: (hints.get(name, Any), ...) for name in names} + return create_model( + f"{method.__name__}OutputModel", __base__=BaseOutputModel, **fields + ) + + +def _output_names(return_type: Type) -> Tuple[str, ...]: + """Names of the task outputs defined by a return type.""" + if dataclasses.is_dataclass(return_type): + return tuple(field.name for field in dataclasses.fields(return_type)) + if issubclass(return_type, tuple) and hasattr(return_type, "_fields"): + return tuple(return_type._fields) + if issubclass(return_type, dict): + # A TypedDict is a dict subclass with annotated keys + return tuple(_type_hints(return_type)) + return tuple() + + +def _input_arguments(method: Callable) -> Dict[str, Any]: + """Task input arguments for the arguments of `method` that cannot be model fields.""" + required: List[str] = list() + optional: List[str] = list() + n_required_positional = 0 + for parameter in inspect.signature(method).parameters.values(): + if parameter.kind is inspect.Parameter.POSITIONAL_ONLY: + # Cannot be a named task input + if parameter.default is inspect.Parameter.empty: + n_required_positional += 1 + continue + if parameter.kind in ( + inspect.Parameter.VAR_POSITIONAL, + inspect.Parameter.VAR_KEYWORD, + ): + continue + if parameter.default is inspect.Parameter.empty: + required.append(parameter.name) + else: + optional.append(parameter.name) + + return { + "input_names": required, + "optional_input_names": optional, + "n_required_positional_inputs": n_required_positional, + } diff --git a/src/ewokscore/task_discovery.py b/src/ewokscore/task_discovery.py index b2d382b..d0b942e 100644 --- a/src/ewokscore/task_discovery.py +++ b/src/ewokscore/task_discovery.py @@ -14,6 +14,7 @@ from ewoksutils.import_utils import qualname from .entry_points import entry_points +from .methodtask import get_method_task from .task import Task @@ -156,9 +157,18 @@ def _iter_method_tasks( if method_name.startswith("_"): continue + task_identifier = qualname(method_qn) + task_class = get_method_task(task_identifier) + output_model = task_class.output_model() yield { "task_type": "method", - **_common_method_task_fields(method_name, method_qn, mod), + **_method_arguments(getattr(mod, method_name)), + "task_identifier": task_identifier, + "output_names": sorted(task_class.output_names()), + "category": task_identifier.split(".")[0], + "description": task_class.__doc__, + "input_model": None, + "output_model": qualname(output_model) if output_model else None, } diff --git a/src/ewokscore/tests/function_task/__init__.py b/src/ewokscore/tests/function_task/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/src/ewokscore/tests/function_task/shared_methods.py b/src/ewokscore/tests/function_task/shared_methods.py new file mode 100644 index 0000000..adb0ac9 --- /dev/null +++ b/src/ewokscore/tests/function_task/shared_methods.py @@ -0,0 +1,22 @@ +from pydantic import BaseModel + +from ...model import BaseOutputModel + + +class RangeInfo(BaseOutputModel): + extent: float + minimum: float + maximum: float + + +def range_info(a: float, b: float) -> RangeInfo: + """Information on the range between two numbers.""" + return RangeInfo(extent=abs(b - a), minimum=min(a, b), maximum=max(a, b)) + + +class PlainModel(BaseModel): + total: int + + +def plain_model_method(a: int, b: int = 1) -> PlainModel: + return PlainModel(total=a + b) diff --git a/src/ewokscore/tests/function_task/test_inputs.py b/src/ewokscore/tests/function_task/test_inputs.py new file mode 100644 index 0000000..3495777 --- /dev/null +++ b/src/ewokscore/tests/function_task/test_inputs.py @@ -0,0 +1,111 @@ +import pytest +from ewoksutils.import_utils import qualname + +from ...methodtask import get_method_task +from .shared_methods import plain_model_method + + +def test_task_inputs_from_signature(): + task_class = get_method_task(qualname(plain_model_method)) + + assert set(task_class.required_input_names()) == {"a"} + assert set(task_class.optional_input_names()) == {"b"} + assert set(task_class.output_names()) == {"total"} + + +def signature_named(a: int, b: int = 1) -> int: + return a + b + + +def signature_var_args(a, *args, b=None, c=3, **kw): + return a + + +def signature_keyword_only(a, *args, b): + return a + + +def signature_var_kwargs(a, b=1, **kw): + return a + + +def signature_positional_only(a, b, /, c, d=1): + return a + + +def signature_positional_only_default(a, b=1, /, c=2): + return a + + +def signature_no_arguments(): + return None + + +@pytest.mark.parametrize( + "method,required,optional,n_positional,has_model", + [ + (signature_named, {"a"}, {"b"}, 0, True), + (signature_var_args, {"a"}, {"b", "c"}, 0, False), + (signature_keyword_only, {"a", "b"}, set(), 0, False), + (signature_var_kwargs, {"a"}, {"b"}, 0, False), + (signature_positional_only, {"c"}, {"d"}, 2, False), + (signature_positional_only_default, set(), {"c"}, 1, False), + (signature_no_arguments, set(), set(), 0, True), + ], + ids=lambda value: getattr(value, "__name__", None), +) +def test_task_inputs_of_signature(method, required, optional, n_positional, has_model): + task_class = get_method_task(qualname(method)) + + assert set(task_class.required_input_names()) == required + assert set(task_class.optional_input_names()) == optional + assert task_class.n_required_positional_inputs() == n_positional + assert (task_class.input_model() is not None) == has_model + + +def positional_only_method(a, b, /, c, d=1): + return a + b + c + d + + +def test_positional_only_inputs(varinfo): + task = get_method_task(qualname(positional_only_method))( + inputs={0: 1, 1: 2, "c": 3}, varinfo=varinfo + ) + task.execute() + assert task.get_output_values() == {"return_value": 7} + + +def test_missing_positional_only_input(varinfo): + task_class = get_method_task(qualname(positional_only_method)) + with pytest.raises(Exception, match="positional argument"): + task_class(inputs={0: 1, "c": 3}, varinfo=varinfo) + + +def variadic_method(*args): + return sum(args) + + +def test_variadic_inputs(varinfo): + task = get_method_task(qualname(variadic_method))( + inputs={0: 3, 1: 5}, varinfo=varinfo + ) + task.execute() + assert task.get_output_values() == {"return_value": 8} + + +def var_kwargs_method(a, b=1, **kw): + return a + b + sum(kw.values()) + + +def test_var_kwargs_inputs(varinfo): + task = get_method_task(qualname(var_kwargs_method))( + inputs={"a": 2, "b": 3, "extra": 4}, varinfo=varinfo + ) + task.execute() + assert task.get_output_values() == {"return_value": 9} + + +def test_missing_required_input(varinfo): + task_class = get_method_task(qualname(var_kwargs_method)) + with pytest.raises(Exception, match=r"Missing inputs.*'a'"): + task_class(inputs={"b": 3}, varinfo=varinfo) diff --git a/src/ewokscore/tests/function_task/test_outputs.py b/src/ewokscore/tests/function_task/test_outputs.py new file mode 100644 index 0000000..d1663c0 --- /dev/null +++ b/src/ewokscore/tests/function_task/test_outputs.py @@ -0,0 +1,139 @@ +from collections import namedtuple +from dataclasses import dataclass +from typing import NamedTuple +from typing import TypedDict + +import pytest +from ewoksutils.import_utils import qualname + +from ...methodtask import get_method_task +from .shared_methods import PlainModel +from .shared_methods import RangeInfo +from .shared_methods import plain_model_method +from .shared_methods import range_info + + +def test_task_outputs_from_return_model(varinfo): + task = get_method_task(qualname(range_info))( + inputs={"a": 15, "b": 10}, varinfo=varinfo + ) + task.execute() + + assert task.get_output_values() == {"extent": 5, "minimum": 10, "maximum": 15} + + +def test_task_outputs_from_plain_model(varinfo): + task = get_method_task(qualname(plain_model_method))( + inputs={"a": 2}, varinfo=varinfo + ) + task.execute() + + assert task.get_output_values() == {"total": 3} + + +def no_return_type(a: int, b: int = 0): + return a + b + + +def test_single_task_output(varinfo): + task_class = get_method_task(qualname(no_return_type)) + assert set(task_class.output_names()) == {"return_value"} + + task = task_class(inputs={"a": 2, "b": 3}, varinfo=varinfo) + task.execute() + + assert task.get_output_values() == {"return_value": 5} + + +class NamedTupleResult(NamedTuple): + total: int + count: int + + +def namedtuple_method(a: int, b: int) -> NamedTupleResult: + return NamedTupleResult(total=a + b, count=2) + + +CollectionsResult = namedtuple("CollectionsResult", ["total", "count"]) + + +def collections_namedtuple_method(a: int, b: int) -> CollectionsResult: + return CollectionsResult(total=a + b, count=2) + + +@dataclass +class DataclassResult: + total: int + count: int + + +def dataclass_method(a: int, b: int) -> DataclassResult: + return DataclassResult(total=a + b, count=2) + + +class TypedDictResult(TypedDict): + total: int + count: int + + +def typeddict_method(a: int, b: int) -> TypedDictResult: + return TypedDictResult(total=a + b, count=2) + + +@pytest.mark.parametrize( + "method", + [ + namedtuple_method, + collections_namedtuple_method, + dataclass_method, + typeddict_method, + ], + ids=lambda method: method.__name__, +) +def test_task_outputs_from_return_type(method, varinfo): + task_class = get_method_task(qualname(method)) + assert set(task_class.output_names()) == {"total", "count"} + + task = task_class(inputs={"a": 2, "b": 3}, varinfo=varinfo) + task.execute() + + assert task.get_output_values() == {"total": 5, "count": 2} + + +def mapping_method(a: float, b: float) -> RangeInfo: + return {"extent": abs(b - a), "minimum": min(a, b), "maximum": max(a, b)} + + +def namedtuple_mapping_method(a: int, b: int) -> "NamedTupleResult": + return {"total": a + b, "count": 2} + + +def collections_mapping_method(a: int, b: int) -> "CollectionsResult": + return {"total": a + b, "count": 2} + + +def dataclass_mapping_method(a: int, b: int) -> "DataclassResult": + return {"total": a + b, "count": 2} + + +def plain_model_mapping_method(a: int, b: int = 1) -> PlainModel: + return {"total": a + b} + + +@pytest.mark.parametrize( + "method,expected", + [ + (mapping_method, {"extent": 5, "minimum": 10, "maximum": 15}), + (plain_model_mapping_method, {"total": 25}), + (namedtuple_mapping_method, {"total": 25, "count": 2}), + (collections_mapping_method, {"total": 25, "count": 2}), + (dataclass_mapping_method, {"total": 25, "count": 2}), + ], + ids=lambda value: getattr(value, "__name__", None), +) +def test_task_outputs_from_mapping(method, expected, varinfo): + """A function may return a mapping instead of an instance of its return type.""" + task = get_method_task(qualname(method))(inputs={"a": 15, "b": 10}, varinfo=varinfo) + task.execute() + + assert task.get_output_values() == expected diff --git a/src/ewokscore/tests/function_task/test_task_class.py b/src/ewokscore/tests/function_task/test_task_class.py new file mode 100644 index 0000000..9dc5e11 --- /dev/null +++ b/src/ewokscore/tests/function_task/test_task_class.py @@ -0,0 +1,17 @@ +from ewoksutils.import_utils import qualname + +from ...methodtask import get_method_task +from .shared_methods import RangeInfo +from .shared_methods import range_info + + +def test_task_class_of_a_function(): + task_class = get_method_task(qualname(range_info)) + + assert task_class is get_method_task(qualname(range_info)) + assert task_class.__name__ == "RangeInfoTask" + assert task_class.__doc__ == "Information on the range between two numbers." + + +def test_function_is_callable(): + assert range_info(15, 10) == RangeInfo(extent=5, minimum=10, maximum=15) diff --git a/src/ewokscore/tests/function_task/test_validation.py b/src/ewokscore/tests/function_task/test_validation.py new file mode 100644 index 0000000..6fd688f --- /dev/null +++ b/src/ewokscore/tests/function_task/test_validation.py @@ -0,0 +1,140 @@ +import sys + +if sys.version_info < (3, 9): + from typing_extensions import Annotated +else: + from typing import Annotated + +import numpy +import pytest +from ewoksutils.import_utils import qualname +from pydantic import ConfigDict +from pydantic import Field +from pydantic import ValidationError +from pydantic import validate_call + +from ...methodtask import get_method_task +from ...model import BaseOutputModel +from .shared_methods import RangeInfo +from .shared_methods import range_info + + +def invalid_output_method(a: float, b: float) -> RangeInfo: + return {"extent": "not a number", "minimum": a, "maximum": b} + + +def test_task_input_validation(varinfo): + task = get_method_task(qualname(range_info))( + inputs={"a": "not a number", "b": 10}, varinfo=varinfo + ) + with pytest.raises(Exception, match="Invalid task inputs"): + task.execute(raise_on_error=True) + + +def test_task_output_validation(varinfo): + task = get_method_task(qualname(invalid_output_method))( + inputs={"a": 15, "b": 10}, varinfo=varinfo + ) + with pytest.raises(Exception, match="Invalid task outputs"): + task.execute(raise_on_error=True) + + +class DescribedResult(BaseOutputModel): + total: int = Field(..., description="Sum of the two numbers") + + +def annotated_method( + a: Annotated[int, Field(description="First number")], + b: Annotated[int, Field(description="Second number", ge=0)] = 0, +) -> DescribedResult: + return DescribedResult(total=a + b) + + +def field_default_method( + a: int = Field(..., description="First number"), + b: int = Field(0, description="Second number", ge=0), +) -> DescribedResult: + return DescribedResult(total=a + b) + + +@pytest.mark.parametrize( + "method", [annotated_method, field_default_method], ids=lambda m: m.__name__ +) +def test_described_task_inputs_and_outputs(method, varinfo): + task_class = get_method_task(qualname(method)) + + assert set(task_class.required_input_names()) == {"a"} + assert set(task_class.optional_input_names()) == {"b"} + + input_fields = task_class.input_model().model_fields + assert input_fields["a"].description == "First number" + assert input_fields["b"].description == "Second number" + assert task_class.output_model().model_fields["total"].description == ( + "Sum of the two numbers" + ) + + task = task_class(inputs={"a": 2, "b": 3}, varinfo=varinfo) + task.execute() + assert task.get_output_values() == {"total": 5} + + +@pytest.mark.parametrize( + "method", [annotated_method, field_default_method], ids=lambda m: m.__name__ +) +def test_task_input_constraint(method, varinfo): + task = get_method_task(qualname(method))(inputs={"a": 2, "b": -1}, varinfo=varinfo) + with pytest.raises(Exception, match="Invalid task inputs"): + task.execute(raise_on_error=True) + + +@validate_call(validate_return=True) +def validated_method(a: float, b: float = 0) -> RangeInfo: + return RangeInfo(extent=abs(b - a), minimum=min(a, b), maximum=max(a, b)) + + +@validate_call(config=ConfigDict(arbitrary_types_allowed=True), validate_return=True) +def validated_arbitrary_method(data: numpy.ndarray) -> int: + return data.size + + +def test_validate_call_task_signature(): + task_class = get_method_task(qualname(validated_method)) + + assert set(task_class.required_input_names()) == {"a"} + assert set(task_class.optional_input_names()) == {"b"} + assert set(task_class.output_names()) == {"extent", "minimum", "maximum"} + + +def test_validate_call_task_execution(varinfo): + task = get_method_task(qualname(validated_method))( + inputs={"a": 15, "b": 10}, varinfo=varinfo + ) + task.execute() + + assert task.get_output_values() == {"extent": 5, "minimum": 10, "maximum": 15} + + +def test_validate_call_arbitrary_type(varinfo): + task = get_method_task(qualname(validated_arbitrary_method))( + inputs={"data": numpy.arange(4)}, varinfo=varinfo + ) + task.execute() + + assert task.get_output_values() == {"return_value": 4} + + +def test_validate_call_direct_call(): + assert validated_method(15, 10) == RangeInfo(extent=5, minimum=10, maximum=15) + + with pytest.raises(ValidationError): + validated_method("not a number") + + +@validate_call(validate_return=True) +def invalid_return_method(a: float) -> RangeInfo: + return {"extent": "not a number", "minimum": a, "maximum": a} + + +def test_validate_call_direct_call_return(): + with pytest.raises(ValidationError): + invalid_return_method(1)