diff --git a/src/dve/core_engine/backends/base/rules.py b/src/dve/core_engine/backends/base/rules.py index a3ed6a4..acfe079 100644 --- a/src/dve/core_engine/backends/base/rules.py +++ b/src/dve/core_engine/backends/base/rules.py @@ -387,90 +387,83 @@ 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. Processes recursively: removes orphans at each level, then processes children. """ - def process_node( - node: HierarchyNode, - orph_messages: Messages | None = None, - processed: bool = False, - ): + def process_node(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}" - ) - 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, - ), + if no_orphs > 0: + 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( + 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 +471,7 @@ def process_node( entities.update(entities) - return [], processed + return [], entity_issues_found def identify_and_remove_missing_mandatory_groups( self, @@ -486,20 +479,15 @@ 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, - ) -> 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 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 +513,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", - ) - ] + _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 - return processed - - 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/implementations/duckdb/rules.py b/src/dve/core_engine/backends/implementations/duckdb/rules.py index 867b262..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].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.filter(f"entity_name = '{config.entity_name}'") - ) + 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 c157b3b..ee5a6bc 100644 --- a/src/dve/pipeline/pipeline.py +++ b/src/dve/pipeline/pipeline.py @@ -643,14 +643,14 @@ 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, @@ -658,15 +658,28 @@ def apply_business_rules( # pylint: disable=R0914,R0915 ) # 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 +690,41 @@ def apply_business_rules( # pylint: disable=R0914,R0915 entity_name, ), ) - 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 - ) + + entity_manager.entities[entity_name] = self.step_implementations.read_parquet( # type: ignore + 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( + 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