Skip to content
Merged
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
108 changes: 35 additions & 73 deletions src/dve/core_engine/backends/base/rules.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,7 +17,6 @@
)
from dve.core_engine.backends.base.core import get_entity_type
from dve.core_engine.backends.exceptions import render_error
from dve.core_engine.backends.metadata.reporting import ReportingConfig
from dve.core_engine.backends.metadata.rules import (
AbstractStep,
Aggregation,
Expand All @@ -36,7 +35,6 @@
Notification,
OneToOneJoin,
OrphanIdentification,
OrphanRemoval,
ParentMetadata,
RenameEntity,
Rule,
Expand All @@ -47,7 +45,6 @@
)
from dve.core_engine.backends.types import Entities, EntityType, StageSuccessful
from dve.core_engine.configuration.v1.hierarchy import EntityHierarchy, HierarchyNode
from dve.core_engine.constants import ORPHANED_RECORD_ENTITY_NAME
from dve.core_engine.exceptions import CriticalProcessingError
from dve.core_engine.loggers import get_logger
from dve.core_engine.message import FeedbackMessage
Expand Down Expand Up @@ -317,7 +314,7 @@
@abstractmethod
def identify_orphans(
self, entities: Entities, *, config: OrphanIdentification
) -> tuple[Messages, int]:
) -> Iterable:
"""Identify records in an entity which don't have at least one corresponding
match in the target. A new boolean column will be added to `entity` ('IsOrphaned')
indicating whether the condition matched.
Expand All @@ -330,18 +327,6 @@
"""
raise NotImplementedError

@abstractmethod
def remove_orphans(self, entities: Entities, *, config: OrphanRemoval) -> Iterator:
"""
Remove orphaned records from an entity based on the orphans found in
identify_orphans method. Returns a generator objects with the records removed
for generating feedback messages from.

This may not be implemented by some backends.

"""
raise NotImplementedError

@abstractmethod
def check_mandatory_group(self, entities: Entities, *, config: GroupIdentification) -> Iterator:
"""
Expand Down Expand Up @@ -395,83 +380,60 @@
Processes recursively: removes orphans at each level, then processes children.
"""

def process_node(node: HierarchyNode):
def process_node(node: HierarchyNode) -> bool:
"""Identify orphans and remove in a given node"""
issues_found: bool = False
if node.parent_entity is None:
return issues_found

self.logger.info(f"Identifying orphans in {node.entity_name}")
self.logger.info(f"Checking for orphan records 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()
)

_, 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}")
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,
),
location = list(node.join_fields.values())[0]

Check warning on line 395 in src/dve/core_engine/backends/base/rules.py

View check run for this annotation

SonarQubeCloud / SonarCloud Code Analysis

Replace "list(...)[0]" with "next(iter(...))" to avoid materializing the entire iterable.

See more on https://sonarcloud.io/project/issues?id=NHSDigital_data-validation-engine&issues=AaDxsnsp_GVctkdHK2LI&open=AaDxsnsp_GVctkdHK2LI&pullRequest=170
with BackgroundMessageWriter(
working_directory=working_directory,
dve_stage=self.__stage_name__,
key_fields=key_fields,
logger=self.logger,
) as msg_writer:
_orph_records = 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,
),
)
_messages = [
FeedbackMessage(
entity=node.entity_name,
record=record, # type: ignore
error_location=location,
error_message=template_object(
node.missing_parent_id_error_message, record
),
failure_type="record",
error_type="record",
error_code=node.missing_parent_id_error_code,
reporting_field=location,
category="Parent Missing",
)
# 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=template_object(
node.missing_parent_id_error_message, record
),
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
]
)
for record in _orph_records
]
msg_writer.write_queue.put(_messages)

return issues_found
return len(_messages) > 0

entity_issues_found: dict[EntityName, bool] = {}

for tree in entity_hierarchy.entity_trees.values():
for node in tree.iterate_root_down():
entity_issues_found[node.entity_name] = process_node(node)

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

entities.update(entities)

return [], entity_issues_found
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@

"""Helper objects for duckdb data contract implementation"""

import itertools
from collections.abc import Generator, Iterator
from dataclasses import is_dataclass
from datetime import date, datetime, time
Expand Down Expand Up @@ -99,8 +100,11 @@ def __call__(self):

def table_exists(connection: DuckDBPyConnection, table_name: str) -> bool:
"""check if a table exists in a given DuckDBPyConnection"""
return table_name in map(lambda x: x[0], connection.sql("SHOW TABLES").fetchall())
return table_name in get_all_existing_ddb_tables(connection)

def get_all_existing_ddb_tables(connection: DuckDBPyConnection) -> tuple[str]:
"""Fetch all tables available ina given duckdb connection"""
return tuple(itertools.chain.from_iterable(connection.sql("SHOW TABLES").fetchall()))

def relation_is_empty(relation: DuckDBPyRelation) -> bool:
"""Check if a duckdb relation is empty"""
Expand Down
49 changes: 13 additions & 36 deletions src/dve/core_engine/backends/implementations/duckdb/rules.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
"""Business rule definitions for duckdb backend"""
# pylint: disable=R0801
from collections.abc import Callable, Iterator
from collections.abc import Callable, Iterable, Iterator
from typing import get_type_hints
from uuid import uuid4

Expand Down Expand Up @@ -53,11 +53,10 @@
Notification,
OneToOneJoin,
OrphanIdentification,
OrphanRemoval,
SemiJoin,
TableUnion,
)
from dve.core_engine.constants import ORPHANED_RECORD_ENTITY_NAME, RECORD_INDEX_COLUMN_NAME
from dve.core_engine.constants import RECORD_INDEX_COLUMN_NAME
from dve.core_engine.functions import implementations as functions
from dve.core_engine.message import FeedbackMessage
from dve.core_engine.templating import template_object
Expand Down Expand Up @@ -385,7 +384,7 @@ def identify_orphans(
entities: DuckDBEntities,
*,
config: OrphanIdentification,
) -> tuple[Messages, int]:
) -> Iterable:
"""Identify records in an entity which don't have at least one corresponding
match in the target. A new boolean column will be added to `entity` ('IsOrphaned')
indicating whether the condition matched.
Expand All @@ -401,56 +400,34 @@ def identify_orphans(

if relation_is_empty(source_rel):
self.logger.info(f"{config.entity_name} is empty. Skipping orphan check.")
return [], 0
return []

match_name = f"matched_{uuid4().hex}"
target_rel = target_rel.select(
StarExpression(exclude=[]), ConstantExpression(1).alias(match_name)
).set_alias(config.target_name)

pk, _fk = config.join_condition.split("=")

orphaned_rel: DuckDBPyRelation = (
source_rel.join(target_rel, condition=config.join_condition, how="left")
.aggregate(
f"{config.entity_name}.{RECORD_INDEX_COLUMN_NAME}, {config.entity_name}.{config.id}, coalesce(count({match_name}), 0)==0 AS IsOrphaned" # pylint: disable=C0301
)
.filter("IsOrphaned")
.select(
RECORD_INDEX_COLUMN_NAME,
ConstantExpression(config.entity_name).alias("entity_name"),
ConstantExpression(pk.strip().rsplit(".")[1]).alias("pk"),
ColumnExpression(config.id).alias("pk_value"), # type: ignore
)
.unique("*")
.select(RECORD_INDEX_COLUMN_NAME)
.set_alias("orphan")
)
_orph_records: tuple[int] = orphaned_rel.count(RECORD_INDEX_COLUMN_NAME).fetchone() # type: ignore # pylint: disable=C0301
if _orph_records:
_no_orphans = _orph_records[0]
if entities.get(ORPHANED_RECORD_ENTITY_NAME) is not None:
entities[ORPHANED_RECORD_ENTITY_NAME] = entities[ORPHANED_RECORD_ENTITY_NAME].union(
orphaned_rel
)
else:
entities[ORPHANED_RECORD_ENTITY_NAME] = orphaned_rel
else:
_no_orphans = 0
self.logger.info(f"Found {_no_orphans} orphaned records in {config.entity_name}.")

return [], _no_orphans
if relation_is_empty(orphaned_rel):
self.logger.info(
f"Found 0 orphan records between {config.entity_name} and {config.target_name}"
)
return []

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")
)
message_rel = (
entities[config.entity_name]
.set_alias(config.entity_name)
.join(
orphan_rel,
orphaned_rel,
f"{config.entity_name}.{RECORD_INDEX_COLUMN_NAME} = orphan.{RECORD_INDEX_COLUMN_NAME}", # pylint: disable=C0301
"semi",
)
Expand All @@ -459,7 +436,7 @@ def remove_orphans(self, entities: DuckDBEntities, *, config: OrphanRemoval) ->
entities[config.entity_name]
.set_alias(config.entity_name)
.join(
orphan_rel,
orphaned_rel,
f"{config.entity_name}.{RECORD_INDEX_COLUMN_NAME} = orphan.{RECORD_INDEX_COLUMN_NAME}", # pylint: disable=C0301
"anti",
)
Expand Down
10 changes: 0 additions & 10 deletions src/dve/core_engine/backends/implementations/spark/rules.py
Original file line number Diff line number Diff line change
Expand Up @@ -43,7 +43,6 @@
Notification,
OneToOneJoin,
OrphanIdentification,
OrphanRemoval,
SelectColumns,
SemiJoin,
TableUnion,
Expand Down Expand Up @@ -378,15 +377,6 @@ def identify_orphans(
entities[config.new_entity_name or config.entity_name] = result
return [], 0

def remove_orphans(
self,
entities: SparkEntities,
*,
config: OrphanRemoval,
) -> Iterator:
# TODO - implement for spark
raise NotImplementedError

def check_mandatory_group(
self, entities: SparkEntities, *, config: GroupIdentification
) -> Iterator:
Expand Down
3 changes: 0 additions & 3 deletions src/dve/core_engine/constants.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,3 @@
CONTRACT_ERROR_VALUE_FIELD_NAME: str = "__error_value"
"""The name of the field that can be used to extract the field value that caused
a pydantic validation error"""

ORPHANED_RECORD_ENTITY_NAME: str = "orphaned_record_tracker"
"""Name of entity to keep track of records where there is a missing parent record"""
3 changes: 2 additions & 1 deletion tests/features/steps/steps_post_pipeline.py
Original file line number Diff line number Diff line change
Expand Up @@ -129,4 +129,5 @@ def check_entity_row_counts(context: Context):
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]
output_df = read_output_parquet(processing_loc, entity_name, "business_rules")
assert expected_count == output_df.shape[0], output_df
Original file line number Diff line number Diff line change
Expand Up @@ -40,7 +40,6 @@
SemiJoin,
TableUnion,
)
from dve.core_engine.constants import ORPHANED_RECORD_ENTITY_NAME
from dve.core_engine.configuration.v1.hierarchy import (
EntityHierarchy, HierarchyNode
)
Expand Down Expand Up @@ -622,7 +621,7 @@ def test_identify_orphan_record_single_entity(self):
)

rules = DuckDBStepImplementations(connection=cnn)
_msgs = rules.identify_orphans(
msgs = rules.identify_orphans(
mod_entities.entities,
config=OrphanIdentification(
id="flight_id",
Expand All @@ -631,9 +630,7 @@ def test_identify_orphan_record_single_entity(self):
join_condition="passengers.flight_id = flights.flight_id"
)
)
result = mod_entities[ORPHANED_RECORD_ENTITY_NAME]
assert result.count("*").fetchone()[0] == 1 # type: ignore
assert result.select("entity_name").unique("*").count("*").fetchone()[0] == 1 # type: ignore
assert len(list(msgs)) == 1

def test_identify_and_remove_orphans(self):
with duckdb.connect() as cnn:
Expand Down
Loading