diff --git a/CHANGELOG.md b/CHANGELOG.md index 1dcab5d..18a4808 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -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 diff --git a/jvspatial/__init__.py b/jvspatial/__init__.py index 9df818b..285e6cf 100644 --- a/jvspatial/__init__.py +++ b/jvspatial/__init__.py @@ -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 @@ -129,6 +129,8 @@ "Walker", "Root", "GraphContext", + "graph_transaction", + "TransactionUnavailable", # Mixins "DeferredSaveMixin", "deferred_saves_globally_allowed", diff --git a/jvspatial/core/__init__.py b/jvspatial/core/__init__.py index 2aebf6a..ef62b81 100644 --- a/jvspatial/core/__init__.py +++ b/jvspatial/core/__init__.py @@ -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, @@ -84,4 +87,7 @@ "scoped_default_context_async", "graph_context", "async_graph_context", + "graph_transaction", + "async_transaction_context", + "TransactionUnavailable", ] diff --git a/jvspatial/core/context.py b/jvspatial/core/context.py index 12fe2c4..2d6cbeb 100644 --- a/jvspatial/core/context.py +++ b/jvspatial/core/context.py @@ -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 diff --git a/jvspatial/version.py b/jvspatial/version.py index f53f5bf..a364d0e 100644 --- a/jvspatial/version.py +++ b/jvspatial/version.py @@ -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" diff --git a/tests/core/test_graph_transaction.py b/tests/core/test_graph_transaction.py new file mode 100644 index 0000000..37052cb --- /dev/null +++ b/tests/core/test_graph_transaction.py @@ -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 diff --git a/tests/db/test_postgres_integration.py b/tests/db/test_postgres_integration.py index b985ef8..9e12747 100644 --- a/tests/db/test_postgres_integration.py +++ b/tests/db/test_postgres_integration.py @@ -32,6 +32,8 @@ import pytest +from jvspatial.core.entities import Edge, Node + try: import asyncpg except ImportError: # pragma: no cover - dependency gated below @@ -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