From c0e30d7e7a189aa31d08bb99e1095516f26cdd57 Mon Sep 17 00:00:00 2001 From: stevenhsd <56357022+stevenhsd@users.noreply.github.com> Date: Fri, 2 Oct 2026 19:05:27 +0100 Subject: [PATCH 1/4] perf: cache entities during orphan and group rejections to avoid recomputing multiple times --- src/dve/core_engine/backends/base/rules.py | 29 ++++- .../backends/implementations/duckdb/rules.py | 33 +++++- .../backends/implementations/spark/rules.py | 59 ++++----- src/dve/pipeline/duckdb_pipeline.py | 1 + src/dve/pipeline/pipeline.py | 2 + .../test_duckdb/test_rules.py | 112 ++++++++++++++++++ .../test_spark/test_rules.py | 106 +++++++++++++++++ 7 files changed, 297 insertions(+), 45 deletions(-) diff --git a/src/dve/core_engine/backends/base/rules.py b/src/dve/core_engine/backends/base/rules.py index dbe2d9d9..79691a57 100644 --- a/src/dve/core_engine/backends/base/rules.py +++ b/src/dve/core_engine/backends/base/rules.py @@ -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 @@ -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 @@ -132,6 +134,7 @@ def __init__( # pylint: disable=unused-argument ): 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 @@ -425,6 +428,8 @@ def process_node(node: HierarchyNode) -> bool: for record in _orph_records ] msg_writer.write_queue.put(_messages) + + self.cache_entity(node.entity_name, entities) return len(_messages) > 0 @@ -434,8 +439,6 @@ def process_node(node: HierarchyNode) -> bool: 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( @@ -493,6 +496,7 @@ def process_node(node: HierarchyNode) -> bool: 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] = {} @@ -502,8 +506,6 @@ def process_node(node: HierarchyNode) -> bool: 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 @@ -852,3 +854,20 @@ def filter_data_contract_record_rejections( 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): + self._remove_cached_artifact(entity_name) diff --git a/src/dve/core_engine/backends/implementations/duckdb/rules.py b/src/dve/core_engine/backends/implementations/duckdb/rules.py index edb7d710..b1f6ad3e 100644 --- a/src/dve/core_engine/backends/implementations/duckdb/rules.py +++ b/src/dve/core_engine/backends/implementations/duckdb/rules.py @@ -17,7 +17,7 @@ BaseStepImplementations, ColumnAddition, ColumnRemoval, - SelectColumns, + SelectColumns ) from dve.core_engine.backends.exceptions import ConstraintError from dve.core_engine.backends.implementations.duckdb.duckdb_helpers import ( @@ -60,8 +60,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 @duckdb_record_index @@ -377,8 +379,8 @@ def join_header(self, entities: DuckDBEntities, *, config: HeaderJoin) -> Messag ) entities[config.new_entity_name or config.entity_name] = joined_rel - return [] - + return [] + def identify_orphans( self, entities: DuckDBEntities, @@ -432,6 +434,7 @@ def identify_orphans( "semi", ) ) + filtered_rel = ( entities[config.entity_name] .set_alias(config.entity_name) @@ -441,7 +444,7 @@ def identify_orphans( "anti", ) ) - + entities[config.entity_name] = filtered_rel return duckdb_rel_to_dictionaries(message_rel) @@ -477,7 +480,7 @@ def check_mandatory_group( self.logger.info( f"Found {_no_valid_children} records with no valid children in {config.entity_name}." ) # pylint: disable=C0301 - + entities[config.entity_name] = filtered_rel return duckdb_rel_to_dictionaries(missing_children_rel) @@ -582,3 +585,21 @@ def notify(self, entities: DuckDBEntities, *, config: Notification) -> Messages: ) ) return messages + + def cache_entity(self, + entity_name: EntityName, + entities: DuckDBEntities): + + _tmp_name = f"{entity_name}_{uuid4().hex}" + + if entity := entities.get(entity_name): + 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}") diff --git a/src/dve/core_engine/backends/implementations/spark/rules.py b/src/dve/core_engine/backends/implementations/spark/rules.py index 84df52b0..5d4205e5 100644 --- a/src/dve/core_engine/backends/implementations/spark/rules.py +++ b/src/dve/core_engine/backends/implementations/spark/rules.py @@ -50,7 +50,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 @@ -343,39 +343,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 @@ -436,3 +404,26 @@ def notify(self, entities: SparkEntities, *, config: Notification) -> Messages: ) ) return messages + + def cache_entity(self, entity_name: str, entities: SparkEntities): + + if not entity_name 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}") + + diff --git a/src/dve/pipeline/duckdb_pipeline.py b/src/dve/pipeline/duckdb_pipeline.py index b6a643f9..5dcd2822 100644 --- a/src/dve/pipeline/duckdb_pipeline.py +++ b/src/dve/pipeline/duckdb_pipeline.py @@ -65,3 +65,4 @@ def write_file_to_parquet( # type: ignore return super().write_file_to_parquet( submission_file_uri, submission_info, output, DuckDBPyRelation ) + diff --git a/src/dve/pipeline/pipeline.py b/src/dve/pipeline/pipeline.py index 60e266a1..8f1eb23f 100644 --- a/src/dve/pipeline/pipeline.py +++ b/src/dve/pipeline/pipeline.py @@ -755,6 +755,8 @@ def apply_business_rules( # pylint: disable=R0914,R0915 entity_manager.entities[entity_name] = self.step_implementations.read_parquet( # type: ignore final_projection ) + + self.step_implementations.clear_entity_cache() fh.remove_prefix( fh.joinuri( diff --git a/tests/test_core_engine/test_backends/test_implementations/test_duckdb/test_rules.py b/tests/test_core_engine/test_backends/test_implementations/test_duckdb/test_rules.py index ecef834b..aae56967 100644 --- a/tests/test_core_engine/test_backends/test_implementations/test_duckdb/test_rules.py +++ b/tests/test_core_engine/test_backends/test_implementations/test_duckdb/test_rules.py @@ -1,6 +1,9 @@ """Test DuckDB backend steps.""" # pylint: disable=redefined-outer-name,unused-import,line-too-long +import datetime +from io import StringIO +import json import tempfile from pathlib import Path from typing import Iterator, List, Optional, Set, Tuple, Type @@ -884,3 +887,112 @@ def test_read_and_write_nested_parquet(nested_typecast_parquet): "datetimefield": "TIMESTAMP", "subfield": "STRUCT(id BIGINT, substrfield VARCHAR, subarrayfield DATE[])[]", } + +def test_cache_management(): + conn = DUCKDB_STEP_BACKEND.connection + with tempfile.NamedTemporaryFile(mode="w") as tf1, tempfile.NamedTemporaryFile(mode="w") as tf2: + td1 = [ + {"greeting": "hi", "num_one": 2, "num_two": 4, "test_date": datetime.date(2020,5,1), "active": True}, + {"greeting": "bonjour", "num_one": 3, "num_two": 9, "test_date": datetime.date(2025,7,4), "active": False}, + ] + tf1.write(json.dumps(td1, default=str)) + + tf1.seek(0) + + td1_schema = {"greeting": "STRING", "num_one": "BIGINT", "num_two" :"BIGINT", "test_date": "DATE", "active": "BOOLEAN"} + + td2 = [ + {"farewell": "aurevoir", "lots_of_nums": [1,2,3], "nested_field": {"nested_str": "test1", "nested_timestamp": datetime.datetime(2024,3,5,1,2,3)}}, + {"farewell": "bye", "lots_of_nums": [4,5,6,7,8], "nested_field": {"nested_str": "test2", "nested_timestamp": datetime.datetime(2022,8,12,10,12,14)}}, + ] + + tf2.write(json.dumps(td2, default=str)) + + tf2.seek(0) + + td2_schema = {"farewell": "STRING", "lots_of_nums": "BIGINT[]", "nested_field": "STRUCT(nested_str STRING, nested_timestamp TIMESTAMP)"} + + em = EntityManager({}) + em.entities["test_one"] = conn.read_json(tf1.name, columns=td1_schema) + em.entities["test_two"] = conn.read_json(tf2.name, columns=td2_schema) + + DUCKDB_STEP_BACKEND.cache_entity("test_one", em.entities) + + cached_tables = [rw["table_name"] for rw in conn.sql("SELECT table_name from duckdb_tables() WHERE database_name = 'temp'").pl().to_dicts()] + + test_one_temp_name = list(filter(lambda x: x.startswith("test_one_"), cached_tables))[0] + + assert "test_one" in DUCKDB_STEP_BACKEND.entity_cache_tracker + assert DUCKDB_STEP_BACKEND.entity_cache_tracker["test_one"] == test_one_temp_name + assert sorted(em.entities["test_one"].pl().to_dicts(), key=lambda x: x.get("test_date")) == td1 + assert sorted(conn.table(test_one_temp_name).pl().to_dicts(), key=lambda x: x.get("test_date")) == td1 + assert "test_two" not in DUCKDB_STEP_BACKEND.entity_cache_tracker + + DUCKDB_STEP_BACKEND.cache_entity("test_two", em.entities) + DUCKDB_STEP_BACKEND._remove_cached_artifact("test_one") + + cached_tables = [rw["table_name"] for rw in conn.sql("SELECT table_name from duckdb_tables() WHERE database_name = 'temp'").pl().to_dicts()] + test_two_temp_name = list(filter(lambda x: x.startswith("test_two_"), cached_tables))[0] + + assert "test_two" in DUCKDB_STEP_BACKEND.entity_cache_tracker + assert DUCKDB_STEP_BACKEND.entity_cache_tracker["test_two"] in cached_tables + assert sorted(em.entities["test_two"].pl().to_dicts(), key=lambda x: x.get("farewell")) == td2 + assert sorted(conn.table(test_two_temp_name).pl().to_dicts(), key=lambda x: x.get("farewell")) == td2 + assert "test_one" not in DUCKDB_STEP_BACKEND.entity_cache_tracker + + DUCKDB_STEP_BACKEND.clear_entity_cache() + + cached_tables = [rw["table_name"] for rw in conn.sql("SELECT table_name from duckdb_tables() WHERE database_name = 'temp'").pl().to_dicts()] + + assert not DUCKDB_STEP_BACKEND.entity_cache_tracker + assert not any(tbl.startswith("test_one_") for tbl in cached_tables) + assert not any(tbl.startswith("test_two_") for tbl in cached_tables) + +def test_cache_management_with_update(): + conn = DUCKDB_STEP_BACKEND.connection + with tempfile.NamedTemporaryFile(mode="w") as tf, tempfile.NamedTemporaryFile(mode="w") as ef: + td = [ + {"idx": 1, "lots_of_nums": [1,2,3], "nested_field": {"nested_str": "test1", "nested_timestamp": datetime.datetime(2024,3,5,1,2,3)}}, + {"idx": 2, "lots_of_nums": [4,5,6,7,8], "nested_field": {"nested_str": "test2", "nested_timestamp": datetime.datetime(2022,8,12,10,12,14)}}, + ] + + tf.write(json.dumps(td, default=str)) + tf.seek(0) + td_schema = {"idx": "BIGINT", + "lots_of_nums": "BIGINT[]", + "nested_field": "STRUCT(nested_str STRING, nested_timestamp TIMESTAMP)"} + + + em = EntityManager({}) + em.entities["test_df"] = conn.read_json(tf.name, columns=td_schema) + + DUCKDB_STEP_BACKEND.cache_entity("test_df", em.entities) + + cached_tables = [rw["table_name"] for rw in conn.sql("SELECT table_name from duckdb_tables() WHERE database_name = 'temp'").pl().to_dicts()] + + first_temp_name = list(filter(lambda x: x.startswith("test_df_"), cached_tables))[0] + + assert "test_df" in DUCKDB_STEP_BACKEND.entity_cache_tracker + assert DUCKDB_STEP_BACKEND.entity_cache_tracker["test_df"] == first_temp_name + assert sorted(em.entities["test_df"].pl().to_dicts(), key=lambda x: x.get("idx")) == td + + extra_data = [{"idx": 3, "lots_of_nums": [9], "nested_field": {"nested_str": "test3", "nested_timestamp": datetime.datetime(2024,1,9,3,2,1)}}] + ef.write(json.dumps(extra_data, default=str)) + ef.seek(0) + em.entities["test_df"] = em.entities["test_df"].union(conn.read_json(ef.name, columns=td_schema)) + DUCKDB_STEP_BACKEND.cache_entity("test_df", em.entities) + + cached_tables = [rw["table_name"] for rw in conn.sql("SELECT table_name from duckdb_tables() WHERE database_name = 'temp'").pl().to_dicts()] + tables_of_interest = list(filter(lambda x: x.startswith("test_df_"), cached_tables)) + assert len(tables_of_interest) == 1 + assert first_temp_name not in tables_of_interest + second_temp_name = tables_of_interest[0] + assert sorted(em.entities["test_df"].pl().to_dicts(), key=lambda x: x.get("idx")) == td + extra_data + assert sorted(conn.table(second_temp_name).pl().to_dicts(), key=lambda x: x.get("idx")) == td + extra_data + + DUCKDB_STEP_BACKEND.clear_entity_cache() + + cached_tables = [rw["table_name"] for rw in conn.sql("SELECT table_name from duckdb_tables() WHERE database_name = 'temp'").pl().to_dicts()] + + assert not DUCKDB_STEP_BACKEND.entity_cache_tracker + assert not any(tbl.startswith("test_df_") for tbl in cached_tables) diff --git a/tests/test_core_engine/test_backends/test_implementations/test_spark/test_rules.py b/tests/test_core_engine/test_backends/test_implementations/test_spark/test_rules.py index 6e38c7a9..2a79d81b 100644 --- a/tests/test_core_engine/test_backends/test_implementations/test_spark/test_rules.py +++ b/tests/test_core_engine/test_backends/test_implementations/test_spark/test_rules.py @@ -1,6 +1,7 @@ """Test Spark backend steps.""" # pylint: disable=redefined-outer-name,unused-import,line-too-long +import datetime from pathlib import Path from typing import List, Optional, Set, Tuple, Type @@ -9,6 +10,7 @@ from pyspark.sql.functions import col, lit from pyspark.sql.types import ( ArrayType, + BooleanType, DateType, LongType, Row, @@ -852,3 +854,107 @@ def test_read_and_write_nested_parquet(nested_typecast_parquet): ), ] ) + +def test_cache_management(): + spark = SPARK_STEP_BACKEND._spark_session + td1 = [ + {"greeting": "hi", "num_one": 2, "num_two": 4, "test_date": datetime.date(2020,5,1), "active": True}, + {"greeting": "bonjour", "num_one": 3, "num_two": 9, "test_date": datetime.date(2025,7,4), "active": False}, + ] + td1_schema = StructType([ + StructField("greeting", StringType()), + StructField("num_one", LongType()), + StructField("num_two", LongType()), + StructField("test_date", DateType()), + StructField("active", BooleanType())]) + td2 = [ + {"farewell": "aurevoir", "lots_of_nums": [1,2,3], "nested_field": {"nested_str": "test1", "nested_timestamp": datetime.datetime(2024,3,5,1,2,3)}}, + {"farewell": "bye", "lots_of_nums": [4,5,6,7,8], "nested_field": {"nested_str": "test2", "nested_timestamp": datetime.datetime(2022,8,12,10,12,14)}}, + ] + td2_schema = StructType( + [ + StructField("farewell", StringType()), + StructField("lots_of_nums", ArrayType(LongType())), + StructField("nested_field", StructType([StructField("nested_str", StringType()), StructField("nested_timestamp", TimestampType())])) + ] + ) + em = EntityManager({}) + em.entities["test_one"] = spark.createDataFrame(td1, schema=td1_schema) + em.entities["test_two"] = spark.createDataFrame(td2, schema=td2_schema) + + SPARK_STEP_BACKEND.cache_entity("test_one", em.entities) + + cached_tables = (rw.tableName for rw in spark.sql("SHOW TABLES").collect()) + + test_one_temp_name = list(filter(lambda x: x.startswith("test_one_"), cached_tables))[0] + + assert "test_one" in SPARK_STEP_BACKEND.entity_cache_tracker + assert SPARK_STEP_BACKEND.entity_cache_tracker["test_one"] == test_one_temp_name + assert sorted([rw.asDict(True) for rw in em.entities["test_one"].collect()], key=lambda x: x.get("test_date")) == td1 + assert sorted([rw.asDict(True) for rw in spark.table(test_one_temp_name).collect()], key=lambda x: x.get("test_date")) == td1 + assert "test_two" not in SPARK_STEP_BACKEND.entity_cache_tracker + + SPARK_STEP_BACKEND.cache_entity("test_two", em.entities) + SPARK_STEP_BACKEND._remove_cached_artifact("test_one") + + cached_tables = list(rw.tableName for rw in spark.sql("SHOW TABLES").collect()) + test_two_temp_name = list(filter(lambda x: x.startswith("test_two_"), cached_tables))[0] + + assert "test_two" in SPARK_STEP_BACKEND.entity_cache_tracker + assert SPARK_STEP_BACKEND.entity_cache_tracker["test_two"] in cached_tables + assert sorted([rw.asDict(True) for rw in em.entities["test_two"].collect()], key=lambda x: x.get("farewell")) == td2 + assert sorted([rw.asDict(True) for rw in spark.table(test_two_temp_name).collect()], key=lambda x: x.get("farewell")) == td2 + assert "test_one" not in SPARK_STEP_BACKEND.entity_cache_tracker + + SPARK_STEP_BACKEND.clear_entity_cache() + + cached_tables = list(rw.tableName for rw in spark.sql("SHOW TABLES").collect()) + + assert not SPARK_STEP_BACKEND.entity_cache_tracker + assert not any(tbl.startswith("test_one_") for tbl in cached_tables) + assert not any(tbl.startswith("test_two_") for tbl in cached_tables) + +def test_cache_management_with_update(): + spark = SPARK_STEP_BACKEND._spark_session + td = [ + {"idx": 1, "lots_of_nums": [1,2,3], "nested_field": {"nested_str": "test1", "nested_timestamp": datetime.datetime(2024,3,5,1,2,3)}}, + {"idx": 2, "lots_of_nums": [4,5,6,7,8], "nested_field": {"nested_str": "test2", "nested_timestamp": datetime.datetime(2022,8,12,10,12,14)}}, + ] + td_schema = StructType( + [ + StructField("idx", LongType()), + StructField("lots_of_nums", ArrayType(LongType())), + StructField("nested_field", StructType([StructField("nested_str", StringType()), StructField("nested_timestamp", TimestampType())])) + ] + ) + em = EntityManager({}) + em.entities["test_df"] = spark.createDataFrame(td, schema=td_schema) + + SPARK_STEP_BACKEND.cache_entity("test_df", em.entities) + + cached_tables = (rw.tableName for rw in spark.sql("SHOW TABLES").collect()) + + first_temp_name = list(filter(lambda x: x.startswith("test_df_"), cached_tables))[0] + + assert "test_df" in SPARK_STEP_BACKEND.entity_cache_tracker + assert SPARK_STEP_BACKEND.entity_cache_tracker["test_df"] == first_temp_name + assert sorted([rw.asDict(True) for rw in em.entities["test_df"].collect()], key=lambda x: x.get("idx")) == td + + extra_data = [{"idx": 3, "lots_of_nums": [9], "nested_field": {"nested_str": "test3", "nested_timestamp": datetime.datetime(2024,1,9,3,2,1)}}] + em.entities["test_df"] = em.entities["test_df"].union(spark.createDataFrame(extra_data, schema=td_schema)) + SPARK_STEP_BACKEND.cache_entity("test_df", em.entities) + + cached_tables = (rw.tableName for rw in spark.sql("SHOW TABLES").collect()) + tables_of_interest = list(filter(lambda x: x.startswith("test_df_"), cached_tables)) + assert len(tables_of_interest) == 1 + assert first_temp_name not in tables_of_interest + second_temp_name = tables_of_interest[0] + assert sorted([rw.asDict(True) for rw in em.entities["test_df"].collect()], key=lambda x: x.get("idx")) == td + extra_data + assert sorted([rw.asDict(True) for rw in spark.table(second_temp_name).collect()], key=lambda x: x.get("idx")) == td + extra_data + + SPARK_STEP_BACKEND.clear_entity_cache() + + cached_tables = list(rw.tableName for rw in spark.sql("SHOW TABLES").collect()) + + assert not SPARK_STEP_BACKEND.entity_cache_tracker + assert not any(tbl.startswith("test_df_") for tbl in cached_tables) From 447562f2f698d8bf38dc603a98f74a3184fe4c6c Mon Sep 17 00:00:00 2001 From: stevenhsd <56357022+stevenhsd@users.noreply.github.com> Date: Mon, 5 Oct 2026 09:35:37 +0100 Subject: [PATCH 2/4] style: address linting and typing issues --- src/dve/core_engine/backends/base/rules.py | 24 ++++++------- .../implementations/duckdb/duckdb_helpers.py | 4 ++- .../backends/implementations/duckdb/rules.py | 34 ++++++++++--------- .../backends/implementations/spark/rules.py | 15 ++++---- .../core_engine/configuration/v1/__init__.py | 4 +-- src/dve/pipeline/duckdb_pipeline.py | 1 - src/dve/pipeline/pipeline.py | 6 ++-- src/dve/reporting/__init__.py | 1 + 8 files changed, 45 insertions(+), 44 deletions(-) diff --git a/src/dve/core_engine/backends/base/rules.py b/src/dve/core_engine/backends/base/rules.py index 79691a57..2471426c 100644 --- a/src/dve/core_engine/backends/base/rules.py +++ b/src/dve/core_engine/backends/base/rules.py @@ -315,9 +315,7 @@ def join_header(self, entities: Entities, *, config: HeaderJoin) -> Messages: 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. @@ -416,9 +414,7 @@ def process_node(node: HierarchyNode) -> bool: 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, @@ -428,7 +424,7 @@ def process_node(node: HierarchyNode) -> bool: for record in _orph_records ] msg_writer.write_queue.put(_messages) - + self.cache_entity(node.entity_name, entities) return len(_messages) > 0 @@ -854,19 +850,19 @@ def filter_data_contract_record_rejections( 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).""" + 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.""" + 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): diff --git a/src/dve/core_engine/backends/implementations/duckdb/duckdb_helpers.py b/src/dve/core_engine/backends/implementations/duckdb/duckdb_helpers.py index 84be0496..55db0945 100644 --- a/src/dve/core_engine/backends/implementations/duckdb/duckdb_helpers.py +++ b/src/dve/core_engine/backends/implementations/duckdb/duckdb_helpers.py @@ -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: @@ -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() diff --git a/src/dve/core_engine/backends/implementations/duckdb/rules.py b/src/dve/core_engine/backends/implementations/duckdb/rules.py index b1f6ad3e..f4430447 100644 --- a/src/dve/core_engine/backends/implementations/duckdb/rules.py +++ b/src/dve/core_engine/backends/implementations/duckdb/rules.py @@ -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 @@ -17,7 +18,7 @@ BaseStepImplementations, ColumnAddition, ColumnRemoval, - SelectColumns + SelectColumns, ) from dve.core_engine.backends.exceptions import ConstraintError from dve.core_engine.backends.implementations.duckdb.duckdb_helpers import ( @@ -65,6 +66,7 @@ TempTableName = str """temp tables to cache intermediate results""" + @duckdb_get_entity_count @duckdb_record_index @duckdb_write_parquet @@ -379,8 +381,8 @@ def join_header(self, entities: DuckDBEntities, *, config: HeaderJoin) -> Messag ) entities[config.new_entity_name or config.entity_name] = joined_rel - return [] - + return [] + def identify_orphans( self, entities: DuckDBEntities, @@ -434,7 +436,7 @@ def identify_orphans( "semi", ) ) - + filtered_rel = ( entities[config.entity_name] .set_alias(config.entity_name) @@ -444,7 +446,7 @@ def identify_orphans( "anti", ) ) - + entities[config.entity_name] = filtered_rel return duckdb_rel_to_dictionaries(message_rel) @@ -480,7 +482,7 @@ def check_mandatory_group( self.logger.info( f"Found {_no_valid_children} records with no valid children in {config.entity_name}." ) # pylint: disable=C0301 - + entities[config.entity_name] = filtered_rel return duckdb_rel_to_dictionaries(missing_children_rel) @@ -585,21 +587,21 @@ def notify(self, entities: DuckDBEntities, *, config: Notification) -> Messages: ) ) return messages - - def cache_entity(self, - entity_name: EntityName, - entities: DuckDBEntities): - + + 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): + + 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}") diff --git a/src/dve/core_engine/backends/implementations/spark/rules.py b/src/dve/core_engine/backends/implementations/spark/rules.py index 5d4205e5..e9871a10 100644 --- a/src/dve/core_engine/backends/implementations/spark/rules.py +++ b/src/dve/core_engine/backends/implementations/spark/rules.py @@ -1,4 +1,5 @@ """Step implementations in Spark.""" + # pylint: disable=R0801 from collections.abc import Callable, Iterator from typing import Optional @@ -406,12 +407,14 @@ 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 not entity_name 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}") @@ -419,11 +422,9 @@ def cache_entity(self, entity_name: str, entities: SparkEntities): 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}") - - diff --git a/src/dve/core_engine/configuration/v1/__init__.py b/src/dve/core_engine/configuration/v1/__init__.py index 421c56ee..634cc112 100644 --- a/src/dve/core_engine/configuration/v1/__init__.py +++ b/src/dve/core_engine/configuration/v1/__init__.py @@ -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": ""` for this entity.' ) return self diff --git a/src/dve/pipeline/duckdb_pipeline.py b/src/dve/pipeline/duckdb_pipeline.py index 5dcd2822..b6a643f9 100644 --- a/src/dve/pipeline/duckdb_pipeline.py +++ b/src/dve/pipeline/duckdb_pipeline.py @@ -65,4 +65,3 @@ def write_file_to_parquet( # type: ignore return super().write_file_to_parquet( submission_file_uri, submission_info, output, DuckDBPyRelation ) - diff --git a/src/dve/pipeline/pipeline.py b/src/dve/pipeline/pipeline.py index 8f1eb23f..d9db0d4c 100644 --- a/src/dve/pipeline/pipeline.py +++ b/src/dve/pipeline/pipeline.py @@ -755,8 +755,8 @@ def apply_business_rules( # pylint: disable=R0914,R0915 entity_manager.entities[entity_name] = self.step_implementations.read_parquet( # type: ignore final_projection ) - - self.step_implementations.clear_entity_cache() + + self.step_implementations.clear_entity_cache() # type: ignore fh.remove_prefix( fh.joinuri( @@ -774,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), ) ) ) diff --git a/src/dve/reporting/__init__.py b/src/dve/reporting/__init__.py index ab78e113..476e0d0d 100644 --- a/src/dve/reporting/__init__.py +++ b/src/dve/reporting/__init__.py @@ -1,2 +1,3 @@ """Error reports module.""" + # pylint: disable=R0801 From 97fbf6e93f2d5bbc8b0a7e27333b1ee179642027 Mon Sep 17 00:00:00 2001 From: stevenhsd <56357022+stevenhsd@users.noreply.github.com> Date: Mon, 5 Oct 2026 09:44:57 +0100 Subject: [PATCH 3/4] style: address sonarqube issues --- src/dve/core_engine/backends/base/rules.py | 2 +- src/dve/core_engine/backends/implementations/spark/rules.py | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/src/dve/core_engine/backends/base/rules.py b/src/dve/core_engine/backends/base/rules.py index 2471426c..b8e17fa0 100644 --- a/src/dve/core_engine/backends/base/rules.py +++ b/src/dve/core_engine/backends/base/rules.py @@ -865,5 +865,5 @@ def _remove_cached_artifact(self, entity_name: EntityName): 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): + for entity_name in self.entity_cache_tracker: self._remove_cached_artifact(entity_name) diff --git a/src/dve/core_engine/backends/implementations/spark/rules.py b/src/dve/core_engine/backends/implementations/spark/rules.py index e9871a10..5f218ec6 100644 --- a/src/dve/core_engine/backends/implementations/spark/rules.py +++ b/src/dve/core_engine/backends/implementations/spark/rules.py @@ -410,7 +410,7 @@ 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 not entity_name in entities: + if entity_name not in entities: return _tmp_name = f"{entity_name}_{uuid4().hex}" From f3d402ac29bb455420f2ecda706d84a30b75c9ca Mon Sep 17 00:00:00 2001 From: stevenhsd <56357022+stevenhsd@users.noreply.github.com> Date: Tue, 6 Oct 2026 10:45:55 +0100 Subject: [PATCH 4/4] fix: revert change to address sonarqube attempted fix - need iterable copy of dictionary to loop through --- src/dve/core_engine/backends/base/rules.py | 2 +- src/dve/core_engine/backends/implementations/duckdb/rules.py | 2 +- src/dve/pipeline/pipeline.py | 2 +- 3 files changed, 3 insertions(+), 3 deletions(-) diff --git a/src/dve/core_engine/backends/base/rules.py b/src/dve/core_engine/backends/base/rules.py index b8e17fa0..2471426c 100644 --- a/src/dve/core_engine/backends/base/rules.py +++ b/src/dve/core_engine/backends/base/rules.py @@ -865,5 +865,5 @@ def _remove_cached_artifact(self, entity_name: EntityName): def clear_entity_cache(self): """Helper method to remove all artifacts and cache trackers at end of processing.""" - for entity_name in self.entity_cache_tracker: + for entity_name in list(self.entity_cache_tracker): self._remove_cached_artifact(entity_name) diff --git a/src/dve/core_engine/backends/implementations/duckdb/rules.py b/src/dve/core_engine/backends/implementations/duckdb/rules.py index f4430447..24910d33 100644 --- a/src/dve/core_engine/backends/implementations/duckdb/rules.py +++ b/src/dve/core_engine/backends/implementations/duckdb/rules.py @@ -594,7 +594,7 @@ def cache_entity(self, entity_name: EntityName, entities: DuckDBEntities): 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 + 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) diff --git a/src/dve/pipeline/pipeline.py b/src/dve/pipeline/pipeline.py index d9db0d4c..c73d1dfd 100644 --- a/src/dve/pipeline/pipeline.py +++ b/src/dve/pipeline/pipeline.py @@ -756,7 +756,7 @@ def apply_business_rules( # pylint: disable=R0914,R0915 final_projection ) - self.step_implementations.clear_entity_cache() # type: ignore + self.step_implementations.clear_entity_cache() # type: ignore fh.remove_prefix( fh.joinuri(