diff --git a/docs/versionhistory.rst b/docs/versionhistory.rst
index 81f656aec..b027e1c67 100644
--- a/docs/versionhistory.rst
+++ b/docs/versionhistory.rst
@@ -75,6 +75,9 @@ This library adheres to `Semantic Versioning 2.0 `_.
- Fixed an asyncio worker thread race that could raise ``RuntimeError`` when the event
loop closed between checking its state and scheduling the worker result
(`#1265 `_; PR by @hansu650)
+- Fixed concurrent ``aclose_forcefully()`` calls on asyncio socket streams returning
+ before the underlying socket was closed
+ (`#1273 `_; PR by @hansu650)
- Fixed ``CapacityLimiter`` on the asyncio backend over-granting tokens when
``total_tokens`` was raised while the limiter was over-subscribed
(`#1223 `_; PR by @zelinewang)
diff --git a/src/anyio/_backends/_asyncio.py b/src/anyio/_backends/_asyncio.py
index 35b0cf371..b3498f6e4 100644
--- a/src/anyio/_backends/_asyncio.py
+++ b/src/anyio/_backends/_asyncio.py
@@ -1397,15 +1397,8 @@ 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
-
- self._transport.close()
- await sleep(0)
- self._transport.abort()
+ self._transport.abort()
+ await sleep(0)
class _RawSocketMixin:
@@ -1689,7 +1682,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:
@@ -1739,7 +1733,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:
diff --git a/tests/test_sockets.py b/tests/test_sockets.py
index cc661b79a..a91700875 100644
--- a/tests/test_sockets.py
+++ b/tests/test_sockets.py
@@ -37,12 +37,14 @@
from anyio import (
BrokenResourceError,
BusyResourceError,
+ CancelScope,
ClosedResourceError,
EndOfStream,
Event,
TCPConnectable,
TypedAttributeLookupError,
UNIXConnectable,
+ aclose_forcefully,
as_connectable,
connect_tcp,
connect_unix,
@@ -54,6 +56,7 @@
create_unix_datagram_socket,
create_unix_listener,
fail_after,
+ get_cancelled_exc_class,
getaddrinfo,
getnameinfo,
move_on_after,
@@ -508,6 +511,40 @@ 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]
+
+ @pytest.mark.parametrize("anyio_backend", asyncio_params)
+ async def test_aclose_propagates_cancellation(
+ self, server_addr: tuple[str, int]
+ ) -> None:
+ stream = await connect_tcp(*server_addr)
+ cancelled_exc: BaseException | None = None
+
+ with CancelScope() as scope:
+ scope.cancel()
+ try:
+ await stream.aclose()
+ except get_cancelled_exc_class() as exc:
+ cancelled_exc = exc
+ raise
+
+ assert cancelled_exc is not None
+
async def test_receive_after_close(self, server_addr: tuple[str, int]) -> None:
stream = await connect_tcp(*server_addr)
await stream.aclose()
@@ -1764,6 +1801,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:
@@ -1951,6 +1999,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: