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
23 changes: 23 additions & 0 deletions test/test_parameters.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
}


Expand All @@ -37,6 +38,9 @@ def params_yaml():
some_bool:
type: bool
default: false
some_directory:
type: Directory
default: "/some/path"
"""


Expand Down Expand Up @@ -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"},
},
}


Expand All @@ -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",
}


Expand All @@ -148,6 +160,7 @@ 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 == {}
Expand All @@ -160,6 +173,7 @@ def test_parameters_from_code_with_xce_config(expected_vars):
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)
Expand All @@ -175,11 +189,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
Expand Down Expand Up @@ -212,6 +228,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",
}
}


Expand All @@ -221,6 +243,7 @@ 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"},
}


Expand Down
47 changes: 36 additions & 11 deletions xcengine/parameters.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -13,15 +13,15 @@

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"
config: dict[str, Any]

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
Expand All @@ -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()
Expand All @@ -70,22 +81,27 @@ 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 = {}
else:
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),
(
Expand Down Expand Up @@ -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()
}
)
Expand Down Expand Up @@ -223,14 +247,15 @@ 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 {
int: "long",
float: "double",
str: "string",
bool: "boolean",
"Directory": "Directory",
}[type_]
except KeyError:
raise ValueError(f"Unhandled type {type_}")
Expand Down
Loading