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
54 changes: 54 additions & 0 deletions grandcypher/test_types.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,54 @@
import pickle

import networkx as nx

from grandcypher import GrandCypher
from grandcypher.types import AttributeRef, IDRef, EntityRef


def test_attribute_ref_pickle_roundtrip():
ref = AttributeRef('a', 'name')
restored = pickle.loads(pickle.dumps(ref))
assert restored == 'a.name'
assert type(restored) is AttributeRef
assert restored.entity_name == 'a'
assert restored.attribute == 'name'


def test_id_ref_pickle_roundtrip():
ref = IDRef('a')
restored = pickle.loads(pickle.dumps(ref))
assert restored == 'ID(a)'
assert type(restored) is IDRef
assert restored.entity_name == 'a'


def test_entity_ref_pickle_roundtrip():
ref = EntityRef('a')
restored = pickle.loads(pickle.dumps(ref))
assert restored == 'a'
assert type(restored) is EntityRef
assert restored.entity_name == 'a'


def test_query_results_pickle_roundtrip():
host = nx.DiGraph()
host.add_node("x", name="Alice", age=30)
host.add_node("y", name="Bob", age=25)
host.add_edge("x", "y", since=2020)

qry = """
MATCH (A)-[E]->(B)
WHERE A.name == "Alice"
RETURN A, A.name, A.age, ID(A), B
"""

results = GrandCypher(host).run(qry)
restored = pickle.loads(pickle.dumps(results))

assert restored == results
for key in results:
assert type(restored[key]) is type(results[key])
for key in restored:
restored_key_type = type(key)
assert any(type(k) is restored_key_type for k in results)
9 changes: 9 additions & 0 deletions grandcypher/types.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,9 @@ def __new__(cls, entity_name):
instance.entity_name = str(entity_name)
return instance

def __getnewargs__(self):
return (self.entity_name,)


class AttributeRef(str):
"""Node/edge attribute reference, e.g. A.age.
Expand All @@ -25,6 +28,9 @@ def __new__(cls, entity_name, attribute):
instance.attribute = str(attribute)
return instance

def __getnewargs__(self):
return (self.entity_name, self.attribute)


class IDRef(str):
"""Reference to ID(A) in WHERE clauses.
Expand All @@ -36,3 +42,6 @@ def __new__(cls, entity_name):
instance = super().__new__(cls, f"ID({entity_name})")
instance.entity_name = str(entity_name)
return instance

def __getnewargs__(self):
return (self.entity_name,)
Loading