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
37 changes: 26 additions & 11 deletions src/dve/core_engine/backends/base/rules.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,7 @@
from abc import ABC, abstractmethod
from collections import defaultdict
from collections.abc import Iterable, Iterator
from typing import Any, ClassVar, Generic, NoReturn, Optional, TypeVar
from typing import Any, ClassVar, Generic, MutableMapping, NoReturn, Optional, TypeVar
from uuid import uuid4

from typing_extensions import Literal, Protocol, get_type_hints
Expand Down Expand Up @@ -59,6 +59,8 @@
"""A convenience type indicating a mapping from config type to step method."""
Stage = Literal["Pre-filter", "Filter", "Post-filter"]
"""The name of a stage within a rule."""
TempTableName = str
"""temp tables to cache intermediate results"""


class _UnboundStepFunction(Generic[T_contra], Protocol): # pylint: disable=too-few-public-methods
Expand Down Expand Up @@ -132,6 +134,7 @@
):
self.logger = logger or get_logger(type(self).__name__)
"""The `logging.Logger instance for the data contract config."""
self.entity_cache_tracker: MutableMapping[EntityName, TempTableName] = {}

@classmethod
@abstractmethod
Expand Down Expand Up @@ -312,9 +315,7 @@
raise NotImplementedError

@abstractmethod
def identify_orphans(
self, entities: Entities, *, config: OrphanIdentification
) -> Iterable:
def identify_orphans(self, entities: Entities, *, config: OrphanIdentification) -> 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 Down Expand Up @@ -413,9 +414,7 @@
entity=node.entity_name,
record=record, # type: ignore
error_location=location,
error_message=template_object(
node.missing_parent_id_error_message, record
),
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,
Expand All @@ -426,6 +425,8 @@
]
msg_writer.write_queue.put(_messages)

self.cache_entity(node.entity_name, entities)

return len(_messages) > 0

entity_issues_found: dict[EntityName, bool] = {}
Expand All @@ -434,8 +435,6 @@
for node in tree.iterate_root_down():
entity_issues_found[node.entity_name] = process_node(node)

entities.update(entities)

return [], entity_issues_found

def identify_and_remove_missing_mandatory_groups(
Expand Down Expand Up @@ -493,6 +492,7 @@
for record in missing_children_records
]
msg_writer.write_queue.put(_messages)
self.cache_entity(node.parent_entity, entities)
return len(_messages) > 0

entity_issues_found: dict[EntityName, bool] = {}
Expand All @@ -502,8 +502,6 @@
if node.parent_entity and node.mandatory:
entity_issues_found[node.parent_entity] = process_node(node)

# entities.update(entities)

return [], entity_issues_found

# pylint: disable=R0912,R0914
Expand Down Expand Up @@ -852,3 +850,20 @@
def get_entity_count(entity: EntityType) -> int:
"""Method to get count of records in entity"""
raise NotImplementedError()

def cache_entity(self, entity_name: EntityName, entities: Entities):
"""Store the materialised query in memory and update entity to query directly.
If the entity is already cached, the new cache should be created first, then the old one
removed as part of the function (in case the newer cache depends on the older one)."""
raise NotImplementedError()

def _remove_cached_artifact(self, entity_name: EntityName):
"""Delete artifact in memory and clear from the entity cache keeping track.
This should not be used directly as removing artifacts from memory, may lead to some
entities being unable to be processed as their execution plans depend on these artifacts."""
raise NotImplementedError()

def clear_entity_cache(self):
"""Helper method to remove all artifacts and cache trackers at end of processing."""
for entity_name in list(self.entity_cache_tracker):

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

View check run for this annotation

SonarQubeCloud / SonarCloud Code Analysis

Remove this unnecessary `list()` call on an already iterable object.

See more on https://sonarcloud.io/project/issues?id=NHSDigital_data-validation-engine&issues=AaEVmHJ8j9Sw1Y3zD3-p&open=AaEVmHJ8j9Sw1Y3zD3-p&pullRequest=172
self._remove_cached_artifact(entity_name)
Original file line number Diff line number Diff line change
Expand Up @@ -102,10 +102,12 @@ def table_exists(connection: DuckDBPyConnection, table_name: str) -> bool:
"""check if a table exists in a given DuckDBPyConnection"""
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"""
if relation.limit(1).shape[0] > 0:
Expand Down Expand Up @@ -292,7 +294,7 @@ def _ddb_filter_contract_errors(
},
)
.filter(
f"FailureType == 'record' AND Status != 'informational' AND OriginalEntity = '{entity_name}'" # pylint: disable=C0301
f"FailureType == 'record' AND Status != 'informational' AND OriginalEntity = '{entity_name}'" # pylint: disable=C0301
) # pylint: disable=C0301
.select("RecordIndex")
.distinct()
Expand Down
25 changes: 24 additions & 1 deletion src/dve/core_engine/backends/implementations/duckdb/rules.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
"""Business rule definitions for duckdb backend"""

# pylint: disable=R0801
from collections.abc import Callable, Iterable, Iterator
from typing import get_type_hints
Expand Down Expand Up @@ -60,7 +61,10 @@
from dve.core_engine.functions import implementations as functions
from dve.core_engine.message import FeedbackMessage
from dve.core_engine.templating import template_object
from dve.core_engine.type_hints import Messages
from dve.core_engine.type_hints import EntityName, Messages

TempTableName = str
"""temp tables to cache intermediate results"""


@duckdb_get_entity_count
Expand Down Expand Up @@ -432,6 +436,7 @@ def identify_orphans(
"semi",
)
)

filtered_rel = (
entities[config.entity_name]
.set_alias(config.entity_name)
Expand Down Expand Up @@ -582,3 +587,21 @@ def notify(self, entities: DuckDBEntities, *, config: Notification) -> Messages:
)
)
return messages

def cache_entity(self, entity_name: EntityName, entities: DuckDBEntities):
"""Store the materialised query in memory and update entity to query directly.
If the entity is already cached, the new cache should be created first, then the old one
removed as part of the function (in case the newer cache depends on the older one)."""
_tmp_name = f"{entity_name}_{uuid4().hex}"

if entity := entities.get(entity_name): # pylint: disable=W0612
self.connection.sql(f"CREATE OR REPLACE TEMP TABLE {_tmp_name} AS SELECT * FROM entity")
entities[entity_name] = self.connection.table(_tmp_name)

self._remove_cached_artifact(entity_name)

self.entity_cache_tracker[entity_name] = _tmp_name

def _remove_cached_artifact(self, entity_name: EntityName):
if _tbl := self.entity_cache_tracker.pop(entity_name, None):
self.connection.sql(f"DROP TABLE IF EXISTS {_tbl}")
60 changes: 26 additions & 34 deletions src/dve/core_engine/backends/implementations/spark/rules.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
"""Step implementations in Spark."""

# pylint: disable=R0801
from collections.abc import Callable, Iterator
from typing import Optional
Expand Down Expand Up @@ -50,7 +51,7 @@
from dve.core_engine.functions import implementations as functions
from dve.core_engine.message import FeedbackMessage
from dve.core_engine.templating import template_object
from dve.core_engine.type_hints import Messages
from dve.core_engine.type_hints import EntityName, Messages


@spark_get_entity_count
Expand Down Expand Up @@ -343,39 +344,7 @@ def identify_orphans(
self, entities: SparkEntities, *, config: OrphanIdentification
) -> tuple[Messages, int]:
# TODO - adjust this to new setup of identify and remove orphans
source_df: DataFrame = entities[config.entity_name]
source_df = source_df.alias(config.entity_name)
target_df: DataFrame = entities[config.target_name]
target_df = target_df.alias(config.target_name)

key_name = f"key_{uuid4().hex}"
source_df = source_df.withColumn(key_name, sf.expr("uuid()")).alias(config.entity_name)
match_name = f"matched_{uuid4().hex}"
target_df = target_df.withColumn(match_name, lit(1)).alias(config.target_name)

joined_df = (
source_df.join(target_df, on=sf.expr(config.join_condition), how="left")
.groupBy(col(key_name))
.agg(sf.coalesce(sf.sum(col(match_name)) == lit(0), lit(True)).alias("IsOrphaned"))
)

if "IsOrphaned" not in source_df.columns:
result = source_df.join(joined_df, on=[key_name], how="left").drop(key_name)
else:
result = source_df.alias("source").join(
joined_df.alias("joined"),
on=col(f"source.{key_name}") == col(f"joined.{key_name}"),
how="left",
)

columns = {name: col(f"source.{name}") for name in source_df.columns}
columns["IsOrphaned"] = col("source.IsOrphaned") | col("joined.IsOrphaned")
columns.pop(key_name, None)

result = result.select(*[column.alias(name) for name, column in columns.items()])

entities[config.new_entity_name or config.entity_name] = result
return [], 0
raise NotImplementedError

def check_mandatory_group(
self, entities: SparkEntities, *, config: GroupIdentification
Expand Down Expand Up @@ -436,3 +405,26 @@ def notify(self, entities: SparkEntities, *, config: Notification) -> Messages:
)
)
return messages

def cache_entity(self, entity_name: str, entities: SparkEntities):
"""Store the materialised query in memory and update entity to query directly.
If the entity is already cached, the new cache should be created first, then the old one
removed as part of the function (in case the newer cache depends on the older one)."""
if entity_name not in entities:
return

_tmp_name = f"{entity_name}_{uuid4().hex}"

entity = entities[entity_name]
entity.createOrReplaceTempView(_tmp_name)
self.spark_session.sql(f"CACHE TABLE {_tmp_name}")
self.spark_session.sql(f"SELECT count(*) FROM {_tmp_name}")
entity = self.spark_session.table(_tmp_name)
self._remove_cached_artifact(entity_name)
self.entity_cache_tracker[entity_name] = _tmp_name

entities[entity_name] = entity

def _remove_cached_artifact(self, entity_name: EntityName):
if _tbl := self.entity_cache_tracker.pop(entity_name, None):
self.spark_session.sql(f"DROP TABLE IF EXISTS {_tbl}")
4 changes: 2 additions & 2 deletions src/dve/core_engine/configuration/v1/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -134,8 +134,8 @@ def _check_non_root_entities_have_a_defined_parent(self):
"""Check that non root entities have a parent defined."""
if not self.is_root_entity and self.parent_entity is None:
raise ValueError(
'Non-root entity has no defined parent entity. If you intend this to be a root ' \
'entity you must specify `"is_root_entity": true` for the entity. ' \
"Non-root entity has no defined parent entity. If you intend this to be a root "
'entity you must specify `"is_root_entity": true` for the entity. '
'Otherwise you must specify a `"parent_entity": "<EntityName>"` for this entity.'
)
return self
Expand Down
4 changes: 3 additions & 1 deletion src/dve/pipeline/pipeline.py
Original file line number Diff line number Diff line change
Expand Up @@ -756,6 +756,8 @@ def apply_business_rules( # pylint: disable=R0914,R0915
final_projection
)

self.step_implementations.clear_entity_cache() # type: ignore

fh.remove_prefix(
fh.joinuri(
self.processed_files_path, submission_info.submission_id, "temp_business_rules"
Expand All @@ -772,7 +774,7 @@ def apply_business_rules( # pylint: disable=R0914,R0915
self.processed_files_path,
submission_info.submission_id,
"data_contract",
rules.global_variables.get('entity', submission_info.dataset_id)
rules.global_variables.get("entity", submission_info.dataset_id),
)
)
)
Expand Down
1 change: 1 addition & 0 deletions src/dve/reporting/__init__.py
Original file line number Diff line number Diff line change
@@ -1,2 +1,3 @@
"""Error reports module."""

# pylint: disable=R0801
Loading
Loading