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
18 changes: 9 additions & 9 deletions grand/backends/_gremlin.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,10 +4,7 @@

from typing import Hashable, Collection

import pandas as pd
from gremlin_python.structure.graph import Graph
from gremlin_python.process.graph_traversal import __, GraphTraversalSource
from gremlin_python.driver.driver_remote_connection import DriverRemoteConnection

from .backend import Backend

Expand Down Expand Up @@ -164,7 +161,7 @@ def add_edge(self, u: Hashable, v: Hashable, metadata: dict):
try:
self.get_edge_by_id(u, v)
e = self._g.V().has(ID, u).outE().as_("e").inV().has(ID, v).select("e")
except IndexError:
except KeyError:
if not self.has_node(u):
self.add_node(u, {})
if not self.has_node(v):
Expand Down Expand Up @@ -231,17 +228,20 @@ def get_edge_by_id(self, u: Hashable, v: Hashable):
dict: Metadata associated with this edge

"""
return (
properties = (
self._g.V()
.has(ID, u)
.outE()
.as_("e")
.inV()
.has(ID, v)
.select("e")
.properties()
.valueMap()
.toList()
)[0]
)
if not properties:
raise KeyError((u, v))
return _node_to_metadata(properties[0])

def get_node_neighbors(
self, u: Hashable, include_metadata: bool = False
Expand Down Expand Up @@ -287,7 +287,7 @@ def get_node_predecessors(
"""
if include_metadata:
return {
e["source"]: e
e["source"]: _node_to_metadata(e["properties"])
for e in (
self._g.V()
.has(ID, u)
Expand All @@ -299,7 +299,7 @@ def get_node_predecessors(
.toList()
)
}
return self._g.V().out().has(ID, u).values(ID).toList()
return self._g.V().has(ID, u).in_().values(ID).toList()

def get_node_count(self) -> int:
"""
Expand Down
85 changes: 85 additions & 0 deletions grand/backends/test_gremlin.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,85 @@
import importlib

import pytest

pytest.importorskip("gremlin_python")
GremlinBackend = importlib.import_module("grand.backends._gremlin").GremlinBackend


class RecordingTraversal:
def __init__(self, result, calls=None):
self.result = result
self.calls = calls if calls is not None else []

def __getattr__(self, name):
def step(*args):
self.calls.append((name, args))
return self

return step

def toList(self):
self.calls.append(("toList", ()))
return self.result


def test_predecessors_start_from_requested_node():
traversal = RecordingTraversal(["parent"])
backend = GremlinBackend(traversal)

assert backend.get_node_predecessors("child") == ["parent"]
assert [(name, args) for name, args in traversal.calls] == [
("V", ()),
("has", ("__id", "child")),
("in_", ()),
("values", ("__id",)),
("toList", ()),
]


def test_predecessors_with_metadata_return_edge_metadata():
traversal = RecordingTraversal(
[{"source": "parent", "target": "child", "properties": {"weight": 2}}]
)
backend = GremlinBackend(traversal)

assert backend.get_node_predecessors("child", include_metadata=True) == {
"parent": {"weight": 2}
}


@pytest.mark.parametrize("metadata", [{}, {"weight": 2}])
def test_get_edge_returns_metadata_dictionary(metadata):
traversal = RecordingTraversal([metadata])
backend = GremlinBackend(traversal)

assert backend.get_edge_by_id("source", "target") == metadata
assert ("valueMap", ()) in traversal.calls


def test_missing_edge_raises_key_error():
backend = GremlinBackend(RecordingTraversal([]))

with pytest.raises(KeyError):
backend.get_edge_by_id("source", "target")


def test_add_edge_updates_existing_edge():
traversal = RecordingTraversal([{}])
backend = GremlinBackend(traversal)

backend.add_edge("source", "target", {"weight": 2})

assert ("addE", ("__edge",)) not in traversal.calls
assert ("property", ("weight", 2)) in traversal.calls


def test_add_edge_creates_missing_edge(monkeypatch):
traversal = RecordingTraversal([])
backend = GremlinBackend(traversal)
monkeypatch.setattr(backend, "get_edge_by_id", lambda *args: (_ for _ in ()).throw(KeyError()))
monkeypatch.setattr(backend, "has_node", lambda node: True)

backend.add_edge("source", "target", {})

assert ("addE", ("__edge",)) in traversal.calls
Loading