From bb6d11a698748d49d0bae07b5edd3052dfd3363c Mon Sep 17 00:00:00 2001 From: Jordan Matelsky Date: Fri, 24 Jul 2026 16:11:52 -0400 Subject: [PATCH] perf: query DynamoDB adjacency indexes --- grand/backends/_dynamodb.py | 119 ++++++++++++++++++++------------ grand/backends/test_dynamodb.py | 84 +++++++++++++++++++++- 2 files changed, 157 insertions(+), 46 deletions(-) diff --git a/grand/backends/_dynamodb.py b/grand/backends/_dynamodb.py index ff430e9..5a36f68 100644 --- a/grand/backends/_dynamodb.py +++ b/grand/backends/_dynamodb.py @@ -9,6 +9,10 @@ _DEFAULT_DYNAMODB_URL = "http://localhost:4566" +_EDGE_SOURCE_INDEX = "grand_Source" +_EDGE_TARGET_INDEX = "grand_Target" + + def _dynamo_table_exists(table_name: str, client: boto3.client): """ Check to see if the DynamoDB table already exists. @@ -26,21 +30,39 @@ def _create_dynamo_table( primary_key: str, client, read_write_units: Optional[int] = None, + index_keys: tuple[str, ...] = (), ): if read_write_units is not None: raise NotImplementedError("Non-on-demand billing is not currently supported.") - return client.create_table( - TableName=table_name, - KeySchema=[ - {"AttributeName": primary_key, "KeyType": "HASH"}, # Partition key - # {"AttributeName": "title", "KeyType": "RANGE"}, # Sort key + table_kwargs = { + "TableName": table_name, + "KeySchema": [ + {"AttributeName": primary_key, "KeyType": "HASH"}, ], - AttributeDefinitions=[ + "AttributeDefinitions": [ {"AttributeName": primary_key, "AttributeType": "S"}, - # {"AttributeName": "title", "AttributeType": "S"}, + *[ + {"AttributeName": index_key, "AttributeType": "S"} + for index_key in index_keys + ], ], - BillingMode="PAY_PER_REQUEST", + "BillingMode": "PAY_PER_REQUEST", + } + if index_keys: + table_kwargs["GlobalSecondaryIndexes"] = [ + { + "IndexName": index_name, + "KeySchema": [{"AttributeName": index_key, "KeyType": "HASH"}], + "Projection": {"ProjectionType": "ALL"}, + } + for index_name, index_key in zip( + (_EDGE_SOURCE_INDEX, _EDGE_TARGET_INDEX), index_keys + ) + ] + + return client.create_table( + **table_kwargs, ) @@ -107,7 +129,10 @@ def __init__( if not _dynamo_table_exists(self._edge_table_name, self._client): edge_creation_response = _create_dynamo_table( - self._edge_table_name, self._primary_key, self._resource + self._edge_table_name, + self._primary_key, + self._resource, + index_keys=(self._edge_source_key, self._edge_target_key), ) # Await table creation: if edge_creation_response: @@ -173,6 +198,28 @@ def _scan_table(self, table, scan_kwargs: dict = None): done = start_key is None return results + def _query_edges(self, index_name: str, key: str, value: Hashable): + results = [] + query_kwargs = { + "IndexName": index_name, + "KeyConditionExpression": Key(key).eq(str(value)), + } + while True: + response = self._edge_table.query(**query_kwargs) + results.extend(response.get("Items", [])) + start_key = response.get("LastEvaluatedKey") + if start_key is None: + return results + query_kwargs["ExclusiveStartKey"] = start_key + + def _incident_edges(self, u: Hashable): + outgoing = self._query_edges(_EDGE_SOURCE_INDEX, self._edge_source_key, u) + incoming = self._query_edges(_EDGE_TARGET_INDEX, self._edge_target_key, u) + return { + edge[self._primary_key]: edge + for edge in [*outgoing, *incoming] + }.values() + def all_nodes_as_iterable(self, include_metadata: bool = False) -> Collection: """ Get a generator of all of the nodes in this graph. @@ -233,17 +280,17 @@ def add_edge(self, u: Hashable, v: Hashable, metadata: dict): raise KeyError( f"'{self._edge_source_key}' should not be in metadata. I need that for PK!" ) - metadata[self._edge_source_key] = u + metadata[self._edge_source_key] = str(u) if self._edge_target_key in metadata: raise KeyError( f"'{self._edge_target_key}' should not be in metadata. I need that for PK!" ) - metadata[self._edge_target_key] = v + metadata[self._edge_target_key] = str(v) if not self.has_node(u): - self._node_table.put_item(Item={self._primary_key: u}) + self._node_table.put_item(Item={self._primary_key: str(u)}) if not self.has_node(v): - self._node_table.put_item(Item={self._primary_key: v}) + self._node_table.put_item(Item={self._primary_key: str(v)}) response = self._edge_table.put_item(Item=metadata) @@ -318,25 +365,16 @@ def get_node_neighbors( Generator """ + u = str(u) if self._directed: - # Return only edges for which `u` is the source - res = self._scan_table( - self._edge_table, - { - "FilterExpression": Key(self._primary_key).begins_with(f"__{u}__"), - }, + res = self._query_edges( + _EDGE_SOURCE_INDEX, + self._edge_source_key, + u, ) else: - res = self._scan_table( - self._edge_table, - { - "FilterExpression": ( - Key(self._edge_source_key).eq(u) - | Key(self._edge_target_key).eq(u) - ), - }, - ) + res = self._incident_edges(u) if include_metadata: results = {} @@ -375,25 +413,16 @@ def get_node_predecessors( Generator """ + u = str(u) if self._directed: - # Return only edges for which `u` is the target - res = self._scan_table( - self._edge_table, - { - "FilterExpression": Key(self._edge_target_key).eq(u), - }, + res = self._query_edges( + _EDGE_TARGET_INDEX, + self._edge_target_key, + u, ) else: - res = self._scan_table( - self._edge_table, - { - "FilterExpression": ( - Key(self._edge_source_key).eq(u) - | Key(self._edge_target_key).eq(u) - ), - }, - ) + res = self._incident_edges(u) if include_metadata: results = {} @@ -477,8 +506,8 @@ def ingest_from_edgelist_dataframe( batch_writer.put_item( Item={ self._primary_key: f"__{source}__{target}", - self._edge_source_key: source, - self._edge_target_key: target, + self._edge_source_key: str(source), + self._edge_target_key: str(target), **metadata, } ) diff --git a/grand/backends/test_dynamodb.py b/grand/backends/test_dynamodb.py index db5cbd2..d4febe9 100644 --- a/grand/backends/test_dynamodb.py +++ b/grand/backends/test_dynamodb.py @@ -16,14 +16,23 @@ sys.modules["boto3.dynamodb"] = ModuleType("boto3.dynamodb") sys.modules["boto3.dynamodb.conditions"] = conditions -from ._dynamodb import DynamoDBBackend # noqa: E402 +from ._dynamodb import ( # noqa: E402 + _EDGE_SOURCE_INDEX, + _EDGE_TARGET_INDEX, + DynamoDBBackend, + _create_dynamo_table, +) @pytest.fixture def backend(): backend = DynamoDBBackend.__new__(DynamoDBBackend) backend._primary_key = "ID" + backend._edge_source_key = "Source" + backend._edge_target_key = "Target" + backend._directed = True backend._node_table = Mock() + backend._edge_table = Mock() return backend @@ -49,3 +58,76 @@ def test_has_node_propagates_table_errors(backend): with pytest.raises(RuntimeError, match="DynamoDB unavailable"): backend.has_node("node") + + +def test_edge_table_creation_adds_adjacency_indexes(): + resource = Mock() + + _create_dynamo_table( + "edges", + "ID", + resource, + index_keys=("Source", "Target"), + ) + + kwargs = resource.create_table.call_args.kwargs + assert kwargs["AttributeDefinitions"] == [ + {"AttributeName": "ID", "AttributeType": "S"}, + {"AttributeName": "Source", "AttributeType": "S"}, + {"AttributeName": "Target", "AttributeType": "S"}, + ] + assert kwargs["GlobalSecondaryIndexes"] == [ + { + "IndexName": _EDGE_SOURCE_INDEX, + "KeySchema": [{"AttributeName": "Source", "KeyType": "HASH"}], + "Projection": {"ProjectionType": "ALL"}, + }, + { + "IndexName": _EDGE_TARGET_INDEX, + "KeySchema": [{"AttributeName": "Target", "KeyType": "HASH"}], + "Projection": {"ProjectionType": "ALL"}, + }, + ] + + +def test_directed_neighbors_query_source_index_with_pagination(backend): + backend._edge_table.query.side_effect = [ + { + "Items": [{"ID": "ab", "Source": "A", "Target": "B"}], + "LastEvaluatedKey": {"ID": "ab"}, + }, + {"Items": [{"ID": "ac", "Source": "A", "Target": "C"}]}, + ] + + assert set(backend.get_node_neighbors("A")) == {"B", "C"} + first_query, second_query = backend._edge_table.query.call_args_list + assert first_query.kwargs["IndexName"] == _EDGE_SOURCE_INDEX + assert "FilterExpression" not in first_query.kwargs + assert second_query.kwargs["ExclusiveStartKey"] == {"ID": "ab"} + + +def test_directed_predecessors_query_target_index(backend): + backend._edge_table.query.return_value = { + "Items": [{"ID": "ab", "Source": "A", "Target": "B"}] + } + + assert list(backend.get_node_predecessors("B")) == ["A"] + assert backend._edge_table.query.call_args.kwargs["IndexName"] == ( + _EDGE_TARGET_INDEX + ) + + +def test_undirected_neighbors_query_both_indexes_and_deduplicate(backend): + backend._directed = False + edge = {"ID": "aa", "Source": "A", "Target": "A"} + backend._edge_table.query.side_effect = [ + {"Items": [edge, {"ID": "ab", "Source": "A", "Target": "B"}]}, + {"Items": [edge, {"ID": "ca", "Source": "C", "Target": "A"}]}, + ] + + assert list(backend.get_node_neighbors("A")) == ["A", "B", "C"] + assert [item.kwargs["IndexName"] for item in backend._edge_table.query.call_args_list] == [ + _EDGE_SOURCE_INDEX, + _EDGE_TARGET_INDEX, + ] + assert backend._edge_table.scan.call_args_list == []