From 3714c08133409650bee4ba576917f275b8ea889a Mon Sep 17 00:00:00 2001 From: 0x676e67 Date: Mon, 5 Oct 2026 20:46:16 +0000 Subject: [PATCH 1/8] perf(aio): read buffered bodies on the event loop and trim coroutine overhead --- src/client/body/stream.rs | 6 +- src/client/req.rs | 79 +++++++++------- src/client/resp/http.rs | 181 ++++++++++++++++++++++++++++++------- src/client/resp/stream.rs | 27 +++++- src/client/resp/ws.rs | 11 ++- src/coroutine.rs | 35 ++++++- src/coroutine/asyncio.rs | 23 +++-- src/coroutine/awaitable.rs | 43 +++++++-- src/coroutine/scope.rs | 30 ++++-- src/lib.rs | 2 + src/macros.rs | 23 +++++ tests/upload_test.py | 16 ++++ 12 files changed, 377 insertions(+), 99 deletions(-) diff --git a/src/client/body/stream.rs b/src/client/body/stream.rs index 72c31cc4..a78f594c 100644 --- a/src/client/body/stream.rs +++ b/src/client/body/stream.rs @@ -226,7 +226,7 @@ impl PyAsyncStream { fn new(generator: Bound<'_, PyAny>) -> PyResult { static FORWARD: PyOnceLock> = PyOnceLock::new(); let py = generator.py(); - let event_loop = py.import("asyncio")?.call_method0("get_running_loop")?; + let event_loop = coroutine::running_loop(py)?; let forward = FORWARD.get_or_try_init(py, || { PyModule::from_code( py, @@ -266,10 +266,10 @@ async def forward(gen, sender): .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,)) { + let task = match event_loop.call_method1(intern!(py, "create_task"), (&coroutine,)) { Ok(task) => task, Err(err) => { - let _ = coroutine.call_method0("close"); + let _ = coroutine.call_method0(intern!(py, "close")); return Err(err); } }; diff --git a/src/client/req.rs b/src/client/req.rs index bbcdba53..92d7a172 100644 --- a/src/client/req.rs +++ b/src/client/req.rs @@ -5,7 +5,9 @@ use std::{ use futures_util::TryFutureExt; use http::header::COOKIE; -use pyo3::{PyResult, exceptions::asyncio::CancelledError, prelude::*, pybacked::PyBackedStr}; +use pyo3::{ + PyResult, exceptions::asyncio::CancelledError, intern, prelude::*, pybacked::PyBackedStr, +}; use crate::{ client::{ @@ -20,6 +22,7 @@ use crate::{ extractor::Extractor, header::{HeaderMap, OrigHeaderMap}, http::{Method, Version}, + macros::lookup, proxy::Proxy, redirect, }; @@ -186,36 +189,50 @@ impl FromPyObject<'_, '_> for Request { fn extract(ob: Borrowed) -> PyResult { let mut request = Self::default(); - extract_option!(ob, request, emulation); - extract_option!(ob, request, proxy); - extract_option!(ob, request, local_address); - extract_option!(ob, request, local_addresses); - extract_option!(ob, request, interface); - - extract_option!(ob, request, timeout); - extract_option!(ob, request, read_timeout); - - extract_option!(ob, request, version); - extract_option!(ob, request, headers); - extract_option!(ob, request, orig_headers); - extract_option!(ob, request, default_headers); - extract_option!(ob, request, cookies); - extract_option!(ob, request, redirect); - extract_option!(ob, request, cookie_provider); - extract_option!(ob, request, auth); - extract_option!(ob, request, bearer_auth); - extract_option!(ob, request, basic_auth); - extract_option!(ob, request, query); - extract_option!(ob, request, form); - extract_option!(ob, request, json); - extract_option!(ob, request, body); - extract_option!(ob, request, multipart); - - extract_option!(ob, request, gzip); - extract_option!(ob, request, brotli); - extract_option!(ob, request, deflate); - extract_option!(ob, request, zstd); - + // Extracting a body or multipart form can start a generator or take stream parts, so + // they are extracted last, after every other option has been validated. + let py = ob.py(); + let body = lookup(&ob, intern!(py, "body")); + let multipart = lookup(&ob, intern!(py, "multipart")); + let found = usize::from(body.is_some()) + usize::from(multipart.is_some()); + // Common keys first: the scan stops once every given key is found. + extract_options!( + ob, + request, + found, + [ + json, + form, + query, + headers, + timeout, + read_timeout, + cookies, + auth, + bearer_auth, + basic_auth, + emulation, + proxy, + local_address, + local_addresses, + interface, + version, + orig_headers, + default_headers, + redirect, + cookie_provider, + gzip, + brotli, + deflate, + zstd, + ] + ); + if let Some(body) = body { + request.body = body.extract()?; + } + if let Some(multipart) = multipart { + request.multipart = multipart.extract()?; + } Ok(request) } } diff --git a/src/client/resp/http.rs b/src/client/resp/http.rs index 48803db6..5abaa735 100644 --- a/src/client/resp/http.rs +++ b/src/client/resp/http.rs @@ -12,8 +12,8 @@ use futures_util::{ }; use http::response::{Parts, Response as HttpResponse}; use http_body::Body as _; -use http_body_util::{BodyExt, Collected}; -use pyo3::{prelude::*, pybacked::PyBackedStr}; +use http_body_util::{BodyExt, Collected, combinators::Collect}; +use pyo3::{prelude::*, pybacked::PyBackedStr, sync::PyOnceLock}; use wreq::Uri; use super::{ext::ResponseExt, stream::Streamer}; @@ -21,7 +21,7 @@ use crate::{ buffer::PyBuffer, client::{SocketAddr, body::Json, nogil}, cookie::Cookie, - coroutine::{self, Coroutine}, + coroutine::{self, Coroutine, EntersSelf}, error::Error, header::HeaderMap, http::{StatusCode, Version}, @@ -32,8 +32,10 @@ use crate::{ /// A response from a request. /// -/// Body reads are written once as futures ([`Response::read_body`]); the async methods -/// run them on the runtime when awaited and [`BlockingResponse`] waits on the caller. +/// Async reads take the body with [`Response::take_bytes`]: a buffered body of known length +/// up to the read's limit finishes on the event loop, and any other body is collected on the +/// runtime by [`Response::collect_later`]. [`BlockingResponse`] waits on the caller through +/// [`Response::read_body`]. #[pyclass(subclass, frozen, str, skip_from_py_object)] pub struct Response { uri: Uri, @@ -62,21 +64,34 @@ enum Body { Released, } +/// A body taken for reading. +enum BodyRead { + /// Read in full now, or failed. + Ready(Result), + /// Not fully buffered yet; collected on the runtime by [`Response::collect_later`]. + Pending(Collect), +} + /// A blocking response from a request. #[pyclass(name = "Response", subclass, frozen, str, skip_from_py_object)] pub struct BlockingResponse(Response); -/// Forbids connection reuse on drop unless disarmed by taking the parts. Held while -/// [`Response::cache_response`] collects the body, so a failed or cancelled read is not -/// pooled. +/// Forbids connection reuse on drop unless disarmed by taking the parts. Held by +/// [`Response::collect_later`] from its creation until the body is read in full, so a failed +/// or cancelled read is not pooled. struct RecycleGuard(Option); // ===== impl Response ===== impl Response { - /// Bodies up to this size are read on a blocking caller without first releasing the GIL. + /// Bodies up to this size are read on a blocking caller without first releasing the GIL; + /// `bytes()` and `text()` read and decode them on the event loop when already buffered. const READ_ATTACHED: u64 = 64 * 1024; + /// Buffered JSON bodies up to this size are parsed on the event loop; larger ones are + /// parsed on the runtime. + const JSON_ATTACHED: u64 = 8 * 1024; + /// Create a new [`Response`] instance. pub fn new(response: wreq::Response, runtime: Runtime) -> Self { let uri = response.uri().clone(); @@ -126,26 +141,70 @@ impl Response { } }; drop(slot); - let parts = self.parts.clone(); + self.collect_later(stream.collect()) + .map_ok(|(parts, bytes)| wreq::Response::from(HttpResponse::from_parts(parts, bytes))) + .boxed() + } + + /// Take the body for reading. Cached bytes are shared at once, and a body of known + /// length up to `limit` that is already buffered is read now on the calling thread; + /// overlapping reads and `stream()` fail with [`Error::Memory`] while a read runs. + fn take_bytes(&self, limit: u64) -> BodyRead { + let mut slot = self.slot(); + let body = match mem::replace(&mut *slot, Body::Taken) { + Body::Unread(body) => body, + Body::Cached(bytes) => { + *slot = Body::Cached(bytes.clone()); + return BodyRead::Ready(Ok(bytes)); + } + other => { + *slot = other; + return BodyRead::Ready(Err(Error::Memory)); + } + }; + drop(slot); + let attached = body.size_hint().exact().is_some_and(|len| len <= limit); + let mut collect = body.collect(); + if attached { + // A read timeout starts a timer on its first poll, which needs the runtime. + let ready = { + let _runtime = self.runtime.handle().enter(); + (&mut collect).now_or_never() + }; + match ready { + Some(Ok(collected)) => { + let bytes = collected.to_bytes(); + cache(&self.body, &bytes); + return BodyRead::Ready(Ok(bytes)); + } + Some(Err(err)) => { + self.forbid_recycle(); + return BodyRead::Ready(Err(Error::Library(err))); + } + None => {} + } + } + BodyRead::Pending(collect) + } + + /// Collect a taken body, caching its bytes and returning them with a copy of the head. + /// The connection stays out of the pool unless the body is read in full, including + /// when the future is dropped before its first poll. + fn collect_later( + &self, + collect: Collect, + ) -> impl Future> + Send + 'static { + let mut guard = RecycleGuard(Some(self.parts.clone())); let body = self.body.clone(); async move { - // Keep the connection out of the pool unless the body is read in full. - let mut guard = RecycleGuard(Some(parts)); - let bytes = stream - .collect() + let bytes = collect .await .map(Collected::to_bytes) .map_err(Error::Library)?; let parts = guard.0.take().ok_or(Error::Memory)?; - // A release during the read wins over caching. - let mut slot = lock(&body); - if let Body::Taken = *slot { - *slot = Body::Cached(bytes.clone()); - } - drop(slot); - Ok(wreq::Response::from(HttpResponse::from_parts(parts, bytes))) + cache(&body, &bytes); + Ok((parts, bytes)) } - .boxed() } /// Take the unread body for a [`Streamer`]; fails with [`Error::Memory`] otherwise, @@ -172,10 +231,13 @@ impl Response { self.cache_response().and_then(read).map_err(Into::into) } - /// Read the body on the runtime once awaited; the body is taken on first await. + /// Read and decode the body once awaited; the body is taken on first await. A body up + /// to `limit` that is already buffered is decoded on the event loop; otherwise the read + /// continues on the runtime. fn read<'py, F, Fut, T>( slf: Bound<'py, Self>, qualname: &'static str, + limit: u64, read: F, ) -> PyResult> where @@ -187,7 +249,30 @@ impl Response { let slf = slf.unbind(); coroutine::local(py, qualname, async move { let this = slf.get(); - coroutine::run(this.runtime.clone(), this.read_body(read)).await + let (mut fut, attached) = match this.take_bytes(limit) { + BodyRead::Ready(bytes) => { + let bytes = bytes?; + let attached = bytes.len() as u64 <= limit; + (read(this.build_response(bytes)).boxed(), attached) + } + BodyRead::Pending(collect) => { + let fut = this.collect_later(collect).and_then(|(parts, bytes)| { + read(wreq::Response::from(HttpResponse::from_parts(parts, bytes))) + }); + (fut.boxed(), false) + } + }; + if attached { + // Decoding a buffered body finishes at once. + let ready = { + let _runtime = this.runtime.handle().enter(); + (&mut fut).now_or_never() + }; + if let Some(output) = ready { + return output.map_err(Into::into); + } + } + coroutine::run(this.runtime.clone(), fut.map_err(Into::into)).await }) } @@ -233,6 +318,14 @@ fn lock(body: &Mutex) -> MutexGuard<'_, Body> { body.lock().unwrap_or_else(PoisonError::into_inner) } +/// Cache the bytes of a finished read; a release during the read wins over caching. +fn cache(body: &Mutex, bytes: &Bytes) { + let mut slot = lock(body); + if let Body::Taken = *slot { + *slot = Body::Cached(bytes.clone()); + } +} + #[pymethods] impl Response { /// Get the URL of the response. @@ -330,19 +423,38 @@ impl Response { slf: Bound<'_, Self>, encoding: Option, ) -> PyResult> { - Self::read(slf, "Response.text", |resp| { + Self::read(slf, "Response.text", Self::READ_ATTACHED, |resp| { ResponseExt::text(resp, encoding) }) } /// Get the JSON content of the response. pub fn json(slf: Bound<'_, Self>) -> PyResult> { - Self::read(slf, "Response.json", ResponseExt::json::) + Self::read( + slf, + "Response.json", + Self::JSON_ATTACHED, + ResponseExt::json::, + ) } /// Read the body as a read-only memoryview, retaining its data after the response closes. pub fn bytes(slf: Bound<'_, Self>) -> PyResult> { - Self::read(slf, "Response.bytes", ResponseExt::bytes) + let py = slf.py(); + let slf = slf.unbind(); + coroutine::local(py, "Response.bytes", async move { + let this = slf.get(); + match this.take_bytes(Self::READ_ATTACHED) { + BodyRead::Ready(bytes) => bytes.map(PyBuffer::from).map_err(Into::into), + BodyRead::Pending(collect) => { + let read = this + .collect_later(collect) + .map_ok(|(_, bytes)| PyBuffer::from(bytes)) + .map_err(Into::into); + coroutine::run(this.runtime.clone(), read).await + } + } + }) } /// Discard the retained body and mark its connection as non-reusable. @@ -355,7 +467,7 @@ impl Response { let slf = slf.unbind(); coroutine::local(py, "Response.close", async move { slf.get().discard(); - Ok(()) + Ok(None::<()>) }) } } @@ -378,11 +490,18 @@ impl Response { let slf = slf.unbind(); coroutine::local(py, "Response.__aexit__", async move { slf.get().destroy(); - Ok(()) + Ok(None::<()>) }) } } +impl EntersSelf for Response { + fn native_aenter() -> &'static PyOnceLock> { + static NATIVE: PyOnceLock> = PyOnceLock::new(); + &NATIVE + } +} + impl Display for Response { fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { write!( @@ -415,8 +534,8 @@ impl Drop for RecycleGuard { // ===== impl BlockingResponse ===== impl BlockingResponse { - /// Read the body with `read` on the calling thread, like [`Response::read`] on the - /// runtime. A small body is read without releasing the GIL unless it must wait. + /// Read the body with `read` on the calling thread. A small body is read without + /// releasing the GIL unless it must wait. fn read(&self, py: Python, read: F) -> PyResult where F: FnOnce(wreq::Response) -> Fut + Send + 'static, diff --git a/src/client/resp/stream.rs b/src/client/resp/stream.rs index ba8b90b9..2fc160b2 100644 --- a/src/client/resp/stream.rs +++ b/src/client/resp/stream.rs @@ -13,6 +13,7 @@ use std::{ use bytes::Bytes; use futures_util::FutureExt; +use http_body::Body as _; use http_body_util::BodyExt; use pyo3::{exceptions::PyRuntimeError, prelude::*}; use tokio::sync::{ @@ -83,6 +84,9 @@ impl Streamer { /// Buffered frames returned before `__anext__` yields to the event loop once. const YIELD_EVERY: usize = 8; + /// An async reader reads a buffered body of known length up to this size at once. + const READ_ATTACHED: u64 = 64 * 1024; + /// Create a new [`Streamer`] instance. #[inline] pub fn new(resp: wreq::Response, runtime: Runtime) -> Streamer { @@ -156,17 +160,24 @@ impl Streamer { } } - /// Read a frame the body already holds, before any read-ahead task starts, so a - /// synchronous reader of a short body never waits on another thread. - fn ready_frame(&self) -> Option> { + /// Read a frame the body already holds, before any read-ahead task starts, so a reader + /// of a short body never waits on another thread. With `limit`, only a body of known + /// length up to it is read here, so an async reader never decodes on the event loop; + /// `end` builds the error that ends iteration. + fn ready_frame(&self, limit: Option, end: fn() -> Error) -> Option> { let mut state = self.reader.lock(); let State::Idle(resp) = &mut *state else { return None; }; + if let Some(limit) = limit + && !resp.size_hint().exact().is_some_and(|len| len <= limit) + { + return None; + } let _runtime = self.runtime.handle().enter(); let Some(frame) = Self::poll_ready(resp)? else { *state = State::Closed; - return Some(Err(Error::StopIteration.into())); + return Some(Err(end().into())); }; if frame.is_err() { *state = State::Closed; @@ -264,7 +275,7 @@ impl Streamer { fn __next__(&self, py: Python) -> PyResult { // Frames already received are returned without releasing the GIL. - if let Some(frame) = self.ready_frame() { + if let Some(frame) = self.ready_frame(None, || Error::StopIteration) { return frame; } nogil::run(py, &self.runtime, self.next(|| Error::StopIteration)) @@ -297,6 +308,12 @@ impl Streamer { let slf = slf.unbind(); coroutine::local(py, "Streamer.__anext__", async move { let this = slf.get(); + // A short body already buffered is read without starting the read-ahead task. + if let Some(frame) = + this.ready_frame(Some(Self::READ_ATTACHED), || Error::StopAsyncIteration) + { + return frame; + } // 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 { diff --git a/src/client/resp/ws.rs b/src/client/resp/ws.rs index cc64fb8a..d83afb43 100644 --- a/src/client/resp/ws.rs +++ b/src/client/resp/ws.rs @@ -7,13 +7,13 @@ use std::{ }; use msg::Message; -use pyo3::prelude::*; +use pyo3::{prelude::*, sync::PyOnceLock}; use wreq::{header::HeaderValue, ws::WebSocketResponse}; use crate::{ client::{SocketAddr, nogil}, cookie::Cookie, - coroutine::{self, Coroutine}, + coroutine::{self, Coroutine, EntersSelf}, extractor::Text, header::HeaderMap, http::{StatusCode, Version}, @@ -180,6 +180,13 @@ impl WebSocket { } } +impl EntersSelf for WebSocket { + fn native_aenter() -> &'static PyOnceLock> { + static NATIVE: PyOnceLock> = PyOnceLock::new(); + &NATIVE + } +} + impl Display for WebSocket { fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { write!(f, "<{} [{}] >", stringify!(WebSocket), self.status.0) diff --git a/src/coroutine.rs b/src/coroutine.rs index 8dea17dc..048b5183 100644 --- a/src/coroutine.rs +++ b/src/coroutine.rs @@ -20,10 +20,13 @@ use std::{ task::Poll, }; -use pyo3::{IntoPyObjectExt, exceptions::PyRuntimeError, prelude::*}; +use pyo3::{ + IntoPyObjectExt, PyTypeInfo, exceptions::PyRuntimeError, intern, prelude::*, sync::PyOnceLock, +}; use tokio_util::task::AbortOnDropHandle; use self::asyncio::Port; +pub(crate) use self::asyncio::running_loop; pub use self::awaitable::Coroutine; use crate::runtime::Runtime; @@ -56,6 +59,20 @@ where Bound::new(py, coroutine(qualname, fut)) } +/// A result type whose native `__aenter__` returns the object itself, so `async with` on a +/// [`managed`] coroutine enters it without awaiting `__aenter__`. +pub trait EntersSelf: PyTypeInfo { + /// The native `__aenter__`, recorded by [`record_enters_self`]. + fn native_aenter() -> &'static PyOnceLock>; +} + +/// Record `T`'s native `__aenter__` at module init, before user code can replace it. +pub fn record_enters_self(py: Python<'_>) -> PyResult<()> { + let aenter = T::type_object(py).getattr(intern!(py, "__aenter__"))?; + let _ = T::native_aenter().set(py, aenter.unbind()); + Ok(()) +} + /// Like [`local`], but `async with` may enter the coroutine directly, as `async with await` /// would: its result must be an async context manager, whose `__aenter__` is awaited too. #[inline] @@ -66,9 +83,21 @@ pub fn managed<'py, F, T>( ) -> PyResult> where F: Future> + Send + 'static, - T: for<'a> IntoPyObject<'a>, + T: EntersSelf + for<'a> IntoPyObject<'a>, { - Bound::new(py, coroutine(qualname, fut).managed()) + Bound::new(py, coroutine(qualname, fut).managed(enters_self::)) +} + +/// Whether `object` is exactly a `T` whose class still has the native `__aenter__`: a +/// subclass or a replaced class attribute may return something else. +fn enters_self(object: &Bound<'_, PyAny>) -> bool { + let py = object.py(); + let ty = object.get_type(); + ty.is(T::type_object(py)) + && T::native_aenter().get(py).is_some_and(|native| { + ty.getattr(intern!(py, "__aenter__")) + .is_ok_and(|aenter| aenter.is(native)) + }) } /// A coroutine that returns `value` without suspending. diff --git a/src/coroutine/asyncio.rs b/src/coroutine/asyncio.rs index c03de570..9ce6e5ae 100644 --- a/src/coroutine/asyncio.rs +++ b/src/coroutine/asyncio.rs @@ -67,15 +67,7 @@ impl Port { /// Return the running loop and its port, opening the port on first use. 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, || { - py.import("asyncio")? - .getattr("get_running_loop") - .map(Bound::unbind) - })? - .bind(py) - .call0()?; + let event_loop = running_loop(py)?; // A loop drops its drain or keeper when closed or freed, so a reused address // never matches. @@ -318,6 +310,19 @@ impl Drop for Keeper { } } +/// The running asyncio loop, through a cached `asyncio.get_running_loop`. +pub(crate) fn running_loop(py: Python<'_>) -> PyResult> { + static GET_RUNNING_LOOP: PyOnceLock> = PyOnceLock::new(); + GET_RUNNING_LOOP + .get_or_try_init(py, || { + py.import("asyncio")? + .getattr("get_running_loop") + .map(Bound::unbind) + })? + .bind(py) + .call0() +} + #[cfg(unix)] fn socket_pair() -> io::Result<(Socket, Socket)> { let (tx, rx) = Socket::pair()?; diff --git a/src/coroutine/awaitable.rs b/src/coroutine/awaitable.rs index 10a46534..195c22dc 100644 --- a/src/coroutine/awaitable.rs +++ b/src/coroutine/awaitable.rs @@ -10,12 +10,16 @@ use std::{ use futures_util::{FutureExt, future::BoxFuture}; use pyo3::{ PyTraverseError, PyVisit, - exceptions::{PyRuntimeError, PyStopIteration}, - intern, + exceptions::{PyBaseException, PyRuntimeError, PyStopIteration}, + ffi, intern, prelude::*, + types::PyTuple, }; -use super::{Port, scope::Scope}; +use super::{ + Port, + scope::{EntersSelf, Scope}, +}; /// An awaitable driving a Rust future on the asyncio event loop thread. /// @@ -79,8 +83,8 @@ impl Coroutine { } /// Let `async with` enter the coroutine; see [`coroutine::managed`](super::managed). - pub(super) fn managed(mut self) -> Self { - self.scope = Scope::Ready; + pub(super) fn managed(mut self, enters_self: EntersSelf) -> Self { + self.scope = Scope::Ready(enters_self); self } @@ -170,8 +174,11 @@ impl Coroutine { slf } - fn __next__(&mut self, py: Python<'_>) -> PyResult> { - self.step(py, None).and_then(Step::into_result) + fn __next__(&mut self, py: Python<'_>) -> PyResult>> { + match self.step(py, None)? { + Step::Yield(value) => Ok(Some(value)), + Step::Return(value) => stop_iteration(value.into_bound(py)), + } } fn __traverse__(&self, visit: PyVisit<'_>) -> Result<(), PyTraverseError> { @@ -206,6 +213,28 @@ impl Step { } } +/// End `__next__` with `value`, skipping PyO3's lazy `StopIteration`. +/// +/// A `NULL` return without an error means `None`. Other values are raised with +/// `PyErr_SetObject`, as CPython does for generators: 3.11 keeps `(StopIteration, value)` +/// unnormalized for `_PyGen_FetchStopIterationValue`, while 3.12+ builds the instance at once. +/// Tuples and exceptions keep PyO3's path, since CPython would unpack or adopt them. +fn stop_iteration(value: Bound<'_, PyAny>) -> PyResult>> { + if value.is_none() { + return Ok(None); + } + if value.is_instance_of::() || value.is_instance_of::() { + return Err(PyStopIteration::new_err((value.unbind(),))); + } + // SAFETY: the thread is attached and `value` is a valid object for the call; + // `PyErr_SetObject` takes its own references to the type and the value. + #[allow(unsafe_code)] + unsafe { + ffi::PyErr_SetObject(ffi::PyExc_StopIteration, value.as_ptr()) + }; + Ok(None) +} + /// Fail if `waiter`, the future a task waits on for this coroutine, is still pending. pub(super) fn ensure_done(waiter: &Bound<'_, PyAny>) -> PyResult<()> { if waiter diff --git a/src/coroutine/scope.rs b/src/coroutine/scope.rs index 133940b6..85d5af36 100644 --- a/src/coroutine/scope.rs +++ b/src/coroutine/scope.rs @@ -12,16 +12,19 @@ use pyo3::{ use super::awaitable::{Coroutine, Step, ensure_done}; +/// Whether a result's `__aenter__` returns the result itself, so it can be skipped. +pub(super) type EntersSelf = fn(&Bound<'_, PyAny>) -> bool; + /// The `async with` state of a coroutine. pub(super) enum Scope { /// `async with` is unsupported. Unsupported, /// Not yet awaited; `async with` may enter it. - Ready, + Ready(EntersSelf), /// Awaited, entered or exited already. Spent, /// Awaited by `async with` until the result, an async context manager, is ready. - Entering, + Entering(EntersSelf), /// Awaiting the manager's `__aenter__` through `delegate`. Like `async with`, its /// bound `__aexit__` is looked up before entering. Opening { @@ -53,7 +56,7 @@ impl Scope { /// Mark a first await, after which `async with` can no longer enter. #[inline] pub(super) fn start(&mut self) { - if let Scope::Ready = self { + if let Scope::Ready(_) = self { *self = Scope::Spent; } } @@ -88,7 +91,7 @@ impl Coroutine { /// Finish with `value`, first entering it when awaited by `async with`. pub(super) fn complete(&mut self, py: Python<'_>, value: Py) -> PyResult { match mem::replace(&mut self.scope, Scope::Spent) { - Scope::Entering => return self.open(py, value), + Scope::Entering(enters_self) => return self.open(py, value, enters_self), Scope::Opening { exit, .. } => self.scope = Scope::Entered(exit), scope => self.scope = scope, } @@ -97,9 +100,20 @@ impl Coroutine { } /// Await the manager's `__aenter__`, as `async with` would. - fn open(&mut self, py: Python<'_>, manager: Py) -> PyResult { + fn open( + &mut self, + py: Python<'_>, + manager: Py, + enters_self: EntersSelf, + ) -> PyResult { self.finish(); let manager = manager.into_bound(py); + // Its `__aenter__` would return the manager itself: enter without awaiting it. + if enters_self(&manager) { + let exit = manager.getattr(intern!(py, "__aexit__"))?; + self.scope = Scope::Entered(exit.unbind()); + return Ok(Step::Return(manager.unbind())); + } // Results have no instance dict, so binding on the instance matches `async with`. let (Ok(enter), Ok(exit)) = ( manager.getattr(intern!(py, "__aenter__")), @@ -175,7 +189,7 @@ impl Coroutine { /// Finish without entering, dropping any context manager being entered. pub(super) fn abandon(&mut self) { self.finish(); - if let Scope::Entering | Scope::Opening { .. } = self.scope { + if let Scope::Entering(_) | Scope::Opening { .. } = self.scope { self.scope = Scope::Spent; } } @@ -200,8 +214,8 @@ impl Coroutine { Scope::Unsupported => Err(PyTypeError::new_err( "coroutine does not support the asynchronous context manager protocol", )), - Scope::Ready if pending => { - slf.scope = Scope::Entering; + Scope::Ready(enters_self) if pending => { + slf.scope = Scope::Entering(enters_self); Ok(slf) } _ => Err(PyRuntimeError::new_err( diff --git a/src/lib.rs b/src/lib.rs index 809c965c..f165edd4 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -321,6 +321,8 @@ fn wreq(py: Python, m: &Bound<'_, PyModule>) -> PyResult<()> { m.add_class::()?; m.add_class::()?; m.add_class::()?; + coroutine::record_enters_self::(py)?; + coroutine::record_enters_self::(py)?; m.add_class::()?; m.add_class::()?; m.add_class::()?; diff --git a/src/macros.rs b/src/macros.rs index 6326aff3..ef545316 100644 --- a/src/macros.rs +++ b/src/macros.rs @@ -14,6 +14,29 @@ pub(crate) fn lookup<'py>( } } +/// Keys left to look up: a dict's length, or unbounded for other mappings. +pub(crate) fn remaining(ob: &Bound<'_, PyAny>) -> usize { + ob.cast::().map_or(usize::MAX, |dict| dict.len()) +} + +/// Like [`extract_option!`] for several fields, but stop looking up keys once every key +/// of a dict has been found, counting `$found` keys already looked up by the caller. Unknown +/// keys are never counted, so they still leave the scan complete and stay ignored. +macro_rules! extract_options { + ($ob:expr, $params:expr, $found:expr, [$($field:ident),* $(,)?]) => {{ + let mut remaining = $crate::macros::remaining(&$ob).saturating_sub($found); + $( + if remaining > 0 + && let Some(value) = + $crate::macros::lookup(&$ob, pyo3::intern!($ob.py(), stringify!($field))) + { + $params.$field = value.extract()?; + remaining -= 1; + } + )* + }}; +} + macro_rules! extract_option { ($ob:expr, $params:expr, $field:ident) => { if let Some(value) = diff --git a/tests/upload_test.py b/tests/upload_test.py index c753e90a..8848a00b 100644 --- a/tests/upload_test.py +++ b/tests/upload_test.py @@ -88,6 +88,22 @@ async def chunks(): await asyncio.wait_for(closed.wait(), 5) +@pytest.mark.asyncio +async def test_invalid_option_does_not_start_the_body(): + started = asyncio.Event() + + async def chunks(): + started.set() + yield b"first" + + async with local_server() as (url, _), wreq.Client(proxies=[]) as client: + # Options are validated before the body generator is consumed. + with pytest.raises(TypeError): + await client.post(url, body=chunks(), timeout=5) + await asyncio.sleep(0.1) + assert not started.is_set() + + @pytest.mark.asyncio @pytest.mark.parametrize("action", ["cancel", "early_response"]) async def test_abandoned_upload_to_stalled_peer(action): From bbe0ecfc325b10bb92c7c6fe4f54969b97e7b862 Mon Sep 17 00:00:00 2001 From: 0x676e67 Date: Mon, 5 Oct 2026 21:56:37 +0000 Subject: [PATCH 2/8] perf(aio): keep HTTP/2 reads off the event loop and parse JSON in one pass --- python/wreq/wreq.py | 8 +- src/client/body/json.rs | 80 ++++++++++++++++- src/client/resp.rs | 15 ++++ src/client/resp/ext.rs | 13 +-- src/client/resp/http.rs | 175 +++++++++++++++---------------------- src/client/resp/stream.rs | 16 ++-- src/coroutine.rs | 35 ++++---- src/coroutine/awaitable.rs | 36 ++------ src/coroutine/scope.rs | 10 +-- src/macros.rs | 24 ++--- tests/response_test.py | 94 ++++++++++++++++++++ tests/upload_test.py | 8 +- 12 files changed, 322 insertions(+), 192 deletions(-) diff --git a/python/wreq/wreq.py b/python/wreq/wreq.py index 94c3c914..04109214 100644 --- a/python/wreq/wreq.py +++ b/python/wreq/wreq.py @@ -454,7 +454,7 @@ async def close(self) -> None: """ async def __aenter__(self) -> Any: ... - async def __aexit__(self, _exc_type: Any, _exc_value: Any, _traceback: Any) -> Any: + async def __aexit__(self, _exc_type: Any, _exc_value: Any, _traceback: Any) -> None: r""" Release the body without forbidding reuse: a fully read connection returns to the pool, while an unread HTTP/1 body drains or closes its connection. @@ -528,7 +528,7 @@ async def close( """ async def __aenter__(self) -> Any: ... - async def __aexit__(self, _exc_type: Any, _exc_value: Any, _traceback: Any) -> Any: + async def __aexit__(self, _exc_type: Any, _exc_value: Any, _traceback: Any) -> None: r""" Close the WebSocket connection without a close code or reason, unless already closed. """ @@ -1134,7 +1134,7 @@ class _RequestCoroutine(Coroutine[Any, Any, _T]): """ async def __aenter__(self) -> _T: ... - async def __aexit__(self, _exc_type: Any, _exc_value: Any, _traceback: Any) -> Any: ... + async def __aexit__(self, _exc_type: Any, _exc_value: Any, _traceback: Any) -> None: ... class Client: @@ -1458,7 +1458,7 @@ async def main(): ... async def __aenter__(self) -> Any: ... - async def __aexit__(self, _exc_type: Any, _exc_value: Any, _traceback: Any) -> Any: + async def __aexit__(self, _exc_type: Any, _exc_value: Any, _traceback: Any) -> None: r""" Close the client like `close()`: cancel pending requests and reject new ones. """ diff --git a/src/client/body/json.rs b/src/client/body/json.rs index 30e9298f..79e96467 100644 --- a/src/client/body/json.rs +++ b/src/client/body/json.rs @@ -1,10 +1,15 @@ +use std::fmt; + use indexmap::IndexMap; use pyo3::{FromPyObject, prelude::*, pybacked::PyBackedStr}; -use serde::{Deserialize, Deserializer, Serialize, Serializer}; +use serde::{ + Deserialize, Deserializer, Serialize, Serializer, + de::{Error, MapAccess, SeqAccess, Visitor}, +}; /// Represents a JSON value for HTTP requests. /// Supports objects, arrays, numbers, strings, booleans, and null. -#[derive(FromPyObject, IntoPyObject, Serialize, Deserialize)] +#[derive(FromPyObject, IntoPyObject, Serialize)] #[serde(untagged)] pub enum Json { Object(IndexMap), @@ -55,3 +60,74 @@ impl<'de> Deserialize<'de> for JsonString { String::deserialize(deserializer).map(JsonString::RustString) } } + +impl<'de> Deserialize<'de> for Json { + #[inline] + fn deserialize(deserializer: D) -> Result + where + D: Deserializer<'de>, + { + deserializer.deserialize_any(JsonVisitor) + } +} + +/// Builds a [`Json`] in one pass; an untagged derive would buffer every value and then try +/// each variant in turn. +struct JsonVisitor; + +impl<'de> Visitor<'de> for JsonVisitor { + type Value = Json; + + fn expecting(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.write_str("a JSON value") + } + + fn visit_bool(self, v: bool) -> Result { + Ok(Json::Boolean(v)) + } + + // Integers outside `isize` become floats, as the untagged variant order did. + fn visit_i64(self, v: i64) -> Result { + Ok(isize::try_from(v).map_or(Json::Float(v as f64), Json::Number)) + } + + fn visit_u64(self, v: u64) -> Result { + Ok(isize::try_from(v).map_or(Json::Float(v as f64), Json::Number)) + } + + fn visit_f64(self, v: f64) -> Result { + Ok(Json::Float(v)) + } + + fn visit_str(self, v: &str) -> Result { + Ok(Json::String(JsonString::RustString(v.to_owned()))) + } + + fn visit_string(self, v: String) -> Result { + Ok(Json::String(JsonString::RustString(v))) + } + + fn visit_unit(self) -> Result { + Ok(Json::Null(None)) + } + + fn visit_none(self) -> Result { + Ok(Json::Null(None)) + } + + fn visit_seq>(self, mut seq: A) -> Result { + let mut values = Vec::new(); + while let Some(value) = seq.next_element()? { + values.push(value); + } + Ok(Json::Array(values)) + } + + fn visit_map>(self, mut map: A) -> Result { + let mut values = IndexMap::new(); + while let Some((key, value)) = map.next_entry()? { + values.insert(key, value); + } + Ok(Json::Object(values)) + } +} diff --git a/src/client/resp.rs b/src/client/resp.rs index e1f1450b..2ca9a0fd 100644 --- a/src/client/resp.rs +++ b/src/client/resp.rs @@ -8,3 +8,18 @@ pub use self::{ stream::Streamer, ws::{BlockingWebSocket, WebSocket, msg::Message}, }; + +/// Buffered bodies of known length up to this size are read on the event loop, and on a +/// blocking caller without first releasing the GIL. +const READ_ATTACHED: u64 = 64 * 1024; + +/// The largest unread body the event loop polls for a response of `version`: `limit` up to +/// HTTP/1.1, else 0. An HTTP/2 stream shares its connection's state behind a lock that the +/// connection task holds while it handles frames, so polling one would stall the loop. +fn loop_limit(version: wreq::Version, limit: u64) -> u64 { + if version <= wreq::Version::HTTP_11 { + limit + } else { + 0 + } +} diff --git a/src/client/resp/ext.rs b/src/client/resp/ext.rs index 6076eed6..f516422a 100644 --- a/src/client/resp/ext.rs +++ b/src/client/resp/ext.rs @@ -1,7 +1,7 @@ use pyo3::pybacked::PyBackedStr; use serde::de::DeserializeOwned; -use crate::{buffer::PyBuffer, error::Error}; +use crate::error::Error; /// Body readers for [`wreq::Response`] that return crate [`Error`]s, shared by the sync /// and async responses. @@ -11,9 +11,6 @@ pub trait ResponseExt { /// Deserialize the body as JSON. async fn json(self) -> Result; - - /// Read the whole body as a read-only buffer. - async fn bytes(self) -> Result; } impl ResponseExt for wreq::Response { @@ -30,12 +27,4 @@ impl ResponseExt for wreq::Response { async fn json(self) -> Result { self.json::().await.map_err(Error::Library) } - - #[inline] - async fn bytes(self) -> Result { - self.bytes() - .await - .map(PyBuffer::from) - .map_err(Error::Library) - } } diff --git a/src/client/resp/http.rs b/src/client/resp/http.rs index 5abaa735..17ef7fe7 100644 --- a/src/client/resp/http.rs +++ b/src/client/resp/http.rs @@ -6,17 +6,14 @@ use std::{ }; use bytes::Bytes; -use futures_util::{ - FutureExt, TryFutureExt, - future::{self, BoxFuture}, -}; +use futures_util::{FutureExt, TryFutureExt, future::Either}; use http::response::{Parts, Response as HttpResponse}; use http_body::Body as _; use http_body_util::{BodyExt, Collected, combinators::Collect}; use pyo3::{prelude::*, pybacked::PyBackedStr, sync::PyOnceLock}; use wreq::Uri; -use super::{ext::ResponseExt, stream::Streamer}; +use super::{READ_ATTACHED, ext::ResponseExt, loop_limit, stream::Streamer}; use crate::{ buffer::PyBuffer, client::{SocketAddr, body::Json, nogil}, @@ -32,10 +29,9 @@ use crate::{ /// A response from a request. /// -/// Async reads take the body with [`Response::take_bytes`]: a buffered body of known length -/// up to the read's limit finishes on the event loop, and any other body is collected on the -/// runtime by [`Response::collect_later`]. [`BlockingResponse`] waits on the caller through -/// [`Response::read_body`]. +/// Reads take the body with [`Response::take_bytes`]: a buffered body of known length up to +/// the read's limit finishes on the caller, the event loop or a blocking thread, and any other +/// body is collected on the runtime by [`Response::collect_later`]. #[pyclass(subclass, frozen, str, skip_from_py_object)] pub struct Response { uri: Uri, @@ -66,8 +62,8 @@ enum Body { /// A body taken for reading. enum BodyRead { - /// Read in full now, or failed. - Ready(Result), + /// Read in full now. + Ready(Bytes), /// Not fully buffered yet; collected on the runtime by [`Response::collect_later`]. Pending(Collect), } @@ -84,10 +80,6 @@ struct RecycleGuard(Option); // ===== impl Response ===== impl Response { - /// Bodies up to this size are read on a blocking caller without first releasing the GIL; - /// `bytes()` and `text()` read and decode them on the event loop when already buffered. - const READ_ATTACHED: u64 = 64 * 1024; - /// Buffered JSON bodies up to this size are parsed on the event loop; larger ones are /// parsed on the runtime. const JSON_ATTACHED: u64 = 8 * 1024; @@ -122,44 +114,20 @@ impl Response { wreq::Response::from(response) } - /// Take the body and return a future that reads it in full, caching the bytes for later - /// reads; a cached body is shared at once. While a first read runs, overlapping reads and - /// `stream()` fail with [`Error::Memory`], as do later ones if it fails or is dropped. - fn cache_response(&self) -> BoxFuture<'static, Result> { - let mut slot = self.slot(); - let stream = match mem::replace(&mut *slot, Body::Taken) { - Body::Unread(stream) => stream, - other => { - let cached = match &other { - Body::Cached(bytes) => Some(bytes.clone()), - _ => None, - }; - *slot = other; - drop(slot); - let response = cached.map(|bytes| self.build_response(bytes)); - return future::ready(response.ok_or(Error::Memory)).boxed(); - } - }; - drop(slot); - self.collect_later(stream.collect()) - .map_ok(|(parts, bytes)| wreq::Response::from(HttpResponse::from_parts(parts, bytes))) - .boxed() - } - /// Take the body for reading. Cached bytes are shared at once, and a body of known /// length up to `limit` that is already buffered is read now on the calling thread; /// overlapping reads and `stream()` fail with [`Error::Memory`] while a read runs. - fn take_bytes(&self, limit: u64) -> BodyRead { + fn take_bytes(&self, limit: u64) -> Result { let mut slot = self.slot(); let body = match mem::replace(&mut *slot, Body::Taken) { Body::Unread(body) => body, Body::Cached(bytes) => { *slot = Body::Cached(bytes.clone()); - return BodyRead::Ready(Ok(bytes)); + return Ok(BodyRead::Ready(bytes)); } other => { *slot = other; - return BodyRead::Ready(Err(Error::Memory)); + return Err(Error::Memory); } }; drop(slot); @@ -175,16 +143,16 @@ impl Response { Some(Ok(collected)) => { let bytes = collected.to_bytes(); cache(&self.body, &bytes); - return BodyRead::Ready(Ok(bytes)); + return Ok(BodyRead::Ready(bytes)); } Some(Err(err)) => { self.forbid_recycle(); - return BodyRead::Ready(Err(Error::Library(err))); + return Err(Error::Library(err)); } None => {} } } - BodyRead::Pending(collect) + Ok(BodyRead::Pending(collect)) } /// Collect a taken body, caching its bytes and returning them with a copy of the head. @@ -222,18 +190,40 @@ impl Response { Ok(self.build_response(body)) } - /// Read the body with `read`: the body is taken now and its bytes cached for later reads. - fn read_body(&self, read: F) -> impl Future> + Send + 'static + /// Take the body and return a future decoding it with `read`, and whether that future + /// finishes at once: the bytes, cached or already buffered, are at most `limit`. Otherwise + /// the body is still arriving or too large to decode on the caller. + fn read_body( + &self, + limit: u64, + read: F, + ) -> Result< + ( + impl Future> + Send + 'static, + bool, + ), + Error, + > where F: FnOnce(wreq::Response) -> Fut + Send + 'static, Fut: Future> + Send + 'static, { - self.cache_response().and_then(read).map_err(Into::into) + Ok(match self.take_bytes(limit)? { + BodyRead::Ready(bytes) => { + let inline = bytes.len() as u64 <= limit; + (Either::Left(read(self.build_response(bytes))), inline) + } + BodyRead::Pending(collect) => { + let read = self.collect_later(collect).and_then(|(parts, bytes)| { + read(wreq::Response::from(HttpResponse::from_parts(parts, bytes))) + }); + (Either::Right(read), false) + } + }) } - /// Read and decode the body once awaited; the body is taken on first await. A body up - /// to `limit` that is already buffered is decoded on the event loop; otherwise the read - /// continues on the runtime. + /// Read and decode the body once awaited; the body is taken on first await. Bytes that + /// [`read_body`](Self::read_body) can decode at once are decoded on the event loop. fn read<'py, F, Fut, T>( slf: Bound<'py, Self>, qualname: &'static str, @@ -249,47 +239,17 @@ impl Response { let slf = slf.unbind(); coroutine::local(py, qualname, async move { let this = slf.get(); - let (mut fut, attached) = match this.take_bytes(limit) { - BodyRead::Ready(bytes) => { - let bytes = bytes?; - let attached = bytes.len() as u64 <= limit; - (read(this.build_response(bytes)).boxed(), attached) - } - BodyRead::Pending(collect) => { - let fut = this.collect_later(collect).and_then(|(parts, bytes)| { - read(wreq::Response::from(HttpResponse::from_parts(parts, bytes))) - }); - (fut.boxed(), false) - } - }; - if attached { - // Decoding a buffered body finishes at once. - let ready = { - let _runtime = this.runtime.handle().enter(); - (&mut fut).now_or_never() - }; - if let Some(output) = ready { - return output.map_err(Into::into); - } + let limit = loop_limit(this.parts.version, limit); + let (read, inline) = this.read_body(limit, read)?; + let read = read.map_err(Into::into); + if inline { + read.await + } else { + coroutine::run(this.runtime.clone(), read).await } - coroutine::run(this.runtime.clone(), fut.map_err(Into::into)).await }) } - /// Whether the body is small enough for a blocking read to finish without first - /// releasing the GIL. An unknown length, as of a chunked or decompressed body, is not: - /// decoding everything already buffered could hold the GIL for long. - fn read_attached(&self) -> bool { - match &*self.slot() { - Body::Unread(body) => body - .size_hint() - .exact() - .is_some_and(|len| len <= Self::READ_ATTACHED), - Body::Cached(bytes) => bytes.len() as u64 <= Self::READ_ATTACHED, - Body::Taken | Body::Released => true, - } - } - /// Keep the connection out of the pool; a rebuilt response shares its reuse flag. fn forbid_recycle(&self) { let mut response = HttpResponse::new(Bytes::new()); @@ -423,7 +383,7 @@ impl Response { slf: Bound<'_, Self>, encoding: Option, ) -> PyResult> { - Self::read(slf, "Response.text", Self::READ_ATTACHED, |resp| { + Self::read(slf, "Response.text", READ_ATTACHED, |resp| { ResponseExt::text(resp, encoding) }) } @@ -444,8 +404,8 @@ impl Response { let slf = slf.unbind(); coroutine::local(py, "Response.bytes", async move { let this = slf.get(); - match this.take_bytes(Self::READ_ATTACHED) { - BodyRead::Ready(bytes) => bytes.map(PyBuffer::from).map_err(Into::into), + match this.take_bytes(loop_limit(this.parts.version, READ_ATTACHED))? { + BodyRead::Ready(bytes) => Ok(PyBuffer::from(bytes)), BodyRead::Pending(collect) => { let read = this .collect_later(collect) @@ -467,7 +427,7 @@ impl Response { let slf = slf.unbind(); coroutine::local(py, "Response.close", async move { slf.get().discard(); - Ok(None::<()>) + Ok(()) }) } } @@ -490,7 +450,7 @@ impl Response { let slf = slf.unbind(); coroutine::local(py, "Response.__aexit__", async move { slf.get().destroy(); - Ok(None::<()>) + Ok(()) }) } } @@ -534,22 +494,21 @@ impl Drop for RecycleGuard { // ===== impl BlockingResponse ===== impl BlockingResponse { - /// Read the body with `read` on the calling thread. A small body is read without - /// releasing the GIL unless it must wait. + /// Read the body with `read` on the calling thread. Bytes that + /// [`Response::read_body`] can decode at once are decoded without releasing the GIL. fn read(&self, py: Python, read: F) -> PyResult where F: FnOnce(wreq::Response) -> Fut + Send + 'static, Fut: Future> + Send + 'static, T: Send, { - let (response, runtime) = (&self.0, &self.0.runtime); - // Decide before the read takes the body. - let attached = response.read_attached(); - let fut = response.read_body(read); - if attached { - nogil::run(py, runtime, fut) + let runtime = &self.0.runtime; + let (read, inline) = self.0.read_body(READ_ATTACHED, read)?; + let read = read.map_err(Into::into); + if inline { + nogil::run(py, runtime, read) } else { - py.detach(|| runtime.handle().block_on(fut)) + py.detach(|| runtime.handle().block_on(read)) } } } @@ -639,7 +598,17 @@ impl BlockingResponse { /// Read the body as a read-only memoryview, retaining its data after the response closes. pub fn bytes(&self, py: Python) -> PyResult { - self.read(py, ResponseExt::bytes) + let response = &self.0; + match response.take_bytes(READ_ATTACHED)? { + BodyRead::Ready(bytes) => Ok(PyBuffer::from(bytes)), + BodyRead::Pending(collect) => { + let read = response + .collect_later(collect) + .map_ok(|(_, bytes)| PyBuffer::from(bytes)) + .map_err(Into::into); + py.detach(|| response.runtime.handle().block_on(read)) + } + } } /// Discard the retained body and mark its connection as non-reusable. diff --git a/src/client/resp/stream.rs b/src/client/resp/stream.rs index 2fc160b2..8c8d621c 100644 --- a/src/client/resp/stream.rs +++ b/src/client/resp/stream.rs @@ -22,6 +22,7 @@ use tokio::sync::{ }; use tokio_util::task::AbortOnDropHandle; +use super::{READ_ATTACHED, loop_limit}; use crate::{ buffer::PyBuffer, client::nogil, @@ -84,9 +85,6 @@ impl Streamer { /// Buffered frames returned before `__anext__` yields to the event loop once. const YIELD_EVERY: usize = 8; - /// An async reader reads a buffered body of known length up to this size at once. - const READ_ATTACHED: u64 = 64 * 1024; - /// Create a new [`Streamer`] instance. #[inline] pub fn new(resp: wreq::Response, runtime: Runtime) -> Streamer { @@ -162,15 +160,18 @@ impl Streamer { /// Read a frame the body already holds, before any read-ahead task starts, so a reader /// of a short body never waits on another thread. With `limit`, only a body of known - /// length up to it is read here, so an async reader never decodes on the event loop; - /// `end` builds the error that ends iteration. + /// length up to it is read here, so an async reader never reads a large or decompressing + /// body on the event loop; `end` builds the error that ends iteration. fn ready_frame(&self, limit: Option, end: fn() -> Error) -> Option> { let mut state = self.reader.lock(); let State::Idle(resp) = &mut *state else { return None; }; if let Some(limit) = limit - && !resp.size_hint().exact().is_some_and(|len| len <= limit) + && !resp + .size_hint() + .exact() + .is_some_and(|len| len <= loop_limit(resp.version(), limit)) { return None; } @@ -309,8 +310,7 @@ impl Streamer { coroutine::local(py, "Streamer.__anext__", async move { let this = slf.get(); // A short body already buffered is read without starting the read-ahead task. - if let Some(frame) = - this.ready_frame(Some(Self::READ_ATTACHED), || Error::StopAsyncIteration) + if let Some(frame) = this.ready_frame(Some(READ_ATTACHED), || Error::StopAsyncIteration) { return frame; } diff --git a/src/coroutine.rs b/src/coroutine.rs index 048b5183..0aa22b90 100644 --- a/src/coroutine.rs +++ b/src/coroutine.rs @@ -15,6 +15,7 @@ mod awaitable; mod scope; use std::{ + any::TypeId, future::{Future, poll_fn}, mem, task::Poll, @@ -54,7 +55,7 @@ pub fn local<'py, F, T>( ) -> PyResult> where F: Future> + Send + 'static, - T: for<'a> IntoPyObject<'a>, + T: for<'a> IntoPyObject<'a> + 'static, { Bound::new(py, coroutine(qualname, fut)) } @@ -83,21 +84,19 @@ pub fn managed<'py, F, T>( ) -> PyResult> where F: Future> + Send + 'static, - T: EntersSelf + for<'a> IntoPyObject<'a>, + T: EntersSelf + for<'a> IntoPyObject<'a> + 'static, { Bound::new(py, coroutine(qualname, fut).managed(enters_self::)) } -/// Whether `object` is exactly a `T` whose class still has the native `__aenter__`: a -/// subclass or a replaced class attribute may return something else. -fn enters_self(object: &Bound<'_, PyAny>) -> bool { - let py = object.py(); - let ty = object.get_type(); - ty.is(T::type_object(py)) - && T::native_aenter().get(py).is_some_and(|native| { - ty.getattr(intern!(py, "__aenter__")) - .is_ok_and(|aenter| aenter.is(native)) - }) +/// Whether `T` still has its native `__aenter__`; a [`managed`] result is always exactly a +/// `T`, but user code may replace the class attribute. +fn enters_self(py: Python<'_>) -> bool { + T::native_aenter().get(py).is_some_and(|native| { + T::type_object(py) + .getattr(intern!(py, "__aenter__")) + .is_ok_and(|aenter| aenter.is(native)) + }) } /// A coroutine that returns `value` without suspending. @@ -126,15 +125,21 @@ pub async fn yield_now() { .await; } -/// A coroutine awaiting `fut` and converting its output to a Python object. +/// A coroutine awaiting `fut` and converting its output to a Python object. Like a sync +/// method, one with no result returns `None`, where `()` would convert to an empty tuple. fn coroutine(qualname: &'static str, fut: F) -> Coroutine where F: Future> + Send + 'static, - T: for<'a> IntoPyObject<'a>, + T: for<'a> IntoPyObject<'a> + 'static, { Coroutine::new(qualname, async move { let value = fut.await?; - Python::attach(|py| value.into_py_any(py)) + Python::attach(|py| { + if TypeId::of::() == TypeId::of::<()>() { + return Ok(py.None()); + } + value.into_py_any(py) + }) }) } diff --git a/src/coroutine/awaitable.rs b/src/coroutine/awaitable.rs index 195c22dc..3d704ce3 100644 --- a/src/coroutine/awaitable.rs +++ b/src/coroutine/awaitable.rs @@ -10,15 +10,14 @@ use std::{ use futures_util::{FutureExt, future::BoxFuture}; use pyo3::{ PyTraverseError, PyVisit, - exceptions::{PyBaseException, PyRuntimeError, PyStopIteration}, - ffi, intern, + exceptions::{PyRuntimeError, PyStopIteration}, + intern, prelude::*, - types::PyTuple, }; use super::{ Port, - scope::{EntersSelf, Scope}, + scope::{EntersSelfFn, Scope}, }; /// An awaitable driving a Rust future on the asyncio event loop thread. @@ -83,7 +82,7 @@ impl Coroutine { } /// Let `async with` enter the coroutine; see [`coroutine::managed`](super::managed). - pub(super) fn managed(mut self, enters_self: EntersSelf) -> Self { + pub(super) fn managed(mut self, enters_self: EntersSelfFn) -> Self { self.scope = Scope::Ready(enters_self); self } @@ -176,8 +175,9 @@ impl Coroutine { fn __next__(&mut self, py: Python<'_>) -> PyResult>> { match self.step(py, None)? { - Step::Yield(value) => Ok(Some(value)), - Step::Return(value) => stop_iteration(value.into_bound(py)), + // A `NULL` return without an error ends iteration with `None`, no `StopIteration`. + Step::Return(value) if value.is_none(py) => Ok(None), + step => step.into_result().map(Some), } } @@ -213,28 +213,6 @@ impl Step { } } -/// End `__next__` with `value`, skipping PyO3's lazy `StopIteration`. -/// -/// A `NULL` return without an error means `None`. Other values are raised with -/// `PyErr_SetObject`, as CPython does for generators: 3.11 keeps `(StopIteration, value)` -/// unnormalized for `_PyGen_FetchStopIterationValue`, while 3.12+ builds the instance at once. -/// Tuples and exceptions keep PyO3's path, since CPython would unpack or adopt them. -fn stop_iteration(value: Bound<'_, PyAny>) -> PyResult>> { - if value.is_none() { - return Ok(None); - } - if value.is_instance_of::() || value.is_instance_of::() { - return Err(PyStopIteration::new_err((value.unbind(),))); - } - // SAFETY: the thread is attached and `value` is a valid object for the call; - // `PyErr_SetObject` takes its own references to the type and the value. - #[allow(unsafe_code)] - unsafe { - ffi::PyErr_SetObject(ffi::PyExc_StopIteration, value.as_ptr()) - }; - Ok(None) -} - /// Fail if `waiter`, the future a task waits on for this coroutine, is still pending. pub(super) fn ensure_done(waiter: &Bound<'_, PyAny>) -> PyResult<()> { if waiter diff --git a/src/coroutine/scope.rs b/src/coroutine/scope.rs index 85d5af36..2288fbaf 100644 --- a/src/coroutine/scope.rs +++ b/src/coroutine/scope.rs @@ -13,18 +13,18 @@ use pyo3::{ use super::awaitable::{Coroutine, Step, ensure_done}; /// Whether a result's `__aenter__` returns the result itself, so it can be skipped. -pub(super) type EntersSelf = fn(&Bound<'_, PyAny>) -> bool; +pub(super) type EntersSelfFn = fn(Python<'_>) -> bool; /// The `async with` state of a coroutine. pub(super) enum Scope { /// `async with` is unsupported. Unsupported, /// Not yet awaited; `async with` may enter it. - Ready(EntersSelf), + Ready(EntersSelfFn), /// Awaited, entered or exited already. Spent, /// Awaited by `async with` until the result, an async context manager, is ready. - Entering(EntersSelf), + Entering(EntersSelfFn), /// Awaiting the manager's `__aenter__` through `delegate`. Like `async with`, its /// bound `__aexit__` is looked up before entering. Opening { @@ -104,12 +104,12 @@ impl Coroutine { &mut self, py: Python<'_>, manager: Py, - enters_self: EntersSelf, + enters_self: EntersSelfFn, ) -> PyResult { self.finish(); let manager = manager.into_bound(py); // Its `__aenter__` would return the manager itself: enter without awaiting it. - if enters_self(&manager) { + if enters_self(py) { let exit = manager.getattr(intern!(py, "__aexit__"))?; self.scope = Scope::Entered(exit.unbind()); return Ok(Step::Return(manager.unbind())); diff --git a/src/macros.rs b/src/macros.rs index ef545316..3cf499bb 100644 --- a/src/macros.rs +++ b/src/macros.rs @@ -19,9 +19,19 @@ pub(crate) fn remaining(ob: &Bound<'_, PyAny>) -> usize { ob.cast::().map_or(usize::MAX, |dict| dict.len()) } +macro_rules! extract_option { + ($ob:expr, $params:expr, $field:ident) => { + if let Some(value) = + $crate::macros::lookup(&$ob, pyo3::intern!($ob.py(), stringify!($field))) + { + $params.$field = value.extract()?; + } + }; +} + /// Like [`extract_option!`] for several fields, but stop looking up keys once every key -/// of a dict has been found, counting `$found` keys already looked up by the caller. Unknown -/// keys are never counted, so they still leave the scan complete and stay ignored. +/// of a dict has been found, counting `$found` keys the caller already looked up. Unknown +/// keys are never found, so they make the scan run to the end and stay ignored. macro_rules! extract_options { ($ob:expr, $params:expr, $found:expr, [$($field:ident),* $(,)?]) => {{ let mut remaining = $crate::macros::remaining(&$ob).saturating_sub($found); @@ -37,16 +47,6 @@ macro_rules! extract_options { }}; } -macro_rules! extract_option { - ($ob:expr, $params:expr, $field:ident) => { - if let Some(value) = - $crate::macros::lookup(&$ob, pyo3::intern!($ob.py(), stringify!($field))) - { - $params.$field = value.extract()?; - } - }; -} - macro_rules! apply_option { (set_if_some, $builder:expr, $option:expr, $method:ident) => { if let Some(value) = $option.take() { diff --git a/tests/response_test.py b/tests/response_test.py index 6a022bc0..4aa7f7ac 100644 --- a/tests/response_test.py +++ b/tests/response_test.py @@ -1,4 +1,5 @@ import asyncio +import json import socket import threading import time @@ -156,6 +157,99 @@ async def test_context_exit_keeps_connection_but_close_forbids_reuse(): assert connections.empty() +# 2**63 overflows `isize`, so it decodes as a float. +JSON_ITEM = {"n": [0, -2, 1.5, 2**63, True, None], "s": "\u00e9\n", "o": {}} + + +@pytest.mark.asyncio +@pytest.mark.parametrize("count", [1, 200, 1200], ids=["inline", "runtime", "large"]) +async def test_reads_decode_buffered_and_cached_bodies(count): + payload = json.dumps([JSON_ITEM] * count).encode() + reply = b"HTTP/1.1 200 OK\r\nConnection: close\r\nContent-Length: %d\r\n\r\n%s" % ( + len(payload), + payload, + ) + # A read timeout starts a timer even when the body is read on the event loop. + options = {"proxies": [], "read_timeout": timedelta(seconds=5)} + + def read_blocking(url): + with wreq.blocking.Client(**options).get(url) as response: + return response.json(), response.text(), bytes(response.bytes()) + + async with local_server() as (url, connections): + async with wreq.Client(**options) as client: + task = asyncio.create_task(client.get(url)) + _, writer = await asyncio.wait_for(connections.get(), 5) + writer.write(reply) + response = await asyncio.wait_for(task, 5) + value = await response.json() + reads = value, await response.text(), bytes(await response.bytes()) + + task = asyncio.create_task(asyncio.to_thread(read_blocking, url)) + _, writer = await asyncio.wait_for(connections.get(), 5) + writer.write(reply) + assert await asyncio.wait_for(task, 5) == reads + + assert reads == ([JSON_ITEM] * count, payload.decode(), payload) + assert isinstance(value[0]["n"][3], float) + + +@pytest.mark.asyncio +async def test_empty_body_and_unit_results(): + async with local_server() as (url, connections): + client = wreq.Client(proxies=[]) + + async def get(): + task = asyncio.create_task(client.get(url)) + _, writer = await asyncio.wait_for(connections.get(), 5) + writer.write(b"HTTP/1.1 200 OK\r\nConnection: close\r\nContent-Length: 0\r\n\r\n") + return await asyncio.wait_for(task, 5) + + response = await get() + assert bytes(await response.bytes()) == b"" + assert await response.text() == "" + assert await response.__aexit__(None, None, None) is None + + response = await get() + streamer = response.stream() + assert [chunk async for chunk in streamer] == [] + assert await streamer.__aexit__(None, None, None) is None + # Coroutines without a result return None, as sync methods do. + coroutine = response.close() + with pytest.raises(StopIteration) as stop: + coroutine.__next__() + assert stop.value.value is None + assert await client.__aexit__(None, None, None) is None + + +@pytest.mark.asyncio +async def test_pending_read_holds_the_body_until_it_ends(): + head = b"HTTP/1.1 200 OK\r\nContent-Length: 10\r\n\r\n" + async with local_server() as (url, connections), wreq.Client(proxies=[]) as client: + task = asyncio.create_task(client.get(url)) + reader, writer = await asyncio.wait_for(connections.get(), 5) + writer.write(head + b"hello") + response = await asyncio.wait_for(task, 5) + first = asyncio.create_task(response.text()) + await asyncio.sleep(0.1) + with pytest.raises(RuntimeError): + await response.bytes() + writer.write(b"world") + assert await asyncio.wait_for(first, 5) == "helloworld" + assert bytes(await response.bytes()) == b"helloworld" + + # The fully read connection is reused; a truncated body fails, and so do later reads. + task = asyncio.create_task(client.get(url)) + await asyncio.wait_for(reader.readuntil(b"\r\n\r\n"), 5) + writer.write(head + b"hello") + response = await asyncio.wait_for(task, 5) + writer.close() + with pytest.raises(wreq.exceptions.DecodingError): + await asyncio.wait_for(response.bytes(), 5) + with pytest.raises(RuntimeError): + await response.text() + + @pytest.mark.asyncio async def test_stream_read_ahead_yields_and_closes_waiting_readers(): stalled = threading.Event() diff --git a/tests/upload_test.py b/tests/upload_test.py index 8848a00b..6fa3442f 100644 --- a/tests/upload_test.py +++ b/tests/upload_test.py @@ -96,10 +96,14 @@ async def chunks(): started.set() yield b"first" + form = wreq.Multipart(wreq.Part("file", iter([b"x"]))) async with local_server() as (url, _), wreq.Client(proxies=[]) as client: - # Options are validated before the body generator is consumed. + # Every option, `zstd` last, is validated before the body or form is taken. with pytest.raises(TypeError): - await client.post(url, body=chunks(), timeout=5) + await client.post(url, body=chunks(), zstd=1) + for _ in range(2): + with pytest.raises(TypeError): + await client.post(url, multipart=form, zstd=1) await asyncio.sleep(0.1) assert not started.is_set() From 967a82e16254e9de698f84a62445fc3a0901defa Mon Sep 17 00:00:00 2001 From: 0x676e67 Date: Mon, 5 Oct 2026 23:00:00 +0000 Subject: [PATCH 3/8] refactor(aio): inline single-use helpers on the read and request paths --- src/client/body/stream.rs | 10 ++-- src/client/req.rs | 24 ++------ src/client/resp/http.rs | 109 +++++++++++++++---------------------- src/coroutine.rs | 24 ++++---- src/coroutine/asyncio.rs | 7 +-- src/coroutine/awaitable.rs | 13 +---- src/coroutine/scope.rs | 20 +++---- src/macros.rs | 26 ++++++--- tests/response_test.py | 3 +- 9 files changed, 93 insertions(+), 143 deletions(-) diff --git a/src/client/body/stream.rs b/src/client/body/stream.rs index a78f594c..bcba68cd 100644 --- a/src/client/body/stream.rs +++ b/src/client/body/stream.rs @@ -266,13 +266,11 @@ async def forward(gen, sender): .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(intern!(py, "create_task"), (&coroutine,)) { - Ok(task) => task, - Err(err) => { + let task = event_loop + .call_method1(intern!(py, "create_task"), (&coroutine,)) + .inspect_err(|_| { let _ = coroutine.call_method0(intern!(py, "close")); - return Err(err); - } - }; + })?; Ok(Self { rx, task: Some((task.unbind(), event_loop.unbind())), diff --git a/src/client/req.rs b/src/client/req.rs index 92d7a172..3b079bd2 100644 --- a/src/client/req.rs +++ b/src/client/req.rs @@ -5,9 +5,7 @@ use std::{ use futures_util::TryFutureExt; use http::header::COOKIE; -use pyo3::{ - PyResult, exceptions::asyncio::CancelledError, intern, prelude::*, pybacked::PyBackedStr, -}; +use pyo3::{PyResult, exceptions::asyncio::CancelledError, prelude::*, pybacked::PyBackedStr}; use crate::{ client::{ @@ -22,7 +20,6 @@ use crate::{ extractor::Extractor, header::{HeaderMap, OrigHeaderMap}, http::{Method, Version}, - macros::lookup, proxy::Proxy, redirect, }; @@ -189,17 +186,11 @@ impl FromPyObject<'_, '_> for Request { fn extract(ob: Borrowed) -> PyResult { let mut request = Self::default(); - // Extracting a body or multipart form can start a generator or take stream parts, so - // they are extracted last, after every other option has been validated. - let py = ob.py(); - let body = lookup(&ob, intern!(py, "body")); - let multipart = lookup(&ob, intern!(py, "multipart")); - let found = usize::from(body.is_some()) + usize::from(multipart.is_some()); - // Common keys first: the scan stops once every given key is found. + // Common keys first. A body or multipart form can start a generator or take stream + // parts, so they are extracted last, after every other option has been validated. extract_options!( ob, request, - found, [ json, form, @@ -225,14 +216,9 @@ impl FromPyObject<'_, '_> for Request { brotli, deflate, zstd, - ] + ], + [body, multipart] ); - if let Some(body) = body { - request.body = body.extract()?; - } - if let Some(multipart) = multipart { - request.multipart = multipart.extract()?; - } Ok(request) } } diff --git a/src/client/resp/http.rs b/src/client/resp/http.rs index 17ef7fe7..9821712a 100644 --- a/src/client/resp/http.rs +++ b/src/client/resp/http.rs @@ -80,10 +80,6 @@ struct RecycleGuard(Option); // ===== impl Response ===== impl Response { - /// Buffered JSON bodies up to this size are parsed on the event loop; larger ones are - /// parsed on the runtime. - const JSON_ATTACHED: u64 = 8 * 1024; - /// Create a new [`Response`] instance. pub fn new(response: wreq::Response, runtime: Runtime) -> Self { let uri = response.uri().clone(); @@ -133,26 +129,25 @@ impl Response { drop(slot); let attached = body.size_hint().exact().is_some_and(|len| len <= limit); let mut collect = body.collect(); - if attached { + let ready = if attached { // A read timeout starts a timer on its first poll, which needs the runtime. - let ready = { - let _runtime = self.runtime.handle().enter(); - (&mut collect).now_or_never() - }; - match ready { - Some(Ok(collected)) => { - let bytes = collected.to_bytes(); - cache(&self.body, &bytes); - return Ok(BodyRead::Ready(bytes)); - } - Some(Err(err)) => { - self.forbid_recycle(); - return Err(Error::Library(err)); - } - None => {} + let _runtime = self.runtime.handle().enter(); + (&mut collect).now_or_never() + } else { + None + }; + match ready { + Some(Ok(collected)) => { + let bytes = collected.to_bytes(); + cache(&self.body, &bytes); + Ok(BodyRead::Ready(bytes)) } + Some(Err(err)) => { + self.forbid_recycle(); + Err(Error::Library(err)) + } + None => Ok(BodyRead::Pending(collect)), } - Ok(BodyRead::Pending(collect)) } /// Collect a taken body, caching its bytes and returning them with a copy of the head. @@ -175,21 +170,6 @@ impl Response { } } - /// Take the unread body for a [`Streamer`]; fails with [`Error::Memory`] otherwise, - /// leaving any cached bytes readable. - fn stream_response(&self) -> Result { - let mut slot = self.slot(); - let body = match mem::replace(&mut *slot, Body::Taken) { - Body::Unread(body) => body, - other => { - *slot = other; - return Err(Error::Memory); - } - }; - drop(slot); - Ok(self.build_response(body)) - } - /// Take the body and return a future decoding it with `read`, and whether that future /// finishes at once: the bytes, cached or already buffered, are at most `limit`. Otherwise /// the body is still arriving or too large to decode on the caller. @@ -197,29 +177,24 @@ impl Response { &self, limit: u64, read: F, - ) -> Result< - ( - impl Future> + Send + 'static, - bool, - ), - Error, - > + ) -> Result<(impl Future> + Send + 'static, bool), Error> where F: FnOnce(wreq::Response) -> Fut + Send + 'static, Fut: Future> + Send + 'static, { - Ok(match self.take_bytes(limit)? { + let (decode, inline) = match self.take_bytes(limit)? { BodyRead::Ready(bytes) => { let inline = bytes.len() as u64 <= limit; (Either::Left(read(self.build_response(bytes))), inline) } - BodyRead::Pending(collect) => { - let read = self.collect_later(collect).and_then(|(parts, bytes)| { + BodyRead::Pending(collect) => ( + Either::Right(self.collect_later(collect).and_then(|(parts, bytes)| { read(wreq::Response::from(HttpResponse::from_parts(parts, bytes))) - }); - (Either::Right(read), false) - } - }) + })), + false, + ), + }; + Ok((decode.map_err(Into::into), inline)) } /// Read and decode the body once awaited; the body is taken on first await. Bytes that @@ -239,9 +214,7 @@ impl Response { let slf = slf.unbind(); coroutine::local(py, qualname, async move { let this = slf.get(); - let limit = loop_limit(this.parts.version, limit); - let (read, inline) = this.read_body(limit, read)?; - let read = read.map_err(Into::into); + let (read, inline) = this.read_body(loop_limit(this.parts.version, limit), read)?; if inline { read.await } else { @@ -370,11 +343,22 @@ impl Response { .map_err(Into::into) } - /// Stream read-only memoryviews and any trailing headers from the body. + /// Stream read-only memoryviews and any trailing headers from the body. Only an unread + /// body can be streamed; cached bytes stay readable. pub fn stream(&self) -> PyResult { - self.stream_response() - .map(|response| Streamer::new(response, self.runtime.clone())) - .map_err(Into::into) + let mut slot = self.slot(); + let body = match mem::replace(&mut *slot, Body::Taken) { + Body::Unread(body) => body, + other => { + *slot = other; + return Err(Error::Memory.into()); + } + }; + drop(slot); + Ok(Streamer::new( + self.build_response(body), + self.runtime.clone(), + )) } /// Get the text content with the response encoding, defaulting to utf-8 when unspecified. @@ -388,14 +372,10 @@ impl Response { }) } - /// Get the JSON content of the response. + /// Get the JSON content of the response. Buffered JSON up to 8 KiB is parsed on the event + /// loop, larger bodies on the runtime. pub fn json(slf: Bound<'_, Self>) -> PyResult> { - Self::read( - slf, - "Response.json", - Self::JSON_ATTACHED, - ResponseExt::json::, - ) + Self::read(slf, "Response.json", 8 * 1024, ResponseExt::json::) } /// Read the body as a read-only memoryview, retaining its data after the response closes. @@ -504,7 +484,6 @@ impl BlockingResponse { { let runtime = &self.0.runtime; let (read, inline) = self.0.read_body(READ_ATTACHED, read)?; - let read = read.map_err(Into::into); if inline { nogil::run(py, runtime, read) } else { diff --git a/src/coroutine.rs b/src/coroutine.rs index 0aa22b90..f6131284 100644 --- a/src/coroutine.rs +++ b/src/coroutine.rs @@ -26,9 +26,9 @@ use pyo3::{ }; use tokio_util::task::AbortOnDropHandle; -use self::asyncio::Port; pub(crate) use self::asyncio::running_loop; pub use self::awaitable::Coroutine; +use self::{asyncio::Port, scope::Scope}; use crate::runtime::Runtime; /// Run `fut` on the runtime once the coroutine named `qualname` is first awaited. @@ -86,17 +86,17 @@ where F: Future> + Send + 'static, T: EntersSelf + for<'a> IntoPyObject<'a> + 'static, { - Bound::new(py, coroutine(qualname, fut).managed(enters_self::)) -} - -/// Whether `T` still has its native `__aenter__`; a [`managed`] result is always exactly a -/// `T`, but user code may replace the class attribute. -fn enters_self(py: Python<'_>) -> bool { - T::native_aenter().get(py).is_some_and(|native| { - T::type_object(py) - .getattr(intern!(py, "__aenter__")) - .is_ok_and(|aenter| aenter.is(native)) - }) + let mut coroutine = coroutine(qualname, fut); + // The result is always exactly a `T`, so it enters itself unless user code replaced + // `T.__aenter__`. + coroutine.scope = Scope::Ready(|py| { + T::native_aenter().get(py).is_some_and(|native| { + T::type_object(py) + .getattr(intern!(py, "__aenter__")) + .is_ok_and(|aenter| aenter.is(native)) + }) + }); + Bound::new(py, coroutine) } /// A coroutine that returns `value` without suspending. diff --git a/src/coroutine/asyncio.rs b/src/coroutine/asyncio.rs index 9ce6e5ae..677dee07 100644 --- a/src/coroutine/asyncio.rs +++ b/src/coroutine/asyncio.rs @@ -314,12 +314,7 @@ impl Drop for Keeper { pub(crate) fn running_loop(py: Python<'_>) -> PyResult> { static GET_RUNNING_LOOP: PyOnceLock> = PyOnceLock::new(); GET_RUNNING_LOOP - .get_or_try_init(py, || { - py.import("asyncio")? - .getattr("get_running_loop") - .map(Bound::unbind) - })? - .bind(py) + .import(py, "asyncio", "get_running_loop")? .call0() } diff --git a/src/coroutine/awaitable.rs b/src/coroutine/awaitable.rs index 3d704ce3..c0e3ee43 100644 --- a/src/coroutine/awaitable.rs +++ b/src/coroutine/awaitable.rs @@ -15,10 +15,7 @@ use pyo3::{ prelude::*, }; -use super::{ - Port, - scope::{EntersSelfFn, Scope}, -}; +use super::{Port, scope::Scope}; /// An awaitable driving a Rust future on the asyncio event loop thread. /// @@ -26,7 +23,7 @@ use super::{ /// the C task fast path work unchanged. A throw or close drops the Rust future, /// which aborts any spawned work. PyO3's borrow flag rejects reentrant polls. /// -/// A coroutine built with [`managed`](Self::managed) also works as `async with`, as +/// A coroutine built with [`managed`](super::managed) also works as `async with`, as /// `async with await` would; the `scope` module holds that state. #[pyclass(module = "wreq")] pub struct Coroutine { @@ -81,12 +78,6 @@ impl Coroutine { } } - /// Let `async with` enter the coroutine; see [`coroutine::managed`](super::managed). - pub(super) fn managed(mut self, enters_self: EntersSelfFn) -> Self { - self.scope = Scope::Ready(enters_self); - self - } - #[inline] pub(super) fn future(&mut self) -> &mut Option>>> { self.future diff --git a/src/coroutine/scope.rs b/src/coroutine/scope.rs index 2288fbaf..419bb5ed 100644 --- a/src/coroutine/scope.rs +++ b/src/coroutine/scope.rs @@ -91,7 +91,12 @@ impl Coroutine { /// Finish with `value`, first entering it when awaited by `async with`. pub(super) fn complete(&mut self, py: Python<'_>, value: Py) -> PyResult { match mem::replace(&mut self.scope, Scope::Spent) { - Scope::Entering(enters_self) => return self.open(py, value, enters_self), + // Its `__aenter__` would return the result itself: enter without awaiting it. + Scope::Entering(enters_self) if enters_self(py) => { + let exit = value.bind(py).getattr(intern!(py, "__aexit__"))?; + self.scope = Scope::Entered(exit.unbind()); + } + Scope::Entering(_) => return self.open(py, value), Scope::Opening { exit, .. } => self.scope = Scope::Entered(exit), scope => self.scope = scope, } @@ -100,20 +105,9 @@ impl Coroutine { } /// Await the manager's `__aenter__`, as `async with` would. - fn open( - &mut self, - py: Python<'_>, - manager: Py, - enters_self: EntersSelfFn, - ) -> PyResult { + fn open(&mut self, py: Python<'_>, manager: Py) -> PyResult { self.finish(); let manager = manager.into_bound(py); - // Its `__aenter__` would return the manager itself: enter without awaiting it. - if enters_self(py) { - let exit = manager.getattr(intern!(py, "__aexit__"))?; - self.scope = Scope::Entered(exit.unbind()); - return Ok(Step::Return(manager.unbind())); - } // Results have no instance dict, so binding on the instance matches `async with`. let (Ok(enter), Ok(exit)) = ( manager.getattr(intern!(py, "__aenter__")), diff --git a/src/macros.rs b/src/macros.rs index 3cf499bb..2c2a0336 100644 --- a/src/macros.rs +++ b/src/macros.rs @@ -14,11 +14,6 @@ pub(crate) fn lookup<'py>( } } -/// Keys left to look up: a dict's length, or unbounded for other mappings. -pub(crate) fn remaining(ob: &Bound<'_, PyAny>) -> usize { - ob.cast::().map_or(usize::MAX, |dict| dict.len()) -} - macro_rules! extract_option { ($ob:expr, $params:expr, $field:ident) => { if let Some(value) = @@ -30,11 +25,19 @@ macro_rules! extract_option { } /// Like [`extract_option!`] for several fields, but stop looking up keys once every key -/// of a dict has been found, counting `$found` keys the caller already looked up. Unknown -/// keys are never found, so they make the scan run to the end and stay ignored. +/// of a dict has been found. `$last` fields are looked up first and extracted after the +/// rest. Unknown keys are never found, so they make the scan run to the end and stay ignored. macro_rules! extract_options { - ($ob:expr, $params:expr, $found:expr, [$($field:ident),* $(,)?]) => {{ - let mut remaining = $crate::macros::remaining(&$ob).saturating_sub($found); + ($ob:expr, $params:expr, [$($field:ident),* $(,)?], [$($last:ident),* $(,)?]) => {{ + $( + let $last = + $crate::macros::lookup(&$ob, pyo3::intern!($ob.py(), stringify!($last))); + )* + // A dict's length bounds the keys left to find; other mappings are scanned in full. + let mut remaining = $ob + .cast::() + .map_or(usize::MAX, |dict| dict.len()) + $(.saturating_sub(usize::from($last.is_some())))*; $( if remaining > 0 && let Some(value) = @@ -44,6 +47,11 @@ macro_rules! extract_options { remaining -= 1; } )* + $( + if let Some(value) = $last { + $params.$last = value.extract()?; + } + )* }}; } diff --git a/tests/response_test.py b/tests/response_test.py index 4aa7f7ac..5ac1d6b2 100644 --- a/tests/response_test.py +++ b/tests/response_test.py @@ -215,9 +215,8 @@ async def get(): assert [chunk async for chunk in streamer] == [] assert await streamer.__aexit__(None, None, None) is None # Coroutines without a result return None, as sync methods do. - coroutine = response.close() with pytest.raises(StopIteration) as stop: - coroutine.__next__() + response.close().__next__() assert stop.value.value is None assert await client.__aexit__(None, None, None) is None From 94c3bda873fa5b70d6fe0b8423f0581865e0f9c7 Mon Sep 17 00:00:00 2001 From: 0x676e67 Date: Mon, 5 Oct 2026 23:00:00 +0000 Subject: [PATCH 4/8] chore(deps): patch netty to skip stale waker wakes on body drop --- Cargo.lock | 6 +++--- Cargo.toml | 2 +- 2 files changed, 4 insertions(+), 4 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index 604e59f5..98ad2764 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -102,7 +102,7 @@ dependencies = [ "regex", "rustc-hash", "shlex", - "syn 3.0.5", + "syn 2.0.119", ] [[package]] @@ -1178,7 +1178,7 @@ checksum = "27b02d87554356db9e9a873add8782d4ea6e3e58ea071a9adb9a2e8ddb884a8b" [[package]] name = "netty" version = "0.2.5" -source = "git+https://github.com/0x676e67/netty?branch=main#f8652ddf5e386d66b7f03cceedbc6b319d849762" +source = "git+https://github.com/0x676e67/netty?branch=demo%2Fskip-stale-drop-wake#27d2077f5ad6022efc267d8b8cfda1f164c6fb48" dependencies = [ "bytes", "futures-channel", @@ -2296,7 +2296,7 @@ version = "0.1.11" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "c2a7b1c03c876122aa43f3020e6c3c3ee5c05081c9a00739faf7503aeba10d22" dependencies = [ - "windows-sys 0.61.2", + "windows-sys 0.52.0", ] [[package]] diff --git a/Cargo.toml b/Cargo.toml index e56a234c..cd9ac929 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -81,7 +81,7 @@ wreq = { git = "https://github.com/0x676e67/wreq", rev = "54fdc7c9e0c1290e2f70f3 btls = { git = "https://github.com/0x676e67/btls", branch = "main" } btls-sys = { git = "https://github.com/0x676e67/btls", branch = "main" } tokio-btls = { git = "https://github.com/0x676e67/btls", branch = "main" } -netty = { git = "https://github.com/0x676e67/netty", branch = "main" } +netty = { git = "https://github.com/0x676e67/netty", branch = "demo/skip-stale-drop-wake" } [profile.release] codegen-units = 1 From fbd93a9b670d6326729dc788642baddfa1471e95 Mon Sep 17 00:00:00 2001 From: 0x676e67 Date: Tue, 6 Oct 2026 00:07:59 +0000 Subject: [PATCH 5/8] chore(deps): track netty main with the stale wake fix --- Cargo.lock | 2 +- Cargo.toml | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index 98ad2764..23ff2ad8 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1178,7 +1178,7 @@ checksum = "27b02d87554356db9e9a873add8782d4ea6e3e58ea071a9adb9a2e8ddb884a8b" [[package]] name = "netty" version = "0.2.5" -source = "git+https://github.com/0x676e67/netty?branch=demo%2Fskip-stale-drop-wake#27d2077f5ad6022efc267d8b8cfda1f164c6fb48" +source = "git+https://github.com/0x676e67/netty?branch=main#663b094fe80389c393880d8d5168c73131629336" dependencies = [ "bytes", "futures-channel", diff --git a/Cargo.toml b/Cargo.toml index cd9ac929..e56a234c 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -81,7 +81,7 @@ wreq = { git = "https://github.com/0x676e67/wreq", rev = "54fdc7c9e0c1290e2f70f3 btls = { git = "https://github.com/0x676e67/btls", branch = "main" } btls-sys = { git = "https://github.com/0x676e67/btls", branch = "main" } tokio-btls = { git = "https://github.com/0x676e67/btls", branch = "main" } -netty = { git = "https://github.com/0x676e67/netty", branch = "demo/skip-stale-drop-wake" } +netty = { git = "https://github.com/0x676e67/netty", branch = "main" } [profile.release] codegen-units = 1 From b75d655c214c87bad5474ab50c19c57cc4390e3b Mon Sep 17 00:00:00 2001 From: 0x676e67 Date: Tue, 6 Oct 2026 00:07:59 +0000 Subject: [PATCH 6/8] perf(upload): end async generator bodies without waiting for channel room --- src/client/body/stream.rs | 72 +++++++++++++++++++++------------------ tests/upload_test.py | 29 ++++++++++++++++ 2 files changed, 67 insertions(+), 34 deletions(-) diff --git a/src/client/body/stream.rs b/src/client/body/stream.rs index bcba68cd..bb051763 100644 --- a/src/client/body/stream.rs +++ b/src/client/body/stream.rs @@ -2,7 +2,10 @@ use std::{ pin::Pin, - sync::{Arc, Mutex, MutexGuard, PoisonError}, + sync::{ + Arc, Mutex, MutexGuard, PoisonError, + atomic::{AtomicBool, Ordering}, + }, task::{Context, Poll}, time::Duration, }; @@ -69,14 +72,19 @@ struct Pump { /// A request body from a Python async generator, forwarded with one chunk of buffering /// by a task on the loop that was running at extraction. Dropping it cancels that task. struct PyAsyncStream { - rx: mpsc::Receiver>, + rx: mpsc::Receiver, + /// Set by [`Sender::finish`], so the channel closing reads as the end of the body. + finished: Arc, task: Option<(Py, Py)>, } /// The channel end given to the forwarding coroutine; awaiting `send` applies upload /// backpressure. Closing it without `finish` fails the body. #[pyclass(frozen)] -struct Sender(Mutex>>>); +struct Sender { + tx: Mutex>>, + finished: Arc, +} // ===== impl PyBytesLike ===== @@ -253,7 +261,7 @@ async def forward(gen, sender): except BaseException as error: await sender.send(error, True) else: - await sender.finish() + sender.finish() ", c"wreq/_async_stream.py", c"wreq._async_stream", @@ -262,9 +270,12 @@ async def forward(gen, sender): .map(Bound::unbind) })?; let (tx, rx) = mpsc::channel(1); - let coroutine = forward - .bind(py) - .call1((generator, Sender(Mutex::new(Some(tx)))))?; + let finished = Arc::new(AtomicBool::new(false)); + let sender = Sender { + tx: Mutex::new(Some(tx)), + finished: finished.clone(), + }; + let coroutine = forward.bind(py).call1((generator, sender))?; // create_task captures the caller's contextvars on the running loop. let task = event_loop .call_method1(intern!(py, "create_task"), (&coroutine,)) @@ -273,6 +284,7 @@ async def forward(gen, sender): })?; Ok(Self { rx, + finished, task: Some((task.unbind(), event_loop.unbind())), }) } @@ -284,17 +296,16 @@ impl Stream for PyAsyncStream { fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { let this = self.get_mut(); match this.rx.poll_recv(cx) { - Poll::Ready(Some(Some(item))) => Poll::Ready(Some(item)), - Poll::Ready(Some(None)) => { - this.rx.close(); - this.task.take(); - Poll::Ready(None) + Poll::Ready(Some(item)) => Poll::Ready(Some(item)), + // Every sender is gone: the end after `finish`. Otherwise the forwarding task was + // cancelled or destroyed, so the body is incomplete. + Poll::Ready(None) + if this.task.take().is_some() && !this.finished.load(Ordering::Acquire) => + { + Poll::Ready(Some(Err(PyRuntimeError::new_err( + "async body generator stopped before it finished", + )))) } - // 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, } @@ -340,41 +351,34 @@ impl Sender { } else { Ok(item.extract()?) }; - let tx = self.sender(); + let tx = self.lock().clone(); // Channel readiness is runtime-independent, so this waits on the Python loop. coroutine::local(py, "Sender.send", async move { Ok(match tx { - Some(tx) => tx.send(Some(item)).await.is_ok(), + Some(tx) => tx.send(item).await.is_ok(), None => false, }) }) } - /// Mark the normal end of the body. - fn finish<'py>(&self, py: Python<'py>) -> PyResult> { + /// Mark the normal end of the body. The last chunk may still be queued: it is read + /// before the closed channel, so finishing never waits for room. + fn finish(&self) { + self.finished.store(true, Ordering::Release); // Python may retain the sender after completion, especially on PyPy. - let tx = self.sender(); - coroutine::local(py, "Sender.finish", async move { - Ok(match tx { - Some(tx) => tx.send(None).await.is_ok(), - None => false, - }) - }) + self.close(); } /// 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(); + let tx = self.lock().take(); drop(tx); } } impl Sender { #[inline] - fn sender(&self) -> Option>> { - self.0 - .lock() - .unwrap_or_else(PoisonError::into_inner) - .clone() + fn lock(&self) -> MutexGuard<'_, Option>> { + self.tx.lock().unwrap_or_else(PoisonError::into_inner) } } diff --git a/tests/upload_test.py b/tests/upload_test.py index 6fa3442f..8bed63c8 100644 --- a/tests/upload_test.py +++ b/tests/upload_test.py @@ -66,6 +66,35 @@ async def chunks(): assert closed.is_set() +@pytest.mark.asyncio +@pytest.mark.parametrize("size", [0, 1 << 24], ids=["empty", "queued_tail"]) +async def test_finish_delivers_queued_chunks(size): + # The peer reads nothing until forwarding has finished. A first chunk far above netty's + # ~408 KiB write buffer stalls the upload, so b"tail" is still queued when `finish` + # drops the sender; it and the end chunk must still be sent. + parts = (b"x" * size, b"tail") if size else () + yielded = asyncio.Event() + + async def chunks(): + for part in parts: + yield part + yielded.set() + + async with local_server() as (url, connections), wreq.Client(proxies=[]) as client: + task = asyncio.create_task(client.post(url, body=chunks())) + reader, writer = await asyncio.wait_for(connections.get(), 5) + await asyncio.wait_for(yielded.wait(), 5) + forwarding = [ + t for t in asyncio.all_tasks() if t.get_coro().__qualname__ == "forward" + ] + # `finish` does not wait for room, so forwarding ends before the peer reads. + await asyncio.wait_for(asyncio.gather(*forwarding), 5) + assert await asyncio.wait_for(read_chunked(reader), 10) == b"".join(parts) + writer.write(b"HTTP/1.1 200 OK\r\nContent-Length: 0\r\n\r\n") + await writer.drain() + await (await asyncio.wait_for(task, 5)).close() + + @pytest.mark.asyncio @pytest.mark.parametrize("failure", ["exception", "type", "cancelled"]) async def test_upload_errors(failure): From c816f085be9064f366716b268b4264ca9050982b Mon Sep 17 00:00:00 2001 From: 0x676e67 Date: Tue, 6 Oct 2026 00:07:59 +0000 Subject: [PATCH 7/8] style(tests): format response tests with ruff --- tests/response_test.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/tests/response_test.py b/tests/response_test.py index 5ac1d6b2..24aadfb5 100644 --- a/tests/response_test.py +++ b/tests/response_test.py @@ -202,7 +202,9 @@ async def test_empty_body_and_unit_results(): async def get(): task = asyncio.create_task(client.get(url)) _, writer = await asyncio.wait_for(connections.get(), 5) - writer.write(b"HTTP/1.1 200 OK\r\nConnection: close\r\nContent-Length: 0\r\n\r\n") + writer.write( + b"HTTP/1.1 200 OK\r\nConnection: close\r\nContent-Length: 0\r\n\r\n" + ) return await asyncio.wait_for(task, 5) response = await get() From 29f92a363a6720274bab9993d2d7989625532629 Mon Sep 17 00:00:00 2001 From: 0x676e67 Date: Tue, 6 Oct 2026 01:46:37 +0000 Subject: [PATCH 8/8] docs: correct read placement notes and mark json impl sections --- docs/source/guide/advanced.md | 9 +++++---- src/client/body/json.rs | 12 +++++++++--- src/client/resp/http.rs | 7 ++++--- 3 files changed, 18 insertions(+), 10 deletions(-) diff --git a/docs/source/guide/advanced.md b/docs/source/guide/advanced.md index 89767806..caf48503 100644 --- a/docs/source/guide/advanced.md +++ b/docs/source/guide/advanced.md @@ -119,10 +119,11 @@ client = Client(runtime=runtime) With `work_steal=False`, workers use independent single-thread Tokio runtimes. Each client is assigned one worker for its lifetime; requests, response reads, -streams and WebSocket operations use that worker. With multiple workers, newly -created clients select a worker randomly and keep that selection. This is not -CPU pinning. Sharing the same `Runtime` between clients is supported, and -`client.runtime` returns the shared runtime object. +streams and WebSocket operations use that worker. Async reads of small HTTP/1 +bodies that have already arrived finish on the event loop thread instead. With +multiple workers, newly created clients select a worker randomly and keep that +selection. This is not CPU pinning. Sharing the same `Runtime` between clients +is supported, and `client.runtime` returns the shared runtime object. `workers=None` uses the available CPU parallelism, or 1 if it cannot be determined. Custom runtimes start their threads during construction, before any client is diff --git a/src/client/body/json.rs b/src/client/body/json.rs index 79e96467..dcdb91ca 100644 --- a/src/client/body/json.rs +++ b/src/client/body/json.rs @@ -29,6 +29,12 @@ pub enum JsonString { RustString(String), } +/// Builds a [`Json`] in one pass; an untagged derive would buffer every value and then try +/// each variant in turn. +struct JsonVisitor; + +// ===== impl JsonString ===== + impl FromPyObject<'_, '_> for JsonString { type Error = PyErr; @@ -61,6 +67,8 @@ impl<'de> Deserialize<'de> for JsonString { } } +// ===== impl Json ===== + impl<'de> Deserialize<'de> for Json { #[inline] fn deserialize(deserializer: D) -> Result @@ -71,9 +79,7 @@ impl<'de> Deserialize<'de> for Json { } } -/// Builds a [`Json`] in one pass; an untagged derive would buffer every value and then try -/// each variant in turn. -struct JsonVisitor; +// ===== impl JsonVisitor ===== impl<'de> Visitor<'de> for JsonVisitor { type Value = Json; diff --git a/src/client/resp/http.rs b/src/client/resp/http.rs index 9821712a..42332175 100644 --- a/src/client/resp/http.rs +++ b/src/client/resp/http.rs @@ -64,7 +64,8 @@ enum Body { enum BodyRead { /// Read in full now. Ready(Bytes), - /// Not fully buffered yet; collected on the runtime by [`Response::collect_later`]. + /// Not read now: still arriving, over the read's limit or of unknown length; collected + /// on the runtime by [`Response::collect_later`]. Pending(Collect), } @@ -372,9 +373,9 @@ impl Response { }) } - /// Get the JSON content of the response. Buffered JSON up to 8 KiB is parsed on the event - /// loop, larger bodies on the runtime. + /// Get the JSON content of the response. pub fn json(slf: Bound<'_, Self>) -> PyResult> { + // Buffered HTTP/1 JSON up to 8 KiB is parsed on the event loop; see `loop_limit`. Self::read(slf, "Response.json", 8 * 1024, ResponseExt::json::) }