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
15 changes: 15 additions & 0 deletions framework/db/simple_module_db/session.py
Original file line number Diff line number Diff line change
Expand Up @@ -148,6 +148,9 @@ def _configure_sqlite(
facade, and ``connect`` fires on the DBAPI connection underneath it. All
three PRAGMAs must be re-issued per connection except ``journal_mode``,
which is a property of the file; re-issuing it is cheap and idempotent.

Also registers a ``savepoint`` listener that opens the transaction before an
outermost SAVEPOINT (pysqlite would otherwise let its RELEASE commit; GH #350).
"""
apply_wal = wal and _is_file_database(database_url)

Expand All @@ -163,3 +166,15 @@ def _set_sqlite_pragmas(dbapi_connection, _connection_record) -> None: # type:
cursor.execute("PRAGMA journal_mode=WAL")
finally:
cursor.close()

@event.listens_for(engine.sync_engine, "savepoint")
def _begin_before_outermost_savepoint(conn, _name) -> None: # type: ignore[misc]
# pysqlite only emits BEGIN lazily before DML, so a SAVEPOINT that is
# the first statement would *start* the transaction and its RELEASE
# would COMMIT it, leaving the outer rollback nothing to undo (GH #350).
# Open the transaction ourselves, only in that case. SQLAlchemy's
# blanket recipe (isolation_level=None + BEGIN on every transaction)
# was rejected: it turns every read into a snapshot, so a later write
# fails with SQLITE_BUSY_SNAPSHOT, which ignores busy_timeout.
if not getattr(conn.connection.driver_connection, "in_transaction", False):
conn.exec_driver_sql("BEGIN")
75 changes: 75 additions & 0 deletions framework/db/tests/test_db_sqlite_savepoint.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,75 @@
"""GH #350: an outermost begin_nested() must not commit on RELEASE on SQLite."""

from __future__ import annotations

import pytest
from simple_module_db.session import init_db
from sqlalchemy import text


@pytest.fixture(params=["memory", "file"])
async def db_state(request, tmp_path):
url = (
"sqlite+aiosqlite:///:memory:"
if request.param == "memory"
else f"sqlite+aiosqlite:///{tmp_path / 'sp.db'}"
)
state = init_db(url)
async with state.engine.begin() as conn:
await conn.execute(text("create table t (id integer primary key)"))
try:
yield state
finally:
await state.engine.dispose()


async def _count(state, where: str = "1=1") -> int:
async with state.session_factory() as db:
return (await db.execute(text(f"select count(*) from t where {where}"))).scalar()


async def test_outermost_savepoint_is_rolled_back_with_the_outer_transaction(db_state):
async with db_state.session_factory() as db:
async with db.begin_nested(): # no DML before it
await db.execute(text("insert into t values (1)"))
await db.rollback()
assert await _count(db_state) == 0


async def test_savepoint_after_write_still_rolls_back(db_state):
async with db_state.session_factory() as db:
await db.execute(text("update t set id = id"))
async with db.begin_nested():
await db.execute(text("insert into t values (2)"))
await db.rollback()
assert await _count(db_state, "id = 2") == 0


async def test_committed_work_with_savepoint_persists(db_state):
async with db_state.session_factory() as db:
async with db.begin_nested():
await db.execute(text("insert into t values (3)"))
await db.commit()
assert await _count(db_state, "id = 3") == 1


async def test_failed_savepoint_rolls_back_only_itself(db_state):
from sqlalchemy.exc import IntegrityError

async with db_state.session_factory() as db:
await db.execute(text("insert into t values (4)"))
with pytest.raises(IntegrityError):
async with db.begin_nested():
await db.execute(text("insert into t values (5)"))
await db.execute(text("insert into t values (4)"))
await db.commit()
assert await _count(db_state, "id = 4") == 1
assert await _count(db_state, "id = 5") == 0


async def test_nested_savepoints_do_not_double_begin(db_state):
async with db_state.session_factory() as db:
async with db.begin_nested(), db.begin_nested():
await db.execute(text("insert into t values (6)"))
await db.rollback()
assert await _count(db_state, "id = 6") == 0
Loading