Skip to content

Commit 309d166

Browse files
committed
Fix close_all loop-switch connection leak
1 parent 8c12adc commit 309d166

2 files changed

Lines changed: 27 additions & 2 deletions

File tree

tests/test_connection.py

Lines changed: 23 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -290,6 +290,29 @@ async def test_close_all_without_discard(mocked_db_config, conn_handler):
290290
assert conn_handler._storage == {"default": conn_1, "other": conn_2}
291291

292292

293+
@pytest.mark.asyncio
294+
@patch("tortoise.connection.ConnectionHandler._create_connection")
295+
@patch("tortoise.connection.ConnectionHandler.db_config", new_callable=PropertyMock)
296+
async def test_close_all_does_not_reconnect_on_loop_change(
297+
mocked_db_config, mocked_create_connection, conn_handler
298+
):
299+
stale_conn = Mock(_check_loop=Mock(return_value=False))
300+
stale_conn.close = AsyncMock()
301+
conn_handler._storage = {"default": stale_conn}
302+
conn_handler._db_config = {"default": {}}
303+
mocked_db_config.return_value = {"default": {}}
304+
305+
with warnings.catch_warnings(record=True) as w:
306+
warnings.simplefilter("always")
307+
await conn_handler.close_all()
308+
309+
stale_conn.close.assert_awaited_once()
310+
mocked_create_connection.assert_not_called()
311+
assert conn_handler._storage == {}
312+
loop_warnings = [x for x in w if issubclass(x.category, TortoiseLoopSwitchWarning)]
313+
assert loop_warnings == []
314+
315+
293316
# --- Event loop validation tests ---
294317

295318

tortoise/connection.py

Lines changed: 4 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -273,11 +273,13 @@ async def close_all(self, discard: bool = True) -> None:
273273
# Handle case where connections were never initialized (e.g., init failed)
274274
if self._db_config is None:
275275
return
276-
tasks = [conn.close() for conn in self.all()]
276+
storage = self._copy_storage()
277+
tasks = [storage[alias].close() for alias in self.db_config if alias in storage]
277278
await asyncio.gather(*tasks)
278279
if discard:
279280
for alias in self.db_config:
280-
self.discard(alias)
281+
if alias in storage:
282+
self.discard(alias)
281283

282284

283285
class _ConnectionsProxy:

0 commit comments

Comments
 (0)