Skip to content

Commit 676ac5a

Browse files
committed
fix: ensure missing parent and group rejection records are being removed - add test coverage
1 parent a817759 commit 676ac5a

6 files changed

Lines changed: 196 additions & 116 deletions

File tree

‎src/dve/core_engine/backends/base/rules.py‎

Lines changed: 86 additions & 93 deletions
Original file line numberDiff line numberDiff line change
@@ -387,119 +387,113 @@ def identify_and_remove_orphans(
387387
entities: Entities,
388388
entity_hierarchy: EntityHierarchy,
389389
key_fields: Optional[dict[str, list[str]]] = None,
390-
) -> tuple[Messages, bool]:
390+
) -> tuple[Messages, dict[EntityName, bool]]:
391391
"""
392392
Identifies and removes orphan records by traversing the EntityHierarchy object.
393393
An orphan is a child record whose parent FK does not exist in the parent entity.
394394
Processes recursively: removes orphans at each level, then processes children.
395395
"""
396396

397397
def process_node(
398-
node: HierarchyNode,
399-
orph_messages: Messages | None = None,
400-
processed: bool = False,
398+
node: HierarchyNode
401399
):
402400
"""Identify orphans and remove in a given node"""
401+
issues_found: bool = False
402+
if node.parent_entity is None:
403+
return issues_found
403404

404-
if orph_messages is None:
405-
orph_messages = []
405+
self.logger.info(f"Identifying orphans in {node.entity_name}")
406406

407-
if node.parent_entity is not None:
408-
self.logger.info(f"Identifying orphans in {node.entity_name}")
407+
join_expr = " AND ".join(
408+
f"{node.parent_entity}.{k} = {node.entity_name}.{v}"
409+
for k, v in node.join_fields.items()
410+
)
409411

410-
join_expr = " AND ".join(
411-
f"{node.parent_entity}.{k} = {node.entity_name}.{v}"
412-
for k, v in node.join_fields.items()
413-
)
412+
_, no_orphs = self.identify_orphans(
413+
entities=entities,
414+
config=OrphanIdentification(
415+
id=list(node.join_fields.values())[0],
416+
entity_name=node.entity_name,
417+
target_name=node.parent_entity,
418+
join_condition=join_expr,
419+
),
420+
)
414421

415-
_, no_orphs = self.identify_orphans(
416-
entities=entities,
417-
config=OrphanIdentification(
418-
id=list(node.join_fields.values())[0],
419-
entity_name=node.entity_name,
420-
target_name=node.parent_entity,
421-
join_condition=join_expr,
422-
),
422+
if no_orphs > 0:
423+
self.logger.info(
424+
f"Removing records with missing parent from {node.entity_name}"
423425
)
424-
425-
if no_orphs > 0:
426-
self.logger.info(
427-
f"Removing records with missing parent from {node.entity_name}"
428-
)
429-
processed = True
430-
location = list(node.join_fields.values())[0]
431-
with BackgroundMessageWriter(
432-
working_directory=working_directory,
433-
dve_stage=self.__stage_name__,
434-
key_fields=key_fields,
435-
logger=self.logger,
436-
) as msg_writer:
437-
_orph_records = self.remove_orphans(
438-
entities=entities,
439-
config=OrphanRemoval(
440-
entity_name=node.entity_name,
441-
reporting=ReportingConfig(
442-
emit="record_failure",
443-
code=node.missing_parent_id_error_code,
444-
message=node.missing_parent_id_error_message,
445-
location=location,
446-
),
426+
issues_found = True
427+
location = list(node.join_fields.values())[0]
428+
with BackgroundMessageWriter(
429+
working_directory=working_directory,
430+
dve_stage=self.__stage_name__,
431+
key_fields=key_fields,
432+
logger=self.logger,
433+
) as msg_writer:
434+
_orph_records = self.remove_orphans(
435+
entities=entities,
436+
config=OrphanRemoval(
437+
entity_name=node.entity_name,
438+
reporting=ReportingConfig(
439+
emit="record_failure",
440+
code=node.missing_parent_id_error_code,
441+
message=node.missing_parent_id_error_message,
442+
location=location,
447443
),
448-
)
449-
# moved to batch the write - risky if large number of
450-
msg_writer.write_queue.put(
451-
[
452-
FeedbackMessage(
453-
entity=node.entity_name,
454-
record=record, # type: ignore
455-
error_location=location,
456-
error_message=node.missing_parent_id_error_message,
457-
failure_type="record",
458-
error_type="record",
459-
error_code=node.missing_parent_id_error_code,
460-
reporting_field=location,
461-
category="Parent Missing",
462-
)
463-
for record in _orph_records
464-
]
465-
)
444+
),
445+
)
446+
# moved to batch the write - risky if large number of
447+
msg_writer.write_queue.put(
448+
[
449+
FeedbackMessage(
450+
entity=node.entity_name,
451+
record=record, # type: ignore
452+
error_location=location,
453+
error_message=node.missing_parent_id_error_message,
454+
failure_type="record",
455+
error_type="record",
456+
error_code=node.missing_parent_id_error_code,
457+
reporting_field=location,
458+
category="Parent Missing",
459+
)
460+
for record in _orph_records
461+
]
462+
)
466463

467-
return processed
464+
return issues_found
468465

469-
processed = False
466+
entity_issues_found: dict[EntityName, bool] = {}
470467

471468
for tree in entity_hierarchy.entity_trees.values():
472469
for node in tree.iterate_root_down():
473-
processed = process_node(node)
470+
entity_issues_found[node.entity_name] = process_node(node)
474471

475472
_orph_rel = entities.get(ORPHANED_RECORD_ENTITY_NAME)
476473
if _orph_rel is not None:
477474
del entities[ORPHANED_RECORD_ENTITY_NAME]
478475

479476
entities.update(entities)
480477

481-
return [], processed
478+
return [], entity_issues_found
482479

483480
def identify_and_remove_missing_mandatory_groups(
484481
self,
485482
working_directory: URI,
486483
entities: Entities,
487484
entity_hierarchy: EntityHierarchy,
488485
key_fields: Optional[dict[str, list[str]]] = None,
489-
) -> tuple[Messages, bool]:
486+
) -> tuple[Messages, dict[EntityName, bool]]:
490487
"""
491488
Identify that an entity with a mandatory key has at least one valid child record.
492489
"""
493490

494491
def process_node(
495-
node: HierarchyNode,
496-
processed: bool = False,
492+
node: HierarchyNode
497493
) -> bool:
498494
"""Identify at least one valid child for a mandatory entity at a given node."""
499495
if node.parent_entity is None or not node.mandatory:
500-
return processed
501-
502-
processed = True
496+
return False
503497

504498
self.logger.info(
505499
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(
525519
join_condition=join_expr,
526520
),
527521
)
528-
for record in missing_children_records:
529-
msg_writer.write_queue.put(
530-
[
531-
FeedbackMessage(
532-
entity=node.parent_entity,
533-
record=record, # type: ignore
534-
error_location=location,
535-
error_message=node.no_valid_records_error_message,
536-
failure_type="record",
537-
error_type="record",
538-
error_code=node.no_valid_records_error_code,
539-
reporting_field=location,
540-
category="Children missing",
541-
)
542-
]
543-
)
544-
545-
return processed
522+
_messages = [
523+
FeedbackMessage(
524+
entity=node.parent_entity,
525+
record=record, # type: ignore
526+
error_location=location,
527+
error_message=node.no_valid_records_error_message,
528+
failure_type="record",
529+
error_type="record",
530+
error_code=node.no_valid_records_error_code,
531+
reporting_field=location,
532+
category="Children missing",
533+
)
534+
for record in missing_children_records
535+
]
536+
msg_writer.write_queue.put(_messages)
537+
return len(_messages) > 0
546538

547-
processed = False
539+
entity_issues_found: dict[EntityName, bool] = {}
548540

549541
for tree in entity_hierarchy.entity_trees.values():
550542
for node in tree.iterate_lowest_descendent_up():
551-
processed = process_node(node, processed)
543+
if node.parent_entity and node.mandatory:
544+
entity_issues_found[node.parent_entity] = process_node(node)
552545

553-
entities.update(entities)
546+
#entities.update(entities)
554547

555-
return [], processed
548+
return [], entity_issues_found
556549

557550
# pylint: disable=R0912,R0914
558551
def apply_sync_filters(

‎src/dve/core_engine/backends/base/utilities.py‎

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -3,7 +3,7 @@
33
import warnings
44
from collections import deque
55
from collections.abc import Sequence
6-
from typing import Optional
6+
from typing import Iterator, Optional
77

88
import pyarrow # type: ignore
99
import pyarrow.parquet as pq # type: ignore
@@ -144,3 +144,5 @@ def check_if_parquet_file(file_location: URI) -> bool:
144144
return True
145145
except (pyarrow.ArrowInvalid, pyarrow.ArrowIOError):
146146
return False
147+
148+

‎src/dve/core_engine/backends/implementations/duckdb/rules.py‎

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -434,7 +434,7 @@ def identify_orphans(
434434

435435
def remove_orphans(self, entities: DuckDBEntities, *, config: OrphanRemoval) -> Iterator:
436436
"""Method to remove identified orphans in the orphan tracker entity."""
437-
orphan_rel = entities[ORPHANED_RECORD_ENTITY_NAME].set_alias("orphan")
437+
orphan_rel = entities[ORPHANED_RECORD_ENTITY_NAME].filter(f"entity_name = '{config.entity_name}'").set_alias("orphan")
438438
filtered_rel = (
439439
entities[config.entity_name]
440440
.set_alias(config.entity_name)
@@ -448,7 +448,7 @@ def remove_orphans(self, entities: DuckDBEntities, *, config: OrphanRemoval) ->
448448
entities[config.entity_name] = filtered_rel
449449

450450
return duckdb_rel_to_dictionaries(
451-
orphan_rel.filter(f"entity_name = '{config.entity_name}'")
451+
orphan_rel
452452
)
453453

454454
def check_mandatory_group(

‎src/dve/pipeline/pipeline.py‎

Lines changed: 42 additions & 19 deletions
Original file line numberDiff line numberDiff line change
@@ -643,30 +643,39 @@ def apply_business_rules( # pylint: disable=R0914,R0915
643643
projected
644644
)
645645

646-
_, orph_or_group = self.step_implementations.identify_and_remove_orphans( # type: ignore
646+
_, orph_issues_1 = self.step_implementations.identify_and_remove_orphans( # type: ignore
647647
working_directory,
648648
entity_manager.entities,
649649
entity_hierarchy,
650650
key_fields,
651651
)
652+
652653

653-
_, orph_or_group = self.step_implementations.identify_and_remove_missing_mandatory_groups( # type: ignore
654+
_, grp_issues_1 = self.step_implementations.identify_and_remove_missing_mandatory_groups( # type: ignore
654655
working_directory,
655656
entity_manager.entities,
656657
entity_hierarchy,
657658
key_fields,
658659
)
660+
659661

660662
# Perform a second time incase the mandatory groups result in new orphans
661-
_, orph_or_group = self.step_implementations.identify_and_remove_orphans( # type: ignore
663+
_, orph_issues_2 = self.step_implementations.identify_and_remove_orphans( # type: ignore
662664
working_directory,
663665
entity_manager.entities,
664666
entity_hierarchy,
665667
key_fields,
666668
)
669+
670+
entity_issues: dict[EntityName, bool] = {
671+
entity: any(
672+
val for val in (orph_issues_1.get(entity, False), grp_issues_1.get(entity, False), orph_issues_2.get(entity, False)))
673+
for entity in orph_issues_1.keys()
674+
}
667675

676+
unchanged_entities: list[EntityName] = []
668677
for entity_name, entity in entity_manager.entities.items():
669-
if orph_or_group:
678+
if entity_issues.get(entity_name, False):
670679
self._logger.info(f"Writing {entity_name} out to disk.")
671680
final_projection = self._step_implementations.write_parquet( # type: ignore
672681
entity,
@@ -677,27 +686,41 @@ def apply_business_rules( # pylint: disable=R0914,R0915
677686
entity_name,
678687
),
679688
)
689+
690+
entity_manager.entities[entity_name] = self.step_implementations.read_parquet( # type: ignore
691+
final_projection
692+
)
680693
else:
681-
self._logger.info(f"Moving {entity_name} from temp_business_rules to business_rules")
682-
final_projection = fh.move_resource(
683-
source_uri=fh.joinuri(
684-
self.processed_files_path,
685-
submission_info.submission_id,
686-
"temp_business_rules",
687-
entity_name
688-
),
689-
target_uri=fh.joinuri(
690-
self.processed_files_path,
691-
submission_info.submission_id,
692-
"business_rules",
693-
entity_name
694-
)
695-
)
694+
unchanged_entities.append(entity_name)
695+
696+
for entity_name in unchanged_entities:
697+
self._logger.info(f"Moving {entity_name} from temp_business_rules to business_rules")
698+
final_projection = fh.move_resource(
699+
source_uri=fh.joinuri(
700+
self.processed_files_path,
701+
submission_info.submission_id,
702+
"temp_business_rules",
703+
entity_name
704+
),
705+
target_uri=fh.joinuri(
706+
self.processed_files_path,
707+
submission_info.submission_id,
708+
"business_rules",
709+
entity_name
710+
),
711+
overwrite=True
712+
)
696713

697714
entity_manager.entities[entity_name] = self.step_implementations.read_parquet( # type: ignore
698715
final_projection
699716
)
700717

718+
719+
fh.remove_prefix(fh.joinuri(
720+
self.processed_files_path,
721+
submission_info.submission_id,
722+
"temp_business_rules"))
723+
701724
submission_status.number_of_records = self.get_entity_count(
702725
entity=entity_manager.entities[f"""Original{rules.global_variables.get(
703726
'entity',

0 commit comments

Comments
 (0)