diff --git a/framework/db/simple_module_db/session.py b/framework/db/simple_module_db/session.py index 6971732f..3012a30c 100644 --- a/framework/db/simple_module_db/session.py +++ b/framework/db/simple_module_db/session.py @@ -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) @@ -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") diff --git a/framework/db/tests/test_db_sqlite_savepoint.py b/framework/db/tests/test_db_sqlite_savepoint.py new file mode 100644 index 00000000..e4b5ffd3 --- /dev/null +++ b/framework/db/tests/test_db_sqlite_savepoint.py @@ -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