diff --git a/src/databricks/labs/dqx/check_funcs.py b/src/databricks/labs/dqx/check_funcs.py index 52441ff60..1399555e1 100644 --- a/src/databricks/labs/dqx/check_funcs.py +++ b/src/databricks/labs/dqx/check_funcs.py @@ -2995,7 +2995,12 @@ def apply(df: DataFrame, spark: SparkSession, ref_dfs: dict[str, DataFrame]) -> ref_df = _get_ref_df(ref_df_name, ref_table, ref_dfs, spark) # map type columns must be skipped as they cannot be compared with eqNullSafe - map_type_columns = {field.name for field in df.schema.fields if isinstance(field.dataType, types.MapType)} + map_type_columns = { + field.name + for schema in (df.schema, ref_df.schema) + for field in schema.fields + if isinstance(field.dataType, types.MapType) + } # columns to compare: present in both df and ref_df, not in PK, not excluded, not map type compare_columns = [ diff --git a/tests/integration/test_dataset_checks.py b/tests/integration/test_dataset_checks.py index a8ce29920..749ae498e 100644 --- a/tests/integration/test_dataset_checks.py +++ b/tests/integration/test_dataset_checks.py @@ -2831,6 +2831,26 @@ def test_dataset_compare_ref_as_table_and_skip_map_col(spark: SparkSession, set_ assertDataFrameEqual(actual, expected) +@pytest.mark.parametrize( + "schema, value, ref_schema, ref_value", + [ + ("id int, value string", "source", "id int, value map", {"key": "reference"}), + ("id int, value map", {"key": "source"}, "id int, value string", "reference"), + ], +) +def test_dataset_compare_skips_map_col_from_either_schema( + spark: SparkSession, schema: str, value: Any, ref_schema: str, ref_value: Any +): + df = spark.createDataFrame([[1, value]], schema) + ref_df = spark.createDataFrame([[1, ref_value]], ref_schema) + condition, apply = compare_datasets(columns=["id"], ref_columns=["id"], ref_df_name="ref_df") + + actual = apply(df, spark, {"ref_df": ref_df}).select(*df.columns, condition) + expected = spark.createDataFrame([[1, value, None]], f"{schema}, {get_column_name_or_alias(condition)} string") + + assertDataFrameEqual(actual, expected) + + def test_dataset_compare_with_no_columns_to_compare_and_check_missing(spark: SparkSession): schema = "id long" diff --git a/tests/unit/test_row_checks.py b/tests/unit/test_row_checks.py index 60e79cfff..53b0e2560 100644 --- a/tests/unit/test_row_checks.py +++ b/tests/unit/test_row_checks.py @@ -1,5 +1,10 @@ +import re +from collections.abc import Callable from typing import cast + import pytest +from pyspark.sql import Column + from databricks.labs.dqx.utils import get_column_name_or_alias from databricks.labs.dqx.check_funcs import ( is_equal_to, @@ -9,9 +14,11 @@ is_not_greater_than, is_not_less_than, is_in_list, + is_not_in_list, is_not_null_and_is_in_list, is_aggr_not_greater_than, has_valid_string_case, + has_json_keys, is_ipv4_address_in_cidr, is_ipv6_address_in_cidr, is_valid_national_id, @@ -335,3 +342,87 @@ def test_is_valid_language_code_unsupported_code_format(): def test_is_valid_language_code_case_insensitive_auto_name(): result = is_valid_language_code("a", case_sensitive=False) assert get_column_name_or_alias(result) == "a_is_not_a_valid_language_code" + + +@pytest.mark.parametrize( + "check_func, kwargs, expected_error, expected_message", + [ + ( + is_not_null_and_is_in_list, + {"column": "a", "allowed": None}, + MissingParameterError, + "allowed list is not provided.", + ), + (is_in_list, {"column": "a", "allowed": None}, MissingParameterError, "allowed list is not provided."), + ( + is_in_list, + {"column": "a", "allowed": "not_a_list"}, + InvalidParameterError, + "allowed parameter must be a list", + ), + (is_not_in_list, {"column": "a", "forbidden": None}, MissingParameterError, "forbidden list is not provided."), + ( + is_not_in_list, + {"column": "a", "forbidden": "not_a_list"}, + InvalidParameterError, + "forbidden parameter must be a list", + ), + (is_not_in_list, {"column": "a", "forbidden": []}, InvalidParameterError, "forbidden list must not be empty."), + ( + is_equal_to, + {"column": "a", "value": 1, "abs_tolerance": -1.0}, + InvalidParameterError, + "tolerances if provided must be non-negative", + ), + ( + is_equal_to, + {"column": "a", "value": 1, "rel_tolerance": -1.0}, + InvalidParameterError, + "tolerances if provided must be non-negative", + ), + ( + is_not_equal_to, + {"column": "a", "value": 1, "abs_tolerance": -1.0}, + InvalidParameterError, + "tolerances if provided must be non-negative", + ), + ( + is_not_equal_to, + {"column": "a", "value": 1, "rel_tolerance": -1.0}, + InvalidParameterError, + "tolerances if provided must be non-negative", + ), + ( + is_ipv4_address_in_cidr, + {"column": "a", "cidr_block": 123}, + InvalidParameterError, + "'cidr_block' must be a string", + ), + ( + is_ipv6_address_in_cidr, + {"column": "a", "cidr_block": 123}, + InvalidParameterError, + "'cidr_block' must be a string", + ), + ( + has_json_keys, + {"column": "a", "keys": []}, + InvalidParameterError, + "The 'keys' parameter must be a non-empty list of strings.", + ), + ( + has_json_keys, + {"column": "a", "keys": ["valid", 1]}, + InvalidParameterError, + "All keys must be of type string.", + ), + ], +) +def test_row_check_rejects_invalid_arguments( + check_func: Callable[..., Column], + kwargs: dict[str, object], + expected_error: type[Exception], + expected_message: str, +): + with pytest.raises(expected_error, match=re.escape(expected_message)): + check_func(**kwargs)