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
64 changes: 40 additions & 24 deletions grand/backends/_sqlbackend.py
Original file line number Diff line number Diff line change
Expand Up @@ -711,6 +711,9 @@ def out_degrees(self, nbunch=None):

"""

if not self._directed:
return self._undirected_degrees(nbunch)

if nbunch is None:
where_clause = None
elif isinstance(nbunch, (list, tuple)):
Expand All @@ -721,18 +724,11 @@ def out_degrees(self, nbunch=None):
# single node:
where_clause = self._edge_table.c[self._edge_source_key] == str(nbunch)

if self._directed:
query = (
select(self._edge_table.c[self._edge_source_key], func.count())
.select_from(self._edge_table)
.group_by(self._edge_table.c[self._edge_source_key])
)
else:
query = (
select(self._edge_table.c[self._edge_source_key], func.count())
.select_from(self._edge_table)
.group_by(self._edge_table.c[self._edge_source_key])
)
query = (
select(self._edge_table.c[self._edge_source_key], func.count())
.select_from(self._edge_table)
.group_by(self._edge_table.c[self._edge_source_key])
)

if where_clause is not None:
query = query.where(where_clause)
Expand All @@ -755,6 +751,9 @@ def in_degrees(self, nbunch=None):

"""

if not self._directed:
return self._undirected_degrees(nbunch)

if nbunch is None:
where_clause = None
elif isinstance(nbunch, (list, tuple)):
Expand All @@ -765,18 +764,11 @@ def in_degrees(self, nbunch=None):
# single node:
where_clause = self._edge_table.c[self._edge_target_key] == str(nbunch)

if self._directed:
query = (
select(self._edge_table.c[self._edge_target_key], func.count())
.select_from(self._edge_table)
.group_by(self._edge_table.c[self._edge_target_key])
)
else:
query = (
select(self._edge_table.c[self._edge_target_key], func.count())
.select_from(self._edge_table)
.group_by(self._edge_table.c[self._edge_target_key])
)
query = (
select(self._edge_table.c[self._edge_target_key], func.count())
.select_from(self._edge_table)
.group_by(self._edge_table.c[self._edge_target_key])
)

if where_clause is not None:
query = query.where(where_clause)
Expand All @@ -787,6 +779,30 @@ def in_degrees(self, nbunch=None):
return results.get(nbunch, 0)
return results

def _undirected_degrees(self, nbunch=None):
endpoints = select(
self._edge_table.c[self._edge_source_key].label("node")
).union_all(
select(self._edge_table.c[self._edge_target_key].label("node"))
).subquery()
query = select(endpoints.c.node, func.count()).group_by(endpoints.c.node)

requested = None
if isinstance(nbunch, (list, tuple)):
requested = list(nbunch)
query = query.where(
endpoints.c.node.in_([str(node) for node in requested])
)
elif nbunch is not None:
query = query.where(endpoints.c.node == str(nbunch))

results = {row[0]: row[1] for row in self._connection.execute(query)}
if requested is not None:
return {node: results.get(str(node), 0) for node in requested}
if nbunch is not None:
return results.get(str(nbunch), 0)
return results

def ingest_from_edgelist_dataframe(
self, edgelist: pd.DataFrame, source_column: str, target_column: str
) -> dict:
Expand Down
28 changes: 28 additions & 0 deletions grand/backends/test_sql_transactions.py
Original file line number Diff line number Diff line change
Expand Up @@ -327,3 +327,31 @@ def test_legacy_sql_edge_remains_readable(tmp_path):
assert backend.get_edge_by_id("A", "B") == {"old": True}
assert backend.get_edge_count() == 1
backend.close()


def test_undirected_sql_degrees_aggregate_both_edge_orientations(tmp_path):
backend = SQLBackend(
db_url=f"sqlite:///{tmp_path / 'graph.db'}", directed=False
)
backend.add_edges_from([("A", "B"), ("C", "A")])

expected = {"A": 2, "B": 1, "C": 1}
assert backend.out_degrees() == expected
assert backend.in_degrees() == expected
backend.close()


def test_undirected_sql_degrees_include_requested_zero_degree_nodes(tmp_path):
backend = SQLBackend(
db_url=f"sqlite:///{tmp_path / 'graph.db'}", directed=False
)
backend.add_nodes_from([("D", {})])
backend.add_edges_from([("A", "B"), ("C", "A")])

expected = {"A": 2, "B": 1, "C": 1, "D": 0}
assert backend.out_degrees(["A", "B", "C", "D"]) == expected
assert backend.in_degrees(["A", "B", "C", "D"]) == expected
assert backend.out_degrees("D") == 0
assert backend.in_degrees("D") == 0
assert backend.out_degrees([1]) == {1: 0}
backend.close()
Loading