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
6 changes: 5 additions & 1 deletion hamilton/function_modifiers/validation.py
Original file line number Diff line number Diff line change
Expand Up @@ -186,7 +186,11 @@ def __init__(self, *validators: dq_base.DataValidator, target_: base.TargetType
self.validators = list(validators)

def get_validators(self, node_to_validate: node.Node) -> list[dq_base.DataValidator]:
return self.validators
return [
validator
for validator in self.validators
if validator.applies_to(node_to_validate.type)
]


class check_output(BaseDataValidationDecorator):
Expand Down
25 changes: 25 additions & 0 deletions tests/function_modifiers/test_validation.py
Original file line number Diff line number Diff line change
Expand Up @@ -128,6 +128,31 @@ def fn(input: pd.Series) -> pd.Series:
)


def test_check_output_custom_only_uses_applicable_validators():
applicable_validator = SampleDataValidator1(equal_to=10, importance="warn")
incompatible_validator = SampleDataValidator2(dataset_length=1, importance="warn")
decorator = check_output_custom(applicable_validator, incompatible_validator)

def fn(input: int) -> int:
return input

validators = decorator.get_validators(node.Node.from_fn(fn))

assert validators == [applicable_validator]


def test_check_output_custom_uses_no_incompatible_validators():
decorator = check_output_custom(
SampleDataValidator2(dataset_length=1, importance="warn"),
SampleDataValidator3(dtype=np.int64, importance="warn"),
)

def fn(input: int) -> int:
return input

assert decorator.get_validators(node.Node.from_fn(fn)) == []


def test_check_output_custom_node_transform_duplicate():
"""You should be able to pass in the same validator twice; IRL it would be different args."""
decorator = check_output_custom(
Expand Down
Loading