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
206 changes: 126 additions & 80 deletions grand/backends/_sqlbackend.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,4 @@
from contextlib import contextmanager
from typing import Hashable, Generator
import time

Expand Down Expand Up @@ -57,6 +58,7 @@ def __init__(
sqlalchemy_kwargs = sqlalchemy_kwargs or {}
self._engine = sqlalchemy.create_engine(db_url, **sqlalchemy_kwargs)
self._connection = self._engine.connect()
self._transaction_depth = 0
self._metadata = sqlalchemy.MetaData()

# Create nodes table
Expand Down Expand Up @@ -102,6 +104,34 @@ def __init__(
tindex = Index("edge_target", target_column)
tindex.create(self._engine, checkfirst=True)

@contextmanager
def _mutation(self):
if self._transaction_depth:
yield
return
try:
yield
self._connection.commit()
except Exception:
self._connection.rollback()
raise

@contextmanager
def transaction(self):
"""Group multiple mutations into one atomic commit."""
outermost = self._transaction_depth == 0
self._transaction_depth += 1
try:
yield self
if outermost:
self._connection.commit()
except Exception:
if outermost:
self._connection.rollback()
raise
finally:
self._transaction_depth -= 1

def is_directed(self) -> bool:
"""
Return True if the backend graph is directed.
Expand Down Expand Up @@ -138,20 +168,21 @@ def add_node(self, node_name: Hashable, metadata: dict) -> Hashable:
Hashable: The ID of this node, as inserted

"""
if self.has_node(node_name):
existing_metadata = self.get_node_by_id(node_name)
existing_metadata.update(metadata)
self._connection.execute(
self._node_table.update().where(
self._node_table.c[self._primary_key] == str(node_name)
),
parameters={"_metadata": existing_metadata},
)
else:
self._connection.execute(
self._node_table.insert(),
parameters={self._primary_key: node_name, "_metadata": metadata},
)
with self._mutation():
if self.has_node(node_name):
existing_metadata = self.get_node_by_id(node_name)
existing_metadata.update(metadata)
self._connection.execute(
self._node_table.update().where(
self._node_table.c[self._primary_key] == str(node_name)
),
parameters={"_metadata": existing_metadata},
)
else:
self._connection.execute(
self._node_table.insert(),
parameters={self._primary_key: node_name, "_metadata": metadata},
)
return node_name

def _insert_empty_node_if_missing(self, node_name: Hashable) -> None:
Expand Down Expand Up @@ -189,7 +220,8 @@ def add_nodes_from(self, nodes_for_adding, **attr):
for node, metadata in nodes_for_adding
]

self._connection.execute(self._node_table.insert(), nodes)
with self._mutation():
self._connection.execute(self._node_table.insert(), nodes)

def _upsert_node(self, node_name: Hashable, metadata: dict) -> Hashable:
"""
Expand Down Expand Up @@ -225,20 +257,19 @@ def remove_node(self, u: Hashable) -> None:
u (Hashable): id of the node
"""

# Remove nodes
statement = delete(self._node_table).where(
self._node_table.c[self._primary_key] == str(u)
)
self._connection.execute(statement)
with self._mutation():
statement = delete(self._node_table).where(
self._node_table.c[self._primary_key] == str(u)
)
self._connection.execute(statement)

# Remove edges for node
statement = delete(self._edge_table).where(
or_(
self._edge_table.c[self._edge_source_key] == str(u),
self._edge_table.c[self._edge_target_key] == str(u)
statement = delete(self._edge_table).where(
or_(
self._edge_table.c[self._edge_source_key] == str(u),
self._edge_table.c[self._edge_target_key] == str(u),
)
)
)
self._connection.execute(statement)
self._connection.execute(statement)

def all_nodes_as_iterable(self, include_metadata: bool = False) -> Generator:
"""
Expand Down Expand Up @@ -302,29 +333,42 @@ def add_edge(self, u: Hashable, v: Hashable, metadata: dict):
"""
pk = f"__{u}__{v}"

self._insert_empty_node_if_missing(u)
self._insert_empty_node_if_missing(v)

try:
self._connection.execute(
self._edge_table.insert(),
parameters={
self._primary_key: pk,
self._edge_source_key: u,
self._edge_target_key: v,
"_metadata": metadata,
},
)
except sqlalchemy.exc.IntegrityError:
# Edge already exists, perform the update:
existing_metadata = self.get_edge_by_id(u, v)
existing_metadata.update(metadata)
self._connection.execute(
self._edge_table.update().where(
self._edge_table.c[self._primary_key] == pk
),
parameters={"_metadata": existing_metadata},
)
with self._mutation():
self._insert_empty_node_if_missing(u)
self._insert_empty_node_if_missing(v)

if self._transaction_depth:
self._connection.execute(
self._edge_table.insert(),
parameters={
self._primary_key: pk,
self._edge_source_key: u,
self._edge_target_key: v,
"_metadata": metadata,
},
)
return pk

try:
with self._connection.begin_nested():
self._connection.execute(
self._edge_table.insert(),
parameters={
self._primary_key: pk,
self._edge_source_key: u,
self._edge_target_key: v,
"_metadata": metadata,
},
)
except sqlalchemy.exc.IntegrityError:
existing_metadata = self.get_edge_by_id(u, v)
existing_metadata.update(metadata)
self._connection.execute(
self._edge_table.update().where(
self._edge_table.c[self._primary_key] == pk
),
parameters={"_metadata": existing_metadata},
)

return pk

Expand All @@ -339,7 +383,8 @@ def add_edges_from(self, ebunch_to_add, **attr):
for u, v, metadata in ebunch_to_add
]

self._connection.execute(self._edge_table.insert(), edges)
with self._mutation():
self._connection.execute(self._edge_table.insert(), edges)

def all_edges_as_iterable(self, include_metadata: bool = False) -> Generator:
"""
Expand Down Expand Up @@ -679,39 +724,39 @@ def ingest_from_edgelist_dataframe(
for source, target, metadata in zip(sources, targets, edge_metadata)
]

if edge_rows:
self._connection.execute(self._edge_table.insert(), edge_rows)
with self._mutation():
if edge_rows:
self._connection.execute(self._edge_table.insert(), edge_rows)

edge_toc = time.time() - edge_tic
edge_toc = time.time() - edge_tic

# now ingest nodes:
node_tic = time.time()
nodes = pd.unique(
pd.concat(
[edgelist[source_column], edgelist[target_column]],
ignore_index=True,
node_tic = time.time()
nodes = pd.unique(
pd.concat(
[edgelist[source_column], edgelist[target_column]],
ignore_index=True,
)
)
)

node_rows = [
{
self._primary_key: str(node),
"_metadata": {},
}
for node in nodes
]
node_rows = [
{
self._primary_key: str(node),
"_metadata": {},
}
for node in nodes
]

if node_rows:
node_insert = self._node_table.insert()
if self._engine.dialect.name == "sqlite":
node_insert = node_insert.prefix_with("OR IGNORE")
self._connection.execute(node_insert, node_rows)
elif self._engine.dialect.name in {"mysql", "mariadb"}:
node_insert = node_insert.prefix_with("IGNORE")
self._connection.execute(node_insert, node_rows)
else:
for node in nodes:
self._insert_empty_node_if_missing(node)
if node_rows:
node_insert = self._node_table.insert()
if self._engine.dialect.name == "sqlite":
node_insert = node_insert.prefix_with("OR IGNORE")
self._connection.execute(node_insert, node_rows)
elif self._engine.dialect.name in {"mysql", "mariadb"}:
node_insert = node_insert.prefix_with("IGNORE")
self._connection.execute(node_insert, node_rows)
else:
for node in nodes:
self._insert_empty_node_if_missing(node)

return {
"node_count": len(nodes),
Expand All @@ -721,7 +766,8 @@ def ingest_from_edgelist_dataframe(
}

def commit(self):
self._connection.commit()
if self._connection.in_transaction():
self._connection.commit()

def close(self):
self._connection.close()
12 changes: 8 additions & 4 deletions grand/backends/test_backends.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,5 @@
from contextlib import nullcontext

import pytest
import os
import pandas as pd
Expand Down Expand Up @@ -480,10 +482,12 @@ def test_node_addition_performance(backend):
def test_get_density_performance(backend):
backend, kwargs = backend
G = Graph(backend=backend(directed=True, **kwargs))
for i in range(1000):
G.nx.add_node(i)
for i in range(1000 - 1):
G.nx.add_edge(i, i + 1)
transaction = getattr(G.backend, "transaction", None)
with transaction() if transaction else nullcontext():
for i in range(1000):
G.nx.add_node(i)
for i in range(1000 - 1):
G.nx.add_edge(i, i + 1)
assert nx.density(G.nx) <= 0.005


Expand Down
92 changes: 92 additions & 0 deletions grand/backends/test_sql_transactions.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,92 @@
import pytest

sqlalchemy = pytest.importorskip("sqlalchemy")

from ._sqlbackend import SQLBackend # noqa: E402


def test_mutations_persist_without_explicit_commit(tmp_path):
db_url = f"sqlite:///{tmp_path / 'graph.db'}"
backend = SQLBackend(db_url=db_url, directed=True)

backend.add_edge("A", "B", {"weight": 1})
backend.close()

reopened = SQLBackend(db_url=db_url, directed=True)
assert set(reopened.all_nodes_as_iterable()) == {"A", "B"}
assert reopened.get_edge_by_id("A", "B") == {"weight": 1}
reopened.close()


def test_remove_node_persists_without_explicit_commit(tmp_path):
db_url = f"sqlite:///{tmp_path / 'graph.db'}"
backend = SQLBackend(db_url=db_url, directed=True)
backend.add_edge("A", "B", {})

backend.remove_node("A")
backend.close()

reopened = SQLBackend(db_url=db_url, directed=True)
assert not reopened.has_node("A")
assert not reopened.has_edge("A", "B")
reopened.close()


def test_existing_edge_update_persists_without_explicit_commit(tmp_path):
db_url = f"sqlite:///{tmp_path / 'graph.db'}"
backend = SQLBackend(db_url=db_url, directed=True)
backend.add_edge("A", "B", {"weight": 1, "kind": "old"})

backend.add_edge("A", "B", {"weight": 2})
backend.close()

reopened = SQLBackend(db_url=db_url, directed=True)
assert reopened.get_edge_by_id("A", "B") == {"weight": 2, "kind": "old"}
reopened.close()


def test_transaction_context_commits_grouped_mutations(tmp_path):
db_url = f"sqlite:///{tmp_path / 'graph.db'}"
backend = SQLBackend(db_url=db_url, directed=True)

with backend.transaction():
backend.add_node("A", {})
backend.add_node("B", {})
backend.add_edge("A", "B", {})
backend.close()

reopened = SQLBackend(db_url=db_url, directed=True)
assert reopened.has_edge("A", "B")
reopened.close()


def test_transaction_context_rolls_back_grouped_mutations(tmp_path):
backend = SQLBackend(db_url=f"sqlite:///{tmp_path / 'graph.db'}")

with pytest.raises(RuntimeError, match="abort"):
with backend.transaction():
backend.add_node("A", {})
raise RuntimeError("abort")

assert not backend.has_node("A")
backend.close()


def test_add_edge_rolls_back_created_nodes_when_edge_insert_fails(tmp_path):
backend = SQLBackend(db_url=f"sqlite:///{tmp_path / 'graph.db'}", directed=True)
original_execute = backend._connection.execute

def fail_edge_insert(statement, *args, **kwargs):
if getattr(statement, "table", None) is backend._edge_table:
raise RuntimeError("edge insert failed")
return original_execute(statement, *args, **kwargs)

backend._connection.execute = fail_edge_insert

with pytest.raises(RuntimeError, match="edge insert failed"):
backend.add_edge("A", "B", {})

backend._connection.execute = original_execute
assert not backend.has_node("A")
assert not backend.has_node("B")
backend.close()
Loading