Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
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,71 @@ 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("relevantVariable", query)
self.assertIn("linkedMember", query)


class TestProvenanceSummaryGenerator(unittest.TestCase):
Expand Down
15 changes: 15 additions & 0 deletions pipeline/workflow/aggregation-helper/aggregation/deleter.py
Original file line number Diff line number Diff line change
Expand Up @@ -86,6 +86,7 @@ def _execute_delete(table_name: str, sql: str, extra_desc: str) -> int:
logging.error(f"Failed to execute partitioned DML for deletions: {e}")
raise


def delete_stat_var_group_edges(self) -> int:
"""Deletes all generated StatVarGroup edges across all provenances in Spanner."""
prefix = f"{get_provenance_prefix(self.is_base_dc)}generated/"
Expand Down Expand Up @@ -115,3 +116,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:

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Why do we need this? Dont we already delete all Edges tied to the provenance?

"""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