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
5 changes: 5 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
@@ -1,3 +1,8 @@
## unreleased

* Do not wait for a connection closure


## v0.0.1 (2025-07-07)

* Initial release
6 changes: 5 additions & 1 deletion asyncpg_lock/guard.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,7 @@ class AdvisoryLockGuard:
"__reconnect_delay",
"__after_acquire_delay",
"__reacquire_delay",
"__tasks",
)

def __init__(
Expand All @@ -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,
Expand All @@ -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")
Expand Down
31 changes: 26 additions & 5 deletions tests/conftest.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand All @@ -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,
Expand All @@ -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:
Expand Down
16 changes: 16 additions & 0 deletions tests/test_guard.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
Loading