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
14 changes: 5 additions & 9 deletions asyncpg/connect_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
131 changes: 131 additions & 0 deletions tests/test_connect.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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
):
Expand Down
Loading