diff --git a/lib/ramble/ramble/expander.py b/lib/ramble/ramble/expander.py index 51ffa5e12..cf3deebba 100644 --- a/lib/ramble/ramble/expander.py +++ b/lib/ramble/ramble/expander.py @@ -776,6 +776,9 @@ def expand_var( if var is None or var == "None": return None if typed else "None" + if isinstance(var, (int, float, bool)): + return var if typed else str(var) + passthrough_setting = allow_passthrough # If disable_passthrough is set, override allow_passthrough from caller diff --git a/lib/ramble/ramble/schema/__init__.py b/lib/ramble/ramble/schema/__init__.py index 32b9269f6..66b5a8a5c 100644 --- a/lib/ramble/ramble/schema/__init__.py +++ b/lib/ramble/ramble/schema/__init__.py @@ -41,9 +41,19 @@ def _deprecated_properties(validator, deprecated, instance, schema): yield jsonschema.ValidationError(msg) - return jsonschema.validators.extend( + ValidatorClass = jsonschema.validators.extend( jsonschema.Draft4Validator, {"deprecatedProperties": _deprecated_properties} ) + import spack.util.spack_yaml as syaml + + # Add syaml_bool to the accepted boolean types for validation + boolean_types = ValidatorClass.DEFAULT_TYPES.get("boolean", bool) + if not isinstance(boolean_types, tuple): + boolean_types = (boolean_types,) + ValidatorClass.DEFAULT_TYPES["boolean"] = boolean_types + (syaml.syaml_bool,) + + return ValidatorClass + Validator = llnl.util.lang.Singleton(_make_validator) diff --git a/lib/ramble/ramble/test/end_to_end/env_var_builtin.py b/lib/ramble/ramble/test/end_to_end/env_var_builtin.py index c68afd915..4759776c2 100644 --- a/lib/ramble/ramble/test/end_to_end/env_var_builtin.py +++ b/lib/ramble/ramble/test/end_to_end/env_var_builtin.py @@ -324,3 +324,43 @@ def test_auto_env_vars(make_workspace_from_config, mock_applications, mock_modif assert "MY_AUTO_ENV_VAR_WL_DEFAULTS" not in data assert "MY_AUTO_ENV_VAR_WG" not in data assert "OBJ_AUTO_ENV_VAR" not in data + + +def test_env_var_bool_case(mock_applications, make_workspace_from_config): + test_config = """ +ramble: + config: + shell: bash + variables: + mpi_command: 'mpirun -n {n_ranks} -ppn {processes_per_node}' + batch_submit: 'batch_submit {execute_experiment}' + partition: 'part1' + processes_per_node: '16' + n_threads: '1' + applications: + interleved-env-vars: + workloads: + test_wl: + experiments: + simple_test: + variables: + n_nodes: 1 + env_vars: + set: + MY_VAR: TRUE + software: + packages: {} + environments: {} +""" + ws, ws_name = make_workspace_from_config(test_config) + + workspace("setup", "--dry-run", global_args=["-w", ws_name]) + + experiment_root = ws.experiment_dir + exp_dir = os.path.join(experiment_root, "interleved-env-vars", "test_wl", "simple_test") + exp_script = os.path.join(exp_dir, "execute_experiment") + + with open(exp_script, encoding="utf-8") as f: + data = f.read() + + assert "MY_VAR=TRUE" in data diff --git a/lib/ramble/ramble/variants.py b/lib/ramble/ramble/variants.py index 19e72abf2..83263b0f7 100644 --- a/lib/ramble/ramble/variants.py +++ b/lib/ramble/ramble/variants.py @@ -15,6 +15,8 @@ import ramble.util.colors as color from ramble.expander import Expander +from spack.util.spack_yaml import syaml_bool + reserved_variants = { "modifier", "package_manager", @@ -223,7 +225,7 @@ def experiment_variant(self, name: str, value: Any): default_var = self.default_variants[name] # If the default value is a boolean, convert the experiment value to a boolean - if default_var and isinstance(default_var.default, bool): + if default_var and isinstance(default_var.default, (bool, syaml_bool)): if isinstance(value, str): value = value.lower() == "true" @@ -413,7 +415,7 @@ def as_set(self, expander: Optional[Expander] = None) -> set: if callable(values): is_valid = values(val) else: - is_valid = str(val) in [str(v) for v in values] + is_valid = str(val).lower() in [str(v).lower() for v in values] if not is_valid: raise RambleVariantError( f"When defining variant {name} the value {val} is not valid.\n" @@ -461,7 +463,7 @@ def __init__( self._definitions = self._build_definitions() def _build_definitions(self) -> tuple: - if isinstance(self.default, bool): + if isinstance(self.default, (bool, syaml_bool)): val_str = str(self.default) return ( self._definition, @@ -477,7 +479,7 @@ def copy(self): def format_value(self, value: Any) -> str: """Format a value for this variant into Spack-like syntax""" - if isinstance(self.default, bool): + if isinstance(self.default, (bool, syaml_bool)): prefix = "+" if value else "~" return f"{prefix}{self.name}" else: diff --git a/lib/ramble/spack/util/spack_yaml.py b/lib/ramble/spack/util/spack_yaml.py index bcaffd4b0..cb636108a 100644 --- a/lib/ramble/spack/util/spack_yaml.py +++ b/lib/ramble/spack/util/spack_yaml.py @@ -50,10 +50,25 @@ class syaml_int(int): __repr__ = int.__repr__ +class syaml_bool(int): + def __new__(cls, val, string_val=None): + obj = super(syaml_bool, cls).__new__(cls, val) + obj.string_val = string_val + return obj + + def __repr__(self): + return self.string_val if self.string_val else ('True' if self else 'False') + + def __str__(self): + return self.string_val if self.string_val else ('True' if self else 'False') + + + #: mapping from syaml type -> primitive type syaml_types = { syaml_str: str, syaml_int: int, + syaml_bool: bool, syaml_dict: dict, syaml_list: list, } @@ -121,6 +136,13 @@ class OrderedLineLoader(RoundTripLoader): # and fill in with mappings later. We preserve this behavior. # + + def construct_yaml_bool(self, node): + value = super(OrderedLineLoader, self).construct_yaml_bool(node) + b = syaml_bool(value, string_val=node.value) + mark(b, node) + return b + def construct_yaml_str(self, node): value = super(OrderedLineLoader, self).construct_yaml_str(node) # There is no specific marker to indicate that we are parsing a key, @@ -156,6 +178,8 @@ def construct_yaml_map(self, node): # register above new constructors +OrderedLineLoader.add_constructor( + 'tag:yaml.org,2002:bool', OrderedLineLoader.construct_yaml_bool) OrderedLineLoader.add_constructor( 'tag:yaml.org,2002:map', OrderedLineLoader.construct_yaml_map) OrderedLineLoader.add_constructor( @@ -184,6 +208,12 @@ def represent_data(self, data): result.value = syaml_str("null") return result + + def represent_bool(self, data): + if hasattr(data, 'string_val') and data.string_val: + return self.represent_scalar('tag:yaml.org,2002:bool', data.string_val) + return super(OrderedLineDumper, self).represent_bool(bool(data)) + def represent_str(self, data): if hasattr(data, 'override') and data.override: data = data + ':' @@ -198,12 +228,19 @@ def ignore_aliases(self, _data): """Make the dumper NEVER print YAML aliases.""" return True + def represent_bool(self, data): + if hasattr(data, 'string_val') and data.string_val: + return self.represent_scalar('tag:yaml.org,2002:bool', data.string_val) + return super(SafeDumper, self).represent_bool(bool(data)) + # Make our special objects look like normal YAML ones. RoundTripDumper.add_representer(syaml_dict, RoundTripDumper.represent_dict) RoundTripDumper.add_representer(syaml_list, RoundTripDumper.represent_list) RoundTripDumper.add_representer(syaml_int, RoundTripDumper.represent_int) +RoundTripDumper.add_representer(syaml_bool, RoundTripDumper.represent_bool) RoundTripDumper.add_representer(syaml_str, RoundTripDumper.represent_str) +OrderedLineDumper.add_representer(syaml_bool, OrderedLineDumper.represent_bool) OrderedLineDumper.add_representer(syaml_str, OrderedLineDumper.represent_str)