Skip to content
Open
Show file tree
Hide file tree
Changes from 3 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
3 changes: 3 additions & 0 deletions docs/versionhistory.rst
Original file line number Diff line number Diff line change
Expand Up @@ -52,6 +52,9 @@ This library adheres to `Semantic Versioning 2.0 <http://semver.org/>`_.
which triggers ``PytestRemovedIn10Warning`` on ``pytest>=9.2`` and crashes pytest at
startup when ``filterwarnings = error`` is configured
(`#1271 <https://github.com/agronholm/anyio/issues/1271>`_; PR by @matthewfeickert)
- Fixed concurrent ``aclose_forcefully()`` calls on asyncio socket streams returning
before the underlying socket was closed
(`#1273 <https://github.com/agronholm/anyio/issues/1273>`_; PR by @hansu650)

**4.14.2**

Expand Down
27 changes: 18 additions & 9 deletions src/anyio/_backends/_asyncio.py
Original file line number Diff line number Diff line change
Expand Up @@ -1250,6 +1250,7 @@ class StreamProtocol(asyncio.Protocol):
write_event: asyncio.Event
exception: Exception | None = None
is_at_eof: bool = False
is_connection_lost: bool = False

def connection_made(self, transport: asyncio.BaseTransport) -> None:
self.read_queue = deque()
Expand All @@ -1259,6 +1260,7 @@ def connection_made(self, transport: asyncio.BaseTransport) -> None:
cast(asyncio.Transport, transport).set_write_buffer_limits(0)

def connection_lost(self, exc: Exception | None) -> None:
self.is_connection_lost = True
if exc:
self.exception = exc

Expand Down Expand Up @@ -1393,15 +1395,20 @@ async def send_eof(self) -> None:

async def aclose(self) -> None:
self._closed = True
if not self._transport.is_closing():
try:
self._transport.write_eof()
except OSError:
pass
try:
if not self._transport.is_closing():
try:
self._transport.write_eof()
except OSError:
pass

self._transport.close()
await sleep(0)
self._transport.close()
await sleep(0)
Comment thread
hansu650 marked this conversation as resolved.
Outdated
finally:
self._transport.abort()
if not self._protocol.is_connection_lost:
with CancelScope(shield=True):
await sleep(0)
Comment thread
hansu650 marked this conversation as resolved.
Outdated


class _RawSocketMixin:
Expand Down Expand Up @@ -1685,7 +1692,8 @@ async def aclose(self) -> None:
if not self._transport.is_closing():
self._transport.close()

await self._protocol.closed_event.wait()
with CancelScope(shield=True):
await self._protocol.closed_event.wait()

async def receive(self) -> tuple[bytes, IPSockAddrType]:
with self._receive_guard:
Expand Down Expand Up @@ -1735,7 +1743,8 @@ async def aclose(self) -> None:
if not self._transport.is_closing():
self._transport.close()

await self._protocol.closed_event.wait()
with CancelScope(shield=True):
await self._protocol.closed_event.wait()

async def receive(self) -> bytes:
with self._receive_guard:
Expand Down
45 changes: 45 additions & 0 deletions tests/test_sockets.py
Original file line number Diff line number Diff line change
Expand Up @@ -43,6 +43,7 @@
TCPConnectable,
TypedAttributeLookupError,
UNIXConnectable,
aclose_forcefully,
as_connectable,
connect_tcp,
connect_unix,
Expand Down Expand Up @@ -507,6 +508,23 @@ async def interrupt() -> None:
with pytest.raises(ClosedResourceError):
await stream.receive()

async def test_concurrent_aclose_forcefully_waits_for_socket_close(
self, server_addr: tuple[str, int]
) -> None:
stream = await connect_tcp(*server_addr)
raw_socket = stream.extra(SocketAttribute.raw_socket)
file_descriptors: list[int] = []

async def close_stream() -> None:
await aclose_forcefully(stream)
file_descriptors.append(raw_socket.fileno())

async with create_task_group() as task_group:
task_group.start_soon(close_stream)
task_group.start_soon(close_stream)

assert file_descriptors == [-1, -1]

async def test_receive_after_close(self, server_addr: tuple[str, int]) -> None:
stream = await connect_tcp(*server_addr)
await stream.aclose()
Expand Down Expand Up @@ -1746,6 +1764,17 @@ async def handle(stream: SocketStream) -> None:
@pytest.mark.network
@pytest.mark.usefixtures("check_asyncio_bug")
class TestUDPSocket:
async def test_aclose_forcefully_waits_for_fd_release(
self, family: AnyIPAddressFamily
) -> None:
host = "127.0.0.1" if family == socket.AF_INET else "::1"
udp = await create_udp_socket(family=family, local_host=host)
raw_socket = udp.extra(SocketAttribute.raw_socket)

await aclose_forcefully(udp)

assert raw_socket.fileno() == -1

async def test_aclose_waits_for_fd_release(
self, family: AnyIPAddressFamily, free_udp_port: int
) -> None:
Expand Down Expand Up @@ -1933,6 +1962,22 @@ async def test_from_socket_wrong_socket_type(
@pytest.mark.network
@pytest.mark.usefixtures("check_asyncio_bug")
class TestConnectedUDPSocket:
async def test_aclose_forcefully_waits_for_fd_release(
self, family: AnyIPAddressFamily
) -> None:
host = "127.0.0.1" if family == socket.AF_INET else "::1"
peer = socket.socket(family, socket.SOCK_DGRAM)
peer.bind((host, 0))
try:
udp = await create_connected_udp_socket(*peer.getsockname()[:2])
raw_socket = udp.extra(SocketAttribute.raw_socket)

await aclose_forcefully(udp)

assert raw_socket.fileno() == -1
finally:
peer.close()

async def test_aclose_waits_for_fd_release(
self, family: AnyIPAddressFamily, free_udp_port: int
) -> None:
Expand Down
Loading