diff --git a/docs/source/getting-started/quickstart.md b/docs/source/getting-started/quickstart.md index 79213db3..e7679a17 100644 --- a/docs/source/getting-started/quickstart.md +++ b/docs/source/getting-started/quickstart.md @@ -81,12 +81,46 @@ data = await response.json() print(data) ``` +### Binary data + +`response.bytes()` returns a read-only `memoryview`, not a `bytes` object. The view shares Rust-owned data without a copy into Python bytes. Reading a complete response can still allocate memory to combine body chunks. + +```python +import hashlib + +view = await response.bytes() +await response.close() +print(view.readonly) # True; closing the response does not invalidate the view +print(hashlib.sha256(view).hexdigest()) # Reads the buffer directly +``` + +The blocking API returns the same type, without `await`. Stream data frames, WebSocket binary fields, header names and values, and peer certificates also return read-only memoryviews. Each view retains its backing data even after the source object is closed or deleted. + +Use views directly with APIs that accept the buffer protocol, such as `file.write(view)` or `hashlib.sha256(view)`. For text, `str(view, "utf-8")` decodes into a string without an intermediate `bytes` object. + +#### Copying data + +Only convert when you need an independent `bytes` object or an API requires one: + +```python +data = bytes(view) # Copies the data; view.tobytes() also copies +view.release() +``` + +This changes the binary return type. `memoryview` has no `.decode()` method or byte-string concatenation. When you finish using a view, you can call `view.release()`; this does not release other views or slices sharing the data. Input types are unchanged. + +Built-in `bytes` and `str` inputs can share their storage. Their subclasses are copied from the actual contents to avoid hidden reference cycles; deleting a view releases its ownership normally, without requiring an explicit `release()`. + +When passing a view back to wreq's binary inputs (`body`, `Part`, `Message` constructors, or `CertStore`), convert it with `bytes(view)`. These inputs do not treat a memoryview as binary data. + ### Response headers Response headers are available as a [HeaderMap](../api/header/?h=HeaerMap#wreq.header.HeaderMap) object: ```python -print(response.headers.get("content-type")) +content_type = response.headers.get("content-type") +if content_type is not None: + print(str(content_type, "ascii")) # application/json ``` diff --git a/docs/source/guide/basic.md b/docs/source/guide/basic.md index f6f9df46..793de9b5 100644 --- a/docs/source/guide/basic.md +++ b/docs/source/guide/basic.md @@ -140,13 +140,17 @@ headers.append("Accept", "application/json") headers.append("Accept", "text/html") # Retrieve a single value -print(headers.get("Content-Type")) +content_type = headers.get("Content-Type") +if content_type is not None: + print(str(content_type, "ascii")) # application/json # Retrieve all values for a multi-value header -print(list(headers.get_all("Accept"))) +print([str(value, "ascii") for value in headers.get_all("Accept")]) # ['application/json', 'text/html'] ``` + +Header names and values are read-only `memoryview` objects. Decode text with `str(view, encoding)` without first copying it into `bytes`. Pass the `HeaderMap` to any request method via the `headers` argument: @@ -161,16 +165,21 @@ response = await wreq.get("https://httpbin.org/headers", headers=headers) For large responses, you can read the body incrementally instead of loading it all into memory at once. Use `resp.stream()` as an async iterator: ```python -from wreq import Client +import sys + +from wreq import Client, HeaderMap async def main(): client = Client() response = await client.get("https://httpbin.org/stream/10") async for chunk in response.stream(): - print(chunk.decode("utf-8")) + if isinstance(chunk, memoryview): + sys.stdout.buffer.write(chunk) + elif isinstance(chunk, HeaderMap): + print("Trailers:", chunk) ``` -Each `chunk` is a `bytes` object. Decode it to a string only if you know the response body is text. +Data chunks are read-only `memoryview` objects; trailer frames are `HeaderMap` objects. The example writes data directly through the buffer protocol. Each view stays valid after the stream is closed. See [Binary data](../getting-started/quickstart.md#binary-data) for copying and releasing views. --- diff --git a/docs/source/guide/blocking.md b/docs/source/guide/blocking.md index c8e4901f..a65cc3d9 100644 --- a/docs/source/guide/blocking.md +++ b/docs/source/guide/blocking.md @@ -166,6 +166,8 @@ if __name__ == "__main__": ### Streaming Response ```python +import sys + from wreq.blocking import Client @@ -175,9 +177,14 @@ def main(): with resp: with resp.stream() as streamer: for chunk in streamer: - print(chunk) + if isinstance(chunk, memoryview): + sys.stdout.buffer.write(chunk) + else: + print("Trailers:", chunk) if __name__ == "__main__": main() ``` + +Data chunks are read-only `memoryview` objects that stay valid after the stream is closed. `resp.bytes()` returns the same type. Pass views directly to APIs that accept the buffer protocol; use `bytes(view)` or `view.tobytes()` only when you need a copy. diff --git a/docs/source/guide/websocket.md b/docs/source/guide/websocket.md index db617203..35ac74f9 100644 --- a/docs/source/guide/websocket.md +++ b/docs/source/guide/websocket.md @@ -4,6 +4,8 @@ - HTTP/1.1 WebSocket - HTTP/2 WebSocket +`Message.data`, `.binary`, `.ping`, and `.pong` return read-only `memoryview` objects when present. The views remain valid after the message is deleted or the connection is closed. Use `bytes(view)` when a consumer requires a `bytes` object; `Message.text` still returns a string. + ### HTTP/1.1 WebSocket Connection ```python diff --git a/examples/blocking/stream.py b/examples/blocking/stream.py index 17dfccb4..a3bda692 100644 --- a/examples/blocking/stream.py +++ b/examples/blocking/stream.py @@ -1,3 +1,4 @@ +import sys import time import wreq @@ -8,7 +9,10 @@ def main(): with wreq.blocking.get("https://httpbin.io/stream/20") as resp: with resp.stream() as streamer: for chunk in streamer: - print(chunk) + if isinstance(chunk, memoryview): + sys.stdout.buffer.write(chunk) + else: + print("Trailers:", chunk) time.sleep(0.1) diff --git a/examples/header_map.py b/examples/header_map.py index 2bc43d1c..1bf038af 100644 --- a/examples/header_map.py +++ b/examples/header_map.py @@ -9,9 +9,11 @@ # Add Accept header (second value) headers.insert("Accept", "text/html") # Get all values for 'Accept' header - print("All Accept:", list(headers.get_all("Accept"))) + print("All Accept:", [str(value, "ascii") for value in headers.get_all("Accept")]) # Get the value for 'Content-Type' header - print("Content-Type:", headers.get("Content-Type")) + content_type = headers.get("Content-Type") + if content_type is not None: + print("Content-Type:", str(content_type, "ascii")) # Print total number of values in the map print("len (all values):", headers.len()) # Print number of unique keys in the map diff --git a/examples/request.py b/examples/request.py index dc39f134..857e6a45 100644 --- a/examples/request.py +++ b/examples/request.py @@ -12,13 +12,15 @@ async def main(): print("Cookies: ", resp.cookies) print("Content-Length: ", resp.content_length) print("Remote Address: ", resp.remote_addr) - print("Headers set-cookie: ", resp.headers["set-cookie"]) + set_cookie = resp.headers["set-cookie"] + if set_cookie is not None: + print("Headers set-cookie: ", str(set_cookie, "latin-1")) - for key in resp.headers: - print(key) + for key in resp.headers.keys(): + print(str(key, "ascii")) for key, value in resp.headers: - print(f"{key}: {value}") + print(f"{str(key, 'ascii')}: {str(value, 'latin-1')}") for cookie in resp.cookies: print(cookie) diff --git a/examples/stream.py b/examples/stream.py index 42ace271..8f7f10ce 100644 --- a/examples/stream.py +++ b/examples/stream.py @@ -1,4 +1,5 @@ import asyncio +import sys import wreq from wreq import Response @@ -8,7 +9,10 @@ async def main(): async with resp: async with resp.stream() as streamer: async for chunk in streamer: - print(chunk) + if isinstance(chunk, memoryview): + sys.stdout.buffer.write(chunk) + else: + print("Trailers:", chunk) await asyncio.sleep(0.1) diff --git a/python/wreq/blocking.py b/python/wreq/blocking.py index 4a30eb75..ddc7e794 100644 --- a/python/wreq/blocking.py +++ b/python/wreq/blocking.py @@ -89,7 +89,7 @@ def raise_for_status(self) -> None: def stream(self) -> Streamer: r""" - Get the response into a `Streamer` of `bytes` from the body. + Stream the body as read-only memoryviews, with HeaderMap frames for trailers. """ ... @@ -104,9 +104,11 @@ def json(self) -> Any: Get the JSON content of the response. """ - def bytes(self) -> bytes: + def bytes(self) -> memoryview: r""" - Get the bytes content of the response. + Read the body as a read-only memoryview without copying it into Python bytes. + The view remains valid after the response is closed or deleted. + Use bytes(view) or view.tobytes() for a copy; view.release() releases this view. """ ... diff --git a/python/wreq/header.py b/python/wreq/header.py index 22291ae1..cb1060bd 100644 --- a/python/wreq/header.py +++ b/python/wreq/header.py @@ -27,9 +27,12 @@ class HeaderMap: The implementation follows HTTP/1.1 specifications for header handling and provides both dictionary-like access and specialized methods for HTTP header manipulation. + + Header names and values are returned as read-only memoryviews. + Each view retains its backing data even if the map is changed or deleted. """ - def __getitem__(self, key: str) -> bytes | None: + def __getitem__(self, key: str) -> memoryview | None: """Get the first value for a header name (case-insensitive).""" ... @@ -49,7 +52,7 @@ def __len__(self) -> int: """Return the total number of header values (not unique names).""" ... - def __iter__(self) -> Iterator[Tuple[bytes, bytes]]: + def __iter__(self) -> Iterator[Tuple[memoryview, memoryview]]: """Iterate all header(name, value) pairs, including duplicates for multiple values.""" ... @@ -140,7 +143,7 @@ def remove(self, key: str) -> None: """ ... - def get(self, key: str, default: bytes | None = None) -> bytes | None: + def get(self, key: str, default: bytes | None = None) -> memoryview | None: r""" Get the first value for a header name with optional default. @@ -153,15 +156,15 @@ def get(self, key: str, default: bytes | None = None) -> bytes | None: default: Value to return if header doesn't exist Returns: - The first header value as bytes, or the default value + A read-only view of the first header value, or of the default value """ ... - def get_all(self, key: str) -> Iterator[bytes]: + def get_all(self, key: str) -> list[memoryview]: r""" Get all values for a header name. - Returns an iterator over all values associated with the header name. + Returns a list of read-only views of all values associated with the header name. This is useful for headers that can have multiple values, such as Set-Cookie, Accept-Encoding, or custom headers. @@ -169,25 +172,25 @@ def get_all(self, key: str) -> Iterator[bytes]: key: The header name (case-insensitive) Returns: - An iterator over all header values + A list of read-only header value views """ ... - def values(self) -> Iterator[bytes]: + def values(self) -> list[memoryview]: """ - Iterate over all header values. + Get all header values. Returns: - An iterator over all header values as bytes. + A list of read-only header value views. """ ... - def keys(self) -> Iterator[bytes]: + def keys(self) -> list[memoryview]: """ - Iterate over unique header names. + Get all unique header names. Returns: - An iterator over unique header names as bytes. + A list of read-only header name views. """ ... @@ -247,6 +250,8 @@ class OrigHeaderMap: The map stores a mapping between the case-insensitive (standard) header name and the original case-sensitive header name as it appeared in the HTTP message. + Iteration returns pairs of read-only memoryviews that retain their backing data. + Example: If an HTTP message included the following headers: @@ -277,7 +282,7 @@ def __init__( """ ... - def __iter__(self) -> Iterator[Tuple[bytes, bytes]]: + def __iter__(self) -> Iterator[Tuple[memoryview, memoryview]]: """ Returns an iterator over the (standard_name, original_name) pairs. diff --git a/python/wreq/tls.py b/python/wreq/tls.py index fa7b72c1..cab2857d 100644 --- a/python/wreq/tls.py +++ b/python/wreq/tls.py @@ -446,8 +446,8 @@ class TlsInfo: Information about the established TLS connection. """ - def peer_certificate(self) -> bytes | None: + def peer_certificate(self) -> memoryview | None: """ - Get the DER encoded leaf certificate of the peer. + Get a read-only memoryview of the peer's DER-encoded leaf certificate. """ ... diff --git a/python/wreq/wreq.py b/python/wreq/wreq.py index 744ee9b2..31155aa2 100644 --- a/python/wreq/wreq.py +++ b/python/wreq/wreq.py @@ -175,11 +175,13 @@ def __init__( class Message: r""" A WebSocket message. + + Binary fields return read-only memoryviews that retain their backing data. """ - data: bytes | None + data: memoryview | None r""" - Returns the data of the message as bytes. + Returns the data of the message as a read-only memoryview. """ text: str | None @@ -187,17 +189,17 @@ class Message: Returns the text content of the message if it is a text message. """ - binary: bytes | None + binary: memoryview | None r""" Returns the binary data of the message if it is a binary message. """ - ping: bytes | None + ping: memoryview | None r""" Returns the ping data of the message if it is a ping message. """ - pong: bytes | None + pong: memoryview | None r""" Returns the pong data of the message if it is a pong message. """ @@ -274,18 +276,20 @@ def __str__(self) -> str: ... class Streamer: r""" A stream response. - An asynchronous iterator yielding data chunks (bytes) or HTTP trailers (HeaderMap) from the response stream. + An asynchronous iterator yielding read-only memoryviews or HTTP trailers (HeaderMap) from the response stream. Used to stream response content and receive HTTP trailers if present. Implemented in the `stream` method of the `Response` class. Can be used in an asynchronous for loop in Python. - When streaming a response, each iteration yields either a bytes object (for body data) or a HeaderMap (for HTTP trailers, if the server sends them). + When streaming a response, each iteration yields either a memoryview (for body data) or a HeaderMap (for HTTP trailers, if the server sends them). This allows you to access HTTP/1.1 or HTTP/2 trailers in addition to the main body. + Data views retain their backing data after the stream is closed. # Examples ```python import asyncio + import sys import wreq from wreq import Method, Emulation, HeaderMap @@ -293,8 +297,8 @@ async def main(): resp = await wreq.get("https://example.com/stream-with-trailers") async with resp.stream() as streamer: async for chunk in streamer: - if isinstance(chunk, bytes): - print("Chunk: ", chunk) + if isinstance(chunk, memoryview): + sys.stdout.buffer.write(chunk) elif isinstance(chunk, HeaderMap): print("Trailers: ", chunk) await asyncio.sleep(0.1) @@ -305,11 +309,11 @@ async def main(): """ def __iter__(self) -> "Streamer": ... - def __next__(self) -> bytes | HeaderMap: ... + def __next__(self) -> memoryview | HeaderMap: ... def __enter__(self) -> Any: ... def __exit__(self, _exc_type: Any, _exc_value: Any, _traceback: Any) -> None: ... - async def __aiter__(self) -> "Streamer": ... - async def __anext__(self) -> bytes | HeaderMap: ... + def __aiter__(self) -> "Streamer": ... + async def __anext__(self) -> memoryview | HeaderMap: ... async def __aenter__(self) -> Any: ... async def __aexit__( self, _exc_type: Any, _exc_value: Any, _traceback: Any @@ -401,7 +405,7 @@ def raise_for_status(self) -> None: def stream(self) -> Streamer: r""" - Get the response into a `Streamer` of `bytes` from the body. + Stream the body as read-only memoryviews, with HeaderMap frames for trailers. """ ... @@ -416,9 +420,11 @@ async def json(self) -> Any: Get the JSON content of the response. """ - async def bytes(self) -> bytes: + async def bytes(self) -> memoryview: r""" - Get the bytes content of the response. + Read the body as a read-only memoryview without copying it into Python bytes. + The view remains valid after the response is closed or deleted. + Use bytes(view) or view.tobytes() for a copy; view.release() releases this view. """ ... diff --git a/src/buffer.rs b/src/buffer.rs index b6c66221..b44d6f25 100644 --- a/src/buffer.rs +++ b/src/buffer.rs @@ -18,29 +18,26 @@ use std::os::raw::c_int; use bytes::Bytes; -use pyo3::{ffi, prelude::*}; +use pyo3::{exceptions::PyOverflowError, ffi, prelude::*, types::PyMemoryView}; use wreq::header::{HeaderCaseName, HeaderName, HeaderValue}; -/// [`PyBuffer`] enables zero-copy conversion of Rust [`Bytes`] to Python bytes. +/// Exposes owned Rust bytes as a read-only Python memoryview without copying. pub struct PyBuffer(BufferView); #[pyclass(frozen, skip_from_py_object)] struct BufferView(Bytes); -// ===== PyBuffer ===== +// ===== impl PyBuffer ===== impl<'a> IntoPyObject<'a> for PyBuffer { - type Target = PyAny; + type Target = PyMemoryView; type Output = Bound<'a, Self::Target>; type Error = PyErr; #[inline(always)] fn into_pyobject(self, py: Python<'a>) -> Result { let buffer = self.0.into_pyobject(py)?; - #[allow(unsafe_code)] - unsafe { - Bound::from_owned_ptr_or_err(py, ffi::PyBytes_FromObject(buffer.as_ptr())) - } + PyMemoryView::from(buffer.as_any()) } } @@ -85,10 +82,12 @@ impl From for PyBuffer { } } -// ===== BufferView ===== +// ===== impl BufferView ===== #[pymethods] impl BufferView { + /// # Safety + /// `view` must be a writable `Py_buffer` supplied by the Python buffer protocol. #[allow(unsafe_code)] unsafe fn __getbuffer__( slf: PyRef, @@ -96,13 +95,16 @@ impl BufferView { flags: c_int, ) -> PyResult<()> { let bytes = &slf.0; + let len = ffi::Py_ssize_t::try_from(bytes.len()) + .map_err(|_| PyOverflowError::new_err("buffer length exceeds Python's maximum size"))?; + // SAFETY: PyO3 supplies a valid output buffer. FillInfo retains the exporter, + // which owns immutable Bytes for the lifetime of this read-only buffer. let ret = unsafe { - // Fill the Py_buffer struct with information about the buffer ffi::PyBuffer_FillInfo( view, slf.as_ptr() as *mut _, bytes.as_ptr() as *mut _, - bytes.len() as _, + len, 1, flags, ) @@ -113,3 +115,43 @@ impl BufferView { Ok(()) } } + +#[cfg(test)] +mod tests { + use pyo3::{ + buffer::PyBuffer as PythonBuffer, + types::{PyBytes, PyString}, + }; + + use super::*; + use crate::extractor::{BytesInput, StrInput}; + + #[test] + fn memoryview_shares_owned_bytes() { + Python::initialize(); + Python::attach(|py| { + let bytes = Bytes::from(vec![0, 1, 255]); + let ptr = bytes.as_ptr(); + let view = PyBuffer::from(bytes).into_pyobject(py).unwrap(); + let buffer = PythonBuffer::::get(view.as_any()).unwrap(); + + assert_eq!(buffer.buf_ptr().cast_const().cast::(), ptr); + assert!(buffer.readonly()); + assert_eq!(buffer.to_vec(py).unwrap(), [0, 1, 255]); + + let binary = PyBytes::new(py, b"builtin bytes"); + let input = binary.extract::().unwrap(); + assert_eq!(input.0.as_ptr(), binary.as_bytes().as_ptr()); + let view = PyBuffer::from(input.0).into_pyobject(py).unwrap(); + let buffer = PythonBuffer::::get(view.as_any()).unwrap(); + assert_eq!( + buffer.buf_ptr().cast_const().cast::(), + binary.as_bytes().as_ptr() + ); + + let text = PyString::new(py, "builtin text"); + let input = text.extract::().unwrap(); + assert_eq!(input.0.as_ptr(), text.to_str().unwrap().as_ptr()); + }); + } +} diff --git a/src/client/body.rs b/src/client/body.rs index d7fe67ac..6b4ecf6c 100644 --- a/src/client/body.rs +++ b/src/client/body.rs @@ -5,24 +5,20 @@ mod json; pub mod multipart; mod stream; -use bytes::Bytes; -use pyo3::{ - FromPyObject, PyResult, - prelude::*, - pybacked::{PyBackedBytes, PyBackedStr}, -}; +use pyo3::{FromPyObject, PyResult, prelude::*}; pub use self::{ form::Form, json::Json, stream::{PyStream, Streamer}, }; +use crate::extractor::{BytesInput, StrInput}; /// Represents the body of an HTTP request. #[derive(FromPyObject)] pub enum Body { - Text(PyBackedStr), - Bytes(PyBackedBytes), + Text(StrInput), + Bytes(BytesInput), Form(Form), Json(Json), Stream(PyStream), @@ -41,8 +37,8 @@ impl TryFrom for wreq::Body { .map_err(crate::Error::Json) .map(wreq::Body::from) .map_err(Into::into), - Body::Text(s) => Ok(wreq::Body::from(Bytes::from_owner(s))), - Body::Bytes(bytes) => Ok(wreq::Body::from(Bytes::from_owner(bytes))), + Body::Text(s) => Ok(wreq::Body::from(s.0)), + Body::Bytes(bytes) => Ok(wreq::Body::from(bytes.0)), Body::Stream(stream) => Ok(wreq::Body::wrap_stream(stream)), } } diff --git a/src/client/body/multipart.rs b/src/client/body/multipart.rs index 51f561e5..34f05280 100644 --- a/src/client/body/multipart.rs +++ b/src/client/body/multipart.rs @@ -1,14 +1,14 @@ use std::path::PathBuf; -use bytes::Bytes; -use pyo3::{ - prelude::*, - pybacked::{PyBackedBytes, PyBackedStr}, - types::PyTuple, -}; +use pyo3::{prelude::*, types::PyTuple}; use wreq::{Body, multipart}; -use crate::{client::body::PyStream, error::Error, header::HeaderMap}; +use crate::{ + client::body::PyStream, + error::Error, + extractor::{BytesInput, StrInput}, + header::HeaderMap, +}; /// A multipart form for a request. #[pyclass(subclass)] @@ -20,8 +20,8 @@ pub struct Multipart { /// The data for a part value of a multipart form. #[derive(FromPyObject)] pub enum Value { - Text(PyBackedStr), - Bytes(PyBackedBytes), + Text(StrInput), + Bytes(BytesInput), File(PathBuf), Stream(PyStream), } @@ -44,12 +44,12 @@ impl Multipart { /// Creates a new multipart. #[new] #[pyo3(signature = (*parts))] - pub fn new(py: Python, parts: &Bound) -> PyResult { + pub fn new(parts: &Bound) -> PyResult { let mut new_parts = Vec::with_capacity(parts.len()); for part in parts { let part = part.cast::()?; let mut part = part.borrow_mut(); - new_parts.push(part.try_clone(py)?); + new_parts.push(part.try_clone()?); } Ok(Self { @@ -88,14 +88,14 @@ impl FromPyObject<'_, '_> for Multipart { // ===== impl Value ===== impl Value { - fn try_clone(&self, py: Python) -> Option { + fn try_clone(&self) -> Option { match self { Value::Text(text) => { - let text = text.clone_ref(py); + let text = text.clone(); Some(Value::Text(text)) } Value::Bytes(bytes) => { - let bytes = bytes.clone_ref(py); + let bytes = bytes.clone(); Some(Value::Bytes(bytes)) } Value::File(path) => { @@ -125,14 +125,14 @@ impl Part { let value = self .value .as_ref() - .and_then(|value| value.try_clone(py)) + .and_then(Value::try_clone) .or_else(|| self.value.take()) .ok_or_else(|| Error::Memory)?; py.detach(move || { let mut inner = match value { - Value::Text(text) => multipart::Part::stream(Bytes::from_owner(text)), - Value::Bytes(bytes) => multipart::Part::stream(Bytes::from_owner(bytes)), + Value::Text(text) => multipart::Part::stream(text.0), + Value::Bytes(bytes) => multipart::Part::stream(bytes.0), Value::File(path) => pyo3_async_runtimes::tokio::get_runtime() .block_on(multipart::Part::file(path)) .map_err(Error::from)?, @@ -161,11 +161,11 @@ impl Part { }) } - fn try_clone(&mut self, py: Python) -> PyResult { + fn try_clone(&mut self) -> PyResult { if let Some(part) = self .value .as_ref() - .and_then(|value| value.try_clone(py)) + .and_then(Value::try_clone) .map(|value| self.with_value(value)) { return Ok(part); diff --git a/src/client/body/stream.rs b/src/client/body/stream.rs index 179aea10..d7b8d9e6 100644 --- a/src/client/body/stream.rs +++ b/src/client/body/stream.rs @@ -7,18 +7,14 @@ use std::{ use bytes::Bytes; use futures_util::{FutureExt, Stream, StreamExt, stream::BoxStream}; use http_body_util::BodyExt; -use pyo3::{ - coroutine::CancelHandle, - intern, - prelude::*, - pybacked::{PyBackedBytes, PyBackedStr}, -}; +use pyo3::{coroutine::CancelHandle, intern, prelude::*}; use tokio::{sync::Mutex, task::JoinHandle}; use crate::{ buffer::PyBuffer, client::nogil::NoGIL, error::{self, Error}, + extractor::{BytesInput, StrInput}, header::HeaderMap, }; @@ -33,11 +29,11 @@ enum PyStreamSource { /// A bytes-like object that can be extracted from Python. #[derive(FromPyObject)] pub enum PyBytesLike { - Bytes(PyBackedBytes), - String(PyBackedStr), + Bytes(BytesInput), + String(StrInput), } -/// A bytes-like object that can be into Python. +/// A response frame exposed as a read-only memoryview or a header map. #[derive(IntoPyObject)] pub enum Frame { Bytes(PyBuffer), @@ -50,7 +46,7 @@ pub struct PyStream { pending: Pending, } -/// A bytes stream response. +/// A response stream yielding read-only memoryviews and any trailing headers. #[derive(Clone)] #[pyclass(subclass, frozen, skip_from_py_object)] pub struct Streamer(Arc>>); @@ -183,8 +179,8 @@ impl From for Bytes { #[inline] fn from(value: PyBytesLike) -> Self { match value { - PyBytesLike::Bytes(b) => Bytes::from_owner(b), - PyBytesLike::String(s) => Bytes::from_owner(s), + PyBytesLike::Bytes(b) => b.0, + PyBytesLike::String(s) => s.0, } } } diff --git a/src/client/resp/http.rs b/src/client/resp/http.rs index 6e0cb882..1831b822 100644 --- a/src/client/resp/http.rs +++ b/src/client/resp/http.rs @@ -209,7 +209,7 @@ impl Response { .map_err(Into::into) } - /// Get the response into a `Stream` of `Bytes` from the body. + /// Stream read-only memoryviews and any trailing headers from the body. pub fn stream(&self) -> PyResult { self.stream_response() .map(Streamer::new) @@ -239,7 +239,7 @@ impl Response { NoGIL::new(fut, cancel).await } - /// Get the bytes content of the response. + /// Read the body as a read-only memoryview, retaining its data after the response closes. pub async fn bytes(&self, #[pyo3(cancel_handle)] cancel: CancelHandle) -> PyResult { let fut = self .cache_response() @@ -366,7 +366,7 @@ impl BlockingResponse { self.0.raise_for_status() } - /// Get the response into a `Stream` of `Bytes` from the body. + /// Stream read-only memoryviews and any trailing headers from the body. #[inline] pub fn stream(&self) -> PyResult { self.0.stream() @@ -397,7 +397,7 @@ impl BlockingResponse { }) } - /// Get the bytes content of the response. + /// Read the body as a read-only memoryview, retaining its data after the response closes. pub fn bytes(&self, py: Python) -> PyResult { py.detach(|| { let fut = self diff --git a/src/client/resp/ws.rs b/src/client/resp/ws.rs index cf4a4394..03e5a8c3 100644 --- a/src/client/resp/ws.rs +++ b/src/client/resp/ws.rs @@ -4,7 +4,7 @@ pub mod msg; use std::{fmt::Display, time::Duration}; use msg::Message; -use pyo3::{coroutine::CancelHandle, prelude::*, pybacked::PyBackedStr}; +use pyo3::{coroutine::CancelHandle, prelude::*}; use tokio::sync::mpsc; use wreq::{ header::HeaderValue, @@ -15,6 +15,7 @@ use crate::{ client::{SocketAddr, nogil::NoGIL}, cookie::Cookie, error::Error, + extractor::StrInput, header::HeaderMap, http::{StatusCode, Version}, }; @@ -136,7 +137,7 @@ impl WebSocket { &self, #[pyo3(cancel_handle)] cancel: CancelHandle, code: Option, - reason: Option, + reason: Option, ) -> PyResult<()> { let tx = self.cmd.clone(); NoGIL::new(cmd::close(tx, code, reason), cancel).await @@ -243,12 +244,7 @@ impl BlockingWebSocket { /// Close the WebSocket connection. #[pyo3(signature = (code=None, reason=None))] - pub fn close( - &self, - py: Python, - code: Option, - reason: Option, - ) -> PyResult<()> { + pub fn close(&self, py: Python, code: Option, reason: Option) -> PyResult<()> { py.detach(|| { pyo3_async_runtimes::tokio::get_runtime().block_on(cmd::close( self.0.cmd.clone(), diff --git a/src/client/resp/ws/cmd.rs b/src/client/resp/ws/cmd.rs index ea00330e..6f381cf7 100644 --- a/src/client/resp/ws/cmd.rs +++ b/src/client/resp/ws/cmd.rs @@ -7,9 +7,8 @@ use std::time::Duration; -use bytes::Bytes; use futures_util::{SinkExt, StreamExt, TryStreamExt}; -use pyo3::{prelude::*, pybacked::PyBackedStr}; +use pyo3::prelude::*; use tokio::sync::{ mpsc::{UnboundedReceiver, UnboundedSender}, oneshot::{self, Sender}, @@ -19,6 +18,7 @@ use super::{ Error, Message, Utf8Bytes, ws::{self, WebSocket}, }; +use crate::extractor::StrInput; /// Commands for WebSocket operations. pub enum Command { @@ -40,7 +40,7 @@ pub enum Command { /// Close the WebSocket connection. /// /// Contains an optional close code, optional reason, and a oneshot sender for the result. - Close(Option, Option, Sender>), + Close(Option, Option, Sender>), } /// The main background task that processes incoming [`Command`]s and interacts with the WebSocket. @@ -97,7 +97,7 @@ pub async fn task(ws: WebSocket, mut cmd: UnboundedReceiver) { .map(ws::message::CloseCode::from) .unwrap_or(ws::message::CloseCode::NORMAL); let reason = reason - .map(Bytes::from_owner) + .map(|reason| reason.0) .map(Utf8Bytes::try_from) .transpose(); @@ -157,7 +157,7 @@ pub async fn send_all(cmd: UnboundedSender, messages: Vec) -> pub async fn close( cmd: UnboundedSender, code: Option, - reason: Option, + reason: Option, ) -> PyResult<()> { send_command(cmd, |tx| Command::Close(code, reason, tx)).await? } diff --git a/src/client/resp/ws/msg.rs b/src/client/resp/ws/msg.rs index ceb5bfd2..100faa69 100644 --- a/src/client/resp/ws/msg.rs +++ b/src/client/resp/ws/msg.rs @@ -9,26 +9,27 @@ use std::fmt::Debug; -use bytes::Bytes; -use pyo3::{ - prelude::*, - pybacked::{PyBackedBytes, PyBackedStr}, -}; +use pyo3::prelude::*; use wreq::ws::message::{self, CloseCode, CloseFrame, Utf8Bytes}; -use crate::{buffer::PyBuffer, client::body::Json, error::Error}; +use crate::{ + buffer::PyBuffer, + client::body::Json, + error::Error, + extractor::{BytesInput, StrInput}, +}; /// An enum representing either a bytes message or a JSON message. #[derive(FromPyObject)] pub enum BytesLike { - Bytes(PyBackedBytes), + Bytes(BytesInput), Json(Json), } /// An enum representing either a text message or a JSON message. #[derive(FromPyObject)] pub enum TextLike { - Text(PyBackedStr), + Text(StrInput), Json(Json), } @@ -39,7 +40,7 @@ pub struct Message(pub message::Message); #[pymethods] impl Message { - /// Returns the data of the message as bytes. + /// Returns the message data as a read-only memoryview. #[getter] pub fn data(&self) -> Option { let bytes = match &self.0 { @@ -62,7 +63,7 @@ impl Message { } } - /// Returns the binary data of the message if it is a binary message. + /// Returns a read-only memoryview if this is a binary message. #[getter] pub fn binary(&self) -> Option { if let message::Message::Binary(data) = &self.0 { @@ -72,7 +73,7 @@ impl Message { } } - /// Returns the ping data of the message if it is a ping message. + /// Returns a read-only memoryview if this is a ping message. #[getter] pub fn ping(&self) -> Option { if let message::Message::Ping(data) = &self.0 { @@ -82,7 +83,7 @@ impl Message { } } - /// Returns the pong data of the message if it is a pong message. + /// Returns a read-only memoryview if this is a pong message. #[getter] pub fn pong(&self) -> Option { if let message::Message::Pong(data) = &self.0 { @@ -118,9 +119,7 @@ impl Message { py.detach(|| match like { TextLike::Text(text) => { // If the string is not valid UTF-8, this will panic. - let msg = message::Message::text( - Utf8Bytes::try_from(Bytes::from_owner(text)).expect("valid UTF-8"), - ); + let msg = message::Message::text(Utf8Bytes::try_from(text.0).expect("valid UTF-8")); Ok(Self(msg)) } TextLike::Json(json) => message::Message::text_from_json(&json) @@ -135,7 +134,7 @@ impl Message { #[pyo3(signature = (like))] pub fn from_binary(py: Python, like: BytesLike) -> PyResult { py.detach(|| match like { - BytesLike::Bytes(bytes) => Ok(Self(message::Message::binary(Bytes::from_owner(bytes)))), + BytesLike::Bytes(bytes) => Ok(Self(message::Message::binary(bytes.0))), BytesLike::Json(json) => message::Message::binary_from_json(&json) .map(Message) .map_err(Error::Library) @@ -146,23 +145,23 @@ impl Message { /// Creates a new ping message. #[staticmethod] #[pyo3(signature = (data))] - pub fn from_ping(data: PyBackedBytes) -> Self { - Self(message::Message::ping(Bytes::from_owner(data))) + pub fn from_ping(data: BytesInput) -> Self { + Self(message::Message::ping(data.0)) } /// Creates a new pong message. #[staticmethod] #[pyo3(signature = (data))] - pub fn from_pong(data: PyBackedBytes) -> Self { - Self(message::Message::pong(Bytes::from_owner(data))) + pub fn from_pong(data: BytesInput) -> Self { + Self(message::Message::pong(data.0)) } /// Creates a new close message. #[staticmethod] #[pyo3(signature = (code, reason=None))] - pub fn from_close(code: u16, reason: Option) -> Self { + pub fn from_close(code: u16, reason: Option) -> Self { let reason = reason - .map(Bytes::from_owner) + .map(|reason| reason.0) .and_then(|b| Utf8Bytes::try_from(b).ok()) .unwrap_or_else(|| Utf8Bytes::from_static("Goodbye")); let msg = message::Message::close(CloseFrame { diff --git a/src/cookie.rs b/src/cookie.rs index 328c673a..7112bba3 100644 --- a/src/cookie.rs +++ b/src/cookie.rs @@ -5,7 +5,7 @@ use cookie::{Cookie as RawCookie, Expiration, ParseError, time::Duration}; use pyo3::{prelude::*, pybacked::PyBackedStr, types::PyDict}; use wreq::header::{self, HeaderMap, HeaderValue}; -use crate::error::Error; +use crate::{error::Error, extractor::StrInput}; define_enum!( /// The Cookie SameSite attribute. @@ -189,8 +189,8 @@ impl FromPyObject<'_, '_> for Cookies { type Error = PyErr; fn extract(ob: Borrowed) -> PyResult { - if let Ok(cookie) = ob.extract::() { - return HeaderValue::from_maybe_shared(Bytes::from_owner(cookie)) + if let Ok(cookie) = ob.extract::() { + return HeaderValue::from_maybe_shared(cookie.0) .map(|cookie| Cookies(vec![cookie])) .map_err(Error::from) .map_err(Into::into); diff --git a/src/extractor.rs b/src/extractor.rs index 229e920c..6cf11077 100644 --- a/src/extractor.rs +++ b/src/extractor.rs @@ -1,10 +1,60 @@ use std::net::{IpAddr, Ipv4Addr, Ipv6Addr}; -use pyo3::{FromPyObject, prelude::*}; +use bytes::Bytes; +use pyo3::{ + FromPyObject, + prelude::*, + pybacked::{PyBackedBytes, PyBackedStr}, + types::{PyBytes, PyString}, +}; /// A generic extractor for various types. pub struct Extractor(pub T); +/// Byte input with no hidden references to Python subclasses. +#[derive(Clone)] +pub struct BytesInput(pub Bytes); + +/// UTF-8 input with no hidden references to Python subclasses. +#[derive(Clone)] +pub struct StrInput(pub Bytes); + +// ===== impl BytesInput ===== + +impl FromPyObject<'_, '_> for BytesInput { + type Error = PyErr; + + fn extract(ob: Borrowed) -> PyResult { + let value = ob.extract::()?; + // Subclasses can reference an exported view; Bytes owners are invisible to GC. + let bytes = if ob.is_instance_of::() && !ob.is_exact_instance_of::() { + Bytes::copy_from_slice(value.as_ref()) + } else { + Bytes::from_owner(value) + }; + Ok(Self(bytes)) + } +} + +// ===== impl StrInput ===== + +impl FromPyObject<'_, '_> for StrInput { + type Error = PyErr; + + fn extract(ob: Borrowed) -> PyResult { + let value = ob.extract::()?; + let bytes = if ob.is_exact_instance_of::() { + Bytes::from_owner(value) + } else { + // Copy the actual UTF-8 payload, not a subclass's __str__ result. + Bytes::copy_from_slice(value.as_bytes()) + }; + Ok(Self(bytes)) + } +} + +// ===== impl Extractor ===== + impl FromPyObject<'_, '_> for Extractor<(Option, Option)> { type Error = PyErr; diff --git a/src/header.rs b/src/header.rs index bf8794f1..057aed64 100644 --- a/src/header.rs +++ b/src/header.rs @@ -1,19 +1,22 @@ -use bytes::Bytes; use pyo3::{ prelude::*, - pybacked::{PyBackedBytes, PyBackedStr}, + pybacked::PyBackedStr, types::{PyDict, PyIterator, PyList}, }; use wreq::header::{self, HeaderName, HeaderValue}; -use crate::{buffer::PyBuffer, error::Error}; +use crate::{ + buffer::PyBuffer, + error::Error, + extractor::{BytesInput, StrInput}, +}; -/// A HTTP header map. +/// An HTTP header map whose names and values are exposed as read-only memoryviews. #[derive(Clone)] #[pyclass(subclass, str, skip_from_py_object)] pub struct HeaderMap(pub header::HeaderMap); -/// A HTTP original header map. +/// An HTTP original header map whose iterator exposes read-only memoryviews. #[derive(Clone)] #[pyclass(subclass, str, skip_from_py_object)] pub struct OrigHeaderMap(pub header::OrigHeaderMap); @@ -43,8 +46,8 @@ impl HeaderMap { }; let value = match value - .extract::() - .map(Bytes::from_owner) + .extract::() + .map(|value| value.0) .map(HeaderValue::from_maybe_shared) { Ok(Ok(v)) => v, @@ -58,7 +61,7 @@ impl HeaderMap { HeaderMap(headers) } - /// Returns a reference to the value associated with the key. + /// Returns a read-only memoryview of the value associated with the key. /// /// If there are multiple values associated with the key, then the first one /// is returned. Use `get_all` to get all values associated with a given @@ -68,12 +71,12 @@ impl HeaderMap { &self, py: Python<'py>, key: PyBackedStr, - default: Option, + default: Option, ) -> Option { py.detach(|| { self.0.get::<&str>(key.as_ref()).cloned().or_else(|| { match default - .map(Bytes::from_owner) + .map(|value| value.0) .map(HeaderValue::from_maybe_shared) { Some(Ok(v)) => Some(v), @@ -84,7 +87,7 @@ impl HeaderMap { .map(PyBuffer::from) } - /// Returns a view of all values associated with a key. + /// Returns a list of read-only memoryviews for the values associated with a key. #[pyo3(signature = (key))] fn get_all<'py>(&self, py: Python<'py>, key: PyBackedStr) -> Vec { py.detach(|| { @@ -99,11 +102,11 @@ impl HeaderMap { /// Insert a key-value pair into the header map. #[pyo3(signature = (key, value))] - fn insert(&mut self, py: Python, key: PyBackedStr, value: PyBackedStr) { + fn insert(&mut self, py: Python, key: PyBackedStr, value: StrInput) { py.detach(|| { if let (Ok(name), Ok(value)) = ( HeaderName::from_bytes(key.as_bytes()), - HeaderValue::from_maybe_shared(Bytes::from_owner(value)), + HeaderValue::from_maybe_shared(value.0), ) { self.0.insert(name, value); } @@ -112,11 +115,11 @@ impl HeaderMap { /// Append a key-value pair to the header map. #[pyo3(signature = (key, value))] - fn append(&mut self, py: Python, key: PyBackedStr, value: PyBackedStr) { + fn append(&mut self, py: Python, key: PyBackedStr, value: StrInput) { py.detach(|| { if let (Ok(name), Ok(value)) = ( HeaderName::from_bytes(key.as_bytes()), - HeaderValue::from_maybe_shared(Bytes::from_owner(value)), + HeaderValue::from_maybe_shared(value.0), ) { self.0.append(name, value); } @@ -137,7 +140,7 @@ impl HeaderMap { py.detach(|| self.0.contains_key::<&str>(key.as_ref())) } - /// An iterator visiting all keys. + /// Returns a list of read-only memoryviews for all keys. #[inline] fn keys<'py>(&self, py: Python<'py>) -> Vec { py.detach(|| { @@ -149,7 +152,7 @@ impl HeaderMap { }) } - /// An iterator visiting all values. + /// Returns a list of read-only memoryviews for all values. #[inline] fn values<'py>(&self, py: Python<'py>) -> Vec { py.detach(|| { @@ -201,7 +204,7 @@ impl HeaderMap { } #[inline] - fn __setitem__(&mut self, py: Python, key: PyBackedStr, value: PyBackedStr) { + fn __setitem__(&mut self, py: Python, key: PyBackedStr, value: StrInput) { self.insert(py, key, value); } @@ -253,9 +256,8 @@ impl FromPyObject<'_, '_> for HeaderMap { }; let value = { - let value = value.extract::()?; - HeaderValue::from_maybe_shared(Bytes::from_owner(value)) - .map_err(Error::from)? + let value = value.extract::()?; + HeaderValue::from_maybe_shared(value.0).map_err(Error::from)? }; headers.insert(name, value); @@ -282,8 +284,8 @@ impl OrigHeaderMap { // and we want to prevent Python's garbage collector from managing it. if let Some(init) = init { for name in init.iter() { - let name = match name.extract::() { - Ok(n) => Bytes::from_owner(n), + let name = match name.extract::() { + Ok(name) => name.0, _ => continue, }; @@ -304,8 +306,8 @@ impl OrigHeaderMap { /// updated, though; this matters for types that can be `==` without being /// identical. #[inline] - pub fn insert(&mut self, value: PyBackedStr) -> bool { - self.0.insert(Bytes::from_owner(value)) + pub fn insert(&mut self, value: StrInput) -> bool { + self.0.insert(value.0) } /// Extends the map with all entries from another [`OrigHeaderMap`], preserving order. @@ -354,8 +356,8 @@ impl FromPyObject<'_, '_> for OrigHeaderMap { header::OrigHeaderMap::with_capacity(list.len()), |mut headers, name| { let name = { - let name = name.extract::()?; - Bytes::from_owner(name) + let name = name.extract::()?; + name.0 }; headers.insert(name); Ok(headers) diff --git a/src/proxy.rs b/src/proxy.rs index 1277277b..d93a994a 100644 --- a/src/proxy.rs +++ b/src/proxy.rs @@ -1,8 +1,7 @@ -use bytes::Bytes; use pyo3::{prelude::*, pybacked::PyBackedStr}; use wreq::header::HeaderValue; -use crate::{error::Error, header::HeaderMap}; +use crate::{error::Error, extractor::StrInput, header::HeaderMap}; /// A builder for `Proxy`. #[derive(Default)] @@ -14,7 +13,7 @@ struct Builder { password: Option, // Optional custom HTTP authentication header. - custom_http_auth: Option, + custom_http_auth: Option, /// Optional custom HTTP headers for the proxy. custom_http_headers: Option, @@ -114,7 +113,7 @@ fn create_proxy<'py>( // Convert the custom HTTP auth string to a header value. if let Some(Ok(custom_http_auth)) = builder .custom_http_auth - .map(Bytes::from_owner) + .map(|value| value.0) .map(HeaderValue::from_maybe_shared) { proxy = proxy.custom_http_auth(custom_http_auth); diff --git a/src/tls.rs b/src/tls.rs index e81a323b..59e3d2a3 100644 --- a/src/tls.rs +++ b/src/tls.rs @@ -426,7 +426,7 @@ pub struct TlsInfo(pub wreq::tls::TlsInfo); #[pymethods] impl TlsInfo { - /// Get the DER encoded leaf certificate of the peer. + /// Get a read-only memoryview of the peer's DER-encoded leaf certificate. #[inline] pub fn peer_certificate(&self) -> Option { self.0 diff --git a/tests/buffer_test.py b/tests/buffer_test.py new file mode 100644 index 00000000..b78f1f30 --- /dev/null +++ b/tests/buffer_test.py @@ -0,0 +1,251 @@ +import asyncio +import datetime +import gc +import threading +import weakref +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer + +import pytest + +import wreq +from wreq import Message, Multipart, Part, Version, blocking +from wreq.header import HeaderMap, OrigHeaderMap + + +def assert_readonly_view(view, expected): + assert type(view) is memoryview + assert view.readonly + assert view.format == "B" + assert view.ndim == 1 + assert view.itemsize == 1 + assert view.shape == (len(expected),) + assert view == expected + + +def test_header_and_message_views(): + headers = HeaderMap({"X-Buffer": "value"}) + headers.append("X-Buffer", "other") + original = OrigHeaderMap(["X-Buffer"]) + messages = [ + Message.from_text("text"), + Message.from_binary(b"binary"), + Message.from_ping(b"ping"), + Message.from_pong(b"pong"), + Message.from_binary(b""), + ] + assert messages[0].text == "text" + assert type(messages[0].text) is str + round_trip = Message.from_binary(bytes(messages[1].binary)) + assert round_trip.binary == b"binary" + assert headers.get("missing") is None + assert headers.get("missing", b"default") == b"default" + + views = [ + headers.get("X-Buffer"), + headers["X-Buffer"], + headers.get("missing", b"default"), + *headers.get_all("X-Buffer"), + *headers.keys(), + *headers.values(), + *(view for pair in headers for view in pair), + *(view for pair in original for view in pair), + *(message.data for message in messages), + messages[1].binary, + messages[2].ping, + messages[3].pong, + round_trip.binary, + ] + expected = [bytes(view) for view in views] + for view, data in zip(views, expected): + assert_readonly_view(view, data) + assert hash(view) == hash(data) + assert {view: "value"}[data] == "value" + + with pytest.raises(TypeError): + views[0][0] = 0 + + headers.clear() + del headers, original, messages, round_trip + gc.collect() + for view, data in zip(views, expected): + assert view == data + + parent = views[0] + child = parent[1:] + parent.release() + assert child == b"alue" + child.release() + + +def test_subclass_input_cycles_are_collected(): + class Text(str): + def __str__(self): + raise AssertionError("subclass conversion must not be called") + + class Binary(bytes): + def __bytes__(self): + raise AssertionError("subclass conversion must not be called") + + class Marker: + pass + + def collect_garbage(): + # PyPy may need several GC cycles to finalize C-extension buffers. + for _ in range(3): + gc.collect() + + def header_value(source, method): + headers = HeaderMap() + getattr(headers, method)("X-Buffer", source) + return headers["X-Buffer"] + + def original_name(source): + headers = OrigHeaderMap() + headers.insert(source) + return next(iter(headers))[1] + + cases = [ + (Text, "payload", lambda source: Message.from_text(source).data), + (Binary, b"payload", lambda source: Message.from_binary(source).binary), + (Binary, b"payload", lambda source: Message.from_ping(source).ping), + (Binary, b"payload", lambda source: Message.from_pong(source).pong), + (Text, "payload", lambda source: HeaderMap({"X-Buffer": source})["X-Buffer"]), + (Text, "payload", lambda source: header_value(source, "insert")), + (Text, "payload", lambda source: header_value(source, "append")), + (Text, "payload", lambda source: header_value(source, "__setitem__")), + (Binary, b"payload", lambda source: HeaderMap().get("missing", source)), + (Text, "X-Buffer", lambda source: next(iter(OrigHeaderMap([source])))[1]), + (Text, "X-Buffer", original_name), + ] + for input_type, payload, make_view in cases: + source = input_type(payload) + source.marker = Marker() + marker = weakref.ref(source.marker) + view = make_view(source) + expected = payload.encode() if isinstance(payload, str) else payload + assert_readonly_view(view, expected) + source.view = view + del source, view + collect_garbage() + assert marker() is None, make_view + + for input_type, payload, make_owner in [ + (Text, "payload", lambda source: Part("field", source)), + (Binary, b"payload", lambda source: Multipart(Part("field", source))), + (Text, "payload", lambda source: Message.from_close(1000, source)), + ]: + source = input_type(payload) + source.marker = Marker() + marker = weakref.ref(source.marker) + owner = make_owner(source) + source.owner = owner + del source, owner + collect_garbage() + assert marker() is None, make_owner + + +@pytest.fixture +def buffer_http_server(): + class Handler(BaseHTTPRequestHandler): + protocol_version = "HTTP/1.1" + + def setup(self): + super().setup() + self.connection.settimeout(3) + + def do_GET(self): + self.send_response(200) + self.send_header("Connection", "close") + if self.path == "/stream": + self.send_header("Transfer-Encoding", "chunked") + self.send_header("Trailer", "X-Buffer") + self.end_headers() + self.wfile.write( + b"5\r\nhello\r\n6\r\n world\r\n0\r\nX-Buffer: complete\r\n\r\n" + ) + else: + self.send_header("Content-Length", "11") + self.end_headers() + self.wfile.write(b"hello world") + + def log_message(self, *_): + pass + + server = ThreadingHTTPServer(("127.0.0.1", 0), Handler) + thread = threading.Thread( + target=server.serve_forever, kwargs={"poll_interval": 0.05}, daemon=True + ) + thread.start() + try: + yield f"http://127.0.0.1:{server.server_port}" + finally: + server.shutdown() + server.server_close() + thread.join(timeout=5) + assert not thread.is_alive() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("blocking_api", [False, True]) +async def test_response_and_stream_views(buffer_http_server, blocking_api): + client_type = blocking.Client if blocking_api else wreq.Client + client = client_type(no_proxy=True, timeout=datetime.timedelta(seconds=3)) + + async def call(method, *args, **kwargs): + if blocking_api: + return await asyncio.wait_for( + asyncio.to_thread(method, *args, **kwargs), timeout=5 + ) + return await asyncio.wait_for(method(*args, **kwargs), timeout=5) + + def collect_blocking(response): + with response.stream() as stream: + return list(stream) + + async def collect_async(response): + async with response.stream() as stream: + return [frame async for frame in stream] + + try: + response = await call( + client.get, f"{buffer_http_server}/body", version=Version.HTTP_11 + ) + view = await call(response.bytes) + repeated = await call(response.bytes) + assert_readonly_view(view, b"hello world") + assert_readonly_view(repeated, b"hello world") + repeated.release() + child = view[6:] + await call(response.close) + del response, repeated + gc.collect() + assert view == b"hello world" + view.release() + assert child == b"world" + + response = await call( + client.get, f"{buffer_http_server}/stream", version=Version.HTTP_11 + ) + if blocking_api: + frames = await asyncio.wait_for( + asyncio.to_thread(collect_blocking, response), timeout=5 + ) + else: + frames = await asyncio.wait_for(collect_async(response), timeout=5) + await call(response.close) + del response + gc.collect() + data = [frame for frame in frames if isinstance(frame, memoryview)] + trailers = [frame for frame in frames if isinstance(frame, HeaderMap)] + assert data + assert b"".join(data) == b"hello world" + assert len(trailers) == 1 + assert_readonly_view(trailers[0]["X-Buffer"], b"complete") + for frame in data: + expected = bytes(frame) + assert_readonly_view(frame, expected) + child = frame[:] + frame.release() + assert child == expected + finally: + client.close() diff --git a/tests/response_test.py b/tests/response_test.py index ad7c017a..df9d0ba9 100644 --- a/tests/response_test.py +++ b/tests/response_test.py @@ -117,4 +117,6 @@ async def test_peer_certificate(): resp = await client.get("https://www.google.com/anything") async with resp: assert resp.tls_info is not None - assert resp.tls_info.peer_certificate() is not None + certificate = resp.tls_info.peer_certificate() + assert type(certificate) is memoryview + assert certificate.readonly diff --git a/tests/tls_test.py b/tests/tls_test.py index 44e9225f..2517deee 100644 --- a/tests/tls_test.py +++ b/tests/tls_test.py @@ -27,10 +27,11 @@ async def test_badssl_invalid_cert(): peer_der_cert = tls_info.peer_certificate() assert peer_der_cert is not None - assert isinstance(peer_der_cert, bytes) + assert type(peer_der_cert) is memoryview + assert peer_der_cert.readonly assert len(peer_der_cert) > 0 - cert_store = CertStore(der_certs=[peer_der_cert]) + cert_store = CertStore(der_certs=[bytes(peer_der_cert)]) assert cert_store is not None client = wreq.Client(tls_verify=cert_store)