Skip to content
Merged
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
119 changes: 74 additions & 45 deletions grand/backends/_dynamodb.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand All @@ -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,
)


Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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.
Expand Down Expand Up @@ -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)

Expand Down Expand Up @@ -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 = {}
Expand Down Expand Up @@ -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 = {}
Expand Down Expand Up @@ -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,
}
)
Expand Down
84 changes: 83 additions & 1 deletion grand/backends/test_dynamodb.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand All @@ -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 == []
Loading