From 3223d667243b9b85871a8219e59424bc6f8f74df Mon Sep 17 00:00:00 2001 From: slawwan Date: Tue, 6 Oct 2026 12:28:35 +0500 Subject: [PATCH] do not wait for a connection closure --- CHANGELOG.md | 5 +++++ asyncpg_lock/guard.py | 6 +++++- tests/conftest.py | 31 ++++++++++++++++++++++++++----- tests/test_guard.py | 16 ++++++++++++++++ 4 files changed, 52 insertions(+), 6 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 340e3db..40a6eab 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,3 +1,8 @@ +## unreleased + +* Do not wait for a connection closure + + ## v0.0.1 (2025-07-07) * Initial release diff --git a/asyncpg_lock/guard.py b/asyncpg_lock/guard.py index 10692d2..544edc9 100644 --- a/asyncpg_lock/guard.py +++ b/asyncpg_lock/guard.py @@ -23,6 +23,7 @@ class AdvisoryLockGuard: "__reconnect_delay", "__after_acquire_delay", "__reacquire_delay", + "__tasks", ) def __init__( @@ -44,6 +45,7 @@ def __init__( self.__after_acquire_delay = after_acquire_delay self.__reacquire_delay = reacquire_delay self.__connect = connect + self.__tasks: set[asyncio.Task[None]] = set() async def run( self, @@ -70,7 +72,9 @@ async def __ensure_connection_established(self, func: Callable[[asyncpg.Connecti try: await func(connection) finally: - await asyncio.shield(connection.close()) + close_task = asyncio.create_task(connection.close()) + close_task.add_done_callback(self.__tasks.discard) + self.__tasks.add(close_task) except Exception: if failed_attempts > 0: logger.exception("Connection closed or not established") diff --git a/tests/conftest.py b/tests/conftest.py index 46699b7..e8004e3 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -31,6 +31,8 @@ def __init__(self, *, src_port: int, dst_port: int): self.dst_host = "127.0.0.1" self.dst_port = dst_port self.connections = set() + self.links = set() + self.frozen_connections = set() async def start(self) -> None: await asyncio.start_server( @@ -40,23 +42,41 @@ async def start(self) -> None: ) async def drop_connections(self) -> None: + self.links.clear() + self.connections.update(self.frozen_connections) + self.frozen_connections.clear() while self.connections: writer = self.connections.pop() writer.close() if hasattr(writer, "wait_closed"): await writer.wait_closed() - @staticmethod - async def _pipe(reader: asyncio.StreamReader, writer: asyncio.StreamWriter): + async def freeze_connections(self) -> None: + """ + Simulates a silently dead network path: the server side is closed, + while the client side stays open and never receives anything. + """ + while self.links: + client_writer, server_writer = self.links.pop() + self.connections.discard(client_writer) + self.connections.discard(server_writer) + self.frozen_connections.add(client_writer) + self.frozen_connections.add(server_writer) + server_writer.close() + + async def _pipe(self, reader: asyncio.StreamReader, writer: asyncio.StreamWriter): try: while not reader.at_eof(): bytes_read = await reader.read(TcpProxy.MAX_BYTES) + if writer in self.frozen_connections: + continue writer.write(bytes_read) await writer.drain() finally: - writer.close() - if hasattr(writer, "wait_closed"): - await writer.wait_closed() + if writer not in self.frozen_connections: + writer.close() + if hasattr(writer, "wait_closed"): + await writer.wait_closed() async def _handle_client( self, @@ -67,6 +87,7 @@ async def _handle_client( self.connections.add(server_writer) self.connections.add(client_writer) + self.links.add((client_writer, server_writer)) try: async with asyncio.TaskGroup() as tg: diff --git a/tests/test_guard.py b/tests/test_guard.py index c649360..9da3620 100644 --- a/tests/test_guard.py +++ b/tests/test_guard.py @@ -218,6 +218,22 @@ async def test_reacquire_lock_after_disruption( assert not tracker.overlaps +async def test_reacquire_lock_after_silent_disruption( + guard: asyncpg_lock.AdvisoryLockGuard, connector: PgConnector, proxy: TcpProxy +) -> None: + tracker = ExecutionTracker(min_completed_executions=8, max_completed_executions=16) + task = asyncio.create_task(guard.run(LOCK_KEY, tracker)) + try: + await tracker.wait_min_completed() + await proxy.freeze_connections() + await tracker.wait_max_completed() + finally: + await cancel_and_wait(task) + + assert connector.total_open_connections == 2 + assert not tracker.overlaps + + async def test_no_overlapping_execution_for_same_keys( guard: asyncpg_lock.AdvisoryLockGuard, connector: PgConnector ) -> None: