Skip to content
Open
Show file tree
Hide file tree
Changes from 8 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
7 changes: 7 additions & 0 deletions pipeline/workflow/aggregation-helper/.dockerignore
Original file line number Diff line number Diff line change
@@ -0,0 +1,7 @@
.venv
__pycache__
*.pyc
.pytest_cache
.coverage
.git
.gitignore
Original file line number Diff line number Diff line change
Expand Up @@ -236,20 +236,72 @@ def test_run_all(self):
mock_job = MagicMock()
self.mock_executor.execute.return_value = mock_job

jobs = generator.run_all(LinkedEdgeConfig(import_names=["import1", "import2"]))
jobs = generator.run_all(LinkedEdgeConfig(
import_names=["import1", "import2"]
))

self.assertEqual(len(jobs), 3) # Should run 3 queries
self.assertEqual(len(jobs), 3) # Runs the 3 scoped linked edge queries
self.assertEqual(self.mock_executor.execute.call_count, 3)

# Verify queries contain import names and connection id
# Verify queries contain connection id and spanner destination uri
calls = self.mock_executor.execute.call_args_list
for call in calls:
query = call[0][0]
self.assertIn("test-conn", query)
self.assertIn("import1", query)
self.assertIn("import2", query)
self.assertIn("spanner-uri", query)
self.assertIn("dc/base/import1", query) # Since is_base_dc=True

def test_run_topic_list_edges(self):
generator = LinkedEdgeGenerator(self.mock_executor, is_base_dc=True)

mock_job = MagicMock()
self.mock_executor.execute.return_value = mock_job

job = generator.run_topic_list_edges()
self.assertEqual(job, mock_job)
self.mock_executor.execute.assert_called_once()

query = self.mock_executor.execute.call_args[0][0]
self.assertIn("temp_raw_topic_edges", query)
self.assertIn("temp_topic_types", query)
self.assertIn("temp_topic_nodes", query)
self.assertIn("temp_svpg_nodes", query)
self.assertIn("relevantVariableList", query)
self.assertIn("memberList", query)
self.assertIn("STRING_AGG(DISTINCT e.object_id, ',' ORDER BY e.object_id)", query)
self.assertIn("CONCAT(SUBSTR(TRIM(list_value), 1, 16), ':', TO_HEX(SHA256(TRIM(list_value))))", query)
self.assertIn('spanner_options = \'{"table": "Node"}\'', query)
self.assertIn('spanner_options = \'{"table": "Edge"}\'', query)
self.assertIn("dc/base/generated/TopicLists", query)

def test_run_topic_list_edges_not_base_dc(self):
generator = LinkedEdgeGenerator(self.mock_executor, is_base_dc=False)

mock_job = MagicMock()
self.mock_executor.execute.return_value = mock_job

job = generator.run_topic_list_edges()
self.assertEqual(job, mock_job)

query = self.mock_executor.execute.call_args[0][0]
self.assertIn("generated/TopicLists", query)
self.assertNotIn("dc/base/generated/TopicLists", query)

def test_run_linked_member(self):
generator = LinkedEdgeGenerator(self.mock_executor, is_base_dc=True)

mock_job = MagicMock()
self.mock_executor.execute.return_value = mock_job

job = generator.run_linked_member(import_names=["import1"])
self.assertEqual(job, mock_job)

query = self.mock_executor.execute.call_args[0][0]
self.assertIn("temp_topic_types", query)
self.assertIn("temp_topic_nodes", query)
self.assertIn("temp_svpg_nodes", query)
self.assertIn("%/topic/%", query)
self.assertIn("%/svpg/%", query)
self.assertIn("linkedMember", query)


class TestProvenanceSummaryGenerator(unittest.TestCase):
Expand Down
42 changes: 23 additions & 19 deletions pipeline/workflow/aggregation-helper/aggregation/deleter.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,7 +14,6 @@

"""Deletes aggregated data in Spanner using Partitioned DML."""

import concurrent.futures
import logging
from typing import List
from google.cloud import spanner
Expand Down Expand Up @@ -67,24 +66,15 @@ def delete_aggregated_data(self, imports_to_delete: List[str]) -> None:
("KeyValueStore", "DELETE FROM KeyValueStore WHERE type = 'ProvenanceSummary' AND provenance IN UNNEST(@provenances)", "")
]

def _execute_delete(table_name: str, sql: str, extra_desc: str) -> int:
rows = db.execute_partitioned_dml(
sql, params=params, param_types=param_types
)
logging.info(f"Deleted {rows} rows from {table_name} table{extra_desc}.")
return rows

try:
with concurrent.futures.ThreadPoolExecutor(max_workers=len(delete_queries)) as executor:
futures = [
executor.submit(_execute_delete, table, sql, desc)
for table, sql, desc in delete_queries
]
for future in concurrent.futures.as_completed(futures):
future.result() # Propagate any worker thread exceptions to main thread
except Exception as e:
logging.error(f"Failed to execute partitioned DML for deletions: {e}")
raise
for table_name, sql, extra_desc in delete_queries:
try:
rows = db.execute_partitioned_dml(
sql, params=params, param_types=param_types
)
logging.info(f"Deleted {rows} rows from {table_name} table{extra_desc}.")
except Exception as e:
logging.error(f"Failed to execute partitioned DML for deletions on {table_name}: {e}")
raise

def delete_stat_var_group_edges(self) -> int:
"""Deletes all generated StatVarGroup edges across all provenances in Spanner."""
Expand Down Expand Up @@ -115,3 +105,17 @@ def delete_linked_edges(self, imports_to_delete: List[str]) -> int:
rows = self.spanner_database.execute_partitioned_dml(sql, params=params, param_types=param_types)
logging.info(f"Deleted {rows} linked relationship edges for imports: {imports_to_delete}")
return rows

def delete_topic_list_edges(self) -> int:
"""Deletes consolidated topic and peer group list edges from Spanner."""
provenance_name = get_provenance_name("generated/TopicLists", self.is_base_dc)
sql = (
"DELETE FROM Edge "
"WHERE provenance = @provenance "
"AND predicate IN ('relevantVariableList', 'memberList')"
)
params = {"provenance": provenance_name}
param_types = {"provenance": spanner.param_types.STRING}
rows = self.spanner_database.execute_partitioned_dml(sql, params=params, param_types=param_types)
logging.info(f"Deleted {rows} topic and peer group list edges for provenance: {provenance_name}")
return rows
44 changes: 43 additions & 1 deletion pipeline/workflow/aggregation-helper/aggregation/deleter_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -87,12 +87,14 @@ def test_delete_aggregated_data_not_base_dc(self, mock_spanner_client):

@patch('aggregation.deleter.spanner.Client')
def test_delete_aggregated_data_exception_propagates(self, mock_spanner_client):
"""Verifies that an exception raised in a worker thread is re-raised by delete_aggregated_data."""
"""Verifies that an exception raised during partitioned DML is re-raised by delete_aggregated_data."""
mock_db = MagicMock()
mock_db.execute_partitioned_dml.side_effect = RuntimeError("Spanner deletion error")
mock_spanner_client.return_value.instance.return_value.database.return_value = mock_db

deleter = AggregationDeleter("proj", "inst", "db")
with self.assertRaises(RuntimeError):
deleter.delete_aggregated_data(["ImportA"])

@patch('aggregation.deleter.spanner.Client')
def test_delete_stat_var_group_edges(self, mock_spanner_client):
Expand Down Expand Up @@ -127,8 +129,48 @@ def test_delete_linked_edges(self, mock_spanner_client):
self.assertIn("DELETE FROM Edge", sql)
self.assertIn("provenance IN UNNEST(@provenances)", sql)
self.assertIn("'linkedContainedInPlace'", sql)
self.assertIn("'linkedMemberOf'", sql)
self.assertIn("'linkedMember'", sql)
self.assertNotIn("'relevantVariableList'", sql)
self.assertNotIn("'memberList'", sql)
self.assertEqual(params, {"provenances": ["dc/base/generated/ImportA"]})

@patch('aggregation.deleter.spanner.Client')
def test_delete_topic_list_edges_base_dc(self, mock_spanner_client):
mock_db = MagicMock()
mock_spanner_client.return_value.instance.return_value.database.return_value = mock_db

deleter = AggregationDeleter("proj", "inst", "db", is_base_dc=True)
deleter.delete_topic_list_edges()

mock_db.execute_partitioned_dml.assert_called_once()
call_args = mock_db.execute_partitioned_dml.call_args
sql = call_args[0][0]
params = call_args[1]["params"]
self.assertIn("DELETE FROM Edge", sql)
self.assertIn("provenance = @provenance", sql)
self.assertIn("'relevantVariableList'", sql)
self.assertIn("'memberList'", sql)
self.assertEqual(params, {"provenance": "dc/base/generated/TopicLists"})

@patch('aggregation.deleter.spanner.Client')
def test_delete_topic_list_edges_custom_dc(self, mock_spanner_client):
mock_db = MagicMock()
mock_spanner_client.return_value.instance.return_value.database.return_value = mock_db

deleter = AggregationDeleter("proj", "inst", "db", is_base_dc=False)
deleter.delete_topic_list_edges()

mock_db.execute_partitioned_dml.assert_called_once()
call_args = mock_db.execute_partitioned_dml.call_args
sql = call_args[0][0]
params = call_args[1]["params"]
self.assertIn("DELETE FROM Edge", sql)
self.assertIn("provenance = @provenance", sql)
self.assertIn("'relevantVariableList'", sql)
self.assertIn("'memberList'", sql)
self.assertEqual(params, {"provenance": "generated/TopicLists"})


if __name__ == '__main__':
unittest.main()
Original file line number Diff line number Diff line change
Expand Up @@ -302,5 +302,98 @@ def test_linked_edges_multiple_imports(self):
self.assertEqual(len(res_b), 3, "ImportB should have 3 scoped linked edges.")


def test_topic_and_svpg_list_edges(self):
"""Tests materialization of relevantVariableList and memberList edges and literal nodes."""
import_name = 'TopicTest_Import'

# 1. Setup mock topics and SVPGs
self.add_node('dc/topic/Environment', 'Environment', types=['Topic'])
self.add_node('custom/topic/AirQuality', 'Air Quality') # Type provided via typeOf edge
self.add_node('dc/svpg/AgeGroups', 'Age Groups', types=['StatVarPeerGroup'])

self.add_edge('custom/topic/AirQuality', 'typeOf', 'Topic', import_name)
self.add_edge('dc/svpg/AgeGroups', 'typeOf', 'StatVarPeerGroup', import_name)

# Add relevantVariable and member arcs
self.add_edge('dc/topic/Environment', 'relevantVariable', 'Count_Person', import_name)
self.add_edge('dc/topic/Environment', 'relevantVariable', 'dc/topic/Water', import_name)

self.add_edge('custom/topic/AirQuality', 'relevantVariable', 'AirPollution_PM25', import_name)
self.add_edge('custom/topic/AirQuality', 'relevantVariable', 'AirPollution_O3', import_name)

self.add_edge('dc/svpg/AgeGroups', 'member', 'Count_Person_18To64', import_name)
self.add_edge('dc/svpg/AgeGroups', 'member', 'Count_Person_0To17', import_name)

self.flush_to_spanner()

calculations = [
{
"name": "Linked Edges With Topic Lists",
"type": "LINKED_EDGES",
"stage": 1,
"input_imports": [import_name],
"generate_topic_list_edges": True
}
]
res = self.run_orchestrator(calculations=calculations, active_imports=[import_name])
self.assertTrue(res.success)

expected_provenance = f'dc/base/generated/TopicLists' if self.is_base_dc else 'generated/TopicLists'

with self.database.snapshot(multi_use=True) as snapshot:
# Verify Edge records
edge_query = """
SELECT subject_id, predicate, provenance
FROM Edge
WHERE predicate IN ('relevantVariableList', 'memberList')
ORDER BY subject_id
"""
edges = list(snapshot.execute_sql(edge_query))
self.assertEqual(len(edges), 3)
self.assertEqual(tuple(edges[0]), ('custom/topic/AirQuality', 'relevantVariableList', expected_provenance))
self.assertEqual(tuple(edges[1]), ('dc/svpg/AgeGroups', 'memberList', expected_provenance))
self.assertEqual(tuple(edges[2]), ('dc/topic/Environment', 'relevantVariableList', expected_provenance))

# Verify Node records contain the aggregated CSV string values
node_query = """
SELECT n.value
FROM Node n
JOIN Edge e ON n.subject_id = e.object_id
WHERE e.predicate IN ('relevantVariableList', 'memberList')
ORDER BY e.subject_id
"""
nodes = [r[0] for r in snapshot.execute_sql(node_query)]
self.assertEqual(nodes, [
'AirPollution_O3,AirPollution_PM25',
'Count_Person_0To17,Count_Person_18To64',
'Count_Person,dc/topic/Water'
])

def test_topic_list_edges_disabled(self):
"""Tests that disabling generate_topic_list_edges skips topic/SVPG list edge materialization."""
import_name = 'TopicDisabledTest_Import'

self.add_node('dc/topic/Environment', 'Environment', types=['Topic'])
self.add_edge('dc/topic/Environment', 'relevantVariable', 'Count_Person', import_name)
self.flush_to_spanner()

calculations = [
{
"name": "Linked Edges Without Topic Lists",
"type": "LINKED_EDGES",
"stage": 1,
"input_imports": [import_name],
"generate_topic_list_edges": False
}
]
res = self.run_orchestrator(calculations=calculations, active_imports=[import_name])
self.assertTrue(res.success)

with self.database.snapshot() as snapshot:
query = "SELECT count(*) FROM Edge WHERE predicate = 'relevantVariableList'"
count = list(snapshot.execute_sql(query))[0][0]
self.assertEqual(count, 0)


class LinkedEdgeGeneratorCustomDcTest(LinkedEdgeGeneratorIntegrationTest):
is_base_dc = False
Loading