diff --git a/Cargo.lock b/Cargo.lock index 604e59f5..23ff2ad8 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=main#663b094fe80389c393880d8d5168c73131629336" 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/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/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..dcdb91ca 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), @@ -24,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; @@ -55,3 +66,74 @@ impl<'de> Deserialize<'de> for JsonString { String::deserialize(deserializer).map(JsonString::RustString) } } + +// ===== impl Json ===== + +impl<'de> Deserialize<'de> for Json { + #[inline] + fn deserialize(deserializer: D) -> Result + where + D: Deserializer<'de>, + { + deserializer.deserialize_any(JsonVisitor) + } +} + +// ===== impl 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/body/stream.rs b/src/client/body/stream.rs index 72c31cc4..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 ===== @@ -226,7 +234,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, @@ -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,19 +270,21 @@ 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)))))?; - // create_task captures the caller's contextvars on the running loop. - let task = match event_loop.call_method1("create_task", (&coroutine,)) { - Ok(task) => task, - Err(err) => { - let _ = coroutine.call_method0("close"); - return Err(err); - } + 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,)) + .inspect_err(|_| { + let _ = coroutine.call_method0(intern!(py, "close")); + })?; Ok(Self { rx, + finished, task: Some((task.unbind(), event_loop.unbind())), }) } @@ -286,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, } @@ -342,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/src/client/req.rs b/src/client/req.rs index bbcdba53..3b079bd2 100644 --- a/src/client/req.rs +++ b/src/client/req.rs @@ -186,36 +186,39 @@ 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); - + // 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, + [ + 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, + ], + [body, multipart] + ); Ok(request) } } 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 48803db6..42332175 100644 --- a/src/client/resp/http.rs +++ b/src/client/resp/http.rs @@ -6,22 +6,19 @@ 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}; -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}; +use super::{READ_ATTACHED, ext::ResponseExt, loop_limit, stream::Streamer}; 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 +29,9 @@ 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. +/// 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, @@ -62,21 +60,27 @@ enum Body { Released, } +/// A body taken for reading. +enum BodyRead { + /// Read in full now. + Ready(Bytes), + /// Not read now: still arriving, over the read's limit or of unknown length; 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. - const READ_ATTACHED: u64 = 64 * 1024; - /// Create a new [`Response`] instance. pub fn new(response: wreq::Response, runtime: Runtime) -> Self { let uri = response.uri().clone(); @@ -107,75 +111,99 @@ 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> { + /// 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) -> Result { let mut slot = self.slot(); - let stream = match mem::replace(&mut *slot, Body::Taken) { - Body::Unread(stream) => stream, + let body = match mem::replace(&mut *slot, Body::Taken) { + Body::Unread(body) => body, + Body::Cached(bytes) => { + *slot = Body::Cached(bytes.clone()); + return Ok(BodyRead::Ready(bytes)); + } 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(); + return Err(Error::Memory); } }; drop(slot); - let parts = self.parts.clone(); + let attached = body.size_hint().exact().is_some_and(|len| len <= limit); + let mut collect = body.collect(); + let ready = if attached { + // A read timeout starts a timer on its first poll, which needs the runtime. + 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)), + } + } + + /// 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, - /// 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)) - } - - /// 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) + 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) => ( + Either::Right(self.collect_later(collect).and_then(|(parts, bytes)| { + read(wreq::Response::from(HttpResponse::from_parts(parts, bytes))) + })), + false, + ), + }; + Ok((decode.map_err(Into::into), inline)) } - /// 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. 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, + limit: u64, read: F, ) -> PyResult> where @@ -187,24 +215,15 @@ 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 (read, inline) = this.read_body(loop_limit(this.parts.version, limit), read)?; + if inline { + read.await + } else { + coroutine::run(this.runtime.clone(), read).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()); @@ -233,6 +252,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. @@ -317,11 +344,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. @@ -330,19 +368,34 @@ impl Response { slf: Bound<'_, Self>, encoding: Option, ) -> PyResult> { - Self::read(slf, "Response.text", |resp| { + Self::read(slf, "Response.text", 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::) + // 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::) } /// 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(loop_limit(this.parts.version, READ_ATTACHED))? { + BodyRead::Ready(bytes) => Ok(PyBuffer::from(bytes)), + 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. @@ -383,6 +436,13 @@ impl Response { } } +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,22 +475,20 @@ 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. 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)?; + if inline { + nogil::run(py, runtime, read) } else { - py.detach(|| runtime.handle().block_on(fut)) + py.detach(|| runtime.handle().block_on(read)) } } } @@ -520,7 +578,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 ba8b90b9..8c8d621c 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::{ @@ -21,6 +22,7 @@ use tokio::sync::{ }; use tokio_util::task::AbortOnDropHandle; +use super::{READ_ATTACHED, loop_limit}; use crate::{ buffer::PyBuffer, client::nogil, @@ -156,17 +158,27 @@ 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 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 <= loop_limit(resp.version(), 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 +276,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 +309,11 @@ 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(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..f6131284 100644 --- a/src/coroutine.rs +++ b/src/coroutine.rs @@ -15,16 +15,20 @@ mod awaitable; mod scope; use std::{ + any::TypeId, future::{Future, poll_fn}, mem, 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 self::{asyncio::Port, scope::Scope}; use crate::runtime::Runtime; /// Run `fut` on the runtime once the coroutine named `qualname` is first awaited. @@ -51,11 +55,25 @@ 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)) } +/// 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 +84,19 @@ pub fn managed<'py, F, T>( ) -> PyResult> where F: Future> + Send + 'static, - T: for<'a> IntoPyObject<'a>, + T: EntersSelf + for<'a> IntoPyObject<'a> + 'static, { - Bound::new(py, coroutine(qualname, fut).managed()) + 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. @@ -97,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/asyncio.rs b/src/coroutine/asyncio.rs index c03de570..677dee07 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,14 @@ 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 + .import(py, "asyncio", "get_running_loop")? + .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..c0e3ee43 100644 --- a/src/coroutine/awaitable.rs +++ b/src/coroutine/awaitable.rs @@ -23,7 +23,7 @@ use super::{Port, scope::Scope}; /// 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 { @@ -78,12 +78,6 @@ impl Coroutine { } } - /// Let `async with` enter the coroutine; see [`coroutine::managed`](super::managed). - pub(super) fn managed(mut self) -> Self { - self.scope = Scope::Ready; - self - } - #[inline] pub(super) fn future(&mut self) -> &mut Option>>> { self.future @@ -170,8 +164,12 @@ 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)? { + // 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), + } } fn __traverse__(&self, visit: PyVisit<'_>) -> Result<(), PyTraverseError> { diff --git a/src/coroutine/scope.rs b/src/coroutine/scope.rs index 133940b6..419bb5ed 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 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, + Ready(EntersSelfFn), /// Awaited, entered or exited already. Spent, /// Awaited by `async with` until the result, an async context manager, is ready. - Entering, + Entering(EntersSelfFn), /// 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,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 => return self.open(py, value), + // 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, } @@ -175,7 +183,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 +208,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..2c2a0336 100644 --- a/src/macros.rs +++ b/src/macros.rs @@ -24,6 +24,37 @@ macro_rules! extract_option { }; } +/// Like [`extract_option!`] for several fields, but stop looking up keys once every key +/// 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, [$($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) = + $crate::macros::lookup(&$ob, pyo3::intern!($ob.py(), stringify!($field))) + { + $params.$field = value.extract()?; + remaining -= 1; + } + )* + $( + if let Some(value) = $last { + $params.$last = 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..24aadfb5 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,100 @@ 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. + with pytest.raises(StopIteration) as stop: + response.close().__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 c753e90a..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): @@ -88,6 +117,26 @@ 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" + + form = wreq.Multipart(wreq.Part("file", iter([b"x"]))) + async with local_server() as (url, _), wreq.Client(proxies=[]) as client: + # Every option, `zstd` last, is validated before the body or form is taken. + with pytest.raises(TypeError): + 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() + + @pytest.mark.asyncio @pytest.mark.parametrize("action", ["cancel", "early_response"]) async def test_abandoned_upload_to_stalled_peer(action):