diff --git a/tests/conftest.py b/tests/conftest.py index 11af99e..b47ec6d 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -48,8 +48,16 @@ async def db_session(di_container: modern_di.Container) -> typing.AsyncIterator[ try: yield create_session(connection) finally: - if connection.in_transaction(): + try: + if not transaction.is_active: + pytest.fail( + "db_session: the outer test transaction is no longer active at teardown, so this test's writes " + "were not rolled back. Something committed it instead of nesting a savepoint; check for a " + 'session created without join_transaction_mode="create_savepoint".', + pytrace=False, + ) await transaction.rollback() - await connection.close() - await engine.dispose() - di_container.reset_override(ioc.Database.database_engine) + finally: + await connection.close() + await engine.dispose() + di_container.reset_override(ioc.Database.database_engine) diff --git a/tests/test_db_session_fixture.py b/tests/test_db_session_fixture.py new file mode 100644 index 0000000..9040ac0 --- /dev/null +++ b/tests/test_db_session_fixture.py @@ -0,0 +1,43 @@ +import pytest + + +pytest_plugins = ["pytester"] + + +def test_db_session_teardown_fails_when_outer_transaction_was_committed(pytester: pytest.Pytester) -> None: + pytester.makepyfile( + """ + import sqlalchemy as sa + from sqlalchemy.ext.asyncio import AsyncSession + + from tests.conftest import app, db_session, di_container + + captured = {} + + + async def test_commits_outer_transaction(db_session): + captured["connection"] = db_session.bind + session = AsyncSession(db_session.bind, join_transaction_mode="control_fully") + await session.execute(sa.text("SELECT 1")) + await session.commit() + await session.close() + + + def test_connection_was_still_closed(): + assert captured["connection"].closed + """, + ) + + result = pytester.runpytest_inprocess( + "-p", + "no:cacheprovider", + "-o", + "asyncio_mode=auto", + "-o", + "asyncio_default_fixture_loop_scope=function", + "-W", + "error", + ) + + result.assert_outcomes(passed=2, errors=1) + result.stdout.fnmatch_lines(["*outer test transaction is no longer active*create_savepoint*"])