From 07d00b775a9b9b371816959c3866ed86edebc072 Mon Sep 17 00:00:00 2001 From: Elvis Pranskevichus Date: Fri, 2 Oct 2026 11:39:32 -0700 Subject: [PATCH] Terminate rejected target_session_attrs candidates before returning With `target_session_attrs`, `connect()` opens candidate connections in turn until one has the requested attribute. Rejected candidates used to be closed in a detached task. That task could outlive the caller's event loop, producing "Event loop is closed" errors and pending-task warnings, and leaving the connections open. Now every candidate that is not returned is terminated before `connect()` exits, whether the attribute check matched, did not match, raised, or was cancelled. --- asyncpg/connect_utils.py | 14 ++--- tests/test_connect.py | 131 +++++++++++++++++++++++++++++++++++++++ 2 files changed, 136 insertions(+), 9 deletions(-) diff --git a/asyncpg/connect_utils.py b/asyncpg/connect_utils.py index 0c5801d7..e215391b 100644 --- a/asyncpg/connect_utils.py +++ b/asyncpg/connect_utils.py @@ -1295,15 +1295,11 @@ async def _connect(*, loop, connection_class, record_class, **kwargs): if target_attr == SessionAttribute.prefer_standby and candidates: chosen_connection = random.choice(candidates) finally: - - async def _close_candidates(conns, chosen): - await asyncio.gather( - *(c.close() for c in conns if c is not chosen), - return_exceptions=True - ) - if candidates: - asyncio.create_task( - _close_candidates(candidates, chosen_connection)) + # Do not wait for rejected hosts to disconnect: their cleanup must + # neither delay a successful connection nor outlive the event loop. + for c in candidates: + if c is not chosen_connection: + c.terminate() if chosen_connection: return chosen_connection diff --git a/tests/test_connect.py b/tests/test_connect.py index fbac0600..1c25d8d6 100644 --- a/tests/test_connect.py +++ b/tests/test_connect.py @@ -1669,6 +1669,95 @@ async def test_connect_args_validation(self): await asyncpg.connect(**{arg: val}) +class TestConnectionCleanup(tb.TestCase): + + async def test_connect_terminates_rejected_connections(self): + for outcome in ('matched', 'unmatched', 'error', 'cancelled'): + with self.subTest(outcome=outcome): + rejected = unittest.mock.Mock(spec=pg_connection.Connection) + candidate = unittest.mock.Mock(spec=pg_connection.Connection) + results = { + 'matched': True, + 'unmatched': False, + 'error': RuntimeError('attribute check failed'), + 'cancelled': asyncio.CancelledError(), + } + tasks_before = asyncio.all_tasks() + try: + with unittest.mock.patch.object( + connect_utils, '_connect_addr', + side_effect=[rejected, candidate], + ), unittest.mock.patch.object( + connect_utils, '_can_use_connection', + side_effect=[False, results[outcome]], + ): + coro = asyncpg.connect( + host=['first', 'second'], user='postgres', + database='postgres', ssl=False, + target_session_attrs='primary', + ) + if outcome == 'matched': + self.assertIs(await coro, candidate) + else: + expected = { + 'unmatched': + exceptions.TargetServerAttributeNotMatched, + 'error': RuntimeError, + 'cancelled': asyncio.CancelledError, + }[outcome] + with self.assertRaises(expected): + await coro + + rejected.terminate.assert_called_once_with() + rejected.close.assert_not_called() + candidate.close.assert_not_called() + if outcome == 'matched': + candidate.terminate.assert_not_called() + else: + candidate.terminate.assert_called_once_with() + self.assertEqual(asyncio.all_tasks(), tasks_before) + finally: + # Clean up unexpected tasks if a regression assertion + # fails. + pending = asyncio.all_tasks() - tasks_before + for task in pending: + task.cancel() + await asyncio.gather(*pending, return_exceptions=True) + + async def test_rejected_host_cannot_time_out_selected_connection(self): + async def close(): + await asyncio.Future() + + rejected = unittest.mock.Mock(spec=pg_connection.Connection) + rejected.close.side_effect = close + selected = unittest.mock.Mock(spec=pg_connection.Connection) + tasks_before = asyncio.all_tasks() + try: + with unittest.mock.patch.object( + connect_utils, '_connect_addr', + side_effect=[rejected, selected], + ), unittest.mock.patch.object( + connect_utils, '_can_use_connection', + side_effect=[False, True], + ): + con = await asyncpg.connect( + host=['first', 'second'], user='postgres', + database='postgres', ssl=False, + target_session_attrs='primary', timeout=0.1, + ) + self.assertIs(con, selected) + rejected.terminate.assert_called_once_with() + rejected.close.assert_not_called() + selected.terminate.assert_not_called() + selected.close.assert_not_called() + self.assertEqual(asyncio.all_tasks(), tasks_before) + finally: + pending = asyncio.all_tasks() - tasks_before + for task in pending: + task.cancel() + await asyncio.gather(*pending, return_exceptions=True) + + class TestConnection(tb.ConnectedTestCase): async def test_connection_cancelled_during_transport_setup(self): @@ -2785,6 +2874,48 @@ async def test_no_explicit_close_with_debug(self): class TestConnectionAttributes(tb.HotStandbyTestCase): + async def test_rejected_standby_is_terminated_before_primary(self): + primary_spec = self.get_cluster_connection_spec(self.master_cluster) + standby_spec = self.get_cluster_connection_spec(self.standby_cluster) + candidates = [] + terminated = [] + + class Connection(pg_connection.Connection): + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + if not candidates: + self.add_termination_listener(terminated.append) + candidates.append(self) + + con = None + try: + con = await self.connect( + host=[standby_spec['host'], primary_spec['host']], + port=[standby_spec['port'], primary_spec['port']], + connection_class=Connection, target_session_attrs='primary', + timeout=5, + ) + self.assertEqual(len(candidates), 2) + rejected = candidates[0] + self.assertTrue(rejected.is_closed()) + self.assertTrue(rejected._transport.is_closing()) + self.assertIs(con, candidates[1]) + self.assertFalse(con.is_closed()) + self.assertTrue(_get_connected_host(con).endswith( + primary_spec['port'])) + self.assertEqual(await con.fetchval('SELECT 42'), 42) + self.assertEqual(terminated, [rejected]) + + # A later graceful close of the terminated connection is safe. + await rejected.close() + self.assertEqual(await con.fetchval('SELECT 43'), 43) + self.assertEqual(terminated, [rejected]) + finally: + for candidate in candidates: + candidate.terminate() + if con is not None: + await con.close() + async def _run_connection_test( self, connect, target_attribute, expected_port ):