From 676ac5a51ef3fba8d61395e22126e22e11519078 Mon Sep 17 00:00:00 2001 From: stevenhsd <56357022+stevenhsd@users.noreply.github.com> Date: Wed, 23 Sep 2026 00:53:36 +0100 Subject: [PATCH 1/2] fix: ensure missing parent and group rejection records are being removed - add test coverage --- src/dve/core_engine/backends/base/rules.py | 179 +++++++++--------- .../core_engine/backends/base/utilities.py | 4 +- .../backends/implementations/duckdb/rules.py | 4 +- src/dve/pipeline/pipeline.py | 61 ++++-- tests/features/flights.feature | 51 ++++- tests/features/steps/steps_post_pipeline.py | 13 ++ 6 files changed, 196 insertions(+), 116 deletions(-) diff --git a/src/dve/core_engine/backends/base/rules.py b/src/dve/core_engine/backends/base/rules.py index a3ed6a4..9fb9e3b 100644 --- a/src/dve/core_engine/backends/base/rules.py +++ b/src/dve/core_engine/backends/base/rules.py @@ -387,7 +387,7 @@ def identify_and_remove_orphans( entities: Entities, entity_hierarchy: EntityHierarchy, key_fields: Optional[dict[str, list[str]]] = None, - ) -> tuple[Messages, bool]: + ) -> tuple[Messages, dict[EntityName, bool]]: """ Identifies and removes orphan records by traversing the EntityHierarchy object. An orphan is a child record whose parent FK does not exist in the parent entity. @@ -395,82 +395,79 @@ def identify_and_remove_orphans( """ def process_node( - node: HierarchyNode, - orph_messages: Messages | None = None, - processed: bool = False, + node: HierarchyNode ): """Identify orphans and remove in a given node""" + issues_found: bool = False + if node.parent_entity is None: + return issues_found - if orph_messages is None: - orph_messages = [] + self.logger.info(f"Identifying orphans in {node.entity_name}") - if node.parent_entity is not None: - self.logger.info(f"Identifying orphans in {node.entity_name}") + join_expr = " AND ".join( + f"{node.parent_entity}.{k} = {node.entity_name}.{v}" + for k, v in node.join_fields.items() + ) - join_expr = " AND ".join( - f"{node.parent_entity}.{k} = {node.entity_name}.{v}" - for k, v in node.join_fields.items() - ) + _, no_orphs = self.identify_orphans( + entities=entities, + config=OrphanIdentification( + id=list(node.join_fields.values())[0], + entity_name=node.entity_name, + target_name=node.parent_entity, + join_condition=join_expr, + ), + ) - _, no_orphs = self.identify_orphans( - entities=entities, - config=OrphanIdentification( - id=list(node.join_fields.values())[0], - entity_name=node.entity_name, - target_name=node.parent_entity, - join_condition=join_expr, - ), + if no_orphs > 0: + self.logger.info( + f"Removing records with missing parent from {node.entity_name}" ) - - if no_orphs > 0: - self.logger.info( - f"Removing records with missing parent from {node.entity_name}" - ) - processed = True - location = list(node.join_fields.values())[0] - with BackgroundMessageWriter( - working_directory=working_directory, - dve_stage=self.__stage_name__, - key_fields=key_fields, - logger=self.logger, - ) as msg_writer: - _orph_records = self.remove_orphans( - entities=entities, - config=OrphanRemoval( - entity_name=node.entity_name, - reporting=ReportingConfig( - emit="record_failure", - code=node.missing_parent_id_error_code, - message=node.missing_parent_id_error_message, - location=location, - ), + issues_found = True + location = list(node.join_fields.values())[0] + with BackgroundMessageWriter( + working_directory=working_directory, + dve_stage=self.__stage_name__, + key_fields=key_fields, + logger=self.logger, + ) as msg_writer: + _orph_records = self.remove_orphans( + entities=entities, + config=OrphanRemoval( + entity_name=node.entity_name, + reporting=ReportingConfig( + emit="record_failure", + code=node.missing_parent_id_error_code, + message=node.missing_parent_id_error_message, + location=location, ), - ) - # moved to batch the write - risky if large number of - msg_writer.write_queue.put( - [ - FeedbackMessage( - entity=node.entity_name, - record=record, # type: ignore - error_location=location, - error_message=node.missing_parent_id_error_message, - failure_type="record", - error_type="record", - error_code=node.missing_parent_id_error_code, - reporting_field=location, - category="Parent Missing", - ) - for record in _orph_records - ] - ) + ), + ) + # moved to batch the write - risky if large number of + msg_writer.write_queue.put( + [ + FeedbackMessage( + entity=node.entity_name, + record=record, # type: ignore + error_location=location, + error_message=node.missing_parent_id_error_message, + failure_type="record", + error_type="record", + error_code=node.missing_parent_id_error_code, + reporting_field=location, + category="Parent Missing", + ) + for record in _orph_records + ] + ) - return processed + return issues_found - processed = False + entity_issues_found: dict[EntityName, bool] = {} for tree in entity_hierarchy.entity_trees.values(): for node in tree.iterate_root_down(): - processed = process_node(node) + entity_issues_found[node.entity_name] = process_node(node) _orph_rel = entities.get(ORPHANED_RECORD_ENTITY_NAME) if _orph_rel is not None: @@ -478,7 +475,7 @@ def process_node( entities.update(entities) - return [], processed + return [], entity_issues_found def identify_and_remove_missing_mandatory_groups( self, @@ -486,20 +483,17 @@ def identify_and_remove_missing_mandatory_groups( entities: Entities, entity_hierarchy: EntityHierarchy, key_fields: Optional[dict[str, list[str]]] = None, - ) -> tuple[Messages, bool]: + ) -> tuple[Messages, dict[EntityName, bool]]: """ Identify that an entity with a mandatory key has at least one valid child record. """ def process_node( - node: HierarchyNode, - processed: bool = False, + node: HierarchyNode ) -> bool: """Identify at least one valid child for a mandatory entity at a given node.""" if node.parent_entity is None or not node.mandatory: - return processed - - processed = True + return False self.logger.info( f"Identifying that mandatory entity `{node.parent_entity}` has at least 1 valid child record" # pylint: disable=C0301 @@ -525,34 +519,33 @@ def process_node( join_condition=join_expr, ), ) - for record in missing_children_records: - msg_writer.write_queue.put( - [ - FeedbackMessage( - entity=node.parent_entity, - record=record, # type: ignore - error_location=location, - error_message=node.no_valid_records_error_message, - failure_type="record", - error_type="record", - error_code=node.no_valid_records_error_code, - reporting_field=location, - category="Children missing", - ) - ] - ) - - return processed + _messages = [ + FeedbackMessage( + entity=node.parent_entity, + record=record, # type: ignore + error_location=location, + error_message=node.no_valid_records_error_message, + failure_type="record", + error_type="record", + error_code=node.no_valid_records_error_code, + reporting_field=location, + category="Children missing", + ) + for record in missing_children_records + ] + msg_writer.write_queue.put(_messages) + return len(_messages) > 0 - processed = False + entity_issues_found: dict[EntityName, bool] = {} for tree in entity_hierarchy.entity_trees.values(): for node in tree.iterate_lowest_descendent_up(): - processed = process_node(node, processed) + if node.parent_entity and node.mandatory: + entity_issues_found[node.parent_entity] = process_node(node) - entities.update(entities) + #entities.update(entities) - return [], processed + return [], entity_issues_found # pylint: disable=R0912,R0914 def apply_sync_filters( diff --git a/src/dve/core_engine/backends/base/utilities.py b/src/dve/core_engine/backends/base/utilities.py index f55bc88..aa99f63 100644 --- a/src/dve/core_engine/backends/base/utilities.py +++ b/src/dve/core_engine/backends/base/utilities.py @@ -3,7 +3,7 @@ import warnings from collections import deque from collections.abc import Sequence -from typing import Optional +from typing import Iterator, Optional import pyarrow # type: ignore import pyarrow.parquet as pq # type: ignore @@ -144,3 +144,5 @@ def check_if_parquet_file(file_location: URI) -> bool: return True except (pyarrow.ArrowInvalid, pyarrow.ArrowIOError): return False + + diff --git a/src/dve/core_engine/backends/implementations/duckdb/rules.py b/src/dve/core_engine/backends/implementations/duckdb/rules.py index 867b262..64abb79 100644 --- a/src/dve/core_engine/backends/implementations/duckdb/rules.py +++ b/src/dve/core_engine/backends/implementations/duckdb/rules.py @@ -434,7 +434,7 @@ def identify_orphans( def remove_orphans(self, entities: DuckDBEntities, *, config: OrphanRemoval) -> Iterator: """Method to remove identified orphans in the orphan tracker entity.""" - orphan_rel = entities[ORPHANED_RECORD_ENTITY_NAME].set_alias("orphan") + orphan_rel = entities[ORPHANED_RECORD_ENTITY_NAME].filter(f"entity_name = '{config.entity_name}'").set_alias("orphan") filtered_rel = ( entities[config.entity_name] .set_alias(config.entity_name) @@ -448,7 +448,7 @@ def remove_orphans(self, entities: DuckDBEntities, *, config: OrphanRemoval) -> entities[config.entity_name] = filtered_rel return duckdb_rel_to_dictionaries( - orphan_rel.filter(f"entity_name = '{config.entity_name}'") + orphan_rel ) def check_mandatory_group( diff --git a/src/dve/pipeline/pipeline.py b/src/dve/pipeline/pipeline.py index c157b3b..8bc85d5 100644 --- a/src/dve/pipeline/pipeline.py +++ b/src/dve/pipeline/pipeline.py @@ -643,30 +643,39 @@ def apply_business_rules( # pylint: disable=R0914,R0915 projected ) - _, orph_or_group = self.step_implementations.identify_and_remove_orphans( # type: ignore + _, orph_issues_1 = self.step_implementations.identify_and_remove_orphans( # type: ignore working_directory, entity_manager.entities, entity_hierarchy, key_fields, ) + - _, orph_or_group = self.step_implementations.identify_and_remove_missing_mandatory_groups( # type: ignore + _, grp_issues_1 = self.step_implementations.identify_and_remove_missing_mandatory_groups( # type: ignore working_directory, entity_manager.entities, entity_hierarchy, key_fields, ) + # Perform a second time incase the mandatory groups result in new orphans - _, orph_or_group = self.step_implementations.identify_and_remove_orphans( # type: ignore + _, orph_issues_2 = self.step_implementations.identify_and_remove_orphans( # type: ignore working_directory, entity_manager.entities, entity_hierarchy, key_fields, ) + + entity_issues: dict[EntityName, bool] = { + entity: any( + val for val in (orph_issues_1.get(entity, False), grp_issues_1.get(entity, False), orph_issues_2.get(entity, False))) + for entity in orph_issues_1.keys() + } + unchanged_entities: list[EntityName] = [] for entity_name, entity in entity_manager.entities.items(): - if orph_or_group: + if entity_issues.get(entity_name, False): self._logger.info(f"Writing {entity_name} out to disk.") final_projection = self._step_implementations.write_parquet( # type: ignore entity, @@ -677,27 +686,41 @@ def apply_business_rules( # pylint: disable=R0914,R0915 entity_name, ), ) + + entity_manager.entities[entity_name] = self.step_implementations.read_parquet( # type: ignore + final_projection + ) else: - self._logger.info(f"Moving {entity_name} from temp_business_rules to business_rules") - final_projection = fh.move_resource( - source_uri=fh.joinuri( - self.processed_files_path, - submission_info.submission_id, - "temp_business_rules", - entity_name - ), - target_uri=fh.joinuri( - self.processed_files_path, - submission_info.submission_id, - "business_rules", - entity_name - ) - ) + unchanged_entities.append(entity_name) + + for entity_name in unchanged_entities: + self._logger.info(f"Moving {entity_name} from temp_business_rules to business_rules") + final_projection = fh.move_resource( + source_uri=fh.joinuri( + self.processed_files_path, + submission_info.submission_id, + "temp_business_rules", + entity_name + ), + target_uri=fh.joinuri( + self.processed_files_path, + submission_info.submission_id, + "business_rules", + entity_name + ), + overwrite=True + ) entity_manager.entities[entity_name] = self.step_implementations.read_parquet( # type: ignore final_projection ) + + fh.remove_prefix(fh.joinuri( + self.processed_files_path, + submission_info.submission_id, + "temp_business_rules")) + submission_status.number_of_records = self.get_entity_count( entity=entity_manager.entities[f"""Original{rules.global_variables.get( 'entity', diff --git a/tests/features/flights.feature b/tests/features/flights.feature index 14572d4..b3e6790 100644 --- a/tests/features/flights.feature +++ b/tests/features/flights.feature @@ -20,6 +20,13 @@ Feature: Pipeline tests using the flights dataset When I run the business rules phase Then there are no file rejections from the business_rules phase And there are no record rejections from the business_rules phase + And the final entities have the following row counts + | entity_name | row_count | + | country | 1 | + | airport | 3 | + | staff | 15 | + | flights | 10 | + | passengers | 25 | When I run the error report phase Then An error report is produced And The statistics entry for the submission shows the following information @@ -54,6 +61,13 @@ Feature: Pipeline tests using the flights dataset | record | StaffHasNoAirport | 15 | | record | FlightHasNoAirport | 10 | | record | PassengerHasNoFlight | 25 | + And the final entities have the following row counts + | entity_name | row_count | + | country | 0 | + | airport | 0 | + | staff | 0 | + | flights | 0 | + | passengers | 0 | When I run the error report phase Then An error report is produced And The statistics entry for the submission shows the following information @@ -83,6 +97,13 @@ Feature: Pipeline tests using the flights dataset | ErrorType | ErrorCode | error_count | | record | FlightIDMissing | 1 | | record | PassengerHasNoFlight | 3 | + And the final entities have the following row counts + | entity_name | row_count | + | country | 1 | + | airport | 3 | + | staff | 15 | + | flights | 9 | + | passengers | 22 | When I run the error report phase Then An error report is produced And The statistics entry for the submission shows the following information @@ -92,7 +113,7 @@ Feature: Pipeline tests using the flights dataset | number_record_rejections | 4 | | number_warnings | 0 | - Scenario: A flights submission with no valid airports record on submission + Scenario: A flights submission with only country id and name submitted Given I submit the flights file only_country_id.xml for processing And A duckdb pipeline is configured with schema file 'flights.dischema.json' And I add initial audit entries for the submission @@ -110,6 +131,13 @@ Feature: Pipeline tests using the flights dataset Then there are errors with the following details and associated error_count from the business_rules phase | ErrorType | ErrorCode | error_count | | record | CountryHasNoAirport | 1 | + And the final entities have the following row counts + | entity_name | row_count | + | country | 0 | + | airport | 0 | + | staff | 0 | + | flights | 0 | + | passengers | 0 | When I run the error report phase Then An error report is produced And The statistics entry for the submission shows the following information @@ -140,6 +168,13 @@ Feature: Pipeline tests using the flights dataset | record | error | PassengerHasNoFlight | 4 | | record | error | AirportHasNoStaff | 1 | | record | error | CountryHasNoAirport | 1 | + And the final entities have the following row counts + | entity_name | row_count | + | country | 0 | + | airport | 0 | + | staff | 0 | + | flights | 0 | + | passengers | 0 | When I run the error report phase Then An error report is produced And The statistics entry for the submission shows the following information @@ -171,6 +206,13 @@ Feature: Pipeline tests using the flights dataset | record | error | AirportHasNoStaff | 1 | | record | error | FlightHasNoAirport | 1 | | record | error | PassengerHasNoFlight | 1 | + And the final entities have the following row counts + | entity_name | row_count | + | country | 1 | + | airport | 1 | + | staff | 1 | + | flights | 1 | + | passengers | 1 | When I run the error report phase Then An error report is produced And The statistics entry for the submission shows the following information @@ -206,6 +248,13 @@ Feature: Pipeline tests using the flights dataset | record | error | StaffHasNoAirport | 1 | | record | error | FlightHasNoAirport | 2 | | record | error | AirportHasNoStaff | 1 | + And the final entities have the following row counts + | entity_name | row_count | + | country | 1 | + | airport | 3 | + | staff | 3 | + | flights | 3 | + | passengers | 2 | When I run the error report phase Then An error report is produced And The statistics entry for the submission shows the following information diff --git a/tests/features/steps/steps_post_pipeline.py b/tests/features/steps/steps_post_pipeline.py index 6c70174..b3663a6 100644 --- a/tests/features/steps/steps_post_pipeline.py +++ b/tests/features/steps/steps_post_pipeline.py @@ -117,3 +117,16 @@ def check_error_aggregates_persisted(context): processing_location = get_processing_location(context) agg_file = Path(processing_location, "audit", "error_aggregates.parquet") assert agg_file.exists() and agg_file.is_file() + +@then("the final entities have the following row counts") +def check_entity_row_counts(context: Context): + processing_loc = get_processing_location(context) + submission_info = get_submission_info(context) + table: Table = context.table + if table is None: + raise ValueError("No table supplied in step") + for row in table: + record = row.as_dict() + entity_name = record["entity_name"] + expected_count = int(record["row_count"]) + assert expected_count == read_output_parquet(processing_loc, entity_name, "business_rules").shape[0] \ No newline at end of file From 431d97ca9e7cdf12b7be639609e1192a44d3e24d Mon Sep 17 00:00:00 2001 From: stevenhsd <56357022+stevenhsd@users.noreply.github.com> Date: Wed, 23 Sep 2026 09:55:39 +0100 Subject: [PATCH 2/2] style: sort linting issues --- src/dve/core_engine/backends/base/rules.py | 36 +++++++---------- .../core_engine/backends/base/utilities.py | 4 +- .../backends/implementations/duckdb/rules.py | 10 +++-- src/dve/pipeline/pipeline.py | 40 ++++++++++--------- 4 files changed, 44 insertions(+), 46 deletions(-) diff --git a/src/dve/core_engine/backends/base/rules.py b/src/dve/core_engine/backends/base/rules.py index 9fb9e3b..acfe079 100644 --- a/src/dve/core_engine/backends/base/rules.py +++ b/src/dve/core_engine/backends/base/rules.py @@ -394,9 +394,7 @@ def identify_and_remove_orphans( Processes recursively: removes orphans at each level, then processes children. """ - def process_node( - node: HierarchyNode - ): + def process_node(node: HierarchyNode): """Identify orphans and remove in a given node""" issues_found: bool = False if node.parent_entity is None: @@ -420,9 +418,7 @@ def process_node( ) if no_orphs > 0: - self.logger.info( - f"Removing records with missing parent from {node.entity_name}" - ) + self.logger.info(f"Removing records with missing parent from {node.entity_name}") issues_found = True location = list(node.join_fields.values())[0] with BackgroundMessageWriter( @@ -488,9 +484,7 @@ def identify_and_remove_missing_mandatory_groups( Identify that an entity with a mandatory key has at least one valid child record. """ - def process_node( - node: HierarchyNode - ) -> bool: + def process_node(node: HierarchyNode) -> bool: """Identify at least one valid child for a mandatory entity at a given node.""" if node.parent_entity is None or not node.mandatory: return False @@ -520,17 +514,17 @@ def process_node( ), ) _messages = [ - FeedbackMessage( - entity=node.parent_entity, - record=record, # type: ignore - error_location=location, - error_message=node.no_valid_records_error_message, - failure_type="record", - error_type="record", - error_code=node.no_valid_records_error_code, - reporting_field=location, - category="Children missing", - ) + FeedbackMessage( + entity=node.parent_entity, + record=record, # type: ignore + error_location=location, + error_message=node.no_valid_records_error_message, + failure_type="record", + error_type="record", + error_code=node.no_valid_records_error_code, + reporting_field=location, + category="Children missing", + ) for record in missing_children_records ] msg_writer.write_queue.put(_messages) @@ -543,7 +537,7 @@ def process_node( if node.parent_entity and node.mandatory: entity_issues_found[node.parent_entity] = process_node(node) - #entities.update(entities) + # entities.update(entities) return [], entity_issues_found diff --git a/src/dve/core_engine/backends/base/utilities.py b/src/dve/core_engine/backends/base/utilities.py index aa99f63..f55bc88 100644 --- a/src/dve/core_engine/backends/base/utilities.py +++ b/src/dve/core_engine/backends/base/utilities.py @@ -3,7 +3,7 @@ import warnings from collections import deque from collections.abc import Sequence -from typing import Iterator, Optional +from typing import Optional import pyarrow # type: ignore import pyarrow.parquet as pq # type: ignore @@ -144,5 +144,3 @@ def check_if_parquet_file(file_location: URI) -> bool: return True except (pyarrow.ArrowInvalid, pyarrow.ArrowIOError): return False - - diff --git a/src/dve/core_engine/backends/implementations/duckdb/rules.py b/src/dve/core_engine/backends/implementations/duckdb/rules.py index 64abb79..c99b152 100644 --- a/src/dve/core_engine/backends/implementations/duckdb/rules.py +++ b/src/dve/core_engine/backends/implementations/duckdb/rules.py @@ -434,7 +434,11 @@ def identify_orphans( def remove_orphans(self, entities: DuckDBEntities, *, config: OrphanRemoval) -> Iterator: """Method to remove identified orphans in the orphan tracker entity.""" - orphan_rel = entities[ORPHANED_RECORD_ENTITY_NAME].filter(f"entity_name = '{config.entity_name}'").set_alias("orphan") + orphan_rel = ( + entities[ORPHANED_RECORD_ENTITY_NAME] + .filter(f"entity_name = '{config.entity_name}'") + .set_alias("orphan") + ) filtered_rel = ( entities[config.entity_name] .set_alias(config.entity_name) @@ -447,9 +451,7 @@ def remove_orphans(self, entities: DuckDBEntities, *, config: OrphanRemoval) -> entities[config.entity_name] = filtered_rel - return duckdb_rel_to_dictionaries( - orphan_rel - ) + return duckdb_rel_to_dictionaries(orphan_rel) def check_mandatory_group( self, entities: DuckDBEntities, *, config: GroupIdentification diff --git a/src/dve/pipeline/pipeline.py b/src/dve/pipeline/pipeline.py index 8bc85d5..ee5a6bc 100644 --- a/src/dve/pipeline/pipeline.py +++ b/src/dve/pipeline/pipeline.py @@ -649,7 +649,6 @@ def apply_business_rules( # pylint: disable=R0914,R0915 entity_hierarchy, key_fields, ) - _, grp_issues_1 = self.step_implementations.identify_and_remove_missing_mandatory_groups( # type: ignore working_directory, @@ -657,7 +656,6 @@ def apply_business_rules( # pylint: disable=R0914,R0915 entity_hierarchy, key_fields, ) - # Perform a second time incase the mandatory groups result in new orphans _, orph_issues_2 = self.step_implementations.identify_and_remove_orphans( # type: ignore @@ -666,12 +664,18 @@ def apply_business_rules( # pylint: disable=R0914,R0915 entity_hierarchy, key_fields, ) - + entity_issues: dict[EntityName, bool] = { entity: any( - val for val in (orph_issues_1.get(entity, False), grp_issues_1.get(entity, False), orph_issues_2.get(entity, False))) + val + for val in ( + orph_issues_1.get(entity, False), + grp_issues_1.get(entity, False), + orph_issues_2.get(entity, False), + ) + ) for entity in orph_issues_1.keys() - } + } unchanged_entities: list[EntityName] = [] for entity_name, entity in entity_manager.entities.items(): @@ -686,13 +690,13 @@ def apply_business_rules( # pylint: disable=R0914,R0915 entity_name, ), ) - + entity_manager.entities[entity_name] = self.step_implementations.read_parquet( # type: ignore - final_projection - ) + final_projection + ) else: unchanged_entities.append(entity_name) - + for entity_name in unchanged_entities: self._logger.info(f"Moving {entity_name} from temp_business_rules to business_rules") final_projection = fh.move_resource( @@ -700,27 +704,27 @@ def apply_business_rules( # pylint: disable=R0914,R0915 self.processed_files_path, submission_info.submission_id, "temp_business_rules", - entity_name + entity_name, ), target_uri=fh.joinuri( self.processed_files_path, submission_info.submission_id, "business_rules", - entity_name + entity_name, ), - overwrite=True + overwrite=True, ) entity_manager.entities[entity_name] = self.step_implementations.read_parquet( # type: ignore final_projection ) - - fh.remove_prefix(fh.joinuri( - self.processed_files_path, - submission_info.submission_id, - "temp_business_rules")) - + fh.remove_prefix( + fh.joinuri( + self.processed_files_path, submission_info.submission_id, "temp_business_rules" + ) + ) + submission_status.number_of_records = self.get_entity_count( entity=entity_manager.entities[f"""Original{rules.global_variables.get( 'entity',