Skip to content

Commit c0e30d7

Browse files
committed
perf: cache entities during orphan and group rejections to avoid recomputing multiple times
1 parent e4dab57 commit c0e30d7

7 files changed

Lines changed: 297 additions & 45 deletions

File tree

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

Lines changed: 24 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -4,7 +4,7 @@
44
from abc import ABC, abstractmethod
55
from collections import defaultdict
66
from collections.abc import Iterable, Iterator
7-
from typing import Any, ClassVar, Generic, NoReturn, Optional, TypeVar
7+
from typing import Any, ClassVar, Generic, MutableMapping, NoReturn, Optional, TypeVar
88
from uuid import uuid4
99

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

6365

6466
class _UnboundStepFunction(Generic[T_contra], Protocol): # pylint: disable=too-few-public-methods
@@ -132,6 +134,7 @@ def __init__( # pylint: disable=unused-argument
132134
):
133135
self.logger = logger or get_logger(type(self).__name__)
134136
"""The `logging.Logger instance for the data contract config."""
137+
self.entity_cache_tracker: MutableMapping[EntityName, TempTableName] = {}
135138

136139
@classmethod
137140
@abstractmethod
@@ -425,6 +428,8 @@ def process_node(node: HierarchyNode) -> bool:
425428
for record in _orph_records
426429
]
427430
msg_writer.write_queue.put(_messages)
431+
432+
self.cache_entity(node.entity_name, entities)
428433

429434
return len(_messages) > 0
430435

@@ -434,8 +439,6 @@ def process_node(node: HierarchyNode) -> bool:
434439
for node in tree.iterate_root_down():
435440
entity_issues_found[node.entity_name] = process_node(node)
436441

437-
entities.update(entities)
438-
439442
return [], entity_issues_found
440443

441444
def identify_and_remove_missing_mandatory_groups(
@@ -493,6 +496,7 @@ def process_node(node: HierarchyNode) -> bool:
493496
for record in missing_children_records
494497
]
495498
msg_writer.write_queue.put(_messages)
499+
self.cache_entity(node.parent_entity, entities)
496500
return len(_messages) > 0
497501

498502
entity_issues_found: dict[EntityName, bool] = {}
@@ -502,8 +506,6 @@ def process_node(node: HierarchyNode) -> bool:
502506
if node.parent_entity and node.mandatory:
503507
entity_issues_found[node.parent_entity] = process_node(node)
504508

505-
# entities.update(entities)
506-
507509
return [], entity_issues_found
508510

509511
# pylint: disable=R0912,R0914
@@ -852,3 +854,20 @@ def filter_data_contract_record_rejections(
852854
def get_entity_count(entity: EntityType) -> int:
853855
"""Method to get count of records in entity"""
854856
raise NotImplementedError()
857+
858+
def cache_entity(self, entity_name: EntityName, entities: Entities):
859+
"""Store the materialised query in memory and update entity to query directly.
860+
If the entity is already cached, the new cache should be created first, then the old one removed
861+
as part of the function (in case the newer cache depends on the older one)."""
862+
raise NotImplementedError()
863+
864+
def _remove_cached_artifact(self, entity_name: EntityName):
865+
"""Delete artifact in memory and clear from the entity cache keeping track.
866+
This should not be used directly as removing artifacts from memory, may lead to some
867+
entities being unable to be processed as their execution plans depend on these artifacts."""
868+
raise NotImplementedError()
869+
870+
def clear_entity_cache(self):
871+
"""Helper method to remove all artifacts and cache trackers at end of processing."""
872+
for entity_name in list(self.entity_cache_tracker):
873+
self._remove_cached_artifact(entity_name)

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

Lines changed: 27 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -17,7 +17,7 @@
1717
BaseStepImplementations,
1818
ColumnAddition,
1919
ColumnRemoval,
20-
SelectColumns,
20+
SelectColumns
2121
)
2222
from dve.core_engine.backends.exceptions import ConstraintError
2323
from dve.core_engine.backends.implementations.duckdb.duckdb_helpers import (
@@ -60,8 +60,10 @@
6060
from dve.core_engine.functions import implementations as functions
6161
from dve.core_engine.message import FeedbackMessage
6262
from dve.core_engine.templating import template_object
63-
from dve.core_engine.type_hints import Messages
63+
from dve.core_engine.type_hints import EntityName, Messages
6464

65+
TempTableName = str
66+
"""temp tables to cache intermediate results"""
6567

6668
@duckdb_get_entity_count
6769
@duckdb_record_index
@@ -377,8 +379,8 @@ def join_header(self, entities: DuckDBEntities, *, config: HeaderJoin) -> Messag
377379
)
378380

379381
entities[config.new_entity_name or config.entity_name] = joined_rel
380-
return []
381-
382+
return []
383+
382384
def identify_orphans(
383385
self,
384386
entities: DuckDBEntities,
@@ -432,6 +434,7 @@ def identify_orphans(
432434
"semi",
433435
)
434436
)
437+
435438
filtered_rel = (
436439
entities[config.entity_name]
437440
.set_alias(config.entity_name)
@@ -441,7 +444,7 @@ def identify_orphans(
441444
"anti",
442445
)
443446
)
444-
447+
445448
entities[config.entity_name] = filtered_rel
446449

447450
return duckdb_rel_to_dictionaries(message_rel)
@@ -477,7 +480,7 @@ def check_mandatory_group(
477480
self.logger.info(
478481
f"Found {_no_valid_children} records with no valid children in {config.entity_name}."
479482
) # pylint: disable=C0301
480-
483+
481484
entities[config.entity_name] = filtered_rel
482485

483486
return duckdb_rel_to_dictionaries(missing_children_rel)
@@ -582,3 +585,21 @@ def notify(self, entities: DuckDBEntities, *, config: Notification) -> Messages:
582585
)
583586
)
584587
return messages
588+
589+
def cache_entity(self,
590+
entity_name: EntityName,
591+
entities: DuckDBEntities):
592+
593+
_tmp_name = f"{entity_name}_{uuid4().hex}"
594+
595+
if entity := entities.get(entity_name):
596+
self.connection.sql(f"CREATE OR REPLACE TEMP TABLE {_tmp_name} AS SELECT * FROM entity")
597+
entities[entity_name] = self.connection.table(_tmp_name)
598+
599+
self._remove_cached_artifact(entity_name)
600+
601+
self.entity_cache_tracker[entity_name] = _tmp_name
602+
603+
def _remove_cached_artifact(self, entity_name: EntityName):
604+
if _tbl := self.entity_cache_tracker.pop(entity_name, None):
605+
self.connection.sql(f"DROP TABLE IF EXISTS {_tbl}")

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

Lines changed: 25 additions & 34 deletions
Original file line numberDiff line numberDiff line change
@@ -50,7 +50,7 @@
5050
from dve.core_engine.functions import implementations as functions
5151
from dve.core_engine.message import FeedbackMessage
5252
from dve.core_engine.templating import template_object
53-
from dve.core_engine.type_hints import Messages
53+
from dve.core_engine.type_hints import EntityName, Messages
5454

5555

5656
@spark_get_entity_count
@@ -343,39 +343,7 @@ def identify_orphans(
343343
self, entities: SparkEntities, *, config: OrphanIdentification
344344
) -> tuple[Messages, int]:
345345
# TODO - adjust this to new setup of identify and remove orphans
346-
source_df: DataFrame = entities[config.entity_name]
347-
source_df = source_df.alias(config.entity_name)
348-
target_df: DataFrame = entities[config.target_name]
349-
target_df = target_df.alias(config.target_name)
350-
351-
key_name = f"key_{uuid4().hex}"
352-
source_df = source_df.withColumn(key_name, sf.expr("uuid()")).alias(config.entity_name)
353-
match_name = f"matched_{uuid4().hex}"
354-
target_df = target_df.withColumn(match_name, lit(1)).alias(config.target_name)
355-
356-
joined_df = (
357-
source_df.join(target_df, on=sf.expr(config.join_condition), how="left")
358-
.groupBy(col(key_name))
359-
.agg(sf.coalesce(sf.sum(col(match_name)) == lit(0), lit(True)).alias("IsOrphaned"))
360-
)
361-
362-
if "IsOrphaned" not in source_df.columns:
363-
result = source_df.join(joined_df, on=[key_name], how="left").drop(key_name)
364-
else:
365-
result = source_df.alias("source").join(
366-
joined_df.alias("joined"),
367-
on=col(f"source.{key_name}") == col(f"joined.{key_name}"),
368-
how="left",
369-
)
370-
371-
columns = {name: col(f"source.{name}") for name in source_df.columns}
372-
columns["IsOrphaned"] = col("source.IsOrphaned") | col("joined.IsOrphaned")
373-
columns.pop(key_name, None)
374-
375-
result = result.select(*[column.alias(name) for name, column in columns.items()])
376-
377-
entities[config.new_entity_name or config.entity_name] = result
378-
return [], 0
346+
raise NotImplementedError
379347

380348
def check_mandatory_group(
381349
self, entities: SparkEntities, *, config: GroupIdentification
@@ -436,3 +404,26 @@ def notify(self, entities: SparkEntities, *, config: Notification) -> Messages:
436404
)
437405
)
438406
return messages
407+
408+
def cache_entity(self, entity_name: str, entities: SparkEntities):
409+
410+
if not entity_name in entities:
411+
return
412+
413+
_tmp_name = f"{entity_name}_{uuid4().hex}"
414+
415+
entity = entities[entity_name]
416+
entity.createOrReplaceTempView(_tmp_name)
417+
self.spark_session.sql(f"CACHE TABLE {_tmp_name}")
418+
self.spark_session.sql(f"SELECT count(*) FROM {_tmp_name}")
419+
entity = self.spark_session.table(_tmp_name)
420+
self._remove_cached_artifact(entity_name)
421+
self.entity_cache_tracker[entity_name] = _tmp_name
422+
423+
entities[entity_name] = entity
424+
425+
def _remove_cached_artifact(self, entity_name: EntityName):
426+
if _tbl := self.entity_cache_tracker.pop(entity_name, None):
427+
self.spark_session.sql(f"DROP TABLE IF EXISTS {_tbl}")
428+
429+

‎src/dve/pipeline/duckdb_pipeline.py‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -65,3 +65,4 @@ def write_file_to_parquet( # type: ignore
6565
return super().write_file_to_parquet(
6666
submission_file_uri, submission_info, output, DuckDBPyRelation
6767
)
68+

‎src/dve/pipeline/pipeline.py‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -755,6 +755,8 @@ def apply_business_rules( # pylint: disable=R0914,R0915
755755
entity_manager.entities[entity_name] = self.step_implementations.read_parquet( # type: ignore
756756
final_projection
757757
)
758+
759+
self.step_implementations.clear_entity_cache()
758760

759761
fh.remove_prefix(
760762
fh.joinuri(

‎tests/test_core_engine/test_backends/test_implementations/test_duckdb/test_rules.py‎

Lines changed: 112 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,9 @@
11
"""Test DuckDB backend steps."""
22

33
# pylint: disable=redefined-outer-name,unused-import,line-too-long
4+
import datetime
5+
from io import StringIO
6+
import json
47
import tempfile
58
from pathlib import Path
69
from typing import Iterator, List, Optional, Set, Tuple, Type
@@ -884,3 +887,112 @@ def test_read_and_write_nested_parquet(nested_typecast_parquet):
884887
"datetimefield": "TIMESTAMP",
885888
"subfield": "STRUCT(id BIGINT, substrfield VARCHAR, subarrayfield DATE[])[]",
886889
}
890+
891+
def test_cache_management():
892+
conn = DUCKDB_STEP_BACKEND.connection
893+
with tempfile.NamedTemporaryFile(mode="w") as tf1, tempfile.NamedTemporaryFile(mode="w") as tf2:
894+
td1 = [
895+
{"greeting": "hi", "num_one": 2, "num_two": 4, "test_date": datetime.date(2020,5,1), "active": True},
896+
{"greeting": "bonjour", "num_one": 3, "num_two": 9, "test_date": datetime.date(2025,7,4), "active": False},
897+
]
898+
tf1.write(json.dumps(td1, default=str))
899+
900+
tf1.seek(0)
901+
902+
td1_schema = {"greeting": "STRING", "num_one": "BIGINT", "num_two" :"BIGINT", "test_date": "DATE", "active": "BOOLEAN"}
903+
904+
td2 = [
905+
{"farewell": "aurevoir", "lots_of_nums": [1,2,3], "nested_field": {"nested_str": "test1", "nested_timestamp": datetime.datetime(2024,3,5,1,2,3)}},
906+
{"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)}},
907+
]
908+
909+
tf2.write(json.dumps(td2, default=str))
910+
911+
tf2.seek(0)
912+
913+
td2_schema = {"farewell": "STRING", "lots_of_nums": "BIGINT[]", "nested_field": "STRUCT(nested_str STRING, nested_timestamp TIMESTAMP)"}
914+
915+
em = EntityManager({})
916+
em.entities["test_one"] = conn.read_json(tf1.name, columns=td1_schema)
917+
em.entities["test_two"] = conn.read_json(tf2.name, columns=td2_schema)
918+
919+
DUCKDB_STEP_BACKEND.cache_entity("test_one", em.entities)
920+
921+
cached_tables = [rw["table_name"] for rw in conn.sql("SELECT table_name from duckdb_tables() WHERE database_name = 'temp'").pl().to_dicts()]
922+
923+
test_one_temp_name = list(filter(lambda x: x.startswith("test_one_"), cached_tables))[0]
924+
925+
assert "test_one" in DUCKDB_STEP_BACKEND.entity_cache_tracker
926+
assert DUCKDB_STEP_BACKEND.entity_cache_tracker["test_one"] == test_one_temp_name
927+
assert sorted(em.entities["test_one"].pl().to_dicts(), key=lambda x: x.get("test_date")) == td1
928+
assert sorted(conn.table(test_one_temp_name).pl().to_dicts(), key=lambda x: x.get("test_date")) == td1
929+
assert "test_two" not in DUCKDB_STEP_BACKEND.entity_cache_tracker
930+
931+
DUCKDB_STEP_BACKEND.cache_entity("test_two", em.entities)
932+
DUCKDB_STEP_BACKEND._remove_cached_artifact("test_one")
933+
934+
cached_tables = [rw["table_name"] for rw in conn.sql("SELECT table_name from duckdb_tables() WHERE database_name = 'temp'").pl().to_dicts()]
935+
test_two_temp_name = list(filter(lambda x: x.startswith("test_two_"), cached_tables))[0]
936+
937+
assert "test_two" in DUCKDB_STEP_BACKEND.entity_cache_tracker
938+
assert DUCKDB_STEP_BACKEND.entity_cache_tracker["test_two"] in cached_tables
939+
assert sorted(em.entities["test_two"].pl().to_dicts(), key=lambda x: x.get("farewell")) == td2
940+
assert sorted(conn.table(test_two_temp_name).pl().to_dicts(), key=lambda x: x.get("farewell")) == td2
941+
assert "test_one" not in DUCKDB_STEP_BACKEND.entity_cache_tracker
942+
943+
DUCKDB_STEP_BACKEND.clear_entity_cache()
944+
945+
cached_tables = [rw["table_name"] for rw in conn.sql("SELECT table_name from duckdb_tables() WHERE database_name = 'temp'").pl().to_dicts()]
946+
947+
assert not DUCKDB_STEP_BACKEND.entity_cache_tracker
948+
assert not any(tbl.startswith("test_one_") for tbl in cached_tables)
949+
assert not any(tbl.startswith("test_two_") for tbl in cached_tables)
950+
951+
def test_cache_management_with_update():
952+
conn = DUCKDB_STEP_BACKEND.connection
953+
with tempfile.NamedTemporaryFile(mode="w") as tf, tempfile.NamedTemporaryFile(mode="w") as ef:
954+
td = [
955+
{"idx": 1, "lots_of_nums": [1,2,3], "nested_field": {"nested_str": "test1", "nested_timestamp": datetime.datetime(2024,3,5,1,2,3)}},
956+
{"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)}},
957+
]
958+
959+
tf.write(json.dumps(td, default=str))
960+
tf.seek(0)
961+
td_schema = {"idx": "BIGINT",
962+
"lots_of_nums": "BIGINT[]",
963+
"nested_field": "STRUCT(nested_str STRING, nested_timestamp TIMESTAMP)"}
964+
965+
966+
em = EntityManager({})
967+
em.entities["test_df"] = conn.read_json(tf.name, columns=td_schema)
968+
969+
DUCKDB_STEP_BACKEND.cache_entity("test_df", em.entities)
970+
971+
cached_tables = [rw["table_name"] for rw in conn.sql("SELECT table_name from duckdb_tables() WHERE database_name = 'temp'").pl().to_dicts()]
972+
973+
first_temp_name = list(filter(lambda x: x.startswith("test_df_"), cached_tables))[0]
974+
975+
assert "test_df" in DUCKDB_STEP_BACKEND.entity_cache_tracker
976+
assert DUCKDB_STEP_BACKEND.entity_cache_tracker["test_df"] == first_temp_name
977+
assert sorted(em.entities["test_df"].pl().to_dicts(), key=lambda x: x.get("idx")) == td
978+
979+
extra_data = [{"idx": 3, "lots_of_nums": [9], "nested_field": {"nested_str": "test3", "nested_timestamp": datetime.datetime(2024,1,9,3,2,1)}}]
980+
ef.write(json.dumps(extra_data, default=str))
981+
ef.seek(0)
982+
em.entities["test_df"] = em.entities["test_df"].union(conn.read_json(ef.name, columns=td_schema))
983+
DUCKDB_STEP_BACKEND.cache_entity("test_df", em.entities)
984+
985+
cached_tables = [rw["table_name"] for rw in conn.sql("SELECT table_name from duckdb_tables() WHERE database_name = 'temp'").pl().to_dicts()]
986+
tables_of_interest = list(filter(lambda x: x.startswith("test_df_"), cached_tables))
987+
assert len(tables_of_interest) == 1
988+
assert first_temp_name not in tables_of_interest
989+
second_temp_name = tables_of_interest[0]
990+
assert sorted(em.entities["test_df"].pl().to_dicts(), key=lambda x: x.get("idx")) == td + extra_data
991+
assert sorted(conn.table(second_temp_name).pl().to_dicts(), key=lambda x: x.get("idx")) == td + extra_data
992+
993+
DUCKDB_STEP_BACKEND.clear_entity_cache()
994+
995+
cached_tables = [rw["table_name"] for rw in conn.sql("SELECT table_name from duckdb_tables() WHERE database_name = 'temp'").pl().to_dicts()]
996+
997+
assert not DUCKDB_STEP_BACKEND.entity_cache_tracker
998+
assert not any(tbl.startswith("test_df_") for tbl in cached_tables)

0 commit comments

Comments
 (0)