Repository navigation
Add a native asyncio network backend #1170
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Open
Kludex
wants to merge
6
commits into
pool-incremental-state
Choose a base branch
from
asyncio-backend
base: pool-incremental-state
Could not load branches
Branch not found: {{ refName }}
Loading
Could not load tags
Nothing to show
Loading
Are you sure you want to change the base?
Some commits from the old base branch may be removed from the timeline,
and old review comments may become outdated.
Open
Changes from all commits
Commits
Show all changes
6 commits
Select commit
Hold shift + click to select a range
bfaa5bb
Add a native asyncio network backend
Kludex 883aefd
Report socket option failures as connection errors
Kludex 6d7231d
Harden TLS upgrades and connection attempts in the asyncio backend
Kludex 9871eef
Tolerate event loops whose connection method has no signature
Kludex 777cd86
Add the original httpx to the benchmark harness
Kludex c6d2d00
Run the sync integration tests once and document the benchmark's refe…
Kludex File filter
Filter by extension
Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
There are no files selected for viewing
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,299 @@ | ||
| from __future__ import annotations | ||
|
|
||
| import asyncio | ||
| import collections | ||
| import inspect | ||
| import ssl | ||
| import typing | ||
|
|
||
| from .._exceptions import ( | ||
| ConnectError, | ||
| ConnectTimeout, | ||
| ReadError, | ||
| ReadTimeout, | ||
| WriteError, | ||
| WriteTimeout, | ||
| ) | ||
| from .base import SOCKET_OPTION, AsyncNetworkBackend, AsyncNetworkStream | ||
|
|
||
| # Stop reading from the socket once this much data is buffered but unread, | ||
| # and start again once the buffer drains below the low-water mark. | ||
| RECEIVE_HIGH_WATER = 256 * 1024 | ||
| RECEIVE_LOW_WATER = 64 * 1024 | ||
|
|
||
| # Stagger connection attempts across resolved addresses, as in RFC 8305. | ||
| HAPPY_EYEBALLS_DELAY = 0.25 | ||
| # The event loop insists on a TLS handshake timeout; this stands in for "none". | ||
| NO_HANDSHAKE_TIMEOUT = 365 * 24 * 60 * 60.0 | ||
|
|
||
| _happy_eyeballs_support: dict[type[asyncio.AbstractEventLoop], bool] = {} | ||
|
|
||
|
|
||
| def _connection_kwargs(loop: asyncio.AbstractEventLoop) -> dict[str, typing.Any]: | ||
| # Not every event loop implements Happy Eyeballs; use it where available. | ||
| supported = _happy_eyeballs_support.get(type(loop)) | ||
| if supported is None: | ||
| try: | ||
| supported = "happy_eyeballs_delay" in inspect.signature(loop.create_connection).parameters | ||
| except (TypeError, ValueError): | ||
| # Some extension-implemented loops expose no signature to inspect. | ||
| supported = False | ||
| _happy_eyeballs_support[type(loop)] = supported | ||
| return {"happy_eyeballs_delay": HAPPY_EYEBALLS_DELAY} if supported else {} | ||
|
|
||
|
|
||
| class _Timeout(Exception): | ||
| """ | ||
| Raised on a waiter future when its deadline passes. | ||
| """ | ||
|
|
||
|
|
||
| def _timeout(waiter: asyncio.Future[None]) -> None: | ||
| if not waiter.done(): | ||
| waiter.set_exception(_Timeout()) | ||
|
|
||
|
|
||
| def _wake(waiter: asyncio.Future[None] | None) -> None: | ||
| if waiter is not None and not waiter.done(): | ||
| waiter.set_result(None) | ||
|
|
||
|
|
||
| class AsyncioStreamProtocol(asyncio.Protocol): | ||
| """ | ||
| Buffers received data for `AsyncioStream`, applying backpressure to the | ||
| transport once too much is buffered, and wakes up pending reads and writes. | ||
| """ | ||
|
|
||
| def __init__(self) -> None: | ||
| self.transport: asyncio.Transport | None = None | ||
| self.chunks: collections.deque[bytes] = collections.deque() | ||
| self.buffered = 0 | ||
| self.reading_paused = False | ||
| self.writing_paused = False | ||
| self.eof = False | ||
| self.closed = False | ||
| self.exception: Exception | None = None | ||
| self.read_waiter: asyncio.Future[None] | None = None | ||
| self.write_waiter: asyncio.Future[None] | None = None | ||
|
|
||
| def connection_made(self, transport: asyncio.BaseTransport) -> None: | ||
| # The transport implements the interface without necessarily subclassing it. | ||
| self.transport = typing.cast(asyncio.Transport, transport) | ||
|
|
||
| def data_received(self, data: bytes) -> None: | ||
| self.chunks.append(data) | ||
| self.buffered += len(data) | ||
| if self.buffered >= RECEIVE_HIGH_WATER and not self.reading_paused: | ||
| assert self.transport is not None | ||
| self.transport.pause_reading() | ||
| self.reading_paused = True | ||
| _wake(self.read_waiter) | ||
|
|
||
| def eof_received(self) -> bool: | ||
| self.eof = True | ||
| _wake(self.read_waiter) | ||
| # Let the transport close: an HTTP peer that has sent a FIN is done. | ||
| return False | ||
|
|
||
| def connection_lost(self, exc: Exception | None) -> None: | ||
| self.closed = True | ||
| self.exception = exc | ||
| _wake(self.read_waiter) | ||
| _wake(self.write_waiter) | ||
|
|
||
| def pause_writing(self) -> None: | ||
| self.writing_paused = True | ||
|
|
||
| def resume_writing(self) -> None: | ||
| self.writing_paused = False | ||
| _wake(self.write_waiter) | ||
|
|
||
|
|
||
| class AsyncioStream(AsyncNetworkStream): | ||
| def __init__(self, transport: asyncio.Transport, protocol: AsyncioStreamProtocol) -> None: | ||
| self._transport = transport | ||
| self._protocol = protocol | ||
|
|
||
| async def read(self, max_bytes: int, timeout: float | None = None) -> bytes: | ||
| protocol = self._protocol | ||
| if not protocol.chunks: | ||
| if protocol.eof or protocol.closed: | ||
| return self._read_at_end() | ||
| try: | ||
| await self._wait("read", timeout) | ||
| except _Timeout: | ||
| raise ReadTimeout("timed out") from None | ||
| if not protocol.chunks: | ||
| return self._read_at_end() | ||
|
|
||
| chunk = protocol.chunks[0] | ||
| if len(chunk) <= max_bytes: | ||
| protocol.chunks.popleft() | ||
| else: | ||
| protocol.chunks[0] = chunk[max_bytes:] | ||
| chunk = chunk[:max_bytes] | ||
| protocol.buffered -= len(chunk) | ||
| if protocol.reading_paused and protocol.buffered <= RECEIVE_LOW_WATER: | ||
| protocol.reading_paused = False | ||
| self._transport.resume_reading() | ||
| return chunk | ||
|
|
||
| def _read_at_end(self) -> bytes: | ||
| # No buffered data and no more coming: a clean EOF reads as empty, | ||
| # a connection dropped by an error is a read error. | ||
| if self._protocol.exception is not None: | ||
| raise ReadError(str(self._protocol.exception)) from self._protocol.exception | ||
| return b"" | ||
|
|
||
| async def write(self, buffer: bytes, timeout: float | None = None) -> None: | ||
| if not buffer: | ||
| return | ||
| protocol = self._protocol | ||
| if protocol.closed or self._transport.is_closing(): | ||
| raise WriteError("Connection closed") | ||
| self._transport.write(buffer) | ||
| if protocol.writing_paused: | ||
| # The transport's send buffer is full; wait for it to drain. | ||
| try: | ||
| await self._wait("write", timeout) | ||
| except _Timeout: | ||
| raise WriteTimeout("timed out") from None | ||
| if protocol.closed: | ||
| raise WriteError(str(protocol.exception or "Connection closed")) | ||
|
|
||
| async def _wait(self, kind: str, timeout: float | None) -> None: | ||
| loop = asyncio.get_running_loop() | ||
| waiter: asyncio.Future[None] = loop.create_future() | ||
| protocol = self._protocol | ||
| if kind == "read": | ||
| protocol.read_waiter = waiter | ||
| else: | ||
| protocol.write_waiter = waiter | ||
| handle = None if timeout is None else loop.call_later(timeout, _timeout, waiter) | ||
| try: | ||
| await waiter | ||
| finally: | ||
| if handle is not None: | ||
| handle.cancel() | ||
| if kind == "read": | ||
| protocol.read_waiter = None | ||
| else: | ||
| protocol.write_waiter = None | ||
|
|
||
| async def aclose(self) -> None: | ||
| if self._protocol.closed: | ||
| return | ||
| self._transport.close() | ||
| # Closing only schedules the socket close on the event loop. Yield once | ||
| # so it runs now, then force it if unsent data is still holding it up. | ||
| await asyncio.sleep(0) | ||
| if not self._protocol.closed: | ||
| self._transport.abort() | ||
|
|
||
| async def start_tls( | ||
| self, | ||
| ssl_context: ssl.SSLContext, | ||
| server_hostname: str | None = None, | ||
| timeout: float | None = None, | ||
| ) -> AsyncNetworkStream: | ||
| protocol = self._protocol | ||
| if protocol.chunks or protocol.eof or protocol.closed: | ||
| # Nothing may arrive before the handshake: anything already buffered | ||
| # is plaintext that must not be mistaken for data received over TLS. | ||
| await self.aclose() | ||
| raise ConnectError("Received unexpected data before the TLS handshake") | ||
|
|
||
| loop = asyncio.get_running_loop() | ||
| # The loop reports its own handshake timeout as a connection error, so | ||
| # the deadline is applied here to raise a timeout, with the loop's own | ||
| # deadline kept out of the way. | ||
| handshake = loop.start_tls( | ||
| self._transport, | ||
| protocol, | ||
| ssl_context, | ||
| server_hostname=server_hostname, | ||
| ssl_handshake_timeout=NO_HANDSHAKE_TIMEOUT if timeout is None else timeout + 1.0, | ||
| ) | ||
| try: | ||
| transport = await asyncio.wait_for(handshake, timeout) | ||
| except (TimeoutError, asyncio.TimeoutError): | ||
| self._transport.close() | ||
| raise ConnectTimeout("timed out") from None | ||
| except (OSError, ssl.SSLError) as exc: | ||
| self._transport.close() | ||
| raise ConnectError(str(exc)) from exc | ||
| if transport is None: # pragma: no cover | ||
| raise ConnectError("TLS handshake failed") | ||
| protocol.transport = transport | ||
| return AsyncioStream(transport, protocol) | ||
|
|
||
| def get_extra_info(self, info: str) -> typing.Any: | ||
| if info == "ssl_object": | ||
| return self._transport.get_extra_info("ssl_object") | ||
| if info == "client_addr": | ||
| return self._transport.get_extra_info("sockname") | ||
| if info == "server_addr": | ||
| return self._transport.get_extra_info("peername") | ||
| if info == "socket": | ||
| return self._transport.get_extra_info("socket") | ||
| if info == "is_readable": | ||
| # The event loop keeps reading while the connection is idle, so a | ||
| # FIN or stray data from the server is already known here without | ||
| # touching the socket. | ||
| protocol = self._protocol | ||
| return bool(protocol.chunks) or protocol.eof or protocol.closed | ||
| return None | ||
|
|
||
|
|
||
| class AsyncioBackend(AsyncNetworkBackend): | ||
| async def connect_tcp( | ||
| self, | ||
| host: str, | ||
| port: int, | ||
| timeout: float | None = None, | ||
| local_address: str | None = None, | ||
| socket_options: typing.Iterable[SOCKET_OPTION] | None = None, | ||
| ) -> AsyncNetworkStream: | ||
| loop = asyncio.get_running_loop() | ||
| local_addr = None if local_address is None else (local_address, 0) | ||
| # By default TCP sockets opened in `asyncio` include TCP_NODELAY. | ||
| connect = loop.create_connection( | ||
| AsyncioStreamProtocol, host, port, local_addr=local_addr, **_connection_kwargs(loop) | ||
| ) | ||
| return await self._connect(connect, timeout, socket_options) | ||
|
|
||
| async def connect_unix_socket( | ||
| self, | ||
| path: str, | ||
| timeout: float | None = None, | ||
| socket_options: typing.Iterable[SOCKET_OPTION] | None = None, | ||
| ) -> AsyncNetworkStream: | ||
| loop = asyncio.get_running_loop() | ||
| connect = loop.create_unix_connection(AsyncioStreamProtocol, path) | ||
| return await self._connect(connect, timeout, socket_options) | ||
|
|
||
| async def _connect( | ||
| self, | ||
| connect: typing.Coroutine[typing.Any, typing.Any, tuple[asyncio.BaseTransport, AsyncioStreamProtocol]], | ||
| timeout: float | None, | ||
| socket_options: typing.Iterable[SOCKET_OPTION] | None, | ||
| ) -> AsyncioStream: | ||
| try: | ||
| transport, protocol = await asyncio.wait_for(connect, timeout) | ||
| except (TimeoutError, asyncio.TimeoutError): | ||
| raise ConnectTimeout("timed out") from None | ||
| except OSError as exc: | ||
| raise ConnectError(str(exc)) from exc | ||
| stream = AsyncioStream(typing.cast(asyncio.Transport, transport), protocol) | ||
| if socket_options: | ||
| sock = transport.get_extra_info("socket") | ||
| try: | ||
| for option in socket_options: | ||
| sock.setsockopt(*option) | ||
| except OSError as exc: | ||
| await stream.aclose() | ||
| raise ConnectError(str(exc)) from exc | ||
| return stream | ||
|
|
||
| async def sleep(self, seconds: float) -> None: | ||
| await asyncio.sleep(seconds) |
Oops, something went wrong.
Add this suggestion to a batch that can be applied as a single commit.
This suggestion is invalid because no changes were made to the code.
Suggestions cannot be applied while the pull request is closed.
Suggestions cannot be applied while viewing a subset of changes.
Only one suggestion per line can be applied in a batch.
Add this suggestion to a batch that can be applied as a single commit.
Applying suggestions on deleted lines is not supported.
You must change the existing code in this line in order to create a valid suggestion.
Outdated suggestions cannot be applied.
This suggestion has been applied or marked resolved.
Suggestions cannot be applied from pending reviews.
Suggestions cannot be applied on multi-line comments.
Suggestions cannot be applied while the pull request is queued to merge.
Suggestion cannot be applied right now. Please check back later.
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
P2: When
--lib httpxis selected through the documented benchmark environment, this import fails because the originalhttpxpackage is not declared or locked by the project. Addhttpxto the benchmark dependency group, or document and provision it separately.Prompt for AI agents
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
Done in c6d2d00, by documenting it:
httpxcannot join the bench group because installing it next tohttpx2breaks the alias tests' type-checks, so the README now shows how to provisionhttpx/punkreqin a separate interpreter passed via--python, and the import site says so.