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
11 changes: 11 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,17 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0

## [Unreleased]

## [0.0.22] - 2026-09-26

### Added

- **`graph_transaction`.** Opens the backend transaction, binds a
`GraphContext` to that handle, and makes it the task-local default so
`Node.create` / `Node.connect` participate. Success commits; any
exception rolls back. Stores without begin/commit/rollback raise
`TransactionUnavailable`. `async_transaction_context` is now an alias —
the old helper opened a transaction but kept writing through the pool.

## [0.0.21] - 2026-09-19

### Added
Expand Down
4 changes: 3 additions & 1 deletion jvspatial/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -61,7 +61,7 @@

# Simplified decorators
from .core.annotations import attribute
from .core.context import GraphContext
from .core.context import GraphContext, TransactionUnavailable, graph_transaction

# Unified entity system
from .core.entities import Edge, Node, Object, Root, Walker
Expand Down Expand Up @@ -129,6 +129,8 @@
"Walker",
"Root",
"GraphContext",
"graph_transaction",
"TransactionUnavailable",
# Mixins
"DeferredSaveMixin",
"deferred_saves_globally_allowed",
Expand Down
6 changes: 6 additions & 0 deletions jvspatial/core/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,11 +7,14 @@

from .context import (
GraphContext,
TransactionUnavailable,
async_graph_context,
async_transaction_context,
clear_default_context,
clear_default_context_global,
get_default_context,
graph_context,
graph_transaction,
reset_default_context,
scoped_default_context,
scoped_default_context_async,
Expand Down Expand Up @@ -84,4 +87,7 @@
"scoped_default_context_async",
"graph_context",
"async_graph_context",
"graph_transaction",
"async_transaction_context",
"TransactionUnavailable",
]
89 changes: 60 additions & 29 deletions jvspatial/core/context.py
Original file line number Diff line number Diff line change
Expand Up @@ -2198,39 +2198,70 @@ async def async_graph_context(database: Optional[Database] = None):
yield ctx


class TransactionUnavailable(RuntimeError):
"""The configured store cannot host a graph transaction."""


def _transaction_database(database: Optional[Any]) -> Any:
"""Unwrap observable/cache wrappers to the adapter that owns connections."""
current = database
seen: set[int] = set()
while current is not None and id(current) not in seen:
seen.add(id(current))
if callable(getattr(current, "begin_transaction", None)):
return current
nested = getattr(current, "inner", None)
if nested is None:
break
current = nested
return database


@asynccontextmanager
async def async_transaction_context(database: Optional[Database] = None):
"""Async context manager for database transactions.
async def graph_transaction(database: Optional[Any] = None):
"""Run Node and Edge writes in one ACID transaction.

Opens the backend transaction, binds a :class:`GraphContext` to that
handle, and makes it the task-local default so ``Node.create`` and
``Node.connect`` participate. Success commits; any exception rolls
back every write in the block.

Captures the transaction object returned by ``begin_transaction()`` and
passes it to ``commit_transaction``/``rollback_transaction`` so that the
MongoDB session handle is not lost between calls.
Requires a backend that exposes ``begin_transaction``,
``commit_transaction``, and ``rollback_transaction`` (Postgres today).

Usage:
async with async_transaction_context(my_db) as ctx:
node = await ctx.create_node(name="Test")
# All operations are automatically committed
async with graph_transaction(db) as ctx:
parent = await Folder.create(name="inbox")
child = await Note.create(title="hello")
await parent.connect(child, edge=Contains)
"""
ctx = GraphContext(database)
txn = None
db = _transaction_database(database)
if db is None:
probe = GraphContext()
db = _transaction_database(probe.database)
required = ("begin_transaction", "commit_transaction", "rollback_transaction")
if not all(callable(getattr(db, name, None)) for name in required):
raise TransactionUnavailable(
"graph_transaction requires a database with begin/commit/rollback"
)
txn = await db.begin_transaction()
ctx = GraphContext(txn)
try:
if hasattr(ctx.database, "begin_transaction"):
txn = await ctx.database.begin_transaction()
yield ctx
if txn is not None and hasattr(ctx.database, "commit_transaction"):
await ctx.database.commit_transaction(txn)
elif txn is None and hasattr(ctx.database, "commit_transaction"):
# Backend's commit_transaction accepts no txn argument (non-Mongo)
try:
await ctx.database.commit_transaction()
except TypeError:
await ctx.database.commit_transaction(None)
except Exception:
if txn is not None and hasattr(ctx.database, "rollback_transaction"):
await ctx.database.rollback_transaction(txn)
elif txn is None and hasattr(ctx.database, "rollback_transaction"):
try:
await ctx.database.rollback_transaction()
except TypeError:
await ctx.database.rollback_transaction(None)
async with scoped_default_context_async(ctx):
yield ctx
await db.commit_transaction(txn)
except BaseException:
await db.rollback_transaction(txn)
raise


@asynccontextmanager
async def async_transaction_context(database: Optional[Any] = None):
"""Alias for :func:`graph_transaction`.

The previous implementation opened a transaction but kept writing
through the pool connection, so ``Node.create`` never joined the
commit. Callers should use :func:`graph_transaction`.
"""
async with graph_transaction(database) as ctx:
yield ctx
2 changes: 1 addition & 1 deletion jvspatial/version.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,4 +9,4 @@
# - MAJOR: Breaking changes
# - MINOR: New features, backward compatible
# - PATCH: Bug fixes, backward compatible
__version__ = "0.0.21"
__version__ = "0.0.22"
93 changes: 93 additions & 0 deletions tests/core/test_graph_transaction.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,93 @@
"""graph_transaction binds Node writes to the backend transaction handle."""

from __future__ import annotations

from typing import Any, Dict, List, Optional
from unittest.mock import AsyncMock

import pytest

from jvspatial.core.context import (
TransactionUnavailable,
get_default_context,
graph_transaction,
)
from jvspatial.core.entities import Edge, Node


class _FakeTxn:
def __init__(self) -> None:
self.saved: List[Dict[str, Any]] = []

async def save(self, collection: str, data: Dict[str, Any]) -> Dict[str, Any]:
self.saved.append({"collection": collection, **data})
return data

async def get(self, collection: str, id: str) -> Optional[Dict[str, Any]]:
return next((row for row in self.saved if row.get("id") == id), None)

async def find(self, collection: str, query: Dict[str, Any], **_kwargs) -> list:
return []

async def delete(self, collection: str, id: str) -> bool:
return False


class _FakeDB:
def __init__(self) -> None:
self.txn = _FakeTxn()
self.committed = False
self.rolled_back = False

async def begin_transaction(self) -> _FakeTxn:
return self.txn

async def commit_transaction(self, transaction: _FakeTxn) -> None:
assert transaction is self.txn
self.committed = True

async def rollback_transaction(self, transaction: _FakeTxn) -> None:
assert transaction is self.txn
self.rolled_back = True


class Folder(Node):
name: str = ""


class Note(Node):
title: str = ""


class Contains(Edge):
pass


@pytest.mark.asyncio
async def test_graph_transaction_writes_through_the_handle_and_commits() -> None:
db = _FakeDB()
async with graph_transaction(db) as ctx:
assert ctx.database is db.txn
assert get_default_context().database is db.txn
await Folder.create(name="inbox")
assert db.committed is True
assert db.rolled_back is False
assert any(row.get("collection") == "node" for row in db.txn.saved)


@pytest.mark.asyncio
async def test_graph_transaction_rolls_back_on_error() -> None:
db = _FakeDB()
with pytest.raises(RuntimeError, match="boom"):
async with graph_transaction(db):
await Folder.create(name="inbox")
raise RuntimeError("boom")
assert db.committed is False
assert db.rolled_back is True


@pytest.mark.asyncio
async def test_graph_transaction_refuses_a_store_without_begin() -> None:
with pytest.raises(TransactionUnavailable):
async with graph_transaction(object()):
pass
42 changes: 42 additions & 0 deletions tests/db/test_postgres_integration.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,6 +32,8 @@

import pytest

from jvspatial.core.entities import Edge, Node

try:
import asyncpg
except ImportError: # pragma: no cover - dependency gated below
Expand Down Expand Up @@ -782,3 +784,43 @@ async def test_hybrid_jsonb_plus_near(self, pg_db: "PostgresDB") -> None:
},
)
assert [r["id"] for r in out] == ["d.x", "d.y"]


class _TxnFolder(Node):
name: str = ""


class _TxnNote(Node):
title: str = ""


class _TxnContains(Edge):
pass


class TestGraphTransaction:
async def test_node_and_edge_roll_back_together(self, pg_db: "PostgresDB") -> None:
from jvspatial.core.context import graph_transaction

with pytest.raises(RuntimeError, match="abort"):
async with graph_transaction(pg_db):
folder = await _TxnFolder.create(name="inbox")
note = await _TxnNote.create(title="hello")
await folder.connect(note, edge=_TxnContains)
raise RuntimeError("abort")
assert await pg_db.find("node", {}) == []
assert await pg_db.find("edge", {}) == []

async def test_node_and_edge_commit_together(self, pg_db: "PostgresDB") -> None:
from jvspatial.core.context import graph_transaction

async with graph_transaction(pg_db):
folder = await _TxnFolder.create(name="inbox")
note = await _TxnNote.create(title="hello")
await folder.connect(note, edge=_TxnContains)
folder_id, note_id = folder.id, note.id
nodes = {row["id"] for row in await pg_db.find("node", {})}
assert folder_id in nodes
assert note_id in nodes
edges = await pg_db.find("edge", {})
assert len(edges) == 1
Loading