From ce6631e7b72c4c603e6f5036bfa7514df50f31db Mon Sep 17 00:00:00 2001 From: Bedram Tamang Date: Fri, 4 Sep 2026 15:49:22 -0700 Subject: [PATCH 1/4] Fix async ORM connection ownership race --- .../masoniteorm/connections/connection.py | 105 +++++++++++------- .../connections/postgres_connection.py | 3 - .../masoniteorm/schema/Blueprint.py | 10 +- .../builder/test_sqlite_builder_pagination.py | 8 ++ .../sqlite/builder/test_sqlite_transaction.py | 34 ++++++ 5 files changed, 112 insertions(+), 48 deletions(-) diff --git a/fastapi_startkit/src/fastapi_startkit/masoniteorm/connections/connection.py b/fastapi_startkit/src/fastapi_startkit/masoniteorm/connections/connection.py index 4ed89f9d..79b8909e 100644 --- a/fastapi_startkit/src/fastapi_startkit/masoniteorm/connections/connection.py +++ b/fastapi_startkit/src/fastapi_startkit/masoniteorm/connections/connection.py @@ -1,4 +1,6 @@ -from typing import List +import asyncio +from dataclasses import dataclass, field +from weakref import WeakKeyDictionary from sqlalchemy import text from sqlalchemy.ext.asyncio import AsyncConnection, AsyncEngine, AsyncTransaction @@ -6,12 +8,35 @@ from fastapi_startkit.masoniteorm.models.builder import QueryBuilder +@dataclass +class _TaskConnectionState: + connection: AsyncConnection | None = None + transactions: list[AsyncTransaction] = field(default_factory=list) + + class Connection: def __init__(self, engine: AsyncEngine, config: dict): self.config = config self.engine: AsyncEngine = engine - self.connection: AsyncConnection | None = None - self.transactions: List[AsyncTransaction] = [] + self._task_states: WeakKeyDictionary[asyncio.Task, _TaskConnectionState] = WeakKeyDictionary() + + def _state(self) -> _TaskConnectionState: + task = asyncio.current_task() + if task is None: + raise RuntimeError("Database operations require an active asyncio task") + state = self._task_states.get(task) + if state is None: + state = _TaskConnectionState() + self._task_states[task] = state + return state + + @property + def connection(self) -> AsyncConnection | None: + return self._state().connection + + @property + def transactions(self) -> list[AsyncTransaction]: + return self._state().transactions def query(self) -> "QueryBuilder": return QueryBuilder( @@ -21,11 +46,11 @@ def query(self) -> "QueryBuilder": ) async def get_connection(self) -> AsyncConnection: - if self.connection is None: - self.connection = await self.engine.connect() + state = self._state() + if state.connection is None: + state.connection = await self.engine.connect() - assert self.connection is not None - return self.connection + return state.connection def get_query_grammar(cls): pass @@ -62,10 +87,12 @@ async def rollback(self) -> None: await self._maybe_cleanup() async def close(self) -> None: - if self.connection is not None: - await self.connection.close() - self.connection = None - self.transactions = [] + states = list(self._task_states.values()) + self._task_states.clear() + for state in states: + if state.connection is not None: + await state.connection.close() + state.transactions.clear() async def reconnect(self) -> None: await self.close() @@ -81,26 +108,25 @@ def sql_alchemy_bindings(query: str, bindings: list | None = None): return (query, params) async def run(self, query: str, bindings: list | None = None): - query, bindings = self.sql_alchemy_bindings(query, bindings) - - conn = await self.get_connection() - result = await conn.execute(text(query), bindings or {}) - - if not self.transactions: - await conn.commit() - - return result + query, params = self.sql_alchemy_bindings(query, bindings) + return await self._execute(query, params) async def execute(self, query: str, bindings: list | None = None): - query, bindings = self.sql_alchemy_bindings(query, bindings) - - conn = await self.get_connection() - result = await conn.execute(text(query), bindings or {}) - - if not self.transactions: - await conn.commit() - - return result + query, params = self.sql_alchemy_bindings(query, bindings) + return await self._execute(query, params) + + async def _execute(self, query: str, bindings: dict): + state = self._state() + if state.transactions: + assert state.connection is not None + return await state.connection.execute(text(query), bindings or {}) + if state.connection is not None: + result = await state.connection.execute(text(query), bindings or {}) + await state.connection.commit() + return result + + async with self.engine.begin() as connection: + return await connection.execute(text(query), bindings or {}) async def insert(self, query: str, bindings: list | None = None) -> int | None: result = await self.execute(query, bindings) @@ -128,24 +154,17 @@ async def select(self, query: str, bindings: list | None = None) -> list[dict]: async def select_one(self, query: str, bindings: list | None = None) -> dict | None: result = await self.run(query, bindings) row = result.fetchone() - result_dict = dict(zip(result.keys(), row)) if row else None - if not self.transactions and self.connection is not None: - await self.connection.commit() - return result_dict + return dict(zip(result.keys(), row)) if row else None async def statement(self, query: str, bindings: list | None = None) -> bool: - query, bindings = self.sql_alchemy_bindings(query, bindings) - - conn = await self.get_connection() - await conn.execute(text(query), bindings or {}) + query, params = self.sql_alchemy_bindings(query, bindings) - # Only commit if NOT inside a transaction - if not self.transactions: - await conn.commit() + await self._execute(query, params) return True async def _maybe_cleanup(self): - if not self.transactions and self.connection: - await self.connection.close() - self.connection = None + state = self._state() + if not state.transactions and state.connection: + await state.connection.close() + state.connection = None diff --git a/fastapi_startkit/src/fastapi_startkit/masoniteorm/connections/postgres_connection.py b/fastapi_startkit/src/fastapi_startkit/masoniteorm/connections/postgres_connection.py index 18ee20ed..43fb32b1 100644 --- a/fastapi_startkit/src/fastapi_startkit/masoniteorm/connections/postgres_connection.py +++ b/fastapi_startkit/src/fastapi_startkit/masoniteorm/connections/postgres_connection.py @@ -10,9 +10,6 @@ class PostgresConnection(Connection): async def insert_get_id(self, query: str, bindings: list | None = None) -> int | None: result = await self.run(query, bindings) row = result.fetchone() - if not self.transactions: - conn = await self.get_connection() - await conn.commit() return row[0] if row is not None else None @classmethod diff --git a/fastapi_startkit/src/fastapi_startkit/masoniteorm/schema/Blueprint.py b/fastapi_startkit/src/fastapi_startkit/masoniteorm/schema/Blueprint.py index 6301abfa..5d67dc89 100644 --- a/fastapi_startkit/src/fastapi_startkit/masoniteorm/schema/Blueprint.py +++ b/fastapi_startkit/src/fastapi_startkit/masoniteorm/schema/Blueprint.py @@ -736,8 +736,14 @@ async def __aexit__(self, exc_type, exc_value, exc_traceback): sql = await self.to_sql() if isinstance(sql, list): - for q in sql: - await self.connection.statement(q, ()) + await self.connection.begin_transaction() + try: + for q in sql: + await self.connection.statement(q, ()) + await self.connection.commit_transaction() + except BaseException: + await self.connection.rollback() + raise return return await self.connection.statement(sql, ()) diff --git a/fastapi_startkit/tests/masoniteorm/sqlite/builder/test_sqlite_builder_pagination.py b/fastapi_startkit/tests/masoniteorm/sqlite/builder/test_sqlite_builder_pagination.py index 67e4791d..4c7879e5 100644 --- a/fastapi_startkit/tests/masoniteorm/sqlite/builder/test_sqlite_builder_pagination.py +++ b/fastapi_startkit/tests/masoniteorm/sqlite/builder/test_sqlite_builder_pagination.py @@ -38,3 +38,11 @@ async def test_simple_paginate(self): self.assertIsInstance(user, User) self.assertIsInstance(paginator.to_json(), str) + + async def test_concurrent_paginate_calls(self): + import asyncio + + paginators = await asyncio.gather(*(User.query().paginate(1) for _ in range(8))) + + self.assertTrue(all(paginator.total for paginator in paginators)) + self.assertTrue(all(paginator.count == 1 for paginator in paginators)) diff --git a/fastapi_startkit/tests/masoniteorm/sqlite/builder/test_sqlite_transaction.py b/fastapi_startkit/tests/masoniteorm/sqlite/builder/test_sqlite_transaction.py index a6afc097..dfafd38d 100644 --- a/fastapi_startkit/tests/masoniteorm/sqlite/builder/test_sqlite_transaction.py +++ b/fastapi_startkit/tests/masoniteorm/sqlite/builder/test_sqlite_transaction.py @@ -1,9 +1,29 @@ +import asyncio + from ...fixtures.model import User from ..fixtures.db import DB from ..test_case import TestCase class TestQueryBuilderTransaction(TestCase): + async def test_concurrent_tasks_own_their_connections_and_transactions(self): + conn = DB.connection("sqlite") + ready = asyncio.Event() + connection_ids = [] + + async def transaction_query(): + await conn.begin_transaction() + connection_ids.append(id(await conn.get_connection())) + if len(connection_ids) == 2: + ready.set() + await ready.wait() + self.assertGreater(await User.query().count(), 0) + await conn.rollback() + + await asyncio.gather(transaction_query(), transaction_query()) + + self.assertEqual(len(set(connection_ids)), 2) + async def test_rollback_undoes_insert(self): conn = DB.connection("sqlite") await conn.begin_transaction() @@ -21,3 +41,17 @@ async def test_commit_persists_insert(self): await conn.commit_transaction() user = await User.where("email", "commit_test@example.com").first() assert user is not None + + async def test_nested_rollback_preserves_outer_transaction(self): + conn = DB.connection("sqlite") + await conn.begin_transaction() + await User.create({"email": "outer@example.com", "name": "Outer", "is_admin": False}) + await conn.begin_transaction() + await User.create({"email": "nested@example.com", "name": "Nested", "is_admin": False}) + await conn.rollback() + + assert await User.where("email", "nested@example.com").first() is None + await conn.commit_transaction() + + assert await User.where("email", "outer@example.com").first() is not None + assert await User.where("email", "nested@example.com").first() is None From 1554fdb9af8cff66f25d493432efdf74323a84ef Mon Sep 17 00:00:00 2001 From: Bedram Tamang Date: Fri, 4 Sep 2026 16:09:48 -0700 Subject: [PATCH 2/4] Refactor async transaction lifecycle --- .../masoniteorm/connections/connection.py | 184 +++++++++++------- .../masoniteorm/schema/Blueprint.py | 7 +- .../masoniteorm/testing/transaction.py | 6 +- .../sqlite/builder/test_sqlite_transaction.py | 48 +++++ 4 files changed, 166 insertions(+), 79 deletions(-) diff --git a/fastapi_startkit/src/fastapi_startkit/masoniteorm/connections/connection.py b/fastapi_startkit/src/fastapi_startkit/masoniteorm/connections/connection.py index 79b8909e..54f5a198 100644 --- a/fastapi_startkit/src/fastapi_startkit/masoniteorm/connections/connection.py +++ b/fastapi_startkit/src/fastapi_startkit/masoniteorm/connections/connection.py @@ -1,44 +1,96 @@ -import asyncio -from dataclasses import dataclass, field -from weakref import WeakKeyDictionary +from __future__ import annotations + +from contextvars import ContextVar, Token +from types import TracebackType +from typing import TYPE_CHECKING from sqlalchemy import text from sqlalchemy.ext.asyncio import AsyncConnection, AsyncEngine, AsyncTransaction from fastapi_startkit.masoniteorm.models.builder import QueryBuilder +if TYPE_CHECKING: + from typing import Self + + +class Transaction: + def __init__(self, owner: Connection): + self.owner = owner + self.connection: AsyncConnection | None = None + self.transaction: AsyncTransaction | None = None + self._token: Token[AsyncConnection | None] | None = None + self._owns_connection = False + + async def __aenter__(self) -> Self: + connection = self.owner.connection + if connection is None: + connection = await self.owner.engine.connect() + self._owns_connection = True + self.connection = connection + self._token = self.owner._connection_context.set(connection) + + try: + if connection.in_transaction(): + self.transaction = await connection.begin_nested() + else: + self.transaction = await connection.begin() + except BaseException: + self.owner._connection_context.reset(self._token) + if self._owns_connection: + await connection.close() + raise + return self + + async def __aexit__( + self, + exc_type: type[BaseException] | None, + exc_value: BaseException | None, + traceback: TracebackType | None, + ) -> None: + assert self.connection is not None + assert self.transaction is not None + assert self._token is not None + try: + await self.transaction.__aexit__(exc_type, exc_value, traceback) + finally: + self.owner._connection_context.reset(self._token) + if self._owns_connection: + await self.connection.close() + + async def commit(self) -> None: + assert self.transaction is not None + await self.transaction.commit() -@dataclass -class _TaskConnectionState: - connection: AsyncConnection | None = None - transactions: list[AsyncTransaction] = field(default_factory=list) + async def rollback(self) -> None: + assert self.transaction is not None + await self.transaction.rollback() class Connection: def __init__(self, engine: AsyncEngine, config: dict): self.config = config self.engine: AsyncEngine = engine - self._task_states: WeakKeyDictionary[asyncio.Task, _TaskConnectionState] = WeakKeyDictionary() - - def _state(self) -> _TaskConnectionState: - task = asyncio.current_task() - if task is None: - raise RuntimeError("Database operations require an active asyncio task") - state = self._task_states.get(task) - if state is None: - state = _TaskConnectionState() - self._task_states[task] = state - return state + self._connection_context: ContextVar[AsyncConnection | None] = ContextVar( + f"masoniteorm_connection_{id(self)}", default=None + ) @property def connection(self) -> AsyncConnection | None: - return self._state().connection + return self._connection_context.get() @property def transactions(self) -> list[AsyncTransaction]: - return self._state().transactions + connection = self.connection + if connection is None: + return [] + nested = connection.get_nested_transaction() + root = connection.get_transaction() + return [transaction for transaction in (root, nested) if transaction is not None] - def query(self) -> "QueryBuilder": + def transaction(self) -> Transaction: + return Transaction(self) + + def query(self) -> QueryBuilder: return QueryBuilder( connection=self, grammar=self.get_query_grammar(), @@ -46,11 +98,7 @@ def query(self) -> "QueryBuilder": ) async def get_connection(self) -> AsyncConnection: - state = self._state() - if state.connection is None: - state.connection = await self.engine.connect() - - return state.connection + return self.connection or await self.engine.connect() def get_query_grammar(cls): pass @@ -59,40 +107,51 @@ def get_post_processor(self): pass async def begin_transaction(self) -> None: - connection = await self.get_connection() - - if not self.transactions: - transaction = await connection.begin() + connection = self.connection + if connection is None: + connection = await self.engine.connect() + self._connection_context.set(connection) + if connection.in_transaction(): + await connection.begin_nested() else: - transaction = await connection.begin_nested() - - self.transactions.append(transaction) + await connection.begin() async def commit_transaction(self) -> None: - if not self.transactions: + connection = self.connection + if connection is None or not connection.in_transaction(): raise RuntimeError("No active transaction to commit") - - transaction = self.transactions.pop() - await transaction.commit() - - await self._maybe_cleanup() + nested = connection.get_nested_transaction() + if nested is not None: + await nested.commit() + else: + transaction = connection.get_transaction() + assert transaction is not None + await transaction.commit() + await self._release_connection(connection) async def rollback(self) -> None: - if not self.transactions: + 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 + await transaction.rollback() + await self._release_connection(connection) - transaction = self.transactions.pop() - await transaction.rollback() - - await self._maybe_cleanup() + async def _release_connection(self, connection: AsyncConnection) -> None: + await connection.close() + if self.connection is connection: + self._connection_context.set(None) async def close(self) -> None: - states = list(self._task_states.values()) - self._task_states.clear() - for state in states: - if state.connection is not None: - await state.connection.close() - state.transactions.clear() + connection = self.connection + if connection is not None: + await connection.close() + self._connection_context.set(None) async def reconnect(self) -> None: await self.close() @@ -116,21 +175,14 @@ async def execute(self, query: str, bindings: list | None = None): return await self._execute(query, params) async def _execute(self, query: str, bindings: dict): - state = self._state() - if state.transactions: - assert state.connection is not None - return await state.connection.execute(text(query), bindings or {}) - if state.connection is not None: - result = await state.connection.execute(text(query), bindings or {}) - await state.connection.commit() - return result - - async with self.engine.begin() as connection: + connection = self.connection + if connection is not None: return await connection.execute(text(query), bindings or {}) + async with self.engine.begin() as operation_connection: + return await operation_connection.execute(text(query), bindings or {}) async def insert(self, query: str, bindings: list | None = None) -> int | None: result = await self.execute(query, bindings) - return getattr(result, "lastrowid", None) async def insert_get_id(self, query: str, bindings: list | None = None) -> int | None: @@ -139,7 +191,6 @@ async def insert_get_id(self, query: str, bindings: list | None = None) -> int | async def update(self, query: str, bindings: list | None = None) -> int: result = await self.execute(query, bindings) - return result.rowcount # type: ignore[return-value] async def delete(self, query: str, bindings: list | None = None) -> int: @@ -148,7 +199,6 @@ async def delete(self, query: str, bindings: list | None = None) -> int: async def select(self, query: str, bindings: list | None = None) -> list[dict]: result = await self.run(query, bindings) - return result.mappings().all() async def select_one(self, query: str, bindings: list | None = None) -> dict | None: @@ -158,13 +208,5 @@ async def select_one(self, query: str, bindings: list | None = None) -> dict | N async def statement(self, query: str, bindings: list | None = None) -> bool: query, params = self.sql_alchemy_bindings(query, bindings) - await self._execute(query, params) - return True - - async def _maybe_cleanup(self): - state = self._state() - if not state.transactions and state.connection: - await state.connection.close() - state.connection = None diff --git a/fastapi_startkit/src/fastapi_startkit/masoniteorm/schema/Blueprint.py b/fastapi_startkit/src/fastapi_startkit/masoniteorm/schema/Blueprint.py index 5d67dc89..f7a17c75 100644 --- a/fastapi_startkit/src/fastapi_startkit/masoniteorm/schema/Blueprint.py +++ b/fastapi_startkit/src/fastapi_startkit/masoniteorm/schema/Blueprint.py @@ -736,14 +736,9 @@ async def __aexit__(self, exc_type, exc_value, exc_traceback): sql = await self.to_sql() if isinstance(sql, list): - await self.connection.begin_transaction() - try: + async with self.connection.transaction(): for q in sql: await self.connection.statement(q, ()) - await self.connection.commit_transaction() - except BaseException: - await self.connection.rollback() - raise return return await self.connection.statement(sql, ()) diff --git a/fastapi_startkit/src/fastapi_startkit/masoniteorm/testing/transaction.py b/fastapi_startkit/src/fastapi_startkit/masoniteorm/testing/transaction.py index cb275e3f..5c8cb03e 100644 --- a/fastapi_startkit/src/fastapi_startkit/masoniteorm/testing/transaction.py +++ b/fastapi_startkit/src/fastapi_startkit/masoniteorm/testing/transaction.py @@ -3,10 +3,12 @@ async def asyncStartTestRun(self): from fastapi_startkit.masoniteorm.models import Model self.connection = Model.db_manager.connection(None) - await self.connection.begin_transaction() + self.transaction = self.connection.transaction() + await self.transaction.__aenter__() async def asyncStopTestRun(self): - await self.connection.rollback() + await self.transaction.rollback() + await self.transaction.__aexit__(None, None, None) class RefreshDatabase(DatabaseTransaction): diff --git a/fastapi_startkit/tests/masoniteorm/sqlite/builder/test_sqlite_transaction.py b/fastapi_startkit/tests/masoniteorm/sqlite/builder/test_sqlite_transaction.py index dfafd38d..e5d56a4c 100644 --- a/fastapi_startkit/tests/masoniteorm/sqlite/builder/test_sqlite_transaction.py +++ b/fastapi_startkit/tests/masoniteorm/sqlite/builder/test_sqlite_transaction.py @@ -1,5 +1,10 @@ import asyncio +from contextvars import Context +from tempfile import NamedTemporaryFile +from sqlalchemy.ext.asyncio import create_async_engine + +from fastapi_startkit.masoniteorm.connections.sqlite_connection import SQliteConnection from ...fixtures.model import User from ..fixtures.db import DB from ..test_case import TestCase @@ -24,6 +29,49 @@ async def transaction_query(): self.assertEqual(len(set(connection_ids)), 2) + async def test_transaction_context_returns_connection_when_cancelled(self): + with NamedTemporaryFile(suffix=".sqlite3") as database: + engine = create_async_engine(f"sqlite+aiosqlite:///{database.name}", pool_size=1, max_overflow=0) + connection = SQliteConnection(engine, {"driver": "sqlite"}) + started = asyncio.Event() + + async def cancelled_transaction(): + async with connection.transaction(): + await connection.statement("CREATE TABLE cancelled (id INTEGER)") + started.set() + await asyncio.Event().wait() + + task = asyncio.create_task(cancelled_transaction()) + await started.wait() + task.cancel() + with self.assertRaises(asyncio.CancelledError): + await task + + self.assertEqual(engine.pool.checkedout(), 0) # type: ignore[attr-defined] + self.assertEqual(await connection.select("SELECT 1 AS value"), [{"value": 1}]) + await engine.dispose() + + async def test_transaction_context_propagates_and_clean_context_is_isolated(self): + conn = DB.connection("sqlite") + async with conn.transaction(): + transaction_connection = await conn.get_connection() + + async def inherited_connection(): + return await conn.get_connection() + + async def clean_connection(): + connection = await conn.get_connection() + try: + return connection + finally: + await connection.close() + + inherited = await asyncio.create_task(inherited_connection()) + isolated = await Context().run(asyncio.create_task, clean_connection()) + + self.assertIs(inherited, transaction_connection) + self.assertIsNot(isolated, transaction_connection) + async def test_rollback_undoes_insert(self): conn = DB.connection("sqlite") await conn.begin_transaction() From b78701f2b9e80f72bedfe1408f259956ed416ad4 Mon Sep 17 00:00:00 2001 From: Bedram Tamang Date: Fri, 4 Sep 2026 16:58:20 -0700 Subject: [PATCH 3/4] test(orm): cover transaction ownership, release and error paths Covers the still-untested paths from the connection ownership rework: Transaction enter-failure cleanup, explicit commit/rollback, nested savepoint contexts, the transactions property, commit/rollback guard errors, reconnect release, PostgresConnection.insert_get_id, and the DatabaseTransaction test harness. --- .../tests/masoniteorm/connections/__init__.py | 0 .../connections/test_postgres_connection.py | 36 ++++++++ .../sqlite/builder/test_sqlite_transaction.py | 91 ++++++++++++++++++- .../sqlite/test_testing_transaction.py | 16 ++++ 4 files changed, 142 insertions(+), 1 deletion(-) create mode 100644 fastapi_startkit/tests/masoniteorm/connections/__init__.py create mode 100644 fastapi_startkit/tests/masoniteorm/connections/test_postgres_connection.py create mode 100644 fastapi_startkit/tests/masoniteorm/sqlite/test_testing_transaction.py diff --git a/fastapi_startkit/tests/masoniteorm/connections/__init__.py b/fastapi_startkit/tests/masoniteorm/connections/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/fastapi_startkit/tests/masoniteorm/connections/test_postgres_connection.py b/fastapi_startkit/tests/masoniteorm/connections/test_postgres_connection.py new file mode 100644 index 00000000..f7e04b0d --- /dev/null +++ b/fastapi_startkit/tests/masoniteorm/connections/test_postgres_connection.py @@ -0,0 +1,36 @@ +from unittest import IsolatedAsyncioTestCase +from unittest.mock import AsyncMock, Mock + +from fastapi_startkit.masoniteorm.connections.postgres_connection import PostgresConnection +from fastapi_startkit.masoniteorm.query.grammars import PostgresGrammar +from fastapi_startkit.masoniteorm.query.processors import PostgresPostProcessor +from fastapi_startkit.masoniteorm.schema.platforms import PostgresPlatform + + +class TestPostgresConnection(IsolatedAsyncioTestCase): + @staticmethod + def _connection_with_result(row): + connection = PostgresConnection(engine=None, config={"driver": "postgres"}) + result = Mock() + result.fetchone.return_value = row + connection.run = AsyncMock(return_value=result) + return connection + + async def test_insert_get_id_returns_first_column_of_returning_row(self): + connection = self._connection_with_result((42,)) + inserted_id = await connection.insert_get_id( + "INSERT INTO users (name) VALUES (?) RETURNING id", ["Joe"] + ) + self.assertEqual(inserted_id, 42) + + async def test_insert_get_id_returns_none_when_no_row(self): + connection = self._connection_with_result(None) + inserted_id = await connection.insert_get_id( + "INSERT INTO users (name) VALUES (?) RETURNING id", ["Joe"] + ) + self.assertIsNone(inserted_id) + + def test_grammar_platform_and_processor_classes(self): + self.assertIs(PostgresConnection.get_query_grammar(), PostgresGrammar) + self.assertIs(PostgresConnection.get_default_platform(), PostgresPlatform) + self.assertIs(PostgresConnection.get_post_processor(), PostgresPostProcessor) diff --git a/fastapi_startkit/tests/masoniteorm/sqlite/builder/test_sqlite_transaction.py b/fastapi_startkit/tests/masoniteorm/sqlite/builder/test_sqlite_transaction.py index e5d56a4c..7a2c02b9 100644 --- a/fastapi_startkit/tests/masoniteorm/sqlite/builder/test_sqlite_transaction.py +++ b/fastapi_startkit/tests/masoniteorm/sqlite/builder/test_sqlite_transaction.py @@ -1,8 +1,9 @@ import asyncio from contextvars import Context from tempfile import NamedTemporaryFile +from unittest import mock -from sqlalchemy.ext.asyncio import create_async_engine +from sqlalchemy.ext.asyncio import AsyncConnection, create_async_engine from fastapi_startkit.masoniteorm.connections.sqlite_connection import SQliteConnection from ...fixtures.model import User @@ -90,6 +91,94 @@ async def test_commit_persists_insert(self): user = await User.where("email", "commit_test@example.com").first() assert user is not None + async def test_transaction_context_nested_savepoint_rolls_back_only_inner(self): + conn = DB.connection("sqlite") + async with conn.transaction(): + await User.create({"email": "ctx_outer@example.com", "name": "Outer", "is_admin": False}) + try: + async with conn.transaction(): + await User.create({"email": "ctx_inner@example.com", "name": "Inner", "is_admin": False}) + raise ValueError("abort inner") + except ValueError: + pass + assert await User.where("email", "ctx_outer@example.com").first() is not None + assert await User.where("email", "ctx_inner@example.com").first() is None + + async def test_transaction_object_commit_persists_insert(self): + conn = DB.connection("sqlite") + transaction = conn.transaction() + async with transaction: + await User.create({"email": "obj_commit@example.com", "name": "Obj Commit", "is_admin": False}) + await transaction.commit() + assert await User.where("email", "obj_commit@example.com").first() is not None + + async def test_transaction_object_rollback_discards_insert(self): + conn = DB.connection("sqlite") + transaction = conn.transaction() + async with transaction: + await User.create({"email": "obj_rollback@example.com", "name": "Obj Rollback", "is_admin": False}) + await transaction.rollback() + assert await User.where("email", "obj_rollback@example.com").first() is None + + async def test_transaction_enter_failure_releases_owned_connection(self): + with NamedTemporaryFile(suffix=".sqlite3") as database: + engine = create_async_engine(f"sqlite+aiosqlite:///{database.name}", pool_size=1, max_overflow=0) + connection = SQliteConnection(engine, {"driver": "sqlite"}) + with mock.patch.object(AsyncConnection, "begin", side_effect=RuntimeError("begin failed")): + with self.assertRaises(RuntimeError): + async with connection.transaction(): + pass + self.assertIsNone(connection.connection) + self.assertEqual(engine.pool.checkedout(), 0) # type: ignore[attr-defined] + await engine.dispose() + + async def test_transactions_property_tracks_root_and_nested(self): + conn = DB.connection("sqlite") + self.assertEqual(conn.transactions, []) + await conn.begin_transaction() + self.assertEqual(len(conn.transactions), 1) + await conn.begin_transaction() + self.assertEqual(len(conn.transactions), 2) + await conn.commit_transaction() + self.assertEqual(len(conn.transactions), 1) + await conn.rollback() + self.assertEqual(conn.transactions, []) + + async def test_nested_commit_is_discarded_by_outer_rollback(self): + conn = DB.connection("sqlite") + await conn.begin_transaction() + # outer DML first: sqlite drivers only issue a real BEGIN before DML, + # so a savepoint opened on a pristine transaction would commit on release + await User.create({"email": "outer_pending@example.com", "name": "Outer Pending", "is_admin": False}) + await conn.begin_transaction() + await User.create({"email": "nested_commit@example.com", "name": "Nested Commit", "is_admin": False}) + await conn.commit_transaction() + await conn.rollback() + + assert await User.where("email", "outer_pending@example.com").first() is None + assert await User.where("email", "nested_commit@example.com").first() is None + + async def test_commit_without_transaction_raises(self): + conn = DB.connection("sqlite") + with self.assertRaises(RuntimeError): + await conn.commit_transaction() + + async def test_rollback_without_transaction_raises(self): + conn = DB.connection("sqlite") + with self.assertRaises(RuntimeError): + await conn.rollback() + + async def test_reconnect_releases_context_connection(self): + with NamedTemporaryFile(suffix=".sqlite3") as database: + engine = create_async_engine(f"sqlite+aiosqlite:///{database.name}", pool_size=1, max_overflow=0) + connection = SQliteConnection(engine, {"driver": "sqlite"}) + await connection.begin_transaction() + self.assertIsNotNone(connection.connection) + await connection.reconnect() + self.assertIsNone(connection.connection) + self.assertEqual(engine.pool.checkedout(), 0) # type: ignore[attr-defined] + await engine.dispose() + async def test_nested_rollback_preserves_outer_transaction(self): conn = DB.connection("sqlite") await conn.begin_transaction() diff --git a/fastapi_startkit/tests/masoniteorm/sqlite/test_testing_transaction.py b/fastapi_startkit/tests/masoniteorm/sqlite/test_testing_transaction.py new file mode 100644 index 00000000..ed0c625d --- /dev/null +++ b/fastapi_startkit/tests/masoniteorm/sqlite/test_testing_transaction.py @@ -0,0 +1,16 @@ +from fastapi_startkit.masoniteorm.testing.transaction import DatabaseTransaction + +from ..fixtures.model import User +from .test_case import TestCase + + +class TestDatabaseTransactionHarness(TestCase): + async def test_start_and_stop_roll_back_writes(self): + harness = DatabaseTransaction() + await harness.asyncStartTestRun() + try: + await User.create({"email": "harness@example.com", "name": "Harness", "is_admin": False}) + assert await User.where("email", "harness@example.com").first() is not None + finally: + await harness.asyncStopTestRun() + assert await User.where("email", "harness@example.com").first() is None From 411929c1a77e57f42c2a0396ef83c8c843114f20 Mon Sep 17 00:00:00 2001 From: Bedram Tamang Date: Fri, 4 Sep 2026 17:09:23 -0700 Subject: [PATCH 4/4] style(tests): format postgres connection tests with ruff --- .../masoniteorm/connections/test_postgres_connection.py | 8 ++------ 1 file changed, 2 insertions(+), 6 deletions(-) diff --git a/fastapi_startkit/tests/masoniteorm/connections/test_postgres_connection.py b/fastapi_startkit/tests/masoniteorm/connections/test_postgres_connection.py index f7e04b0d..e12b684c 100644 --- a/fastapi_startkit/tests/masoniteorm/connections/test_postgres_connection.py +++ b/fastapi_startkit/tests/masoniteorm/connections/test_postgres_connection.py @@ -18,16 +18,12 @@ def _connection_with_result(row): async def test_insert_get_id_returns_first_column_of_returning_row(self): connection = self._connection_with_result((42,)) - inserted_id = await connection.insert_get_id( - "INSERT INTO users (name) VALUES (?) RETURNING id", ["Joe"] - ) + inserted_id = await connection.insert_get_id("INSERT INTO users (name) VALUES (?) RETURNING id", ["Joe"]) self.assertEqual(inserted_id, 42) async def test_insert_get_id_returns_none_when_no_row(self): connection = self._connection_with_result(None) - inserted_id = await connection.insert_get_id( - "INSERT INTO users (name) VALUES (?) RETURNING id", ["Joe"] - ) + inserted_id = await connection.insert_get_id("INSERT INTO users (name) VALUES (?) RETURNING id", ["Joe"]) self.assertIsNone(inserted_id) def test_grammar_platform_and_processor_classes(self):