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
Original file line number Diff line number Diff line change
Expand Up @@ -104,9 +104,34 @@ def validation_base_name(function_name: str):


@typechecked
def validate_validator_combinations(param_name: str, validations_dict: dict):
def validate_validator_combinations(
param_name: str, defined_type: str, validations_dict: dict
):
validation_names = {validation_base_name(name) for name in validations_dict}

scalar_bound_validators = {'bounds', 'gt', 'gt_eq', 'lt', 'lt_eq'}
element_bound_validators = {
'element_bounds',
'lower_element_bounds',
'upper_element_bounds',
}
if array_type(defined_type) and validation_names.intersection(
scalar_bound_validators
):
raise compile_error(
"Parameter {} has array type '{}' but uses a scalar bound validator. "
"Use 'element_bounds/lower_element_bounds/upper_element_bounds' instead.".format(
param_name, defined_type
)
)
if not array_type(defined_type) and validation_names.intersection(
element_bound_validators
):
raise compile_error(
"Parameter {} has scalar type '{}' but uses an element bound validator. "
"Use 'bounds/gt/gt_eq/lt/lt_eq' instead.".format(param_name, defined_type)
)

if 'element_bounds' in validation_names and {
'lower_element_bounds',
'upper_element_bounds',
Expand All @@ -117,9 +142,8 @@ def validate_validator_combinations(param_name: str, validations_dict: dict):
)
)

scalar_bound_validators = {'gt', 'gt_eq', 'lt', 'lt_eq'}
if 'bounds' in validation_names and validation_names.intersection(
scalar_bound_validators
scalar_bound_validators - {'bounds'}
):
raise compile_error(
"Parameter {} cannot combine 'bounds' with scalar bound validators "
Expand Down Expand Up @@ -793,7 +817,7 @@ def preprocess_inputs(language, name, value, nested_name_list):
if is_fixed_type(defined_type):
validations_dict['size_lt<>'] = fixed_type_size(defined_type) + 1

validate_validator_combinations(param_name, validations_dict)
validate_validator_combinations(param_name, defined_type, validations_dict)

for func_name in validations_dict:
args = validations_dict[func_name]
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -21,7 +21,10 @@
from generate_parameter_library_py.generate_cpp_header import run as run_cpp
from generate_parameter_library_py.generate_python_module import run as run_python
from generate_parameter_library_py.generate_markdown import run as run_md
from generate_parameter_library_py.parse_yaml import YAMLSyntaxError
from generate_parameter_library_py.parse_yaml import (
YAMLSyntaxError,
validate_validator_combinations,
)
from generate_parameter_library_py.generate_cpp_header import parse_args


Expand Down Expand Up @@ -100,3 +103,48 @@ def test_parse_valid_parameter_files(yaml_test_file):
set_up(yaml_test_file)
except Exception as e:
assert False, f'failed to parse valid file, reason:{e}'


@pytest.mark.parametrize(
'defined_type,validations,expected_message',
[
(
'double_array',
{'bounds<>': [0.0, 1.0]},
'uses a scalar bound validator',
),
(
'int_array_fixed_03',
{'gt_eq<>': [0]},
'uses a scalar bound validator',
),
(
'double',
{'element_bounds<>': [0.0, 1.0]},
'uses an element bound validator',
),
(
'int',
{'upper_element_bounds<>': [10]},
'uses an element bound validator',
),
],
)
def test_bound_validator_matches_parameter_type(
defined_type, validations, expected_message
):
with pytest.raises(YAMLSyntaxError, match=expected_message):
validate_validator_combinations('test_param', defined_type, validations)


@pytest.mark.parametrize(
'defined_type,validations',
[
('double_array', {'element_bounds<>': [0.0, 1.0]}),
('int_array_fixed_03', {'lower_element_bounds<>': [0]}),
('double', {'bounds<>': [0.0, 1.0]}),
('int', {'gt_eq<>': [0]}),
],
)
def test_bound_validator_accepts_matching_parameter_type(defined_type, validations):
validate_validator_combinations('test_param', defined_type, validations)