diff --git a/src/dve/core_engine/backends/base/rules.py b/src/dve/core_engine/backends/base/rules.py index ee93106..8508189 100644 --- a/src/dve/core_engine/backends/base/rules.py +++ b/src/dve/core_engine/backends/base/rules.py @@ -448,8 +448,7 @@ def process_node(node: HierarchyNode): record=record, # type: ignore error_location=location, error_message=template_object( - node.missing_parent_id_error_message, - record + node.missing_parent_id_error_message, record ), failure_type="record", error_type="record", @@ -522,10 +521,7 @@ def process_node(node: HierarchyNode) -> bool: entity=node.parent_entity, record=record, # type: ignore error_location=location, - error_message=template_object( - node.no_valid_records_error_message, - record - ), + error_message=template_object(node.no_valid_records_error_message, record), failure_type="record", error_type="record", error_code=node.no_valid_records_error_code, @@ -624,6 +620,7 @@ def apply_sync_filters( excluded_columns=filter_column_names, reporting=rule.reporting, parent=rule.parent, + error_on_null=True, ), ) if not success: @@ -655,6 +652,7 @@ def apply_sync_filters( expression=f"NOT ({rule.expression})", reporting=rule.reporting, parent=rule.parent, + error_on_null=True, ), ) if not success: @@ -887,3 +885,8 @@ def filter_data_contract_record_rejections( ): """Method to filter out record rejection errors from the data contract for a given entity""" raise NotImplementedError() + + @staticmethod + def get_entity_count(entity: EntityType) -> int: + """Method to get count of records in entity""" + raise NotImplementedError() diff --git a/src/dve/core_engine/backends/implementations/duckdb/rules.py b/src/dve/core_engine/backends/implementations/duckdb/rules.py index dbbbc54..1723ce5 100644 --- a/src/dve/core_engine/backends/implementations/duckdb/rules.py +++ b/src/dve/core_engine/backends/implementations/duckdb/rules.py @@ -1,5 +1,5 @@ """Business rule definitions for duckdb backend""" - +# pylint: disable=R0801 from collections.abc import Callable, Iterator from typing import get_type_hints from uuid import uuid4 @@ -23,6 +23,7 @@ from dve.core_engine.backends.implementations.duckdb.duckdb_helpers import ( DDBStruct, ddb_filter_contract_errors, + duckdb_get_entity_count, duckdb_read_parquet, duckdb_record_index, duckdb_rel_to_dictionaries, @@ -62,6 +63,7 @@ from dve.core_engine.type_hints import Messages +@duckdb_get_entity_count @duckdb_record_index @duckdb_write_parquet @duckdb_read_parquet @@ -568,7 +570,12 @@ def notify(self, entities: DuckDBEntities, *, config: Notification) -> Messages: """ messages: Messages = [] entity = entities[config.entity_name] - + if config.error_if_expression_null: + if self.get_entity_count(entity.filter(f"({config.expression}) IS NULL")) > 0: + raise ValueError( + f"The filter evaluated for error code {config.reporting.code}" + + f" in entity {config.entity_name} produced some NULL results. Please investigate." # pylint: disable=C0301 + ) matched = entity.filter(config.expression) if config.excluded_columns: matched = matched.select(StarExpression(exclude=config.excluded_columns)) diff --git a/src/dve/core_engine/backends/implementations/spark/rules.py b/src/dve/core_engine/backends/implementations/spark/rules.py index cec1156..717151e 100644 --- a/src/dve/core_engine/backends/implementations/spark/rules.py +++ b/src/dve/core_engine/backends/implementations/spark/rules.py @@ -1,5 +1,5 @@ """Step implementations in Spark.""" - +# pylint: disable=R0801 from collections.abc import Callable, Iterator from typing import Optional from uuid import uuid4 @@ -15,6 +15,7 @@ get_all_registered_udfs, object_to_spark_literal, spark_filter_contract_errors, + spark_get_entity_count, spark_read_parquet, spark_record_index, spark_write_parquet, @@ -53,6 +54,7 @@ from dve.core_engine.type_hints import Messages +@spark_get_entity_count @spark_record_index @spark_write_parquet @spark_read_parquet @@ -411,6 +413,13 @@ def notify(self, entities: SparkEntities, *, config: Notification) -> Messages: messages: Messages = [] entity = entities[config.entity_name] + if config.error_if_expression_null: + if self.get_entity_count(entity.filter(f"({config.expression}) IS NULL")) > 0: + raise ValueError( + f"The filter evaluated for error code {config.reporting.code}" + + f" in entity {config.entity_name} produced some NULL results. Please investigate." # pylint: disable=C0301 + ) + matched = entity.filter(config.expression) if config.excluded_columns: matched = matched.drop(*config.excluded_columns) diff --git a/src/dve/core_engine/backends/metadata/rules.py b/src/dve/core_engine/backends/metadata/rules.py index 1b0121e..a1db7cc 100644 --- a/src/dve/core_engine/backends/metadata/rules.py +++ b/src/dve/core_engine/backends/metadata/rules.py @@ -282,6 +282,8 @@ class Notification(AbstractStep): """Columns to be excluded from the record in the report.""" reporting: ReportingConfig """The reporting information for the filter.""" + error_if_expression_null: bool = False + """Raise error if the results of evaluating the expression passed leads to some NULL results""" def get_required_entities(self) -> set[EntityName]: return {self.entity_name} diff --git a/src/dve/reporting/__init__.py b/src/dve/reporting/__init__.py index 9a93c67..ab78e11 100644 --- a/src/dve/reporting/__init__.py +++ b/src/dve/reporting/__init__.py @@ -1 +1,2 @@ """Error reports module.""" +# pylint: disable=R0801 diff --git a/tests/test_core_engine/test_backends/test_implementations/test_duckdb/test_rules.py b/tests/test_core_engine/test_backends/test_implementations/test_duckdb/test_rules.py index 5447e14..10480a1 100644 --- a/tests/test_core_engine/test_backends/test_implementations/test_duckdb/test_rules.py +++ b/tests/test_core_engine/test_backends/test_implementations/test_duckdb/test_rules.py @@ -820,6 +820,25 @@ def test_planets_notify(planets_rel: DuckDBPyRelation): assert len(messages[0]) == 4 +def test_notify_null_errors(planets_rel: DuckDBPyRelation): + + config = Notification( + entity_name="planets", + expression="CASE WHEN planet=='Mercury' THEN NULL ELSE False END", + excluded_columns=["mass", "diameter"], + reporting=ReportingConfig( + code="TESTNULLERROR", message="this is a test", location="planet, has_ring_system" + ), + error_if_expression_null=True + ) + entities = EntityManager({"planets": planets_rel}) + messages, success = DUCKDB_STEP_BACKEND.evaluate(entities, config=config) + + assert not success + assert len(messages) == 1 + assert messages[0].is_critical + + def test_read_and_write_simple_parquet(simple_typecast_parquet): parquet_uri, data = simple_typecast_parquet entity: DuckDBPyRelation = DUCKDB_STEP_BACKEND.read_parquet(path=parquet_uri) diff --git a/tests/test_core_engine/test_backends/test_implementations/test_spark/test_rules.py b/tests/test_core_engine/test_backends/test_implementations/test_spark/test_rules.py index 654dc08..6e38c7a 100644 --- a/tests/test_core_engine/test_backends/test_implementations/test_spark/test_rules.py +++ b/tests/test_core_engine/test_backends/test_implementations/test_spark/test_rules.py @@ -21,6 +21,7 @@ from dve.core_engine.backends.base.core import EntityManager from dve.core_engine.backends.exceptions import MissingEntity from dve.core_engine.backends.implementations.spark.rules import SparkStepImplementations +from dve.core_engine.backends.metadata.reporting import ReportingConfig from dve.core_engine.backends.metadata.rules import ( Aggregation, AntiJoin, @@ -32,6 +33,7 @@ HeaderJoin, InnerJoin, LeftJoin, + Notification, OneToOneJoin, OrphanIdentification, RenameEntity, @@ -447,6 +449,24 @@ def test_join_can_take_all_cols( expected_rows = sorted(expected_df.collect(), key=lambda row: row.planet) assert actual_rows == expected_rows + +def test_notify_null_errors(planets_df: DataFrame): + + config = Notification( + entity_name="planets", + expression="CASE WHEN planet=='Mercury' THEN NULL ELSE False END", + excluded_columns=["mass", "diameter"], + reporting=ReportingConfig( + code="TESTNULLERROR", message="this is a test", location="planet, has_ring_system" + ), + error_if_expression_null=True + ) + entities = EntityManager({"planets": planets_df}) + messages, success = SPARK_STEP_BACKEND.evaluate(entities, config=config) + + assert not success + assert len(messages) == 1 + assert messages[0].is_critical def test_one_to_one_join_multi_matches_raises(planets_df: DataFrame, satellites_df: DataFrame):