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
163 changes: 91 additions & 72 deletions grandcypher/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -338,6 +338,70 @@ def __str__(self):
return f"type({self.argument})"


class AggregationExpression:
name = None

def __init__(self, argument):
self.argument = argument

def evaluate(self, results, group_keys):
grouped_data = {}
for i, value in enumerate(results[self.argument]):
group = tuple(results[key][i] for key in group_keys if key in results)
grouped_data.setdefault(group, []).append(value)

return {
group: self.aggregate(values) for group, values in grouped_data.items()
}

def aggregate(self, values):
raise NotImplementedError

def __str__(self):
return f"{self.name}({self.argument})"


class Count(AggregationExpression):
name = "COUNT"

def aggregate(self, values):
return len(values)


class Sum(AggregationExpression):
name = "SUM"

def aggregate(self, values):
return sum(value or 0 for value in values)


class Avg(AggregationExpression):
name = "AVG"

def aggregate(self, values):
values = [value or 0 for value in values]
return sum(values) / len(values)


class Max(AggregationExpression):
name = "MAX"

def aggregate(self, values):
return max(value if value is not None else -float("inf") for value in values)


class Min(AggregationExpression):
name = "MIN"

def aggregate(self, values):
return min(value if value is not None else float("inf") for value in values)


_AGGREGATION_FUNCTIONS = {
function.name: function for function in (Count, Sum, Avg, Max, Min)
}


def _evaluate_expression(value, match, host, return_edges, scope=None):
if isinstance(value, ExpressionBase):
return value.evaluate(match, host, return_edges, scope)
Expand Down Expand Up @@ -1079,54 +1143,8 @@ def _lookup(self, data_paths: List[str], offset_limit) -> Dict[str, List]:
result[data_path] = list(ret)[offset_limit]
return result

def _format_aggregation_key(self, func, entity):
return f"{func}({entity})"

def aggregate(self, func, results, entity, group_keys):
# Collect data based on group keys
grouped_data = {}
for i in range(len(results[entity])):
group_tuple = tuple(results[key][i] for key in group_keys if key in results)
if group_tuple not in grouped_data:
grouped_data[group_tuple] = []
grouped_data[group_tuple].append(results[entity][i])


def _collate_data(data, func):
# for ["COUNT", "SUM", "AVG"], we treat None as 0
if func in ["COUNT", "SUM", "AVG"]:
collated_data = [
# label: [
(v or 0)
for v in data
]
elif func in ["MAX", "MIN"]:
collated_data = [
v
for v in data
]


return collated_data

# Apply aggregation function
aggregate_results = {}
for group, data in grouped_data.items():
collated_data = _collate_data(data, func)
if func == "COUNT":
aggregate_results[group] = len(collated_data)
elif func == "SUM":
aggregate_results[group] = sum(collated_data)
elif func == "AVG":
aggregate_results[group] = sum(collated_data) / len(collated_data)
elif func == "MAX":
aggregate_results[group] = max([(d if d is not None else -float("inf")) for d in collated_data])
elif func == "MIN":
aggregate_results[group] = min([(d if d is not None else float("inf")) for d in collated_data ])
# aggregate_results = [v for v in aggregate_results.values()]
return aggregate_results

def returns(self, ignore_limit=False):
aggregation_result_keys = []
data_paths = (
self._return_requests
+ list(self._order_by_attributes)
Expand All @@ -1142,26 +1160,25 @@ def returns(self, ignore_limit=False):
group_keys = [
key
for key in results.keys()
if not any(key.endswith(func[1]) for func in self._aggregate_functions)
if not any(key.endswith(func.argument) for func in self._aggregate_functions)
]

aggregated_results = {}
for func, entity in self._aggregate_functions:
aggregated_data = self.aggregate(func, results, entity, group_keys)
aggregation_keys = None
for func in self._aggregate_functions:
aggregated_data = func.evaluate(results, group_keys)
aggregated_values = list(aggregated_data.values())
aggregated_keys = list(aggregated_data.keys())
func_key = self._format_aggregation_key(func, entity)
current_keys = list(aggregated_data.keys())
if aggregation_keys is None:
aggregation_keys = current_keys
elif current_keys != aggregation_keys:
raise ValueError("Aggregation functions produced different groups")
func_key = str(func)
aggregated_results[func_key] = aggregated_values
self._return_requests.append(func_key)
# TODO: the group_keys is the same for all func
# let's have aggregated keys 1st
# then have aggregated values
# so we don't have to repeat the groups key population here
# for i in range(len(gro up_keys)):
# results[group_keys[i]] = [k[i] for k in aggregated_keys]
aggregation_result_keys.append(func_key)
results.update(aggregated_results)
for i in range(len(group_keys)):
results[group_keys[i]] = [k[i] for k in aggregated_keys]
results[group_keys[i]] = [key[i] for key in aggregation_keys or []]

# update the results with the given alias(es)
results = {self._entity2alias.get(k, k): v for k, v in results.items()}
Expand All @@ -1178,7 +1195,7 @@ def returns(self, ignore_limit=False):
return_requests = [
r if isinstance(r, (EntityRef, AttributeRef, IDRef)) else str(r)
for r in self._return_requests
]
] + aggregation_result_keys

# Only include keys that were asked for in `RETURN` in the final results
results = {
Expand All @@ -1203,7 +1220,11 @@ def returns(self, ignore_limit=False):
def _apply_order_by(self, results):
if self._order_by:
sort_lists = [
(results[field], field, direction)
(
results[self._entity2alias.get(field, field)],
self._entity2alias.get(field, field),
direction,
)
for field, direction in self._order_by
]

Expand Down Expand Up @@ -1606,13 +1627,11 @@ def return_clause(self, clause):
alias = self._extract_alias(item)
item = item.children[0] if isinstance(item, Tree) else item
if isinstance(item, Tree) and item.data == "aggregation_function":
func, entity = self._parse_aggregation_token(item)
func = self._parse_aggregation_token(item)
if alias:
self._executors[-1]._entity2alias[
self._executors[-1]._format_aggregation_key(func, entity)
] = alias
self._executors[-1]._aggregation_attributes.add(entity)
self._executors[-1]._aggregate_functions.append((func, entity))
self._executors[-1]._entity2alias[str(func)] = alias
self._executors[-1]._aggregation_attributes.add(func.argument)
self._executors[-1]._aggregate_functions.append(func)
else:
if isinstance(item, ExpressionBase) and not isinstance(
item, (EntityRef, AttributeRef, IDRef)
Expand Down Expand Up @@ -1644,7 +1663,7 @@ def _parse_aggregation_token(self, item: Tree):
if len(item.children) > 2:
entity += "." + str(item.children[2].children[0].value)

return func, entity
return _AGGREGATION_FUNCTIONS[func](entity)

def _extract_alias(self, item: Tree):
"""
Expand Down Expand Up @@ -1677,9 +1696,9 @@ def order_clause(self, order_clause):
isinstance(item.children[0], Tree)
and item.children[0].data == "aggregation_function"
):
func, entity = self._parse_aggregation_token(item.children[0])
field = self._executors[-1]._format_aggregation_key(func, entity)
self._executors[-1]._order_by_attributes.add(entity)
func = self._parse_aggregation_token(item.children[0])
field = str(func)
self._executors[-1]._order_by_attributes.add(func.argument)
else:
field = str(
item.children[0]
Expand Down
77 changes: 77 additions & 0 deletions grandcypher/test_aggregation_classes.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,77 @@
import networkx as nx
import pytest

from . import Avg, Count, GrandCypher, Max, Min, Sum


@pytest.mark.parametrize(
("aggregation", "expected"),
[
(Count, 3),
(Sum, 3),
(Avg, 1),
(Max, 2),
(Min, 1),
],
)
def test_aggregation_expression_preserves_existing_null_behavior(aggregation, expected):
result = aggregation("value").evaluate({"value": [1, None, 2]}, [])

assert result == {(): expected}


def test_aggregation_expression_groups_values():
result = Sum("value").evaluate(
{"group": ["a", "a", "b"], "value": [1, 2, 5]}, ["group"]
)

assert result == {("a",): 3, ("b",): 5}


def test_aggregation_alias_order_and_multiple_functions():
host = nx.DiGraph()
host.add_nodes_from(
[
("a", {"group": "x", "value": 1}),
("b", {"group": "x", "value": 3}),
("c", {"group": "y", "value": 10}),
]
)

result = GrandCypher(host).run(
"MATCH (n) RETURN n.group, SUM(n.value) AS total, AVG(n.value) AS average "
"ORDER BY SUM(n.value) DESC"
)

assert result == {
"n.group": ["y", "x"],
"total": [10, 4],
"average": [10, 2],
}


def test_aggregation_query_can_be_reused():
host = nx.DiGraph()
host.add_nodes_from([("a", {"group": "x"}), ("b", {"group": "x"})])
grand_cypher = GrandCypher(host)
query = "MATCH (n) RETURN n.group, COUNT(n)"

expected = {"n.group": ["x"], "COUNT(n)": [2]}
assert grand_cypher.run(query) == expected
assert grand_cypher.run(query) == expected


def test_aliased_ordered_aggregation_query_can_be_reused():
host = nx.DiGraph()
host.add_nodes_from(
[("a", {"group": "x", "value": 1}), ("b", {"group": "y", "value": 2})]
)
grand_cypher = GrandCypher(host)
query = (
"MATCH (n) RETURN n.group, SUM(n.value) AS total "
"ORDER BY SUM(n.value) DESC"
)

expected = {"n.group": ["y", "x"], "total": [2, 1]}
assert grand_cypher.run(query) == expected
assert grand_cypher.run(query) == expected
Loading