From c2fe7cbba205817dd367ff3a325698ef5b04c3c5 Mon Sep 17 00:00:00 2001 From: 0x676e67 Date: Sun, 4 Oct 2026 16:59:12 +0000 Subject: [PATCH 1/4] fix(ws): read and write on separate tasks so a pending recv never blocks --- src/client/resp/ws.rs | 6 +- src/client/resp/ws/cmd.rs | 235 ++++++++++++++++++++------------------ tests/websocket_test.py | 90 ++++++++++++--- 3 files changed, 198 insertions(+), 133 deletions(-) diff --git a/src/client/resp/ws.rs b/src/client/resp/ws.rs index 95d94458..13711484 100644 --- a/src/client/resp/ws.rs +++ b/src/client/resp/ws.rs @@ -8,7 +8,6 @@ use std::{ use msg::Message; use pyo3::prelude::*; -use tokio::sync::mpsc; use wreq::{header::HeaderValue, ws::WebSocketResponse}; use crate::{ @@ -44,7 +43,7 @@ pub struct WebSocket { #[pyo3(get)] headers: HeaderMap, protocol: Option, - cmd: mpsc::UnboundedSender, + cmd: cmd::Handle, runtime: Runtime, } @@ -66,8 +65,7 @@ impl WebSocket { ); let websocket = response.into_websocket().await?; let protocol = websocket.protocol().cloned(); - let (cmd, rx) = mpsc::unbounded_channel(); - tokio::spawn(cmd::task(websocket, rx)); + let cmd = cmd::spawn(websocket); Ok(WebSocket { runtime, diff --git a/src/client/resp/ws/cmd.rs b/src/client/resp/ws/cmd.rs index c9e4ca07..286667bf 100644 --- a/src/client/resp/ws/cmd.rs +++ b/src/client/resp/ws/cmd.rs @@ -1,21 +1,25 @@ -//! WebSocket Command Utilities +//! The background tasks behind a [`WebSocket`](super::WebSocket) and the requests Python +//! sends them. //! -//! This module defines the `Command` enum for representing WebSocket operations -//! (send, receive, close) and provides async helpers for sending commands to the -//! WebSocket background task. It enables safe, concurrent, and ergonomic control -//! of WebSocket communication from Python bindings. +//! Reads and writes run on separate tasks, so a pending receive never holds up a send or +//! close. A receive whose caller is gone stops waiting, and a message read just as its +//! caller left is kept for the next receive. Closing also ends the read task. use std::time::Duration; -use futures_util::{SinkExt, StreamExt, TryStreamExt, stream}; +use futures_util::{ + SinkExt, StreamExt, TryStreamExt, + stream::{self, SplitSink, SplitStream}, +}; use pyo3::prelude::*; use tokio::{ sync::{ - mpsc::{UnboundedReceiver, UnboundedSender}, + mpsc::{self, UnboundedReceiver, UnboundedSender}, oneshot::{self, Sender}, }, time, }; +use tokio_util::sync::CancellationToken; use wreq::ws::{ WebSocket, message::{self, CloseCode, CloseFrame, Utf8Bytes}, @@ -24,79 +28,105 @@ use wreq::ws::{ use super::Message; use crate::{error::Error, extractor::Text}; -/// Commands for WebSocket operations. -pub enum Command { - /// Send a WebSocket message. - /// - /// Contains the message to send and a oneshot sender for the result. - Send(Message, Sender>), - - /// Send multiple WebSocket messages. - /// - /// Contains a vector of messages to send and a oneshot sender for the result. - SendMany(Vec, Sender>), +/// The request channels of a WebSocket's read and write tasks. +#[derive(Clone)] +pub struct Handle { + reads: UnboundedSender, + writes: UnboundedSender, +} - /// Receive a WebSocket message. - /// - /// Contains an optional timeout and a oneshot sender for the result. - Recv(Option, Sender>>), +/// A receive with an optional timeout. +struct Read(Option, Sender>>); - /// Close the WebSocket connection. - /// - /// Contains an optional close code, optional reason, and a oneshot sender for the result. +/// A write to the WebSocket. +enum Write { + Send(Message, Sender>), + SendMany(Vec, Sender>), Close(Option, Option, Sender>), } -/// The main background task that processes incoming [`Command`]s and interacts with the WebSocket. -/// -/// Handles sending, receiving, and closing the WebSocket connection based on received commands. -pub async fn task(ws: WebSocket, mut cmd: UnboundedReceiver) { - let (mut writer, mut reader) = ws.split(); - while let Some(command) = cmd.recv().await { - match command { - Command::Send(msg, tx) => { - let res = writer - .send(msg.0) - .await - .map_err(Error::Library) - .map_err(Into::into); +/// Start the read and write tasks of `ws` on the current runtime. +pub fn spawn(ws: WebSocket) -> Handle { + let (writer, reader) = ws.split(); + let (reads, read_rx) = mpsc::unbounded_channel(); + let (writes, write_rx) = mpsc::unbounded_channel(); + let closed = CancellationToken::new(); + tokio::spawn(read(reader, read_rx, closed.clone())); + tokio::spawn(write(writer, write_rx, closed)); + Handle { reads, writes } +} - let _ = tx.send(res); +/// Serve receives in order until the WebSocket is closed or dropped. +async fn read( + mut reader: SplitStream, + mut reads: UnboundedReceiver, + closed: CancellationToken, +) { + let mut unclaimed = None; + loop { + let Read(timeout, mut tx) = tokio::select! { + biased; + _ = closed.cancelled() => return, + read = reads.recv() => match read { + Some(read) => read, + None => return, + }, + }; + if let Some(res) = unclaimed.take() { + if let Err(res) = tx.send(res) { + unclaimed = Some(res); } - Command::SendMany(many_msg, tx) => { - let messages = many_msg.into_iter().map(|m| Ok(m.0)); - let res = writer - .send_all(&mut stream::iter(messages)) + continue; + } + let next = async { + match timeout { + Some(timeout) => time::timeout(timeout, reader.try_next()) .await - .map_err(Error::Library) - .map_err(Into::into); - - let _ = tx.send(res); + .map_err(Error::Timeout), + None => Ok(reader.try_next().await), } - Command::Recv(timeout, tx) => { - let fut = async { - reader - .try_next() - .await - .map(|opt| opt.map(Message)) - .map_err(Error::Library) - .map_err(Into::into) - }; - - if let Some(timeout) = timeout { - match time::timeout(timeout, fut).await { - Ok(res) => { - let _ = tx.send(res); - } - Err(err) => { - let _ = tx.send(Err(Error::Timeout(err).into())); - } - } - } else { - let _ = tx.send(fut.await); + }; + // Reading the stream is cancel-safe: leaving for a gone caller loses nothing. + let res = tokio::select! { + biased; + _ = closed.cancelled() => return, + _ = tx.closed() => continue, + next = next => match next { + Ok(next) => next.map(|msg| msg.map(Message)).map_err(Error::Library), + Err(timeout) => { + let _ = tx.send(Err(timeout.into())); + continue; } + }, + }; + // The caller may have left after the read finished; keep it for the next receive. + if let Err(res) = tx.send(res.map_err(Into::into)) { + unclaimed = Some(res); + } + } +} + +/// Serve writes in order until a close or the WebSocket is dropped, then end the read task. +async fn write( + mut writer: SplitSink, + mut writes: UnboundedReceiver, + closed: CancellationToken, +) { + let _closed = closed.drop_guard(); + while let Some(write) = writes.recv().await { + match write { + // A caller gone before its write starts does not send it. + Write::Send(_, tx) | Write::SendMany(_, tx) if tx.is_closed() => {} + Write::Send(msg, tx) => { + let res = writer.send(msg.0).await.map_err(Error::Library); + let _ = tx.send(res.map_err(Into::into)); + } + Write::SendMany(messages, tx) => { + let mut messages = stream::iter(messages.into_iter().map(|msg| Ok(msg.0))); + let res = writer.send_all(&mut messages).await.map_err(Error::Library); + let _ = tx.send(res.map_err(Into::into)); } - Command::Close(code, reason, tx) => { + Write::Close(code, reason, tx) => { let reason = reason .map(|reason| reason.0) .map(Utf8Bytes::try_from) @@ -114,85 +144,68 @@ pub async fn task(ws: WebSocket, mut cmd: UnboundedReceiver) { let res = writer .send(message::Message::Close(close_frame)) .await - .map_err(Error::Library) - .map_err(Into::into); + .map_err(Error::Library); let _ = writer.close().await; - let _ = tx.send(res); - break; + let _ = tx.send(res.map_err(Into::into)); + return; } } } } -/// Sends a [`Command::Recv`] to the background task and awaits a message from the WebSocket. -/// -/// Returns the received message or an error if the connection is closed or timeout. +/// Receive the next message, or `None` once the peer has closed. #[inline] -pub async fn recv( - cmd: UnboundedSender, - timeout: Option, -) -> PyResult> { - send_command(cmd, |tx| Command::Recv(timeout, tx)) +pub async fn recv(handle: Handle, timeout: Option) -> PyResult> { + request(&handle.reads, |tx| Read(timeout, tx)) .await .ok_or(Error::WebSocketDisconnected)? } -/// Sends a [`Command::Send`] to the background task to transmit a message over the WebSocket. -/// -/// Returns Ok if the message was sent successfully, or an error otherwise. +/// Send a message. #[inline] -pub async fn send(cmd: UnboundedSender, message: Message) -> PyResult<()> { - send_command(cmd, |tx| Command::Send(message, tx)) +pub async fn send(handle: Handle, message: Message) -> PyResult<()> { + request(&handle.writes, |tx| Write::Send(message, tx)) .await .ok_or(Error::WebSocketDisconnected)? } -/// Send as [`Command::SendMany`] to the background task to transmit multiple messages over the -/// WebSocket. -/// -/// Returns Ok if all messages were sent successfully, or an error otherwise. +/// Send messages in order. #[inline] -pub async fn send_all(cmd: UnboundedSender, messages: Vec) -> PyResult<()> { +pub async fn send_all(handle: Handle, messages: Vec) -> PyResult<()> { if messages.is_empty() { return Ok(()); } - send_command(cmd, |tx| Command::SendMany(messages, tx)) + request(&handle.writes, |tx| Write::SendMany(messages, tx)) .await .ok_or(Error::WebSocketDisconnected)? } -/// Sends a [`Command::Close`] to the background task to gracefully close the WebSocket connection. -/// -/// Returns Ok if the connection was closed successfully, or an error otherwise. +/// Send a close frame and close the connection. #[inline] -pub async fn close( - cmd: UnboundedSender, - code: Option, - reason: Option, -) -> PyResult<()> { - send_command(cmd, |tx| Command::Close(code, reason, tx)) +pub async fn close(handle: Handle, code: Option, reason: Option) -> PyResult<()> { + request(&handle.writes, |tx| Write::Close(code, reason, tx)) .await .ok_or(Error::WebSocketDisconnected)? } -/// Closes the WebSocket like [`close`], treating an already closed connection as done, as a -/// context manager exit does. +/// Close like [`close`], treating an already closed connection as done, as a context +/// manager exit does. #[inline] -pub async fn close_on_exit(cmd: UnboundedSender) -> PyResult<()> { - send_command(cmd, |tx| Command::Close(None, None, tx)) +pub async fn close_on_exit(handle: Handle) -> PyResult<()> { + request(&handle.writes, |tx| Write::Close(None, None, tx)) .await .unwrap_or(Ok(())) } -/// Run a command on the background task, or return `None` once it has ended. -async fn send_command( - cmd: UnboundedSender, - make: impl FnOnce(oneshot::Sender) -> Command, +/// Run a request on a task, or return `None` once the task has ended. +async fn request( + requests: &UnboundedSender, + make: impl FnOnce(oneshot::Sender) -> R, ) -> Option { - if cmd.is_closed() { + if requests.is_closed() { return None; } let (tx, rx) = oneshot::channel(); - cmd.send(make(tx)).ok()?; + requests.send(make(tx)).ok()?; rx.await.ok() } diff --git a/tests/websocket_test.py b/tests/websocket_test.py index bc425bf5..4d9ff0e9 100644 --- a/tests/websocket_test.py +++ b/tests/websocket_test.py @@ -11,32 +11,40 @@ GUID = b"258EAFA5-E914-47DA-95CA-C5AB0DC85B11" -async def read_close(reader): - """Read one masked client frame and return its close code and reason.""" - opcode, size = await reader.readexactly(2) - assert opcode == 0x88 +async def read_frame(reader): + """Read one masked client frame and return its opcode and payload.""" + first, size = await reader.readexactly(2) mask = await reader.readexactly(4) - payload = bytes( - b ^ mask[i % 4] for i, b in enumerate(await reader.readexactly(size & 0x7F)) - ) + payload = await reader.readexactly(size & 0x7F) + return first & 0x0F, bytes(b ^ mask[i % 4] for i, b in enumerate(payload)) + + +async def read_close(reader): + """Read one masked client close frame and return its code and reason.""" + opcode, payload = await read_frame(reader) + assert opcode == 0x8 if not payload: return None, "" return struct.unpack("!H", payload[:2])[0], payload[2:].decode() +async def handshake(reader, writer): + head = await reader.readuntil(b"\r\n\r\n") + key = next( + line.split(b":", 1)[1].strip() + for line in head.split(b"\r\n") + if line.lower().startswith(b"sec-websocket-key:") + ) + accept = base64.b64encode(hashlib.sha1(key + GUID).digest()) + writer.write( + b"HTTP/1.1 101 Switching Protocols\r\nUpgrade: websocket\r\n" + b"Connection: Upgrade\r\nSec-WebSocket-Accept: " + accept + b"\r\n\r\n" + ) + + async def serve(reader, writer, frames): try: - head = await reader.readuntil(b"\r\n\r\n") - key = next( - line.split(b":", 1)[1].strip() - for line in head.split(b"\r\n") - if line.lower().startswith(b"sec-websocket-key:") - ) - accept = base64.b64encode(hashlib.sha1(key + GUID).digest()) - writer.write( - b"HTTP/1.1 101 Switching Protocols\r\nUpgrade: websocket\r\n" - b"Connection: Upgrade\r\nSec-WebSocket-Accept: " + accept + b"\r\n\r\n" - ) + await handshake(reader, writer) frames.put_nowait(await read_close(reader)) writer.write(b"\x88\x00") await writer.drain() @@ -92,3 +100,49 @@ def test_blocking_websocket_exit_after_close(): server.close() loop.run_until_complete(asyncio.wait_for(server.wait_closed(), 5)) loop.close() + + +@pytest.mark.asyncio +async def test_websocket_reads_and_writes_do_not_block_each_other(): + received = asyncio.Queue() + release = asyncio.Event() + + async def serve_messages(reader, writer): + async def send_later(): + await release.wait() + for text in (b"first", b"second"): + writer.write(bytes([0x81, len(text)]) + text) + await writer.drain() + + sender = None + try: + await handshake(reader, writer) + sender = asyncio.ensure_future(send_later()) + while True: + received.put_nowait(await read_frame(reader)) + except (asyncio.IncompleteReadError, ConnectionError): + pass + finally: + if sender is not None: + sender.cancel() + writer.close() + + server = await asyncio.start_server(serve_messages, "127.0.0.1", 0) + url = f"ws://127.0.0.1:{server.sockets[0].getsockname()[1]}/" + try: + async with wreq.Client(proxies=[]) as client: + async with client.websocket(url) as ws: + # A receive cancelled by its caller stops waiting and loses nothing. + with pytest.raises(asyncio.TimeoutError): + await asyncio.wait_for(ws.recv(), 0.2) + # A pending receive does not hold up a send. + pending = asyncio.ensure_future(ws.recv()) + await asyncio.wait_for(ws.send(wreq.Message.from_text("ping")), 5) + assert await asyncio.wait_for(received.get(), 5) == (0x1, b"ping") + release.set() + assert (await asyncio.wait_for(pending, 5)).text == "first" + assert (await asyncio.wait_for(ws.recv(), 5)).text == "second" + finally: + release.set() + server.close() + await server.wait_closed() From da9370f330e83aec918fd3d21f9d452a65aa538c Mon Sep 17 00:00:00 2001 From: 0x676e67 Date: Sun, 4 Oct 2026 17:00:54 +0000 Subject: [PATCH 2/4] fix(body): fail async bodies whose forwarding task stops before finishing --- src/client/body/stream.rs | 59 +++++++++++++++++++++++++++++---------- tests/upload_test.py | 52 ++++++++++++++++++++++++++++++++++ 2 files changed, 96 insertions(+), 15 deletions(-) diff --git a/src/client/body/stream.rs b/src/client/body/stream.rs index 9baa5ff8..88c582fa 100644 --- a/src/client/body/stream.rs +++ b/src/client/body/stream.rs @@ -10,7 +10,7 @@ use std::{ use bytes::Bytes; use futures_util::Stream; use pyo3::{ - exceptions::PyStopIteration, + exceptions::{PyRuntimeError, PyStopIteration}, intern, prelude::*, sync::PyOnceLock, @@ -74,9 +74,9 @@ struct PyAsyncStream { } /// The channel end given to the forwarding coroutine; awaiting `send` applies upload -/// backpressure. +/// backpressure. Closing it without `finish` fails the body. #[pyclass(frozen)] -struct Sender(mpsc::Sender>); +struct Sender(Mutex>>>); // ===== impl PyBytesLike ===== @@ -243,8 +243,11 @@ async def forward(gen, sender): if close is not None: await close() except asyncio.CancelledError as error: - # Task cancellation must not wait for space in a retained body. - if not asyncio.current_task().cancelling(): + # Task cancellation must not wait for space in a retained body, and must not leave + # the body waiting while a traceback keeps this frame and its sender alive. + if asyncio.current_task().cancelling(): + sender.close() + else: await sender.send(error, True) raise except BaseException as error: @@ -259,7 +262,9 @@ async def forward(gen, sender): .map(Bound::unbind) })?; let (tx, rx) = mpsc::channel(1); - let coroutine = forward.bind(py).call1((generator, Sender(tx)))?; + let coroutine = forward + .bind(py) + .call1((generator, Sender(Mutex::new(Some(tx)))))?; // create_task captures the caller's contextvars on the running loop. let task = match event_loop.call_method1("create_task", (&coroutine,)) { Ok(task) => task, @@ -287,10 +292,12 @@ impl Stream for PyAsyncStream { this.task.take(); Poll::Ready(None) } - Poll::Ready(_) => { - this.rx.close(); - Poll::Ready(None) - } + // Every sender is gone without `finish`: the forwarding task was cancelled or + // destroyed, so the body is incomplete. + Poll::Ready(None) if this.task.take().is_some() => Poll::Ready(Some(Err( + PyRuntimeError::new_err("async body generator stopped before it finished"), + ))), + Poll::Ready(None) => Poll::Ready(None), Poll::Pending => Poll::Pending, } } @@ -323,7 +330,7 @@ impl Drop for PyAsyncStream { #[pymethods] impl Sender { /// Queue a chunk, or `item` as the error that ends the body. Resolves to `False` - /// once the body is dropped, which stops forwarding. + /// once the body is dropped or the sender closed, which stops forwarding. fn send<'py>( &self, py: Python<'py>, @@ -335,19 +342,41 @@ impl Sender { } else { Ok(item.extract()?) }; - let tx = self.0.clone(); + let tx = self.sender(); // Channel readiness is runtime-independent, so this waits on the Python loop. aio::local(py, "Sender.send", async move { - Ok(tx.send(Some(item)).await.is_ok()) + Ok(match tx { + Some(tx) => tx.send(Some(item)).await.is_ok(), + None => false, + }) }) } /// Mark the normal end of the body. fn finish<'py>(&self, py: Python<'py>) -> PyResult> { // Python may retain the sender after completion, especially on PyPy. - let tx = self.0.clone(); + let tx = self.sender(); aio::local(py, "Sender.finish", async move { - Ok(tx.send(None).await.is_ok()) + Ok(match tx { + Some(tx) => tx.send(None).await.is_ok(), + None => false, + }) }) } + + /// Drop the channel end at once, so the body fails instead of waiting for more. + fn close(&self) { + let tx = self.0.lock().unwrap_or_else(PoisonError::into_inner).take(); + drop(tx); + } +} + +impl Sender { + #[inline] + fn sender(&self) -> Option>> { + self.0 + .lock() + .unwrap_or_else(PoisonError::into_inner) + .clone() + } } diff --git a/tests/upload_test.py b/tests/upload_test.py index cbf79917..c753e90a 100644 --- a/tests/upload_test.py +++ b/tests/upload_test.py @@ -145,3 +145,55 @@ async def chunks(): with pytest.raises(asyncio.CancelledError): await asyncio.wait_for(task, 5) await asyncio.wait_for(closed.wait(), 5) + + +@pytest.mark.asyncio +async def test_cancelled_forwarding_fails_the_upload(): + # Cancelling the task that forwards an async generator body must fail the request, + # not leave it waiting or send the partial body as complete. + first = asyncio.Event() + complete = asyncio.Queue() + + async def serve(reader, writer): + if not server.is_serving(): + writer.close() + return + ended = False + try: + await reader.readuntil(b"\r\n\r\n") + while size := int((await reader.readuntil(b"\r\n")).strip(), 16): + await reader.readexactly(size + 2) + first.set() + await reader.readuntil(b"\r\n") + ended = True + writer.write(b"HTTP/1.1 200 OK\r\nContent-Length: 0\r\n\r\n") + await writer.drain() + except (asyncio.IncompleteReadError, ConnectionError): + pass + finally: + complete.put_nowait(ended) + writer.close() + + async def body(): + yield b"part" + await asyncio.Event().wait() + yield b"rest" + + server = await asyncio.start_server(serve, "127.0.0.1", 0) + url = f"http://127.0.0.1:{server.sockets[0].getsockname()[1]}/" + try: + async with wreq.Client(proxies=[]) as client: + request = asyncio.ensure_future(client.post(url, body=body())) + await asyncio.wait_for(first.wait(), 5) + (forward,) = [ + task + for task in asyncio.all_tasks() + if task.get_coro().__qualname__ == "forward" + ] + forward.cancel() + with pytest.raises(wreq.exceptions.RequestError): + await asyncio.wait_for(request, 5) + assert await asyncio.wait_for(complete.get(), 5) is False + finally: + server.close() + await server.wait_closed() From 9bfaf37437625f40259c6af9df0024f15d93bc27 Mon Sep 17 00:00:00 2001 From: 0x676e67 Date: Sun, 4 Oct 2026 17:02:37 +0000 Subject: [PATCH 3/4] test: resolve the dns override test locally and cover every runtime config --- tests/dns_test.py | 22 ++++++++-------------- tests/runtime_test.py | 19 ++++++++++--------- 2 files changed, 18 insertions(+), 23 deletions(-) diff --git a/tests/dns_test.py b/tests/dns_test.py index 0b491ac0..b1461f4d 100644 --- a/tests/dns_test.py +++ b/tests/dns_test.py @@ -6,9 +6,8 @@ import pytest import wreq -from wreq import Client, blocking +from wreq import blocking from wreq.dns import DnsOptions, LookupIpStrategy -from wreq.exceptions import ConnectionError @pytest.fixture @@ -77,16 +76,11 @@ def request(): @pytest.mark.asyncio -@pytest.mark.flaky(reruns=3, reruns_delay=2) -async def test_dns_resolve_override(): +async def test_dns_resolve_override(dns_http_server): + # An override wins over the real address of a resolvable host. dns_options = DnsOptions(lookup_ip_strategy=LookupIpStrategy.IPV4_ONLY) - dns_options.add_resolve("www.google.com", [IPv4Address("192.168.1.1")]) - client = Client( - dns_options=dns_options, - ) - - try: - await client.get("https://www.google.com") - assert False, "ConnectionError was expected" - except ConnectionError: - pass + dns_options.add_resolve("www.google.com", [IPv4Address("127.0.0.1")]) + async with wreq.Client(dns_options=dns_options, no_proxy=True) as client: + url = f"http://www.google.com:{dns_http_server}/" + async with await asyncio.wait_for(client.get(url), 5) as response: + assert await response.text() == "DNS resolved" diff --git a/tests/runtime_test.py b/tests/runtime_test.py index 5ab1a93d..7e889a06 100644 --- a/tests/runtime_test.py +++ b/tests/runtime_test.py @@ -35,16 +35,17 @@ def test_runtime_configuration(): max_blocking_threads=3, thread_keep_alive=duration, ) + for factory in (wreq.Client, wreq.blocking.Client): + client = factory(runtime=runtime) + alias = client.runtime + with pytest.raises(AttributeError): + client.runtime = runtime + client.close() + del client + # Releasing a client does not close a shared runtime. + other = factory(runtime=alias) + other.close() for factory in (wreq.Client, wreq.blocking.Client): - client = factory(runtime=runtime) - alias = client.runtime - with pytest.raises(AttributeError): - client.runtime = runtime - client.close() - del client - # Releasing a client does not close a shared runtime. - other = factory(runtime=alias) - other.close() client = factory(runtime=None) assert isinstance(client.runtime, Runtime) client.close() From 0bce2981580a70e27eb97def2a632762b2f319cf Mon Sep 17 00:00:00 2001 From: 0x676e67 Date: Sun, 4 Oct 2026 17:06:47 +0000 Subject: [PATCH 4/4] refactor(coroutine): rename aio to coroutine and split out asyncio specifics --- src/client.rs | 14 ++++----- src/client/body/stream.rs | 6 ++-- src/client/resp/http.rs | 12 ++++---- src/client/resp/stream.rs | 10 +++---- src/client/resp/ws.rs | 14 ++++----- src/{aio.rs => coroutine.rs} | 14 +++++---- src/{aio/port.rs => coroutine/asyncio.rs} | 16 ++++++++-- .../coroutine.rs => coroutine/awaitable.rs} | 30 +++++++++---------- src/{aio => }/coroutine/scope.rs | 4 +-- src/lib.rs | 8 ++--- 10 files changed, 71 insertions(+), 57 deletions(-) rename src/{aio.rs => coroutine.rs} (91%) rename src/{aio/port.rs => coroutine/asyncio.rs} (94%) rename src/{aio/coroutine.rs => coroutine/awaitable.rs} (92%) rename src/{aio => }/coroutine/scope.rs (99%) diff --git a/src/client.rs b/src/client.rs index becfd2a5..32324568 100644 --- a/src/client.rs +++ b/src/client.rs @@ -21,8 +21,8 @@ use self::{ resp::{BlockingResponse, BlockingWebSocket, Response, WebSocket}, }; use crate::{ - aio::{self, Coroutine}, cookie::Jar, + coroutine::{self, Coroutine}, dns::{DnsOptions, HickoryResolver, LookupIpStrategy}, emulate::EmulationLike, error::Error, @@ -265,7 +265,7 @@ impl Client { kwds: Option>, ) -> PyResult> { let kwds = kwds.map(Bound::unbind); - aio::managed(py, qualname, self.clone().execute(method, url, kwds)) + coroutine::managed(py, qualname, self.clone().execute(method, url, kwds)) } /// Send a request on the client's runtime, extracting options on first await so an @@ -277,7 +277,7 @@ impl Client { kwds: Option>, ) -> PyResult { let kwds = Python::attach(|py| kwds.map(|kwds| kwds.bind(py).extract()).transpose())?; - aio::run( + coroutine::run( self.runtime.clone(), execute_request(self, method, url, kwds), ) @@ -291,7 +291,7 @@ impl Client { kwds: Option>, ) -> PyResult { let kwds = Python::attach(|py| kwds.map(|kwds| kwds.bind(py).extract()).transpose())?; - aio::run( + coroutine::run( self.runtime.clone(), execute_websocket_request(self, url, kwds), ) @@ -663,14 +663,14 @@ impl Client { kwds: Option>, ) -> PyResult> { let kwds = kwds.map(Bound::unbind); - aio::managed(py, "Client.websocket", self.clone().connect(url, kwds)) + coroutine::managed(py, "Client.websocket", self.clone().connect(url, kwds)) } } #[pymethods] impl Client { fn __aenter__(slf: Bound<'_, Self>) -> PyResult> { - aio::ready("Client.__aenter__", slf) + coroutine::ready("Client.__aenter__", slf) } /// Close the client like `close()`: cancel pending requests and reject new ones. @@ -682,7 +682,7 @@ impl Client { _traceback: Py, ) -> PyResult> { let cancel = self.cancel.clone(); - aio::local(py, "Client.__aexit__", async move { + coroutine::local(py, "Client.__aexit__", async move { cancel.cancel(); Ok(()) }) diff --git a/src/client/body/stream.rs b/src/client/body/stream.rs index 88c582fa..72c31cc4 100644 --- a/src/client/body/stream.rs +++ b/src/client/body/stream.rs @@ -24,7 +24,7 @@ use tokio::{ }; use crate::{ - aio::{self, Coroutine}, + coroutine::{self, Coroutine}, extractor::{Binary, Text}, runtime, }; @@ -344,7 +344,7 @@ impl Sender { }; let tx = self.sender(); // Channel readiness is runtime-independent, so this waits on the Python loop. - aio::local(py, "Sender.send", async move { + coroutine::local(py, "Sender.send", async move { Ok(match tx { Some(tx) => tx.send(Some(item)).await.is_ok(), None => false, @@ -356,7 +356,7 @@ impl Sender { fn finish<'py>(&self, py: Python<'py>) -> PyResult> { // Python may retain the sender after completion, especially on PyPy. let tx = self.sender(); - aio::local(py, "Sender.finish", async move { + coroutine::local(py, "Sender.finish", async move { Ok(match tx { Some(tx) => tx.send(None).await.is_ok(), None => false, diff --git a/src/client/resp/http.rs b/src/client/resp/http.rs index 2c7488e5..48803db6 100644 --- a/src/client/resp/http.rs +++ b/src/client/resp/http.rs @@ -18,10 +18,10 @@ use wreq::Uri; use super::{ext::ResponseExt, stream::Streamer}; use crate::{ - aio::{self, Coroutine}, buffer::PyBuffer, client::{SocketAddr, body::Json, nogil}, cookie::Cookie, + coroutine::{self, Coroutine}, error::Error, header::HeaderMap, http::{StatusCode, Version}, @@ -185,9 +185,9 @@ impl Response { { let py = slf.py(); let slf = slf.unbind(); - aio::local(py, qualname, async move { + coroutine::local(py, qualname, async move { let this = slf.get(); - aio::run(this.runtime.clone(), this.read_body(read)).await + coroutine::run(this.runtime.clone(), this.read_body(read)).await }) } @@ -353,7 +353,7 @@ impl Response { pub fn close(slf: Bound<'_, Self>) -> PyResult> { let py = slf.py(); let slf = slf.unbind(); - aio::local(py, "Response.close", async move { + coroutine::local(py, "Response.close", async move { slf.get().discard(); Ok(()) }) @@ -363,7 +363,7 @@ impl Response { #[pymethods] impl Response { fn __aenter__(slf: Bound<'_, Self>) -> PyResult> { - aio::ready("Response.__aenter__", slf) + coroutine::ready("Response.__aenter__", slf) } /// Release the body without forbidding reuse: a fully read connection returns @@ -376,7 +376,7 @@ impl Response { ) -> PyResult> { let py = slf.py(); let slf = slf.unbind(); - aio::local(py, "Response.__aexit__", async move { + coroutine::local(py, "Response.__aexit__", async move { slf.get().destroy(); Ok(()) }) diff --git a/src/client/resp/stream.rs b/src/client/resp/stream.rs index 0668033d..ba8b90b9 100644 --- a/src/client/resp/stream.rs +++ b/src/client/resp/stream.rs @@ -22,9 +22,9 @@ use tokio::sync::{ use tokio_util::task::AbortOnDropHandle; use crate::{ - aio::{self, Coroutine}, buffer::PyBuffer, client::nogil, + coroutine::{self, Coroutine}, error::Error, header::HeaderMap, runtime::Runtime, @@ -295,20 +295,20 @@ impl Streamer { fn __anext__(slf: Bound<'_, Self>) -> PyResult> { let py = slf.py(); let slf = slf.unbind(); - aio::local(py, "Streamer.__anext__", async move { + coroutine::local(py, "Streamer.__anext__", async move { let this = slf.get(); // Buffered frames complete without suspending; yield to the event loop // periodically so timeouts, cancellation and other tasks run. if this.reader.since_yield.fetch_add(1, Ordering::Relaxed) >= Self::YIELD_EVERY { this.reader.since_yield.store(0, Ordering::Relaxed); - aio::yield_now().await; + coroutine::yield_now().await; } this.next(|| Error::StopAsyncIteration).await }) } fn __aenter__(slf: Bound<'_, Self>) -> PyResult> { - aio::ready("Streamer.__aenter__", slf) + coroutine::ready("Streamer.__aenter__", slf) } /// Release the body and end any pending read; returned views stay valid. @@ -320,7 +320,7 @@ impl Streamer { _traceback: Py, ) -> PyResult> { let reader = self.reader.clone(); - aio::local(py, "Streamer.__aexit__", async move { + coroutine::local(py, "Streamer.__aexit__", async move { reader.close(); Ok(()) }) diff --git a/src/client/resp/ws.rs b/src/client/resp/ws.rs index 13711484..cc64fb8a 100644 --- a/src/client/resp/ws.rs +++ b/src/client/resp/ws.rs @@ -11,9 +11,9 @@ use pyo3::prelude::*; use wreq::{header::HeaderValue, ws::WebSocketResponse}; use crate::{ - aio::{self, Coroutine}, client::{SocketAddr, nogil}, cookie::Cookie, + coroutine::{self, Coroutine}, extractor::Text, header::HeaderMap, http::{StatusCode, Version}, @@ -106,7 +106,7 @@ impl WebSocket { py: Python<'py>, timeout: Option, ) -> PyResult> { - aio::spawn( + coroutine::spawn( py, "WebSocket.recv", &self.runtime, @@ -117,7 +117,7 @@ impl WebSocket { /// Send a message to the WebSocket. #[pyo3(signature = (message))] pub fn send<'py>(&self, py: Python<'py>, message: Message) -> PyResult> { - aio::spawn( + coroutine::spawn( py, "WebSocket.send", &self.runtime, @@ -132,7 +132,7 @@ impl WebSocket { py: Python<'py>, messages: Vec, ) -> PyResult> { - aio::spawn( + coroutine::spawn( py, "WebSocket.send_all", &self.runtime, @@ -148,7 +148,7 @@ impl WebSocket { code: Option, reason: Option, ) -> PyResult> { - aio::spawn( + coroutine::spawn( py, "WebSocket.close", &self.runtime, @@ -160,7 +160,7 @@ impl WebSocket { #[pymethods] impl WebSocket { fn __aenter__(slf: Bound<'_, Self>) -> PyResult> { - aio::ready("WebSocket.__aenter__", slf) + coroutine::ready("WebSocket.__aenter__", slf) } /// Close the WebSocket connection without a close code or reason, unless already closed. @@ -171,7 +171,7 @@ impl WebSocket { _exc_val: Py, _traceback: Py, ) -> PyResult> { - aio::spawn( + coroutine::spawn( py, "WebSocket.__aexit__", &self.runtime, diff --git a/src/aio.rs b/src/coroutine.rs similarity index 91% rename from src/aio.rs rename to src/coroutine.rs index 9e0b19af..8dea17dc 100644 --- a/src/aio.rs +++ b/src/coroutine.rs @@ -1,4 +1,7 @@ -//! Awaitables that drive Rust futures from asyncio tasks. +//! Awaitables that drive Rust futures from Python coroutines. +//! +//! `awaitable` implements the coroutine protocol, `scope` adds `async with` to request +//! coroutines, and `asyncio` holds what is specific to asyncio event loops. //! //! A [`Coroutine`] polls its future on the event loop thread, and [`spawn`] moves //! the work to Tokio. A wake marks the coroutine ready and queues it on the [`Port`] @@ -7,8 +10,9 @@ //! socket is woken with `call_soon_threadsafe`, which does attach the waking thread. //! The loop thread then resolves every queued asyncio future in one batch. -mod coroutine; -mod port; +mod asyncio; +mod awaitable; +mod scope; use std::{ future::{Future, poll_fn}, @@ -19,8 +23,8 @@ use std::{ use pyo3::{IntoPyObjectExt, exceptions::PyRuntimeError, prelude::*}; use tokio_util::task::AbortOnDropHandle; -pub use self::coroutine::Coroutine; -use self::port::Port; +use self::asyncio::Port; +pub use self::awaitable::Coroutine; use crate::runtime::Runtime; /// Run `fut` on the runtime once the coroutine named `qualname` is first awaited. diff --git a/src/aio/port.rs b/src/coroutine/asyncio.rs similarity index 94% rename from src/aio/port.rs rename to src/coroutine/asyncio.rs index 0815c922..c03de570 100644 --- a/src/aio/port.rs +++ b/src/coroutine/asyncio.rs @@ -1,3 +1,6 @@ +//! The asyncio event loop side: the futures a suspended task waits on, and the per-loop +//! [`Port`] that resolves them when Rust work wakes. + use std::{ cell::RefCell, io::{self, ErrorKind, Read, Write}, @@ -10,7 +13,7 @@ use std::{ use pyo3::{PyTraverseError, PyVisit, intern, prelude::*, sync::PyOnceLock}; -use super::coroutine::Slot; +use super::awaitable::Slot; /// Wakes queued for one event loop, delivered by a single bell per batch. pub(crate) struct Port { @@ -53,8 +56,17 @@ thread_local! { // ===== impl Port ===== impl Port { + /// A new asyncio future for a task to wait on, and the port of the loop that runs it. + pub(super) fn waiter(py: Python<'_>) -> PyResult<(Py, Arc)> { + let (event_loop, port) = Self::current(py)?; + let waiter = event_loop.call_method0(intern!(py, "create_future"))?; + // Tasks only accept futures marked as yielded by `await`. + waiter.setattr(intern!(py, "_asyncio_future_blocking"), true)?; + Ok((waiter.unbind(), port)) + } + /// Return the running loop and its port, opening the port on first use. - pub(super) fn current(py: Python<'_>) -> PyResult<(Bound<'_, PyAny>, Arc)> { + fn current(py: Python<'_>) -> PyResult<(Bound<'_, PyAny>, Arc)> { static GET_RUNNING_LOOP: PyOnceLock> = PyOnceLock::new(); let event_loop = GET_RUNNING_LOOP .get_or_try_init(py, || { diff --git a/src/aio/coroutine.rs b/src/coroutine/awaitable.rs similarity index 92% rename from src/aio/coroutine.rs rename to src/coroutine/awaitable.rs index 8e158113..10a46534 100644 --- a/src/aio/coroutine.rs +++ b/src/coroutine/awaitable.rs @@ -1,5 +1,3 @@ -mod scope; - use std::{ future::Future, sync::{ @@ -17,8 +15,7 @@ use pyo3::{ prelude::*, }; -use self::scope::Scope; -use super::Port; +use super::{Port, scope::Scope}; /// An awaitable driving a Rust future on the asyncio event loop thread. /// @@ -35,11 +32,11 @@ pub struct Coroutine { future: Mutex>>>>, slot: Arc, waker: Waker, - scope: Scope, + pub(super) scope: Scope, } /// The outcome of a step: the protocol's `StopIteration` is built only for Python. -enum Step { +pub(super) enum Step { /// Suspend, handing the task an asyncio future to wait on, or `None`. Yield(Py), /// Finish with this value. @@ -81,20 +78,24 @@ impl Coroutine { } } - /// Let `async with` enter the coroutine; see [`aio::managed`](super::managed). + /// Let `async with` enter the coroutine; see [`coroutine::managed`](super::managed). pub(super) fn managed(mut self) -> Self { self.scope = Scope::Ready; self } #[inline] - fn future(&mut self) -> &mut Option>>> { + pub(super) fn future(&mut self) -> &mut Option>>> { self.future .get_mut() .unwrap_or_else(PoisonError::into_inner) } - fn step(&mut self, py: Python<'_>, sent: Option<&Bound<'_, PyAny>>) -> PyResult { + pub(super) fn step( + &mut self, + py: Python<'_>, + sent: Option<&Bound<'_, PyAny>>, + ) -> PyResult { if self.scope.is_opening() { return self.forward(py, sent); } @@ -129,7 +130,7 @@ impl Coroutine { .inspect_err(|_| self.abandon()) } - fn finish(&mut self) { + pub(super) fn finish(&mut self) { *self.future() = None; self.slot.state.store(DONE, Ordering::Release); // A finished coroutine must not keep its loop's port open. @@ -197,7 +198,7 @@ impl Coroutine { // ===== impl Step ===== impl Step { - fn into_result(self) -> PyResult> { + pub(super) fn into_result(self) -> PyResult> { match self { Step::Yield(value) => Ok(value), Step::Return(value) => Err(PyStopIteration::new_err((value,))), @@ -206,7 +207,7 @@ impl Step { } /// Fail if `waiter`, the future a task waits on for this coroutine, is still pending. -fn ensure_done(waiter: &Bound<'_, PyAny>) -> PyResult<()> { +pub(super) fn ensure_done(waiter: &Bound<'_, PyAny>) -> PyResult<()> { if waiter .call_method0(intern!(waiter.py(), "done"))? .is_truthy()? @@ -228,10 +229,7 @@ impl Slot { if self.state.load(Ordering::Acquire) == NOTIFIED { return Ok(py.None()); } - let (event_loop, port) = Port::current(py)?; - let waiter = event_loop.call_method0(intern!(py, "create_future"))?; - waiter.setattr(intern!(py, "_asyncio_future_blocking"), true)?; - let waiter = waiter.unbind(); + let (waiter, port) = Port::waiter(py)?; // Publish the port and waiter before waiting, so a wake that sees WAITING // always reaches the loop that runs this task. *self.lock_port() = Some(port); diff --git a/src/aio/coroutine/scope.rs b/src/coroutine/scope.rs similarity index 99% rename from src/aio/coroutine/scope.rs rename to src/coroutine/scope.rs index 2a7c4a70..133940b6 100644 --- a/src/aio/coroutine/scope.rs +++ b/src/coroutine/scope.rs @@ -10,7 +10,7 @@ use pyo3::{ prelude::*, }; -use super::{Coroutine, Step, ensure_done}; +use super::awaitable::{Coroutine, Step, ensure_done}; /// The `async with` state of a coroutine. pub(super) enum Scope { @@ -309,7 +309,7 @@ impl Delegate { fn close(&self, py: Python<'_>) -> PyResult<()> { match self { - Delegate::Native(coroutine) => coroutine.bind(py).try_borrow_mut()?.close(py), + Delegate::Native(coroutine) => coroutine.bind(py).try_borrow_mut()?.stop(py), Delegate::Foreign(iter) => match iter.bind(py).getattr(intern!(py, "close")) { Ok(close) => close.call0().map(drop), Err(_) => Ok(()), diff --git a/src/lib.rs b/src/lib.rs index cd98ad04..809c965c 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -5,10 +5,10 @@ #[macro_use] mod macros; -mod aio; mod buffer; mod client; mod cookie; +mod coroutine; mod dns; mod emulate; mod error; @@ -65,8 +65,8 @@ mod r#async { use pyo3::{prelude::*, pybacked::PyBackedStr, types::PyDict}; use crate::{ - aio::{self, Coroutine}, client::Client, + coroutine::{self, Coroutine}, http::Method, }; @@ -179,7 +179,7 @@ mod r#async { kwds: Option>, ) -> PyResult> { let kwds = kwds.map(Bound::unbind); - aio::managed(py, "websocket", async move { + coroutine::managed(py, "websocket", async move { Client::default().connect(url, kwds).await }) } @@ -194,7 +194,7 @@ mod r#async { kwds: Option>, ) -> PyResult> { let kwds = kwds.map(Bound::unbind); - aio::managed(py, qualname, async move { + coroutine::managed(py, qualname, async move { Client::default().execute(method, url, kwds).await }) }