Skip to content
Closed
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
7 changes: 6 additions & 1 deletion src/databricks/labs/dqx/check_funcs.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 = [
Expand Down
20 changes: 20 additions & 0 deletions tests/integration/test_dataset_checks.py
Original file line number Diff line number Diff line change
Expand Up @@ -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<string, string>", {"key": "reference"}),
("id int, value map<string, string>", {"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"

Expand Down
91 changes: 91 additions & 0 deletions tests/unit/test_row_checks.py
Original file line number Diff line number Diff line change
@@ -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,
Expand All @@ -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,
Expand Down Expand Up @@ -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)
Loading