diff --git a/CLAUDE.md b/AGENTS.md similarity index 75% rename from CLAUDE.md rename to AGENTS.md index b122cc0a..48346f60 100644 --- a/CLAUDE.md +++ b/AGENTS.md @@ -1,9 +1,8 @@ -# CLAUDE.md +# Repository Guidelines -This file provides guidance to Claude Code (claude.ai/code) when working with code in this repository. +This file provides repository conventions for contributors. ## Project Overview - This is a **monorepo** for the FastAPI Startkit ecosystem — a modular, provider-driven framework for building Python applications with FastAPI. It contains four main components: | Directory | Purpose | Published as | @@ -14,9 +13,7 @@ This is a **monorepo** for the FastAPI Startkit ecosystem — a modular, provide | `application/` | Starter application template | Not published — clone/scaffold target | ### `fastapi_startkit/` — Core Package - -The PyPI package (`fastapi-startkit`, currently v0.13.6). Source lives under `src/fastapi_startkit/`. This is the foundational framework all other components depend on. - +The PyPI package `fastapi-startkit`. Source lives under `src/fastapi_startkit/`. This is the foundational framework all other components depend on. **Do not modify framework code unless explicitly necessary.** Changes to core abstractions (Container, Application, Model, Provider, Facades) can have broad breaking effects on downstream applications. Optional extras are installed with pip/uv extras: @@ -49,53 +46,11 @@ Self-contained apps demonstrating specific features. Each subdirectory is an ind The template users clone when starting a new project. Contains the minimal scaffolding: `artisan` entrypoint, `bootstrap/`, `config/`, `providers/`, `routes/`, and `storage/`. It mirrors a typical project layout and is a uv workspace member of this monorepo. -## Task Tracking (Keera Agent MCP) - -Planned work for this project is tracked in **Keera Agent** via an MCP server running locally at `http://127.0.0.1:4545`. - -### Load tasks at the start of a session - -```bash -# List all tasks for this project -curl -s -X POST http://127.0.0.1:4545/mcp \ - -H "Content-Type: application/json" \ - -d '{"jsonrpc":"2.0","method":"tools/call","params":{"name":"list_tasks","arguments":{"project_path":"/Users/ellite/code/packages/fastapi-startkit-framework/fastapi_startkit"}},"id":1}' -``` - -Or open the Keera Agent UI at: `http://127.0.0.1:4545/framework` - -### Available MCP tools +## Coding Standards -| Tool | Purpose | -|---|---| -| `list_tasks` | List tasks (filter by `status`: pending / in_progress / completed / cancelled) | -| `get_task` | Get full details of a task by numeric ID | -| `create_task` | Create a new task with title, description, acceptance criteria, testing methods, and validation steps | -| `update_task` | Update any field of a task | -| `update_task_status` | Change a task's status | -| `send_message_to_agent` | Send a message to another project's agent | -| `get_agent_messages` | Read messages in this project's agent inbox | - -### MCP JSON-RPC usage - -All calls follow the JSON-RPC 2.0 protocol — `POST http://127.0.0.1:4545/mcp` with `Content-Type: application/json`: - -```bash -# Initialize (once per session) -curl -s -X POST http://127.0.0.1:4545/mcp \ - -H "Content-Type: application/json" \ - -d '{"jsonrpc":"2.0","method":"initialize","params":{"protocolVersion":"2024-11-05","capabilities":{},"clientInfo":{"name":"claude-code","version":"1.0"}},"id":0}' - -# Call a tool -curl -s -X POST http://127.0.0.1:4545/mcp \ - -H "Content-Type: application/json" \ - -d '{"jsonrpc":"2.0","method":"tools/call","params":{"name":"","arguments":{...}},"id":1}' -``` - -The `project_path` for this repo is always: -``` -/Users/ellite/code/packages/fastapi-startkit-framework/fastapi_startkit -``` +1. Do not add explanatory comments or docstrings. Express intent through clear names, arguments, and structure. Behavioural directives required by tooling are allowed. +2. Before implementing a change, write an architectural decision record in `docs/adr/`. Explain the problem, alternatives, chosen design, implementation details, and validation plan. +3. Add each ADR to `docs/adr/AGENTS.md` with its date, title, and a brief explanation of the decision. Keep the index sufficient to identify relevant records without reading every ADR. ## Commands @@ -107,10 +62,10 @@ uv sync cd fastapi_startkit && uv build # Run framework tests -uv run pytest fastapi_startkit/src/fastapi_startkit/tests/ -v +cd fastapi_startkit && uv run pytest tests/ -v # Run a single test file -uv run pytest fastapi_startkit/src/fastapi_startkit/tests/configurations/test_config_merge.py -v +cd fastapi_startkit && uv run pytest tests/core/test_configuration.py -v # Serve the docs locally cd fastapi_startkit.github.io.git && npm run dev @@ -151,7 +106,7 @@ Configured in `pyproject.toml` under `[tool.coverage.*]`: |---|---| | Source tracked | `src/fastapi_startkit/` | | Omitted | `*/tests/*`, `*/migrations/*`, `*/__init__.py`, `*.pyi` | -| Minimum threshold | `fail_under = 40` | +| Minimum threshold | `fail_under = 80` | | HTML output dir | `htmlcov/` | The test suite **fails** if total coverage drops below the `fail_under` threshold. Raise this value in `pyproject.toml` as coverage improves. @@ -246,7 +201,7 @@ Async-first fork of Masonite ORM built on SQLAlchemy async: - `Model` auto-pluralizes table names via `inflection` - `created_at`/`updated_at` managed as `pendulum` Carbon objects - Relationships: `HasOne`, `HasMany`, `BelongsTo`, `BelongsToMany`, `HasOneThrough` -- `AsyncQueryBuilder` provides the chainable query interface +- `QueryBuilder` provides the chainable query interface ### Facades (`facades/`) diff --git a/docs/adr/001-masonite-orm-aftercommit.md b/docs/adr/001-masonite-orm-aftercommit.md new file mode 100644 index 00000000..c7ad9db4 --- /dev/null +++ b/docs/adr/001-masonite-orm-aftercommit.md @@ -0,0 +1,68 @@ +# 001: ORM after-commit callbacks + +Date: 2026-10-01 +Status: Accepted + +## Problem + +Applications often need to run follow-up work, such as sending a notification, only once a database change is durable. Model observers (`created`, `updated`, `saved`) fire while the model operation is still running, so they can fire before the surrounding transaction commits. The ORM offers explicit transactions and transaction context managers, but no way to hook into a commit. + +Repository conventions are also moving into `AGENTS.md`, and ADRs need a concise index. + +## Alternatives + +- **Fire model observers after commit.** This changes the timing of existing observers and does not cover raw queries. +- **Use SQLAlchemy engine events.** Commit events fire before the commit completes and cannot await asynchronous callbacks. +- **Track callbacks in the ORM transaction lifecycle (chosen).** This covers every public transaction API and can await callbacks once the commit has succeeded. + +## Decision + +Add `await connection.after_commit(callback)` and `await DB.after_commit(callback, name=None)`. A callback takes no arguments and may be synchronous or return an awaitable. Callers bind arguments with a closure or `functools.partial`. Non-callables are rejected at registration. + +Behaviour: + +- **Active transaction:** the callback is queued. +- **No transaction (or one the ORM did not start):** the callback runs and is awaited immediately. +- **Outermost commit:** queued callbacks run sequentially in registration order. +- **Nested commit (savepoint):** callbacks merge into the parent and do not run yet. +- **Rollback:** a nested rollback discards only that scope's callbacks; an outer rollback discards everything. +- **Close, reconnect, cancellation, failed commit:** no callbacks survive into later transactions. +- **Callback errors:** the error propagates to the caller. The commit stays durable, remaining callbacks are skipped, and the queue is already cleared. + +The root connection is released before callbacks run, so they can read committed data or open a new transaction, even with a pool of one connection. + +This is an in-process hook, not a durable delivery guarantee. Work that must not be lost should use an outbox. + +## Usage + +```python +from fastapi_startkit.masoniteorm.facades import DB + +async with DB.connection().transaction(): + await DB.table("users").where("id", user_id).update({"verified": True}) + await DB.after_commit(lambda: send_confirmation(user_id)) +``` + +`send_confirmation` runs after the transaction exits successfully and is skipped on rollback. + +## Implementation + +The ORM `Connection` keeps a registry of callback frames, keyed by the SQLAlchemy connection and transaction, with each frame linking to its parent. This follows the existing ContextVar connection propagation: tasks that inherit the context share the current transaction, while independent tasks and named connections use separate physical connections and queues. SQLAlchemy's `connection.info` is avoided because reading it during invalidation can force a DBAPI reconnect and block cleanup. + +A transaction the ORM did not register (for example one begun directly on the raw SQLAlchemy connection) is treated as untracked: callbacks registered inside it run immediately, and an ORM savepoint inside it acts as its own root. + +Frames are created, merged and discarded in `Connection.begin_transaction`, `commit_transaction` and `rollback`, and in the `Transaction` enter, exit, commit and rollback paths. Root batches are detached before they run, and closing a connection clears its frames. + +Repository guidance moves into `AGENTS.md` with corrected wording and test paths, and `docs/adr/AGENTS.md` becomes the ADR index. + +## Validation + +Real SQLite integration tests cover: + +- sync and async callbacks, registration order, and visibility of committed data +- manual, context-manager and transaction-object APIs, including mixed use +- nested and deeply nested savepoints, commit and rollback +- callback errors, failed commits, cancellation, close and reconnect, and invalidated connections +- task isolation, named facade connections, callbacks that start new transactions, and untracked transactions + +Existing SQLite transaction tests, Ruff, and the framework type checker also run. diff --git a/docs/adr/AGENTS.md b/docs/adr/AGENTS.md new file mode 100644 index 00000000..a988f0a9 --- /dev/null +++ b/docs/adr/AGENTS.md @@ -0,0 +1,10 @@ +--- +title: Architectural Decision Records +description: Index of architectural decisions, their dates, and the reasons for each implementation. +--- + +Read this index first, then open the records relevant to the change. Before implementation, add or update an ADR explaining the problem, alternatives, decision, implementation, and validation. Keep this index current. + +| Index | Date | Title | Abstract | +| --- | --- | --- | --- | +| [001](001-masonite-orm-aftercommit.md) | 2026-10-01 | ORM after-commit callbacks | Queue sync and async callbacks on the active database transaction, defer nested callbacks until the outer commit, and discard callbacks on rollback. | diff --git a/fastapi_startkit/src/fastapi_startkit/masoniteorm/connections/connection.py b/fastapi_startkit/src/fastapi_startkit/masoniteorm/connections/connection.py index 485bec1e..a62448ce 100644 --- a/fastapi_startkit/src/fastapi_startkit/masoniteorm/connections/connection.py +++ b/fastapi_startkit/src/fastapi_startkit/masoniteorm/connections/connection.py @@ -1,6 +1,9 @@ from __future__ import annotations +from collections.abc import Callable from contextvars import ContextVar, Token +from dataclasses import dataclass, field +from inspect import isawaitable from types import TracebackType from typing import TYPE_CHECKING @@ -16,6 +19,16 @@ from fastapi_startkit.masoniteorm.schema.platforms.Platform import Platform +AfterCommitCallback = Callable[[], object] + + +@dataclass +class _CallbackFrame: + parent: AsyncTransaction | None + callbacks: list[tuple[int, AfterCommitCallback]] = field(default_factory=list) + next_order: int = 0 + + class Transaction: def __init__(self, owner: Connection): self.owner = owner @@ -32,11 +45,13 @@ async def __aenter__(self) -> Self: self.connection = connection self._token = self.owner._connection_context.set(connection) + parent = connection.get_nested_transaction() or connection.get_transaction() try: if connection.in_transaction(): self.transaction = await connection.begin_nested() else: self.transaction = await connection.begin() + self.owner._register_transaction(connection, self.transaction, parent) except BaseException: self.owner._connection_context.reset(self._token) if self._owns_connection: @@ -53,33 +68,61 @@ async def __aexit__( assert self.connection is not None assert self.transaction is not None assert self._token is not None + callbacks: list[AfterCommitCallback] = [] + was_active = self.transaction.is_active try: await self.transaction.__aexit__(exc_type, exc_value, traceback) + if was_active: + callbacks = self.owner._finish_transaction(self.connection, self.transaction, exc_type is None) + except BaseException: + self.owner._finish_transaction(self.connection, self.transaction, False) + raise finally: + current = self.owner.connection self.owner._connection_context.reset(self._token) + if current is not None and current is not self.connection: + self.owner._connection_context.set(current) if self._owns_connection: - await self.connection.close() + await self.owner._release_connection(self.connection) + await self.owner._run_callbacks(callbacks) async def commit(self) -> None: + assert self.connection is not None assert self.transaction is not None - await self.transaction.commit() + try: + await self.transaction.commit() + except BaseException: + self.owner._finish_transaction(self.connection, self.transaction, False) + raise + callbacks = self.owner._finish_transaction(self.connection, self.transaction, True) + if self._owns_connection and not self.connection.in_transaction(): + await self.owner._release_connection(self.connection) + await self.owner._run_callbacks(callbacks) async def rollback(self) -> None: + assert self.connection is not None assert self.transaction is not None - await self.transaction.rollback() + try: + await self.transaction.rollback() + finally: + self.owner._finish_transaction(self.connection, self.transaction, False) class Connection: def __init__(self, engine: AsyncEngine, config: dict): self.config = config self.engine: AsyncEngine = engine + self._transaction_callbacks: dict[AsyncConnection, dict[AsyncTransaction, _CallbackFrame]] = {} self._connection_context: ContextVar[AsyncConnection | None] = ContextVar( f"masoniteorm_connection_{id(self)}", default=None ) @property def connection(self) -> AsyncConnection | None: - return self._connection_context.get() + connection = self._connection_context.get() + if connection is not None and connection.closed: + return None + return connection @property def transactions(self) -> list[AsyncTransaction]: @@ -115,52 +158,125 @@ def get_post_processor(cls) -> type | None: def get_default_platform(cls) -> type[Platform]: raise NotImplementedError + async def after_commit(self, callback: AfterCommitCallback) -> None: + if not callable(callback): + raise TypeError("after_commit requires a callable") + connection = self.connection + transaction = None + if connection is not None: + transaction = connection.get_nested_transaction() or connection.get_transaction() + if connection is None or transaction is None: + await self._run_callbacks([callback]) + return + frames = self._callback_frames(connection) + frame = frames.get(transaction) + if frame is None: + await self._run_callbacks([callback]) + return + root = frame + while root.parent is not None and root.parent in frames: + root = frames[root.parent] + frame.callbacks.append((root.next_order, callback)) + root.next_order += 1 + + def _callback_frames(self, connection: AsyncConnection) -> dict[AsyncTransaction, _CallbackFrame]: + return self._transaction_callbacks.setdefault(connection, {}) + + def _register_transaction( + self, connection: AsyncConnection, transaction: AsyncTransaction, parent: AsyncTransaction | None + ) -> None: + frames = self._callback_frames(connection) + if parent is None: + frames.clear() + frames[transaction] = _CallbackFrame(parent) + + def _finish_transaction( + self, connection: AsyncConnection, transaction: AsyncTransaction, committed: bool + ) -> list[AfterCommitCallback]: + frames = self._transaction_callbacks.get(connection, {}) + frame = frames.pop(transaction, None) + if frame is None: + return [] + descendants = {transaction} + callbacks = frame.callbacks + for child, child_frame in list(frames.items()): + if child_frame.parent in descendants: + descendants.add(child) + callbacks.extend(child_frame.callbacks) + del frames[child] + if not committed: + if not frames: + self._transaction_callbacks.pop(connection, None) + return [] + if frame.parent is not None and frame.parent in frames: + frames[frame.parent].callbacks.extend(callbacks) + return [] + self._transaction_callbacks.pop(connection, None) + return [callback for _, callback in sorted(callbacks, key=lambda item: item[0])] + + @staticmethod + async def _run_callbacks(callbacks: list[AfterCommitCallback]) -> None: + for callback in callbacks: + result = callback() + if isawaitable(result): + await result + async def begin_transaction(self) -> None: connection = self.connection + owns_connection = connection is None if connection is None: connection = await self.engine.connect() self._connection_context.set(connection) - if connection.in_transaction(): - await connection.begin_nested() - else: - await connection.begin() + parent = connection.get_nested_transaction() or connection.get_transaction() + try: + transaction = await connection.begin_nested() if parent is not None else await connection.begin() + self._register_transaction(connection, transaction, parent) + except BaseException: + if owns_connection: + await self._release_connection(connection) + raise async def commit_transaction(self) -> None: connection = self.connection if connection is None or not connection.in_transaction(): raise RuntimeError("No active transaction to commit") - nested = connection.get_nested_transaction() - if nested is not None: - await nested.commit() - else: - transaction = connection.get_transaction() - assert transaction is not None + transaction = connection.get_nested_transaction() or connection.get_transaction() + assert transaction is not None + try: await transaction.commit() + except BaseException: + self._finish_transaction(connection, transaction, False) + raise + callbacks = self._finish_transaction(connection, transaction, True) + if not connection.in_transaction(): await self._release_connection(connection) + await self._run_callbacks(callbacks) async def rollback(self) -> None: connection = self.connection if connection is None or not connection.in_transaction(): raise RuntimeError("No active transaction to rollback") - nested = connection.get_nested_transaction() - if nested is not None: - await nested.rollback() - else: - transaction = connection.get_transaction() - assert transaction is not None + transaction = connection.get_nested_transaction() or connection.get_transaction() + assert transaction is not None + try: await transaction.rollback() - await self._release_connection(connection) + finally: + self._finish_transaction(connection, transaction, False) + if not connection.in_transaction(): + await self._release_connection(connection) async def _release_connection(self, connection: AsyncConnection) -> None: - await connection.close() - if self.connection is connection: - self._connection_context.set(None) + self._transaction_callbacks.pop(connection, None) + try: + await connection.close() + finally: + if self.connection is connection: + self._connection_context.set(None) async def close(self) -> None: connection = self.connection if connection is not None: - await connection.close() - self._connection_context.set(None) + await self._release_connection(connection) async def reconnect(self) -> None: await self.close() diff --git a/fastapi_startkit/src/fastapi_startkit/masoniteorm/facades/DB.py b/fastapi_startkit/src/fastapi_startkit/masoniteorm/facades/DB.py index 902a1b91..51a3e0bd 100644 --- a/fastapi_startkit/src/fastapi_startkit/masoniteorm/facades/DB.py +++ b/fastapi_startkit/src/fastapi_startkit/masoniteorm/facades/DB.py @@ -3,7 +3,7 @@ from typing import TYPE_CHECKING if TYPE_CHECKING: - from fastapi_startkit.masoniteorm.connections.connection import Connection + from fastapi_startkit.masoniteorm.connections.connection import AfterCommitCallback, Connection from fastapi_startkit.masoniteorm.connections.manager import DatabaseManager from fastapi_startkit.masoniteorm.models.builder import QueryBuilder @@ -64,3 +64,7 @@ async def commit(cls, name: str | None = None) -> None: @classmethod async def rollback(cls, name: str | None = None) -> None: await cls.instance().connection(name).rollback() + + @classmethod + async def after_commit(cls, callback: AfterCommitCallback, name: str | None = None) -> None: + await cls.instance().connection(name).after_commit(callback) diff --git a/fastapi_startkit/tests/masoniteorm/connections/test_after_commit.py b/fastapi_startkit/tests/masoniteorm/connections/test_after_commit.py new file mode 100644 index 00000000..5887ac40 --- /dev/null +++ b/fastapi_startkit/tests/masoniteorm/connections/test_after_commit.py @@ -0,0 +1,422 @@ +import asyncio +from contextvars import Context +from functools import partial +from unittest.mock import patch + +import pytest +from sqlalchemy.exc import PendingRollbackError +from sqlalchemy.ext.asyncio import AsyncConnection, AsyncTransaction, create_async_engine + +from fastapi_startkit.masoniteorm.connections.sqlite_connection import SQliteConnection +from fastapi_startkit.masoniteorm.facades.DB import DB + + +@pytest.fixture +async def connection(tmp_path): + engine = create_async_engine(f"sqlite+aiosqlite:///{tmp_path / 'callbacks.sqlite3'}", pool_size=1, max_overflow=0) + conn = SQliteConnection(engine, {"driver": "sqlite"}) + await conn.statement("CREATE TABLE records (id INTEGER)") + try: + yield conn + finally: + await conn.close() + await engine.dispose() + + +async def test_immediate_sync_async_and_returned_awaitable(connection): + called = [] + + async def callback(): + await asyncio.sleep(0) + called.append("async") + + await connection.after_commit(partial(called.append, "sync")) + await connection.after_commit(callback) + await connection.after_commit(lambda: callback()) + assert called == ["sync", "async", "async"] + + +async def test_rejects_non_callable(connection): + with pytest.raises(TypeError, match="requires a callable"): + await connection.after_commit(None) + + +@pytest.mark.parametrize("mode", ["manual", "context", "object"]) +async def test_commit_releases_connection_and_callbacks_see_durable_data(connection, mode): + called = [] + + async def callback(): + assert connection.connection is None + assert await connection.select("SELECT id FROM records") == [{"id": 1}] + async with connection.transaction(): + await connection.statement("INSERT INTO records VALUES (2)") + called.append("async") + + async def register(): + await connection.statement("INSERT INTO records VALUES (1)") + await connection.after_commit(partial(called.append, "sync")) + await connection.after_commit(callback) + assert called == [] + + if mode == "manual": + await connection.begin_transaction() + await register() + await connection.commit_transaction() + else: + async with connection.transaction() as transaction: + await register() + if mode == "object": + await transaction.commit() + assert called == ["sync", "async"] + assert connection.connection is None + assert await connection.select("SELECT id FROM records ORDER BY id") == [{"id": 1}, {"id": 2}] + + +@pytest.mark.parametrize("mode", ["manual", "context", "object"]) +async def test_rollback_discards_callbacks(connection, mode): + called = [] + if mode == "manual": + await connection.begin_transaction() + await connection.after_commit(partial(called.append, "discard")) + await connection.rollback() + else: + try: + async with connection.transaction() as transaction: + await connection.after_commit(partial(called.append, "discard")) + if mode == "object": + await transaction.rollback() + else: + raise ValueError("abort") + except ValueError: + pass + async with connection.transaction(): + await connection.after_commit(partial(called.append, "next")) + assert called == ["next"] + + +@pytest.mark.parametrize("mode", ["manual", "context", "object"]) +@pytest.mark.parametrize("rollback", [False, True]) +async def test_nested_callbacks_defer_and_rollback_is_scoped(connection, mode, rollback): + called = [] + async with connection.transaction(): + await connection.statement("INSERT INTO records VALUES (1)") + await connection.after_commit(partial(called.append, "before")) + if mode == "manual": + await connection.begin_transaction() + await connection.after_commit(partial(called.append, "nested")) + if rollback: + await connection.rollback() + else: + await connection.commit_transaction() + else: + try: + async with connection.transaction() as transaction: + await connection.after_commit(partial(called.append, "nested")) + if mode == "object": + if rollback: + await transaction.rollback() + else: + await transaction.commit() + elif rollback: + raise ValueError("abort") + except ValueError: + pass + assert called == [] + await connection.after_commit(partial(called.append, "after")) + assert called == (["before", "after"] if rollback else ["before", "nested", "after"]) + + +async def test_deep_savepoints_preserve_order(connection): + called = [] + await connection.begin_transaction() + await connection.statement("INSERT INTO records VALUES (1)") + for value in [1, 2, 3]: + await connection.after_commit(partial(called.append, value)) + await connection.begin_transaction() + await connection.after_commit(partial(called.append, 4)) + for _ in range(3): + await connection.commit_transaction() + assert called == [] + await connection.commit_transaction() + assert called == [1, 2, 3, 4] + + +async def test_outer_rollback_discards_committed_nested_callbacks(connection): + called = [] + await connection.begin_transaction() + await connection.statement("INSERT INTO records VALUES (1)") + async with connection.transaction(): + await connection.after_commit(partial(called.append, "inner")) + await connection.rollback() + async with connection.transaction(): + pass + assert called == [] + + +@pytest.mark.parametrize("explicit", [False, True]) +async def test_ancestor_commit_includes_open_savepoint_callbacks(connection, explicit): + called = [] + async with connection.transaction() as root: + await connection.statement("INSERT INTO records VALUES (1)") + await connection.after_commit(partial(called.append, "root")) + await connection.begin_transaction() + await connection.after_commit(partial(called.append, "child")) + await connection.begin_transaction() + await connection.after_commit(partial(called.append, "grandchild")) + if explicit: + await root.commit() + assert called == ["root", "child", "grandchild"] + + +@pytest.mark.parametrize("mode", ["manual", "context", "object"]) +async def test_callback_error_propagates_after_commit_without_queue_leak(connection, mode): + called = [] + + async def fail(): + raise ValueError("callback failed") + + async def register(): + await connection.statement("INSERT INTO records VALUES (1)") + await connection.after_commit(fail) + await connection.after_commit(partial(called.append, "skip")) + + with pytest.raises(ValueError, match="callback failed"): + if mode == "manual": + await connection.begin_transaction() + await register() + await connection.commit_transaction() + else: + async with connection.transaction() as transaction: + await register() + if mode == "object": + await transaction.commit() + assert connection.connection is None + assert await connection.select("SELECT id FROM records") == [{"id": 1}] + async with connection.transaction(): + await connection.after_commit(partial(called.append, "next")) + assert called == ["next"] + + +@pytest.mark.parametrize("mode", ["manual", "object"]) +async def test_failed_commit_discards_callbacks(connection, mode): + called = [] + with patch.object(AsyncTransaction, "commit", side_effect=RuntimeError("commit failed")): + with pytest.raises(RuntimeError, match="commit failed"): + if mode == "manual": + await connection.begin_transaction() + await connection.after_commit(partial(called.append, "failed")) + await connection.commit_transaction() + else: + async with connection.transaction() as transaction: + await connection.after_commit(partial(called.append, "failed")) + await transaction.commit() + await connection.close() + async with connection.transaction(): + await connection.after_commit(partial(called.append, "next")) + assert called == ["next"] + + +@pytest.mark.parametrize("method", ["close", "reconnect"]) +async def test_close_discards_callbacks(connection, method): + called = [] + await connection.begin_transaction() + await connection.after_commit(partial(called.append, "discard")) + await getattr(connection, method)() + async with connection.transaction(): + await connection.after_commit(partial(called.append, "next")) + assert called == ["next"] + + +async def test_cancelled_transaction_discards_callbacks(connection): + called = [] + ready = asyncio.Event() + + async def worker(): + async with connection.transaction(): + await connection.after_commit(partial(called.append, "discard")) + ready.set() + await asyncio.Event().wait() + + task = asyncio.create_task(worker()) + await ready.wait() + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + async with connection.transaction(): + await connection.after_commit(partial(called.append, "next")) + assert called == ["next"] + + +async def test_inherited_task_joins_parent_and_clean_context_is_isolated(connection): + called = [] + async with connection.transaction(): + await asyncio.create_task(connection.after_commit(partial(called.append, "inherited"))) + await Context().run(asyncio.create_task, connection.after_commit(partial(called.append, "isolated"))) + assert called == ["isolated"] + assert called == ["isolated", "inherited"] + + +async def test_independent_tasks_have_separate_callback_queues(tmp_path): + engine = create_async_engine(f"sqlite+aiosqlite:///{tmp_path / 'tasks.sqlite3'}", pool_size=2) + connection = SQliteConnection(engine, {"driver": "sqlite"}) + ready = asyncio.Event() + committed = asyncio.Event() + called = [] + + async def committing(): + await connection.begin_transaction() + await connection.after_commit(partial(called.append, "commit")) + await ready.wait() + await connection.commit_transaction() + committed.set() + + async def rolling_back(): + await connection.begin_transaction() + await connection.after_commit(partial(called.append, "rollback")) + ready.set() + await committed.wait() + assert called == ["commit"] + await connection.rollback() + + try: + await asyncio.gather(committing(), rolling_back()) + assert called == ["commit"] + finally: + await engine.dispose() + + +async def test_facade_default_and_named_connections(connection, tmp_path): + engine = create_async_engine(f"sqlite+aiosqlite:///{tmp_path / 'other.sqlite3'}") + other = SQliteConnection(engine, {"driver": "sqlite"}) + called = [] + with patch.object(DB, "instance") as instance: + instance.return_value.connection.side_effect = lambda name=None: other if name == "other" else connection + async with connection.transaction(): + async with other.transaction(): + await DB.after_commit(partial(called.append, "default")) + await DB.after_commit(partial(called.append, "other"), name="other") + assert called == [] + assert called == ["other"] + assert called == ["other", "default"] + await engine.dispose() + + +async def test_inherited_task_can_query_after_parent_exits(connection): + ready = asyncio.Event() + + async def child(): + await ready.wait() + await connection.after_commit(lambda: None) + return await connection.select("SELECT 1 AS value") + + async with connection.transaction(): + task = asyncio.create_task(child()) + ready.set() + assert await task == [{"value": 1}] + + +async def test_callback_manual_transaction_survives_old_context_exit(connection): + called = [] + + async def callback(): + await connection.begin_transaction() + await connection.statement("INSERT INTO records VALUES (2)") + await connection.after_commit(partial(called.append, "new")) + + async with connection.transaction() as transaction: + await connection.after_commit(callback) + await transaction.commit() + replacement = connection.connection + assert connection.connection is replacement + assert replacement is not None + await connection.commit_transaction() + assert called == ["new"] + assert await connection.select("SELECT id FROM records") == [{"id": 2}] + + +@pytest.mark.parametrize("mode", ["manual", "context", "object"]) +async def test_invalidated_connection_discards_callbacks_and_can_be_closed(connection, mode): + called = [] + with pytest.raises(PendingRollbackError): + if mode == "manual": + await connection.begin_transaction() + await connection.after_commit(partial(called.append, "discard")) + await connection.connection.invalidate() + await connection.commit_transaction() + else: + async with connection.transaction() as transaction: + await connection.after_commit(partial(called.append, "discard")) + await connection.connection.invalidate() + if mode == "object": + await transaction.commit() + await connection.close() + assert connection.connection is None + async with connection.transaction(): + await connection.after_commit(partial(called.append, "next")) + assert called == ["next"] + + +async def test_unregistered_transaction_runs_callback_immediately(connection): + called = [] + raw = await connection.engine.connect() + connection._connection_context.set(raw) + await raw.begin() + await connection.after_commit(partial(called.append, "immediate")) + assert called == ["immediate"] + await connection.close() + + +async def test_savepoint_in_unregistered_transaction_does_not_raise(connection): + called = [] + raw = await connection.engine.connect() + connection._connection_context.set(raw) + await raw.begin() + async with connection.transaction(): + await connection.after_commit(partial(called.append, "savepoint")) + assert called == ["savepoint"] + await connection.close() + + +async def test_nested_rollback_then_new_registration_keeps_order(connection): + called = [] + async with connection.transaction(): + await connection.after_commit(partial(called.append, 1)) + await connection.begin_transaction() + await connection.after_commit(partial(called.append, "discard")) + await connection.rollback() + await connection.after_commit(partial(called.append, 2)) + assert called == [1, 2] + + +@pytest.mark.parametrize("mode", ["manual", "context"]) +async def test_failed_begin_releases_owned_connection(connection, mode): + with patch.object(AsyncConnection, "begin", side_effect=RuntimeError("begin failed")): + with pytest.raises(RuntimeError, match="begin failed"): + if mode == "manual": + await connection.begin_transaction() + else: + async with connection.transaction(): + pass + assert connection.connection is None + async with connection.transaction(): + await connection.after_commit(lambda: None) + + +async def test_commit_and_rollback_require_active_transaction(connection): + with pytest.raises(RuntimeError, match="No active transaction to commit"): + await connection.commit_transaction() + with pytest.raises(RuntimeError, match="No active transaction to rollback"): + await connection.rollback() + + +async def test_rollback_failure_clears_callbacks(connection): + called = [] + await connection.begin_transaction() + await connection.after_commit(partial(called.append, "discard")) + with patch.object(AsyncTransaction, "rollback", side_effect=RuntimeError("rollback failed")): + with pytest.raises(RuntimeError, match="rollback failed"): + await connection.rollback() + await connection.close() + async with connection.transaction(): + await connection.after_commit(partial(called.append, "next")) + assert called == ["next"] diff --git a/fastapi_startkit/uv.lock b/fastapi_startkit/uv.lock index a7ce2ca7..b82c7c25 100644 --- a/fastapi_startkit/uv.lock +++ b/fastapi_startkit/uv.lock @@ -539,7 +539,7 @@ wheels = [ [[package]] name = "fastapi-startkit" -version = "0.59.0" +version = "0.61.0" source = { editable = "." } dependencies = [ { name = "cleo" },