Skip to content
Open
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
18 changes: 10 additions & 8 deletions ably/transport/websockettransport.py
Original file line number Diff line number Diff line change
Expand Up @@ -257,10 +257,14 @@ async def dispose(self):
if self.idle_timer:
self.idle_timer.cancel()

# Schedule cleanup of cancelled tasks in the background to avoid blocking dispose()
# This prevents deadlock when dispose() is called from within these tasks
# Await the cancelled tasks so none is left unfinalized when dispose() returns. When
# dispose() runs inside one of them, awaiting it would deadlock, so that case is
# cleaned up in the background instead.
if tasks_to_await:
asyncio.create_task(self._cleanup_tasks(tasks_to_await))
if asyncio.current_task() in tasks_to_await:
asyncio.create_task(self._cleanup_tasks(tasks_to_await))
else:
await self._cleanup_tasks(tasks_to_await)

if self.websocket:
try:
Expand All @@ -270,11 +274,9 @@ async def dispose(self):

async def _cleanup_tasks(self, tasks):
"""Wait for cancelled tasks to complete their cleanup."""
for task in tasks:
try:
await task
except Exception:
pass # Ignore all exceptions from cancelled/failed tasks
# return_exceptions ignores what the cancelled or failed tasks raise, but still lets
# this coroutine be cancelled itself
await asyncio.gather(*tasks, return_exceptions=True)

async def close(self):
await self.send({'action': ProtocolMessageAction.CLOSE})
Expand Down
33 changes: 33 additions & 0 deletions test/unit/websockettransport_test.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,4 @@
import asyncio
from unittest.mock import MagicMock, patch

from ably.transport.websockettransport import WebSocketTransport
Expand Down Expand Up @@ -37,3 +38,35 @@ def test_websocket_url_defaults_to_wss_and_443():
def test_websocket_url_defaults_to_ws_and_80_when_tls_disabled():
url = _connect_url(tls=False)
assert url == 'ws://example.com:80?format=json'


# RTN12
async def test_dispose_finishes_cancelled_tasks_before_returning():
transport = WebSocketTransport(MagicMock(), 'example.com', {'format': 'json'})
started = asyncio.Event()

async def read_loop():
started.set()
await asyncio.Event().wait()

transport.read_loop = asyncio.create_task(read_loop())
await started.wait()

await transport.dispose()

assert transport.read_loop.cancelled()


# RTN12
async def test_dispose_called_from_the_read_loop_does_not_deadlock():
transport = WebSocketTransport(MagicMock(), 'example.com', {'format': 'json'})

async def read_loop():
await transport.dispose()

transport.read_loop = asyncio.create_task(read_loop())

done, pending = await asyncio.wait({transport.read_loop}, timeout=1)

assert not pending
assert transport.read_loop.cancelled()