diff --git a/asyncpg/transaction.py b/asyncpg/transaction.py index 562811e6..280a323c 100644 --- a/asyncpg/transaction.py +++ b/asyncpg/transaction.py @@ -146,6 +146,8 @@ async def start(self): await self._connection.execute(query) except BaseException: self._state = TransactionState.FAILED + if con._top_xact is self: + con._top_xact = None raise else: self._state = TransactionState.STARTED diff --git a/tests/test_transaction.py b/tests/test_transaction.py index f84cf7c0..c14a13e2 100644 --- a/tests/test_transaction.py +++ b/tests/test_transaction.py @@ -5,6 +5,8 @@ # the Apache 2.0 License: http://www.apache.org/licenses/LICENSE-2.0 +from unittest import mock + import asyncpg from asyncpg import _testbase as tb @@ -180,6 +182,34 @@ async def test_transaction_within_manual_transaction(self): self.assertIsNone(self.con._top_xact) self.assertFalse(self.con.is_in_transaction()) + async def test_transaction_failed_begin(self): + self.assertIsNone(self.con._top_xact) + + tr = self.con.transaction() + error = asyncpg.PostgresConnectionError( + 'Timed-out waiting to acquire database connection.' + ) + + with mock.patch.object( + type(self.con), 'execute', side_effect=error + ) as execute: + with self.assertRaises(asyncpg.PostgresConnectionError) as caught: + await tr.start() + + self.assertIs(caught.exception, error) + execute.assert_awaited_once_with('BEGIN;') + + self.assertIsNone(self.con._top_xact) + self.assertFalse(self.con.is_in_transaction()) + + # The next transaction must be a real top-level one. + async with self.con.transaction(): + self.assertTrue(self.con.is_in_transaction()) + await self.con.execute('SELECT 1') + + self.assertIsNone(self.con._top_xact) + self.assertFalse(self.con.is_in_transaction()) + async def test_isolation_level(self): await self.con.reset() default_isolation = await self.con.fetchval(