diff --git a/test/test_parameters.py b/test/test_parameters.py index 4275a15..de52767 100644 --- a/test/test_parameters.py +++ b/test/test_parameters.py @@ -19,6 +19,7 @@ def expected_vars(): "some_float": (float, 3.14159), "some_string": (str, "foo"), "some_bool": (bool, False), + "some_directory": ("Directory", "/some/path") } @@ -37,6 +38,9 @@ def params_yaml(): some_bool: type: bool default: false +some_directory: + type: Directory + default: "/some/path" """ @@ -130,6 +134,13 @@ def test_parameters_get_commandline_inputs(notebook_parameters): "doc": "some_bool", "inputBinding": {"prefix": "--some-bool"}, }, + "some_directory": { + "type": "Directory", + "default": "/some/path", + "label": "some_directory", + "doc": "some_directory", + "inputBinding": {"prefix": "--some-directory"}, + }, } @@ -139,6 +150,7 @@ def test_parameters_get_cwl_step_inputs(notebook_parameters): "some_float": "some_float", "some_string": "some_string", "some_bool": "some_bool", + "some_directory": "some_directory", } @@ -148,23 +160,33 @@ def test_parameters_from_code(expected_vars): some_float = 3.14159 some_string = "foo" some_bool = False +some_directory: "EOInput" = "/some/path" """) assert parameters.params == expected_vars assert parameters.config == {} -def test_parameters_from_code_with_xce_config(expected_vars): - xce_config = dict(foo=1, bar="hi!", baz={}) +@pytest.mark.parametrize("config", [ + (dict(foo=1, bar="hi!", baz={}), True), + ("Not a dict", False), + ({42: "Wrong key type"}, False)]) +def test_parameters_from_code_with_xce_config(expected_vars, config): + xce_config, valid = config code = f""" some_int = 42 some_float = 3.14159 some_string = "foo" some_bool = False +some_directory: "EOInput" = "/some/path" {NotebookParameters.config_var_name} = {xce_config!r} """ - parameters = xcengine.parameters.NotebookParameters.from_code(code) - assert parameters.params == expected_vars - assert parameters.config == xce_config + if valid: + parameters = xcengine.parameters.NotebookParameters.from_code(code) + assert parameters.params == expected_vars + assert parameters.config == xce_config + else: + with pytest.raises(TypeError): + xcengine.parameters.NotebookParameters.from_code(code) def test_parameters_from_code_with_setup(expected_vars): @@ -175,11 +197,13 @@ def test_parameters_from_code_with_setup(expected_vars): some_float = 3.14159 some_string = some_uppercase_string.lower() some_bool = not not_some_bool +some_directory: "EOInput" = "/" + "/".join(some_path_components) """, setup_code=""" half_of_some_int = 21 some_uppercase_string = "FOO" not_some_bool = True +some_path_components = ["some", "path"] """, ).params == expected_vars @@ -212,6 +236,12 @@ def test_parameters_get_workflow_inputs(notebook_parameters): "label": "some_bool", "doc": "some_bool", }, + "some_directory": { + "type": "Directory", + "default": "/some/path", + "label": "some_directory", + "doc": "some_directory", + } } @@ -221,9 +251,20 @@ def test_parameters_to_yaml(notebook_parameters): "some_float": {"type": "float", "default": 3.14159}, "some_string": {"type": "str", "default": "foo"}, "some_bool": {"type": "bool", "default": False}, + "some_directory": {"type": "Directory", "default": "/some/path"}, } +def test_parameters_to_yaml_unhandled_type(): + with pytest.raises(TypeError): + # Create empty parameters and modify them afterwards to avoid + # __init__ catching the mistake. + np = xcengine.parameters.NotebookParameters({}) + # Disable the inspection, since this is wrong on purpose. + # noinspection bad-assignment + np.params = {"foo": (42, 42)} + np.to_yaml() + def test_parameters_from_yaml(expected_vars, params_yaml): assert NotebookParameters.from_yaml(params_yaml).params == expected_vars @@ -241,6 +282,15 @@ def test_parameters_from_yaml_with_dataset(): } +def test_parameters_from_yaml_unknown_type(): + with pytest.raises(ValueError): + NotebookParameters.from_yaml(""" +some_input: + type: unsupported + default: null + """) + + def test_parameters_from_file(tmp_path, expected_vars, params_yaml): path = tmp_path / "params.yaml" path.write_text(params_yaml) diff --git a/xcengine/parameters.py b/xcengine/parameters.py index c88f3dc..0992096 100644 --- a/xcengine/parameters.py +++ b/xcengine/parameters.py @@ -2,7 +2,7 @@ import os import pathlib import typing -from typing import Any, ClassVar +from typing import Any, ClassVar, cast import xarray as xr import yaml @@ -13,7 +13,7 @@ class NotebookParameters: - params: dict[str, tuple[type, Any]] + params: dict[str, tuple[type | str, Any]] cwl_params: dict[str, tuple[type | str, Any]] dataset_inputs: list[str] config_var_name: ClassVar[str] = "xcengine_config" @@ -21,7 +21,7 @@ class NotebookParameters: def __init__( self, - params: dict[str, tuple[type, Any]], + params: dict[str, tuple[type | str, Any]], config: dict[str, Any] | None = None, ): self.params = params @@ -46,16 +46,27 @@ def from_code( ) -> "NotebookParameters": variables = cls.extract_variables(code, setup_code) config = variables.pop(cls.config_var_name, (None, None)) - # TODO: throw an error here if config has wrong type - return cls(variables, config[1]) + if config[1] is not None: + if type(config[1]) is not dict: + raise TypeError("Configuration variable must be a dict") + if not all(type(k) is str for k in cast(dict, config[1]).keys()): + raise TypeError("Configuration dict keys must be strings") + return cls(variables, cast(dict[str, Any], config[1])) @classmethod def from_yaml(cls, yaml_content: str | typing.IO) -> "NotebookParameters": input_data = yaml.safe_load(yaml_content) + def convert_type(yaml_spec: str) -> type | str: + match yaml_spec: + case "int" | "float" | "bool" | "str" | "Dataset": + return eval(yaml_spec, globals(), {"Dataset": xr.Dataset}) + case "Directory": + return "Directory" + raise ValueError(f'Unknown type in YAML: "{yaml_spec}"') return cls( { k: ( - eval(v["type"], globals(), {"Dataset": xr.Dataset}), + convert_type(v["type"]), v["default"], ) for k, v in input_data.items() @@ -70,7 +81,7 @@ def from_yaml_file(cls, path: str | os.PathLike) -> "NotebookParameters": @classmethod def extract_variables( cls, code: str, setup_code: str | None = None - ) -> dict[str, tuple[type, Any]]: + ) -> dict[str, tuple[type | str, Any]]: if setup_code is None: locals_: dict[str, object] = {} old_locals = {} @@ -78,14 +89,19 @@ def extract_variables( exec(setup_code, globals(), locals_ := {}) old_locals = locals_.copy() exec(code, globals(), locals_) + annotations = cls.read_annotations(code) new_vars = locals_.keys() - old_locals.keys() new_var_dict = { - k: cls.make_param_tuple(k, locals_[k]) for k in new_vars + k: cls.make_param_tuple(k, locals_[k]) for k in new_vars if not k.startswith("__") } + for k in new_var_dict: + if k in annotations and annotations[k] == "'EOInput'": + old_var = new_var_dict[k] + new_var_dict[k] = ("Directory", old_var[1]) return dict(sorted(new_var_dict.items())) @classmethod - def make_param_tuple(cls, key: str, value: Any) -> tuple[type, Any]: + def make_param_tuple(cls, key: str, value: Any) -> tuple[type | str, Any]: return ( t := type(value), ( @@ -125,9 +141,17 @@ def get_cwl_commandline_input(self, var_name: str) -> dict[str, Any]: } def to_yaml(self) -> str: + def dump_type(type_: type | str) -> str: + match type_: + case type(): + return type_.__name__ + case str(): + return type_ + case _: + raise TypeError(f"Unhandled type {type_} for YAML export") return yaml.safe_dump( { - name: {"type": type_.__name__, "default": default_} + name: {"type": dump_type(type_), "default": default_} for name, (type_, default_) in self.params.items() } ) @@ -223,7 +247,7 @@ def read_staged_in_dataset( return xr.open_dataset(stage_in_path / asset.href) @staticmethod - def cwl_type(type_: type) -> str: + def cwl_type(type_: type | str) -> str: try: # noinspection PyTypeChecker return { @@ -231,6 +255,7 @@ def cwl_type(type_: type) -> str: float: "double", str: "string", bool: "boolean", + "Directory": "Directory", }[type_] except KeyError: raise ValueError(f"Unhandled type {type_}")