Skip to content
Open
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
3 changes: 3 additions & 0 deletions docs/versionhistory.rst
Original file line number Diff line number Diff line change
Expand Up @@ -75,6 +75,9 @@ This library adheres to `Semantic Versioning 2.0 <http://semver.org/>`_.
- 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 <https://github.com/agronholm/anyio/issues/1265>`_; PR by @hansu650)
- 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)
- Fixed ``CapacityLimiter`` on the asyncio backend over-granting tokens when
``total_tokens`` was raised while the limiter was over-subscribed
(`#1223 <https://github.com/agronholm/anyio/pull/1223>`_; PR by @zelinewang)
Expand Down
17 changes: 6 additions & 11 deletions src/anyio/_backends/_asyncio.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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:
Expand Down
64 changes: 64 additions & 0 deletions tests/test_sockets.py
Original file line number Diff line number Diff line change
Expand Up @@ -37,12 +37,14 @@
from anyio import (
BrokenResourceError,
BusyResourceError,
CancelScope,
ClosedResourceError,
EndOfStream,
Event,
TCPConnectable,
TypedAttributeLookupError,
UNIXConnectable,
aclose_forcefully,
as_connectable,
connect_tcp,
connect_unix,
Expand All @@ -54,6 +56,7 @@
create_unix_datagram_socket,
create_unix_listener,
fail_after,
get_cancelled_exc_class,
getaddrinfo,
getnameinfo,
move_on_after,
Expand Down Expand Up @@ -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()
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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:
Expand Down
Loading