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: