diff --git a/src/dve/core_engine/backends/base/rules.py b/src/dve/core_engine/backends/base/rules.py index dbe2d9d..2471426 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 @@ -312,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. @@ -413,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, @@ -426,6 +425,8 @@ def process_node(node: HierarchyNode) -> bool: ] msg_writer.write_queue.put(_messages) + self.cache_entity(node.entity_name, entities) + return len(_messages) > 0 entity_issues_found: dict[EntityName, bool] = {} @@ -434,8 +435,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 +492,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 +502,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 +850,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/duckdb_helpers.py b/src/dve/core_engine/backends/implementations/duckdb/duckdb_helpers.py index 84be049..55db094 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 edb7d71..24910d3 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 @@ -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 @@ -432,6 +436,7 @@ def identify_orphans( "semi", ) ) + filtered_rel = ( entities[config.entity_name] .set_alias(config.entity_name) @@ -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}") diff --git a/src/dve/core_engine/backends/implementations/spark/rules.py b/src/dve/core_engine/backends/implementations/spark/rules.py index 84df52b..5f218ec 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 @@ -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 @@ -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 @@ -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}") diff --git a/src/dve/core_engine/configuration/v1/__init__.py b/src/dve/core_engine/configuration/v1/__init__.py index 421c56e..634cc11 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/pipeline.py b/src/dve/pipeline/pipeline.py index 60e266a..c73d1df 100644 --- a/src/dve/pipeline/pipeline.py +++ b/src/dve/pipeline/pipeline.py @@ -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" @@ -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), ) ) ) diff --git a/src/dve/reporting/__init__.py b/src/dve/reporting/__init__.py index ab78e11..476e0d0 100644 --- a/src/dve/reporting/__init__.py +++ b/src/dve/reporting/__init__.py @@ -1,2 +1,3 @@ """Error reports module.""" + # pylint: disable=R0801 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 ecef834..aae5696 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 6e38c7a..2a79d81 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)