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
1 change: 1 addition & 0 deletions asyncpg/protocol/protocol.pxd
Original file line number Diff line number Diff line change
Expand Up @@ -70,6 +70,7 @@ cdef class BaseProtocol(CoreProtocol):
cdef _on_result__copy_in(self, object waiter)

cdef _handle_waiter_on_connection_lost(self, cause)
cdef _complete_cancel_waiters(self)

cdef _dispatch_result(self)

Expand Down
61 changes: 44 additions & 17 deletions asyncpg/protocol/protocol.pyx
Original file line number Diff line number Diff line change
Expand Up @@ -591,7 +591,15 @@ cdef class BaseProtocol(CoreProtocol):
return not self.closing and self.con_status == CONNECTION_OK

def abort(self):
# Always finish pending cancel waiters. close() sets closing=True
# before awaiting them, so a later abort() must still unblock
# those futures and drop the transport.
self._complete_cancel_waiters()
if self.closing:
if self.transport is not None:
transport = self.transport
self.transport = None
transport.abort()
return
self.closing = True
self._handle_waiter_on_connection_lost(None)
Expand All @@ -604,26 +612,25 @@ cdef class BaseProtocol(CoreProtocol):
return

self.closing = True
timeout = self._get_timeout_impl(timeout)

if self.cancel_sent_waiter is not None:
await self.cancel_sent_waiter
self.cancel_sent_waiter = None

if self.cancel_waiter is not None:
await self.cancel_waiter

if self.waiter is not None:
# If there is a query running, cancel it
self._request_cancel()
await self.cancel_sent_waiter
self.cancel_sent_waiter = None
if self.cancel_waiter is not None:
await self.cancel_waiter
try:
if timeout is None:
await self._drain_cancels()
else:
await asyncio.wait_for(self._drain_cancels(), timeout)
except asyncio.TimeoutError:
# Server never acknowledged the cancel. Abort rather than
# hanging on cancel_waiter until the transport dies.
self._complete_cancel_waiters()
if self.transport is not None:
transport = self.transport
self.transport = None
transport.abort()
return

assert self.waiter is None

timeout = self._get_timeout_impl(timeout)

# Ask the server to terminate the connection and wait for it
# to drop.
self.waiter = self._new_waiter(timeout)
Expand All @@ -637,7 +644,9 @@ cdef class BaseProtocol(CoreProtocol):
pass
finally:
self.waiter = None
self.transport.abort()
if self.transport is not None:
self.transport.abort()
self.transport = None

def _request_cancel(self):
self.cancel_waiter = self.create_future()
Expand Down Expand Up @@ -686,6 +695,21 @@ cdef class BaseProtocol(CoreProtocol):
def _create_future_fallback(self):
return asyncio.Future(loop=self.loop)

cdef _complete_cancel_waiters(self):
if self.cancel_sent_waiter is not None and not self.cancel_sent_waiter.done():
self.cancel_sent_waiter.set_result(None)
self.cancel_sent_waiter = None
if self.cancel_waiter is not None and not self.cancel_waiter.done():
self.cancel_waiter.set_result(None)
self.cancel_waiter = None

async def _drain_cancels(self):
await self._wait_for_cancellation()
if self.waiter is not None:
# If there is a query running, cancel it
self._request_cancel()
await self._wait_for_cancellation()

cdef _handle_waiter_on_connection_lost(self, cause):
if self.waiter is not None and not self.waiter.done():
exc = apg_exc.ConnectionDoesNotExistError(
Expand All @@ -695,6 +719,8 @@ cdef class BaseProtocol(CoreProtocol):
exc.__cause__ = cause
self.waiter.set_exception(exc)
self.waiter = None
self._complete_cancel_waiters()


cdef _set_server_parameter(self, name, val):
self.settings.add_setting(name, val)
Expand Down Expand Up @@ -940,6 +966,7 @@ cdef class BaseProtocol(CoreProtocol):
else:
self.waiter.set_exception(exc)
self.waiter = None
self._complete_cancel_waiters()
else:
# The connection was lost because it was
# terminated or due to another error;
Expand Down
17 changes: 17 additions & 0 deletions tests/test_timeout.py
Original file line number Diff line number Diff line change
Expand Up @@ -152,3 +152,20 @@ async def test_timeout_covers_prepare_01(self):
with self.assertRaises(asyncio.TimeoutError):
meth = getattr(self.con, methname)
await meth('select pg_sleep($1)', 0.2)


class TestCloseTimeoutPendingCancel(tb.ConnectedTestCase):

async def test_close_times_out_pending_cancel(self):
with self.assertRaises(asyncio.TimeoutError):
await self.con.fetch('select pg_sleep(10)', timeout=0.05)

proto = self.con._protocol
proto.cancel_waiter = self.loop.create_future()
proto.cancel_sent_waiter = None

with self.assertRunUnder(1):
await self.con.close(timeout=0.15)

self.assertTrue(self.con.is_closed())