diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 67cf1880..fef4c7ee 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -129,8 +129,10 @@ jobs: *) echo "Expected a pp311 wheel, got: $wheel"; exit 1 ;; esac - name: Run tests - run: .venv-pypy/bin/python -m pytest + timeout-minutes: 15 + run: .venv-pypy/bin/python -u -m pytest -vv -x -o faulthandler_timeout=60 - name: Upload wheel + if: always() uses: actions/upload-artifact@v7 with: name: wheels-linux-x86_64-pypy311 diff --git a/Cargo.lock b/Cargo.lock index 7af90083..cc403a56 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -47,18 +47,6 @@ dependencies = [ "rustversion", ] -[[package]] -name = "async-channel" -version = "2.5.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "924ed96dd52d1b75e9c1a3e6275715fd320f5f9439fb5a4a11fa51f4221158d2" -dependencies = [ - "concurrent-queue", - "event-listener-strategy", - "futures-core", - "pin-project-lite", -] - [[package]] name = "async-compression" version = "0.4.47" @@ -271,15 +259,6 @@ version = "0.4.33" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "6e8ccc4ea9f6acc32d102c0f6d471d11d913ad15f20c04de743374861fa1d414" -[[package]] -name = "concurrent-queue" -version = "2.5.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "4ca0197aee26d1ae37445ee532fefce43251d24cc7c166799f4d46817f1d3973" -dependencies = [ - "crossbeam-utils", -] - [[package]] name = "cookie" version = "0.18.2" @@ -439,26 +418,6 @@ version = "1.0.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "877a4ace8713b0bcf2a4e7eec82529c029f1d0619886d18145fea96c3ffe5c0f" -[[package]] -name = "event-listener" -version = "5.4.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "5a23add41df1562121a9393cb065eab5146a1242410f23a644851e90cfd669d2" -dependencies = [ - "parking", - "pin-project-lite", -] - -[[package]] -name = "event-listener-strategy" -version = "0.5.4" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "8be9f3dfaaffdae2972880079a491a1a8bb7cbed0b8dd7a347f668b4150a3b93" -dependencies = [ - "event-listener", - "pin-project-lite", -] - [[package]] name = "find-msvc-tools" version = "0.1.12" @@ -547,7 +506,6 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b1f9e3d69d39e4862ffed03ed071a76f9a13ba1d9109d355b0f0aa6b15e393c4" dependencies = [ "futures-core", - "futures-sink", ] [[package]] @@ -1237,12 +1195,6 @@ dependencies = [ "syn 2.0.119", ] -[[package]] -name = "parking" -version = "2.2.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f38d5652c16fde515bb1ecef450ab0f6a219d619a7274976324d5e377f7dceba" - [[package]] name = "parking_lot" version = "0.12.5" @@ -1278,6 +1230,20 @@ version = "0.2.17" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "a89322df9ebe1c1578d689c92318e070967d1042b512afbe49518723f4e6d5cd" +[[package]] +name = "pingora-runtime" +version = "0.9.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "89aed2e58b34196682bdc2370c4293f908dece5741c56f5f5ba1018f55de68f2" +dependencies = [ + "log", + "once_cell", + "rand 0.8.8", + "serde", + "thread_local", + "tokio", +] + [[package]] name = "pkg-config" version = "0.3.34" @@ -1350,21 +1316,6 @@ dependencies = [ "pyo3-macros", ] -[[package]] -name = "pyo3-async-runtimes" -version = "0.29.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b3ef68daa7316a3fac65e5e18b2203f010346de1c1c53456811a2624673ab046" -dependencies = [ - "async-channel", - "futures-channel", - "futures-util", - "once_cell", - "pin-project-lite", - "pyo3", - "tokio", -] - [[package]] name = "pyo3-build-config" version = "0.29.2" @@ -1429,13 +1380,24 @@ version = "6.0.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "f8dcc9c7d52a811697d2151c701e0d08956f92b0e24136cf4cf27b57a6a0d9bf" +[[package]] +name = "rand" +version = "0.8.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e058c7de0b26af77780c769414d6257830bb240f3c38477dbc2c16e5f54d6d4c" +dependencies = [ + "libc", + "rand_chacha 0.3.1", + "rand_core 0.6.4", +] + [[package]] name = "rand" version = "0.9.5" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b9ef1d0d795eb7d84685bca4f72f3649f064e6641543d3a8c415898726a57b41" dependencies = [ - "rand_chacha", + "rand_chacha 0.9.0", "rand_core 0.9.5", ] @@ -1450,6 +1412,16 @@ dependencies = [ "rand_core 0.10.1", ] +[[package]] +name = "rand_chacha" +version = "0.3.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e6c10a63a0fa32252be49d21e7709d4d4baf8d231c2dbce1eaa8141b9b127d88" +dependencies = [ + "ppv-lite86", + "rand_core 0.6.4", +] + [[package]] name = "rand_chacha" version = "0.9.0" @@ -1460,6 +1432,15 @@ dependencies = [ "rand_core 0.9.5", ] +[[package]] +name = "rand_core" +version = "0.6.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ec0be4795e2f6a28069bec0b5ff3e2ac9bafc99e6a9a7dc3547996c5c816922c" +dependencies = [ + "getrandom 0.2.17", +] + [[package]] name = "rand_core" version = "0.9.5" @@ -1843,6 +1824,15 @@ dependencies = [ "syn 3.0.5", ] +[[package]] +name = "thread_local" +version = "1.1.10" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1ad99c4c6d32803332c548b1af0540b357b3f5fc0be8f6c6bfe8b2e6ae784070" +dependencies = [ + "cfg-if", +] + [[package]] name = "tikv-jemalloc-sys" version = "0.7.1+5.3.1-0-g81034ce1f1373e37dc865038e1bc8eeecf559ce8" @@ -2475,9 +2465,8 @@ dependencies = [ "http-body-util", "indexmap", "mimalloc", - "pin-project-lite", + "pingora-runtime", "pyo3", - "pyo3-async-runtimes", "serde", "serde_json", "serde_urlencoded", diff --git a/Cargo.toml b/Cargo.toml index 539d6cd5..837f4d70 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -28,7 +28,8 @@ abi3-py313 = ["pyo3/abi3-py313"] abi3-py314 = ["pyo3/abi3-py314"] [dependencies] -tokio = "1.52.2" +pingora-runtime = "0.9.0" +tokio = { version = "1.52.2", features = ["rt-multi-thread", "sync", "time", "net"] } tokio-util = { version = "0.7.18", features = ["rt"] } pyo3 = { version = "0.29.0", features = [ "indexmap", @@ -37,11 +38,6 @@ pyo3 = { version = "0.29.0", features = [ "generate-import-lib", "experimental-async", ] } -pyo3-async-runtimes = { version = "0.29.0", features = [ - "tokio-runtime", - "unstable-streams", -] } -pin-project-lite = "0.2.16" futures-util = { version = "0.3.33", default-features = false } serde = { version = "1.0.228", features = ["derive"] } serde_json = "1" @@ -92,4 +88,5 @@ debug = false incremental = false lto = "fat" opt-level = 3 +panic = "abort" strip = true diff --git a/docs/mkdocs.yml b/docs/mkdocs.yml index 81652186..51ae0e80 100644 --- a/docs/mkdocs.yml +++ b/docs/mkdocs.yml @@ -113,6 +113,7 @@ nav: - Modules: - wreq: api/wreq.md - wreq.blocking: api/blocking.md + - wreq.runtime: api/runtime.md - wreq.header: api/header.md - wreq.cookie: api/cookie.md - wreq.exceptions: api/exceptions.md diff --git a/docs/source/api/runtime.md b/docs/source/api/runtime.md new file mode 100644 index 00000000..dfda7fb5 --- /dev/null +++ b/docs/source/api/runtime.md @@ -0,0 +1,8 @@ +# wreq.runtime + +Runtime configuration for asynchronous and blocking clients. `Runtime` is also +available as `wreq.Runtime`. + +::: wreq.runtime.Runtime + options: + show_root_heading: true diff --git a/docs/source/guide/advanced.md b/docs/source/guide/advanced.md index a7bb3987..f8961f93 100644 --- a/docs/source/guide/advanced.md +++ b/docs/source/guide/advanced.md @@ -8,6 +8,15 @@ Send data using async generators for streaming uploads: +Async upload generators run on the caller's running event loop with its context variables. +Their exceptions fail the request. When an upload ends, the generator is closed; +cancelling or dropping the upload schedules producer cancellation and cleanup on that loop. +Keep the loop running until generator cleanup has finished. This also applies to async multipart parts. +Construct async-generator `Part` objects inside a running event loop; their producers +start at construction, with bounded buffering before the request consumes them. +Use synchronous iterators for blocking uploads. A blocking call on the producer's +event-loop thread prevents async generators from progressing. + ```python import asyncio import wreq @@ -87,6 +96,60 @@ if __name__ == "__main__": asyncio.run(main()) ``` +### Custom runtimes + +Clients share a global multi-thread runtime when `runtime` is omitted or `None`. +It starts on first use. Construct a `Runtime` to +start a separate worker pool for an async or blocking client: + +```python +from datetime import timedelta + +from wreq import Client +from wreq.runtime import Runtime + +runtime = Runtime( + workers=1, + work_steal=False, + thread_name="http-client", + max_blocking_threads=8, + thread_keep_alive=timedelta(seconds=10), +) +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. + +`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 +bound or request is sent. + +`thread_name=None` uses the package name, `wreq-python`, as the thread name. + +`thread_keep_alive` accepts a nonnegative `datetime.timedelta`. +`max_blocking_threads` and `thread_keep_alive` default to Tokio's settings +(512 and 10 seconds). In +no-steal mode these limits apply to **each worker's** blocking pool, not the pool +as a whole. Python async upload generators still run on the caller's event loop. +Standalone multipart file preparation and upload-task cleanup can use the +shared runtime; a dedicated client runtime does not isolate Python's GIL or +every process resource. DNS resolvers are owned by individual clients so their +connections are not shared across runtimes. + +Closing a client cancels pending requests and rejects new requests with +`asyncio.CancelledError`, for both async and blocking APIs. It does not shut down +the runtime or invalidate existing responses and WebSockets. +Clients, responses, streams and active tasks share ownership. Dropping the last +owner automatically releases a custom runtime without synchronously waiting for +its workers; already running blocking work may finish later. The default runtime +is shared for the process lifetime. Zero thread counts, NUL characters in thread +names and negative durations raise `ValueError`. + ### TLS Key Logging Capture TLS keys for debugging with tools like Wireshark: diff --git a/docs/source/guide/blocking.md b/docs/source/guide/blocking.md index a65cc3d9..a4ef7c1c 100644 --- a/docs/source/guide/blocking.md +++ b/docs/source/guide/blocking.md @@ -65,6 +65,27 @@ if __name__ == "__main__": main() ``` +### Custom Runtime + +The blocking client accepts the same `Runtime` as the async client. Without one, +it uses the shared global multi-thread runtime. + +```python +from wreq.blocking import Client +from wreq.runtime import Runtime + +runtime = Runtime(workers=1, work_steal=False) +with Client(runtime=runtime) as client: + with client.get("https://httpbin.io/get") as response: + print(response.text()) +``` + +Network work runs on the selected worker while the calling thread waits. +`client.runtime` is read-only. Closing the client does not shut down a shared +runtime; it cancels pending requests and rejects new ones with +`asyncio.CancelledError`. See [custom runtimes](advanced.md#custom-runtimes) for +configuration and lifetime details. + ### Cookies ```python diff --git a/python/wreq/__init__.py b/python/wreq/__init__.py index f92a9ae4..a3955b41 100644 --- a/python/wreq/__init__.py +++ b/python/wreq/__init__.py @@ -12,6 +12,22 @@ from .dns import * from .redirect import * from .proxy import * +from .runtime import * + +import sys as _sys + +if _sys.implementation.name == "pypy": + from ._compat import _install + + # Creating and closing an unpolled coroutine does not start a request or runtime. + _coroutine = get("") + try: + _install(type(_coroutine)) + finally: + _coroutine.close() + del _coroutine, _install + +del _sys __all__ = ( header.__all__ @@ -24,4 +40,5 @@ + dns.__all__ + redirect.__all__ + proxy.__all__ + + runtime.__all__ ) # type: ignore diff --git a/python/wreq/_compat.py b/python/wreq/_compat.py new file mode 100644 index 00000000..9b610e7d --- /dev/null +++ b/python/wreq/_compat.py @@ -0,0 +1,48 @@ +"""Compatibility for PyPy's legacy exception delegation to PyO3 coroutines.""" + +from types import TracebackType + + +def _install(coroutine_type): + original = vars(coroutine_type)["throw"] + if getattr(original, "_wreq_pypy_throw_compat", False): + return + + # Remove this shim once PyO3 accepts throw(type, value, traceback). + def throw(self, exc, value=None, traceback=None): + if traceback is not None and type(traceback) is not TracebackType: + raise TypeError("throw() third argument must be a traceback object") + if issubclass(type(exc), BaseException): + if value is not None: + raise TypeError("instance exception may not have a separate value") + instance = exc + if traceback is not None: + BaseException.with_traceback(instance, traceback) + elif issubclass(type(exc), type) and issubclass(exc, BaseException): + try: + if issubclass(type(value), BaseException) and issubclass( + type(value), exc + ): + instance = value + elif value is None: + instance = exc() + elif issubclass(type(value), tuple): + instance = exc(*value) + else: + instance = exc(value) + except BaseException as error: + instance = error + own_traceback = BaseException.__traceback__.__get__(error).tb_next + BaseException.with_traceback(instance, own_traceback or traceback) + else: + if not issubclass(type(instance), BaseException): + instance = TypeError( + "exception constructor did not return an exception" + ) + BaseException.with_traceback(instance, traceback) + else: + raise TypeError("exceptions must derive from BaseException") + return original(self, instance) + + throw._wreq_pypy_throw_compat = True + coroutine_type.throw = throw diff --git a/python/wreq/blocking.py b/python/wreq/blocking.py index ddc7e794..3ef7064a 100644 --- a/python/wreq/blocking.py +++ b/python/wreq/blocking.py @@ -19,6 +19,7 @@ from .cookie import Cookie, Jar from .header import HeaderMap from .redirect import History +from .runtime import Runtime from .tls import TlsInfo @@ -213,6 +214,9 @@ class Client: A blocking client for making HTTP requests. """ + runtime: Runtime + """Read-only shared runtime used by this client and its responses.""" + cookie_jar: Jar | None r""" Get the cookie jar used by this client (if enabled/configured). @@ -247,9 +251,8 @@ def __init__( def close(self) -> None: r""" - Closes the client and any associated resources. - - After calling this method, the client should not be used to make further requests. + Cancels pending requests and rejects new ones with `asyncio.CancelledError`. + Existing responses, WebSockets and the shared runtime remain usable. Examples: diff --git a/python/wreq/exceptions.py b/python/wreq/exceptions.py index 68aaa5d8..b50190c0 100644 --- a/python/wreq/exceptions.py +++ b/python/wreq/exceptions.py @@ -29,7 +29,7 @@ class RustPanic(Exception): r""" - A panic occurred in the underlying Rust code. + Compatibility exception; Rust panics are not translated to this type. """ diff --git a/python/wreq/runtime.py b/python/wreq/runtime.py new file mode 100644 index 00000000..1a3c6ab7 --- /dev/null +++ b/python/wreq/runtime.py @@ -0,0 +1,35 @@ +import datetime +from typing import final + +__all__ = ["Runtime"] + + +@final +class Runtime: + """Shared Tokio runtime whose workers start when constructed. + + Clients, responses and active work keep it alive; the last owner releases + it automatically. + No-steal clients keep a fixed worker, without CPU pinning. + """ + + def __init__( + self, + *, + workers: int | None = None, + work_steal: bool = True, + thread_name: str | None = None, + max_blocking_threads: int | None = None, + thread_keep_alive: datetime.timedelta | None = None, + ) -> None: + """Thread counts must be positive; thread_keep_alive is a nonnegative timedelta. + + workers defaults to available CPU parallelism, or 1 if unavailable. + work_steal=True uses a multi-thread pool; False uses independent + single-thread workers. + thread_name defaults to the package name, "wreq-python". + Thread names cannot contain NUL characters. + Blocking-pool settings apply to each worker runtime in no-steal mode. + None preserves Tokio's blocking-pool defaults. + """ + ... diff --git a/python/wreq/wreq.py b/python/wreq/wreq.py index 31155aa2..18441cda 100644 --- a/python/wreq/wreq.py +++ b/python/wreq/wreq.py @@ -26,6 +26,7 @@ from .http2 import Http2Options from .proxy import * from .redirect import History +from .runtime import Runtime from .tls import * @@ -161,6 +162,9 @@ def __init__( r""" Creates a new part. + Construct async-generator parts inside a running event loop. Their producers + start immediately on that loop and are closed after use. + # Arguments - `name` - The name of the part. - `value` - The value of the part, either text, bytes, a file path, or a async or sync stream. @@ -309,12 +313,19 @@ async def main(): """ def __iter__(self) -> "Streamer": ... + def __next__(self) -> memoryview | HeaderMap: ... + def __enter__(self) -> Any: ... + def __exit__(self, _exc_type: Any, _exc_value: Any, _traceback: Any) -> None: ... + 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 ) -> None: ... @@ -514,6 +525,9 @@ def __str__(self) -> str: ... class ClientConfig(TypedDict): + runtime: NotRequired[Runtime | None] + """Runtime for this client and its responses; None uses the shared default.""" + emulation: NotRequired[emulation.Emulation | emulation.Profile] """Emulation config.""" @@ -947,6 +961,8 @@ class Request(TypedDict): ] """ The body to use for the request. + Async generators run on the caller's running event loop. Upload errors fail + the request; cancellation schedules generator cleanup on that loop. """ multipart: NotRequired[Multipart] @@ -1101,6 +1117,9 @@ class Client: A client for making HTTP requests. """ + runtime: Runtime + """Read-only shared runtime used by this client and its responses.""" + cookie_jar: Jar | None r""" Get the cookie jar used by this client (if enabled/configured). @@ -1138,9 +1157,8 @@ async def main(): def close(self) -> None: r""" - Closes the client and any associated resources. - - After calling this method, the client should not be used to make further requests. + Cancels pending requests and rejects new ones with `asyncio.CancelledError`. + Existing responses, WebSockets and the shared runtime remain usable. Examples: diff --git a/src/client.rs b/src/client.rs index e5f28296..c545dda4 100644 --- a/src/client.rs +++ b/src/client.rs @@ -32,7 +32,7 @@ use crate::{ http1::Http1Options, http2::Http2Options, proxy::Proxy, - redirect, + redirect, runtime, tls::{Identity, KeyLog, TlsOptions, TlsVerify, TlsVersion}, }; @@ -59,6 +59,7 @@ impl_print_str!(Display, SocketAddr); /// A builder for `Client`. #[derive(Default)] struct Builder { + runtime: Option, /// The Emulation settings for the client. emulation: Option, /// The user agent to use for the client. @@ -171,6 +172,7 @@ impl FromPyObject<'_, '_> for Builder { fn extract(ob: Borrowed) -> PyResult { let mut builder = Self::default(); + extract_option!(ob, builder, runtime); extract_option!(ob, builder, emulation); extract_option!(ob, builder, user_agent); extract_option!(ob, builder, headers); @@ -229,10 +231,11 @@ impl FromPyObject<'_, '_> for Builder { } /// A client for making HTTP requests. -#[derive(Default, Clone)] +#[derive(Clone)] #[pyclass(subclass, frozen, skip_from_py_object)] pub struct Client { inner: wreq::Client, + runtime: runtime::Runtime, cancel: CancellationToken, raise_for_status: bool, @@ -248,6 +251,19 @@ pub struct BlockingClient(Client); // ====== Client ===== +impl Default for Client { + fn default() -> Self { + let runtime = runtime::get(); + Self { + inner: wreq::Client::default(), + runtime: runtime.clone(), + cancel: CancellationToken::new(), + raise_for_status: false, + cookie_jar: None, + } + } +} + #[pymethods] impl Client { /// Creates a new Client instance. @@ -255,6 +271,10 @@ impl Client { #[pyo3(signature = (**kwds))] fn new(py: Python, kwds: Option) -> PyResult { py.detach(|| { + let runtime = match kwds.as_ref().and_then(|config| config.runtime.as_ref()) { + Some(runtime) => runtime.select()?, + None => runtime::get().clone(), + }; // Create the client builder. let mut builder = wreq::Client::builder(); let mut cookie_jar: Option = None; @@ -464,6 +484,7 @@ impl Client { .build() .map(|inner| Client { inner, + runtime, cancel: CancellationToken::new(), cookie_jar, raise_for_status, @@ -473,12 +494,19 @@ impl Client { }) } - /// Close the client, preventing any new requests. + /// Cancel pending requests and reject new ones with asyncio.CancelledError. + /// Existing responses, WebSockets and the shared runtime remain usable. #[inline] pub fn close(&self) { self.cancel.cancel(); } + /// The runtime used by this client and its responses. + #[getter] + pub fn runtime(&self) -> runtime::Runtime { + self.runtime.clone() + } + /// Make a GET request to the given URL. #[inline(always)] #[pyo3(signature = (url, **kwds))] @@ -585,10 +613,10 @@ impl Client { url: PyBackedStr, kwds: Option, ) -> PyResult { - NoGIL::new_with_token( + NoGIL::with_cancel( + &self.runtime, execute_request(self.clone(), method, url, kwds), cancel, - self.cancel.clone(), ) .await } @@ -602,10 +630,10 @@ impl Client { url: PyBackedStr, kwds: Option, ) -> PyResult { - NoGIL::new_with_token( + NoGIL::with_cancel( + &self.runtime, execute_websocket_request(self.clone(), url, kwds), cancel, - self.cancel.clone(), ) .await } @@ -628,6 +656,12 @@ impl Client { #[pymethods] impl BlockingClient { + /// The runtime used by this client and its responses. + #[getter] + pub fn runtime(&self) -> runtime::Runtime { + self.0.runtime() + } + /// Creates a new blocking Client instance. #[new] #[inline] @@ -643,7 +677,8 @@ impl BlockingClient { self.0.cookie_jar.clone() } - /// Close the client, preventing any new requests. + /// Cancel pending requests and reject new ones with asyncio.CancelledError. + /// Existing responses, WebSockets and the shared runtime remain usable. #[inline] pub fn close(&self) { self.0.close(); @@ -755,9 +790,11 @@ impl BlockingClient { kwds: Option, ) -> PyResult { py.detach(|| { - pyo3_async_runtimes::tokio::get_runtime() - .block_on(execute_request(self.0.clone(), method, url, kwds)) - .map(Into::into) + nogil::block_on( + &self.0.runtime, + execute_request(self.0.clone(), method, url, kwds), + ) + .map(Into::into) }) } @@ -770,9 +807,11 @@ impl BlockingClient { kwds: Option, ) -> PyResult { py.detach(|| { - pyo3_async_runtimes::tokio::get_runtime() - .block_on(execute_websocket_request(self.0.clone(), url, kwds)) - .map(Into::into) + nogil::block_on( + &self.0.runtime, + execute_websocket_request(self.0.clone(), url, kwds), + ) + .map(Into::into) }) } } diff --git a/src/client/body/multipart.rs b/src/client/body/multipart.rs index 34f05280..7648f2c2 100644 --- a/src/client/body/multipart.rs +++ b/src/client/body/multipart.rs @@ -133,9 +133,9 @@ impl Part { let mut inner = match value { 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)?, + Value::File(path) => crate::runtime::get().handle().block_on(async move { + multipart::Part::file(path).await.map_err(Error::from) + })?, Value::Stream(stream) => { let stream = Body::wrap_stream(stream); match self.length { diff --git a/src/client/body/stream.rs b/src/client/body/stream.rs index 68ac7504..9908c53c 100644 --- a/src/client/body/stream.rs +++ b/src/client/body/stream.rs @@ -5,17 +5,23 @@ use std::{ }; use bytes::Bytes; -use futures_util::{FutureExt, Stream, StreamExt, stream::BoxStream}; +use futures_util::{FutureExt, Stream, future::poll_fn}; use http_body_util::BodyExt; -use pyo3::{coroutine::CancelHandle, exceptions::PyStopIteration, intern, prelude::*}; -use tokio::{sync::Mutex, task::JoinHandle}; +use pyo3::{ + coroutine::CancelHandle, exceptions::PyStopIteration, intern, prelude::*, sync::PyOnceLock, +}; +use tokio::{ + sync::{Mutex, mpsc}, + task::JoinHandle, +}; use crate::{ buffer::PyBuffer, - client::nogil::NoGIL, - error::{self, Error}, + client::nogil::{self, NoGIL}, + error::Error, extractor::{BytesInput, StrInput}, header::HeaderMap, + runtime::Runtime, }; type Pending = Option>>>; @@ -23,7 +29,7 @@ type Pending = Option>>>; /// Python stream source. enum PyStreamSource { Sync(Arc>), - Async(Arc>>>), + Async(PyAsyncStream), } /// A bytes-like object that can be extracted from Python. @@ -46,30 +52,28 @@ pub struct PyStream { pending: Pending, } +/// Adapts a Python async generator into a byte stream with bounded buffering. +/// Dropping the stream cancels its producer on the Python event loop. +struct PyAsyncStream { + rx: mpsc::Receiver>>, + task: Option<(Py, Py)>, +} + +#[pyclass(frozen)] +struct Sender(mpsc::Sender>>); + /// A response stream yielding read-only memoryviews and any trailing headers. #[derive(Clone)] #[pyclass(subclass, frozen, skip_from_py_object)] -pub struct Streamer(Arc>>); - -// ===== impl PyStream ===== - -impl From for PyStream { - #[inline] - fn from(inner: PyStreamSource) -> Self { - PyStream { - inner, - pending: None, - } - } -} +pub struct Streamer(Arc>>, Runtime); // ===== impl Streamer ===== impl Streamer { /// Create a new [`Streamer`] instance. #[inline] - pub fn new(resp: wreq::Response) -> Streamer { - Streamer(Arc::new(Mutex::new(Some(resp)))) + pub fn new(resp: wreq::Response, runtime: Runtime) -> Streamer { + Streamer(Arc::new(Mutex::new(Some(resp))), runtime) } async fn next(self, error: fn() -> Error) -> PyResult { @@ -109,10 +113,7 @@ impl Streamer { #[inline] fn __next__(&self, py: Python) -> PyResult { - py.detach(|| { - pyo3_async_runtimes::tokio::get_runtime() - .block_on(self.clone().next(|| Error::StopIteration)) - }) + py.detach(|| nogil::block_on(&self.1, self.clone().next(|| Error::StopIteration))) } #[inline] @@ -139,12 +140,33 @@ impl Streamer { slf } + /// Read the next frame when awaited; returns a coroutine, not a Future. #[inline] fn __anext__<'py>(&self, py: Python<'py>) -> PyResult> { - pyo3_async_runtimes::tokio::future_into_py( + let this = self.clone(); + let cancel = CancelHandle::new(); + // PyO3 0.29 cannot wrap an async __anext__ slot; use its macro constructor. + // Recheck this internal API when upgrading PyO3. + Bound::new( py, - self.clone().next(|| Error::StopAsyncIteration), + pyo3::impl_::coroutine::new_coroutine( + intern!(py, "__anext__"), + Some("Streamer"), + Some(cancel.throw_callback()), + async move { + let runtime = this.1.clone(); + let frame = NoGIL::with_cancel( + &runtime, + this.next(|| Error::StopAsyncIteration), + cancel, + ) + .await?; + // PyO3 polls this coroutine while attached, outside the Tokio task. + Python::attach(|py| frame.into_pyobject(py).map(|obj| obj.unbind())) + }, + ), ) + .map(Bound::into_any) } #[inline] @@ -160,20 +182,17 @@ impl Streamer { _traceback: Py, ) -> PyResult<()> { let this = self.0.clone(); - NoGIL::new( - async move { - if let Some(resp) = this.lock().await.take() { - drop(resp) - } - Ok(()) - }, - CancelHandle::new(), - ) + NoGIL::new(&self.1, async move { + if let Some(resp) = this.lock().await.take() { + drop(resp) + } + Ok(()) + }) .await } } -// ===== PyBytesLike ===== +// ===== impl PyBytesLike ===== impl From for Bytes { #[inline] @@ -187,15 +206,22 @@ impl From for Bytes { // ===== impl PyStream ===== +impl From for PyStream { + #[inline] + fn from(inner: PyStreamSource) -> Self { + PyStream { + inner, + pending: None, + } + } +} + impl FromPyObject<'_, '_> for PyStream { type Error = PyErr; fn extract(ob: Borrowed) -> PyResult { if ob.hasattr(intern!(ob.py(), "asend"))? { - pyo3_async_runtimes::tokio::into_stream_v2(ob.to_owned()) - .map(StreamExt::boxed) - .map(Mutex::new) - .map(Arc::new) + PyAsyncStream::new(ob.to_owned()) .map(PyStreamSource::Async) .map(PyStream::from) } else { @@ -213,39 +239,25 @@ impl Stream for PyStream { fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { let this = self.as_mut().get_mut(); + let ob = match &mut this.inner { + PyStreamSource::Async(stream) => return Pin::new(stream).poll_next(cx), + PyStreamSource::Sync(ob) => ob, + }; let mut pending = match this.pending.take() { Some(pending) => pending, None => { - let runtime = pyo3_async_runtimes::tokio::get_runtime(); - - // Move GIL acquisition to blocking threads to prevent blocking async runtime. - // This is crucial because holding the GIL in async tasks can block the entire - // async executor and cause deadlocks or performance degradation. - match this.inner { - PyStreamSource::Sync(ref ob) => { - let ob = ob.clone(); - runtime.spawn_blocking(move || { - error::attach(|py| match ob.call_method0(py, intern!(py, "__next__")) { - Ok(ob) => Some(ob.extract(py)), - Err(err) if err.is_instance_of::(py) => None, - Err(err) => Some(Err(err)), - }) - .unwrap_or_else(|err| Some(Err(err.into()))) - }) - } - PyStreamSource::Async(ref stream) => { - let stream = stream.clone(); - runtime.spawn(async move { - let ob = stream.lock().await.next().await; - tokio::task::spawn_blocking(move || { - error::attach(|py| ob.map(|ob| ob.extract(py))) - .unwrap_or_else(|err| Some(Err(err.into()))) - }) - .await - .ok()? - }) - } - } + // Acquiring the interpreter must not block a Tokio worker. + let ob = ob.clone(); + tokio::task::spawn_blocking(move || { + Python::try_attach(|py| match ob.call_method0(py, intern!(py, "__next__")) { + Ok(ob) => Some(ob.extract(py)), + Err(err) if err.is_instance_of::(py) => None, + Err(err) => Some(Err(err)), + }) + // Once Python is unavailable, stop reading without creating + // a PyErr that could require another attachment to format. + .flatten() + }) } }; @@ -259,3 +271,144 @@ impl Stream for PyStream { } } } + +// ===== impl PyAsyncStream ===== + +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 forward = FORWARD.get_or_try_init(py, || { + PyModule::from_code( + py, + c"import asyncio + +async def forward(gen, sender): + try: + try: + async for item in gen: + if not await sender.send(item, False): + return + finally: + close = getattr(gen, 'aclose', None) + if close is not None: + await close() + except asyncio.CancelledError as error: + # Task cancellation must not wait for space in a retained body. + if not asyncio.current_task().cancelling(): + await sender.send(error, True) + raise + except BaseException as error: + await sender.send(error, True) + else: + await sender.finish() +", + c"wreq/_async_stream.py", + c"wreq._async_stream", + )? + .getattr("forward") + .map(Bound::unbind) + })?; + let (tx, rx) = mpsc::channel(1); + let coroutine = forward.bind(py).call1((generator, Sender(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); + } + }; + Ok(Self { + rx, + task: Some((task.unbind(), event_loop.unbind())), + }) + } +} + +impl Stream for PyAsyncStream { + type Item = PyResult; + + 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(_) => { + this.rx.close(); + Poll::Ready(None) + } + Poll::Pending => Poll::Pending, + } + } +} + +impl Drop for PyAsyncStream { + fn drop(&mut self) { + self.rx.close(); + if let Some((task, event_loop)) = self.task.take() { + // Body drop can run on Tokio: acquire the interpreter on a blocking thread. + crate::runtime::get().handle().spawn_blocking(move || { + Python::try_attach(|py| { + if let Ok(cancel) = task.bind(py).getattr(intern!(py, "cancel")) { + let _ = event_loop.call_method1( + py, + intern!(py, "call_soon_threadsafe"), + (cancel,), + ); + } + }); + }); + } + } +} + +// ===== impl Sender ===== + +#[pymethods] +impl Sender { + async fn send( + &self, + item: Py, + error: bool, + #[pyo3(cancel_handle)] cancel: CancelHandle, + ) -> PyResult { + let item = Python::attach(|py| { + if error { + Ok(Err(PyErr::from_value(item.into_bound(py)))) + } else { + item.extract(py).map(Ok) + } + })?; + self.send_item(Some(item), cancel).await + } + + async fn finish(&self, #[pyo3(cancel_handle)] cancel: CancelHandle) -> PyResult { + // Python may retain the sender after completion, especially on PyPy. + self.send_item(None, cancel).await + } +} + +impl Sender { + async fn send_item( + &self, + item: Option>, + mut cancel: CancelHandle, + ) -> PyResult { + let item = match self.0.try_send(item) { + Ok(()) => return Ok(true), + Err(mpsc::error::TrySendError::Closed(_)) => return Ok(false), + Err(mpsc::error::TrySendError::Full(item)) => item, + }; + let tx = self.0.clone(); + // Channel readiness is runtime-independent; keep this on the Python loop. + let mut send = std::pin::pin!(tx.send(item)); + tokio::select! { + biased; + exception = poll_fn(|cx| cancel.poll_cancelled(cx)) => { + Err(Python::attach(|py| PyErr::from_value(exception.into_bound(py)))) + } + result = poll_fn(|cx| nogil::poll_with_guard(send.as_mut(), cx)) => Ok(result.is_ok()), + } + } +} diff --git a/src/client/nogil.rs b/src/client/nogil.rs index db537352..8afe8ff1 100644 --- a/src/client/nogil.rs +++ b/src/client/nogil.rs @@ -1,65 +1,57 @@ use std::{ future::Future, pin::Pin, - sync::{Arc, Once}, + sync::Arc, task::{Context, Poll, Wake, Waker}, }; -use pin_project_lite::pin_project; -use pyo3::{ - coroutine::CancelHandle, - exceptions::{PyRuntimeError, asyncio::CancelledError}, - prelude::*, -}; -use tokio_util::{sync::CancellationToken, task::AbortOnDropHandle}; +use pyo3::{coroutine::CancelHandle, exceptions::PyRuntimeError, prelude::*}; +use tokio_util::task::AbortOnDropHandle; -use crate::error; +use crate::runtime::Runtime; -pin_project! { - /// A future that allows Python threads to run while it is being polled or executed. - /// It also handles cancellation and spawns the task in tokio runtime. - pub struct NoGIL { - #[pin] - handle: AbortOnDropHandle>, - cancel: CancelHandle, - } +/// A future that allows Python threads to run while it is being polled or executed. +/// It also handles cancellation and spawns the task in tokio runtime. +pub struct NoGIL { + handle: AbortOnDropHandle>, + cancel: CancelHandle, } +struct GuardedWaker(Waker); + +// ===== impl NoGIL ===== + impl NoGIL where T: Send + 'static, { - /// Create [`NoGIL`] from a future + /// Spawn internal work without a Python cancellation source. #[inline] - pub fn new(fut: Fut, cancel: CancelHandle) -> Self + pub fn new(runtime: &Runtime, fut: Fut) -> Self where Fut: Future> + Send + 'static, { - Self { - handle: AbortOnDropHandle::new(pyo3_async_runtimes::tokio::get_runtime().spawn(fut)), - cancel, - } + Self::with_cancel(runtime, fut, CancelHandle::new()) } - /// Create [`NoGIL`] from a future and a cancellation token + /// Spawn with Python cancellation, keeping the runtime alive until the task ends. #[inline] - pub fn new_with_token( - fut: Fut, - cancel: CancelHandle, - cancel_token: CancellationToken, - ) -> Self + pub fn with_cancel(runtime: &Runtime, fut: Fut, cancel: CancelHandle) -> Self where Fut: Future> + Send + 'static, { - Self::new( - async move { - tokio::select! { - result = fut => result, - _ = cancel_token.cancelled() => Err(CancelledError::new_err("Operation was cancelled: client has been closed")), - } - }, + let owner = runtime.clone(); + Self { + handle: AbortOnDropHandle::new(Python::attach(|py| { + py.detach(|| { + runtime.handle().spawn(async move { + let _owner = owner; + fut.await + }) + }) + })), cancel, - ) + } } } @@ -71,7 +63,7 @@ where #[inline] fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll { - let this = self.project(); + let this = self.get_mut(); // A Python throw must win even when the Tokio task has already finished. if let Poll::Ready(exc) = this.cancel.poll_cancelled(cx) { this.handle.abort(); @@ -80,11 +72,10 @@ where }))); } - let waker = Waker::from(Arc::new(GuardedWaker(cx.waker().clone()))); + let waker = cx.waker(); Python::attach(|py| { py.detach(|| { - let mut cx = Context::from_waker(&waker); - match this.handle.poll(&mut cx) { + match poll_with_guard(Pin::new(&mut this.handle), &mut Context::from_waker(waker)) { Poll::Ready(Ok(result)) => Poll::Ready(result), Poll::Ready(Err(e)) => Poll::Ready(Err(PyRuntimeError::new_err(e.to_string()))), Poll::Pending => Poll::Pending, @@ -94,12 +85,28 @@ where } } -/// Wakes the Python coroutine from Tokio threads. -/// -/// PyO3's coroutine waker calls `Python::attach`, which panics once the interpreter has -/// shut down. Waking inside [`error::attach`] lets it reuse that attachment instead; when -/// Python is gone the wake is dropped and reported once, as no coroutine is left to resume. -struct GuardedWaker(Waker); +/// Run network work on its selected worker; only wait for completion on the caller. +/// The caller must be detached from Python and outside an async Tokio context. +/// Its runtime borrow keeps the workers alive until the join completes. +pub fn block_on(runtime: &Runtime, future: F) -> PyResult +where + F: Future> + Send + 'static, + T: Send + 'static, +{ + let task = runtime.handle().spawn(future); + runtime + .handle() + .block_on(task) + .map_err(|err| PyRuntimeError::new_err(err.to_string()))? +} + +/// Protect Python wakers retained by futures that can be woken from Rust threads. +pub fn poll_with_guard(future: Pin<&mut F>, cx: &mut Context<'_>) -> Poll { + let waker = Waker::from(Arc::new(GuardedWaker(cx.waker().clone()))); + future.poll(&mut Context::from_waker(&waker)) +} + +// ===== impl GuardedWaker ===== impl Wake for GuardedWaker { fn wake(self: Arc) { @@ -107,9 +114,8 @@ impl Wake for GuardedWaker { } fn wake_by_ref(self: &Arc) { - if let Err(err) = error::attach(|_| self.0.wake_by_ref()) { - static REPORTED: Once = Once::new(); - REPORTED.call_once(|| eprintln!("wreq: failed to wake a Python coroutine: {err}")); - } + // PyO3's nested attach reuses this attachment. If Python is unavailable, + // skip the wake instead of invoking its infallible attachment path. + Python::try_attach(|_| self.0.wake_by_ref()); } } diff --git a/src/client/req.rs b/src/client/req.rs index 9ae6b98f..fa8bc77a 100644 --- a/src/client/req.rs +++ b/src/client/req.rs @@ -5,7 +5,7 @@ use std::{ use futures_util::TryFutureExt; use http::header::COOKIE; -use pyo3::{PyResult, prelude::*, pybacked::PyBackedStr}; +use pyo3::{PyResult, exceptions::asyncio::CancelledError, prelude::*, pybacked::PyBackedStr}; use crate::{ client::{ @@ -293,135 +293,143 @@ pub async fn execute_request( where U: AsRef, { - // Create the request builder. - let mut builder = client.inner.request(method.into_ffi(), url.as_ref()); - - if let Some(mut request) = request { - // Emulation options. - apply_option!(set_if_some, builder, request.emulation, emulation); - - // Version options. - apply_option!( - set_if_some_map, - builder, - request.version, - version, - Version::into_ffi - ); - - // Timeout options. - apply_option!(set_if_some, builder, request.timeout, timeout); - apply_option!(set_if_some, builder, request.read_timeout, read_timeout); - - // Network options. - apply_option!(set_if_some_inner, builder, request.proxy, proxy); - apply_option!(set_if_some, builder, request.local_address, local_address); - apply_option!( - set_if_some_tuple_inner, - builder, - request.local_addresses, - local_addresses - ); - - #[cfg(any( - target_os = "android", - target_os = "fuchsia", - target_os = "illumos", - target_os = "ios", - target_os = "linux", - target_os = "macos", - target_os = "solaris", - target_os = "tvos", - target_os = "visionos", - target_os = "watchos", - ))] - apply_option!(set_if_some, builder, request.interface, interface); - - // Headers options. - apply_option!(set_if_some_inner, builder, request.headers, headers); - apply_option!( - set_if_some_inner, - builder, - request.orig_headers, - orig_headers - ); - apply_option!( - set_if_some, - builder, - request.default_headers, - default_headers - ); - - // Cookies options. - apply_option!( - set_if_some_iter_inner_with_key, - builder, - request.cookies, - header, - COOKIE - ); - apply_option!( - set_if_some_inner, - builder, - request.cookie_provider, - cookie_provider - ); - - // Authentication options. - apply_option!( - set_if_some_map_ref, - builder, - request.auth, - auth, - AsRef::::as_ref - ); - apply_option!(set_if_some, builder, request.bearer_auth, bearer_auth); - apply_option!(set_if_some_tuple, builder, request.basic_auth, basic_auth); - - // Allow redirects options. - apply_option!(set_if_some_inner, builder, request.redirect, redirect); - - // Compression options. - apply_option!(set_if_some, builder, request.gzip, gzip); - apply_option!(set_if_some, builder, request.brotli, brotli); - apply_option!(set_if_some, builder, request.deflate, deflate); - apply_option!(set_if_some, builder, request.zstd, zstd); - - // Query options. - apply_option!(set_if_some_ref, builder, request.query, query); - - // Body options. - apply_option!(set_if_some_ref, builder, request.form, form); - apply_option!(set_if_some_ref, builder, request.json, json); - apply_option!( - set_if_some, - builder, - request.multipart.and_then(|form| form.form), - multipart - ); - apply_option!( - set_if_some_map_try, - builder, - request.body, - body, - wreq::Body::try_from - ); + let future = async { + // Create the request builder. + let mut builder = client.inner.request(method.into_ffi(), url.as_ref()); + + if let Some(mut request) = request { + // Emulation options. + apply_option!(set_if_some, builder, request.emulation, emulation); + + // Version options. + apply_option!( + set_if_some_map, + builder, + request.version, + version, + Version::into_ffi + ); + + // Timeout options. + apply_option!(set_if_some, builder, request.timeout, timeout); + apply_option!(set_if_some, builder, request.read_timeout, read_timeout); + + // Network options. + apply_option!(set_if_some_inner, builder, request.proxy, proxy); + apply_option!(set_if_some, builder, request.local_address, local_address); + apply_option!( + set_if_some_tuple_inner, + builder, + request.local_addresses, + local_addresses + ); + + #[cfg(any( + target_os = "android", + target_os = "fuchsia", + target_os = "illumos", + target_os = "ios", + target_os = "linux", + target_os = "macos", + target_os = "solaris", + target_os = "tvos", + target_os = "visionos", + target_os = "watchos", + ))] + apply_option!(set_if_some, builder, request.interface, interface); + + // Headers options. + apply_option!(set_if_some_inner, builder, request.headers, headers); + apply_option!( + set_if_some_inner, + builder, + request.orig_headers, + orig_headers + ); + apply_option!( + set_if_some, + builder, + request.default_headers, + default_headers + ); + + // Cookies options. + apply_option!( + set_if_some_iter_inner_with_key, + builder, + request.cookies, + header, + COOKIE + ); + apply_option!( + set_if_some_inner, + builder, + request.cookie_provider, + cookie_provider + ); + + // Authentication options. + apply_option!( + set_if_some_map_ref, + builder, + request.auth, + auth, + AsRef::::as_ref + ); + apply_option!(set_if_some, builder, request.bearer_auth, bearer_auth); + apply_option!(set_if_some_tuple, builder, request.basic_auth, basic_auth); + + // Allow redirects options. + apply_option!(set_if_some_inner, builder, request.redirect, redirect); + + // Compression options. + apply_option!(set_if_some, builder, request.gzip, gzip); + apply_option!(set_if_some, builder, request.brotli, brotli); + apply_option!(set_if_some, builder, request.deflate, deflate); + apply_option!(set_if_some, builder, request.zstd, zstd); + + // Query options. + apply_option!(set_if_some_ref, builder, request.query, query); + + // Body options. + apply_option!(set_if_some_ref, builder, request.form, form); + apply_option!(set_if_some_ref, builder, request.json, json); + apply_option!( + set_if_some, + builder, + request.multipart.and_then(|form| form.form), + multipart + ); + apply_option!( + set_if_some_map_try, + builder, + request.body, + body, + wreq::Body::try_from + ); + } + + // Send request. + builder + .send() + .await + .and_then(|r| { + if client.raise_for_status { + r.error_for_status() + } else { + Ok(r) + } + }) + .map(|response| Response::new(response, client.runtime.clone())) + .map_err(Error::Library) + .map_err(Into::into) + }; + + tokio::select! { + biased; + _ = client.cancel.cancelled() => Err(CancelledError::new_err("Operation was cancelled: client has been closed")), + result = future => result, } - - // Send request. - builder - .send() - .await - .and_then(|r| { - if client.raise_for_status { - r.error_for_status() - } else { - Ok(r) - } - }) - .map(Response::new) - .map_err(Error::Library) - .map_err(Into::into) } pub async fn execute_websocket_request( @@ -432,123 +440,131 @@ pub async fn execute_websocket_request( where U: AsRef, { - // Create the WebSocket builder. - let mut builder = client.inner.websocket(url.as_ref()); - - if let Some(mut request) = request { - // Emulation options. - apply_option!(set_if_some, builder, request.emulation, emulation); - - // Version options. - apply_option!( - set_if_some_map, - builder, - request.version, - version, - Version::into_ffi - ); - - // Subprotocols options. - apply_option!(set_if_some, builder, request.protocols, protocols); - - // WebSocket config - apply_option!( - set_if_some, - builder, - request.read_buffer_size, - read_buffer_size - ); - apply_option!( - set_if_some, - builder, - request.write_buffer_size, - write_buffer_size - ); - apply_option!( - set_if_some, - builder, - request.max_write_buffer_size, - max_write_buffer_size - ); - apply_option!(set_if_some, builder, request.max_frame_size, max_frame_size); - apply_option!( - set_if_some, - builder, - request.max_message_size, - max_message_size - ); - apply_option!( - set_if_some, - builder, - request.accept_unmasked_frames, - accept_unmasked_frames - ); - - // Network options. - apply_option!(set_if_some_inner, builder, request.proxy, proxy); - apply_option!(set_if_some, builder, request.local_address, local_address); - apply_option!( - set_if_some_tuple_inner, - builder, - request.local_addresses, - local_addresses - ); - #[cfg(any( - target_os = "android", - target_os = "fuchsia", - target_os = "illumos", - target_os = "ios", - target_os = "linux", - target_os = "macos", - target_os = "solaris", - target_os = "tvos", - target_os = "visionos", - target_os = "watchos", - ))] - apply_option!(set_if_some, builder, request.interface, interface); - - // Headers options. - apply_option!(set_if_some_inner, builder, request.headers, headers); - apply_option!( - set_if_some_inner, - builder, - request.orig_headers, - orig_headers - ); - apply_option!( - set_if_some, - builder, - request.default_headers, - default_headers - ); - apply_option!( - set_if_some_iter_inner_with_key, - builder, - request.cookies, - header, - COOKIE - ); - - // Authentication options. - apply_option!( - set_if_some_map_ref, - builder, - request.auth, - auth, - AsRef::::as_ref - ); - apply_option!(set_if_some, builder, request.bearer_auth, bearer_auth); - apply_option!(set_if_some_tuple, builder, request.basic_auth, basic_auth); - - // Query options. - apply_option!(set_if_some_ref, builder, request.query, query); + let future = async { + // Create the WebSocket builder. + let mut builder = client.inner.websocket(url.as_ref()); + + if let Some(mut request) = request { + // Emulation options. + apply_option!(set_if_some, builder, request.emulation, emulation); + + // Version options. + apply_option!( + set_if_some_map, + builder, + request.version, + version, + Version::into_ffi + ); + + // Subprotocols options. + apply_option!(set_if_some, builder, request.protocols, protocols); + + // WebSocket config + apply_option!( + set_if_some, + builder, + request.read_buffer_size, + read_buffer_size + ); + apply_option!( + set_if_some, + builder, + request.write_buffer_size, + write_buffer_size + ); + apply_option!( + set_if_some, + builder, + request.max_write_buffer_size, + max_write_buffer_size + ); + apply_option!(set_if_some, builder, request.max_frame_size, max_frame_size); + apply_option!( + set_if_some, + builder, + request.max_message_size, + max_message_size + ); + apply_option!( + set_if_some, + builder, + request.accept_unmasked_frames, + accept_unmasked_frames + ); + + // Network options. + apply_option!(set_if_some_inner, builder, request.proxy, proxy); + apply_option!(set_if_some, builder, request.local_address, local_address); + apply_option!( + set_if_some_tuple_inner, + builder, + request.local_addresses, + local_addresses + ); + #[cfg(any( + target_os = "android", + target_os = "fuchsia", + target_os = "illumos", + target_os = "ios", + target_os = "linux", + target_os = "macos", + target_os = "solaris", + target_os = "tvos", + target_os = "visionos", + target_os = "watchos", + ))] + apply_option!(set_if_some, builder, request.interface, interface); + + // Headers options. + apply_option!(set_if_some_inner, builder, request.headers, headers); + apply_option!( + set_if_some_inner, + builder, + request.orig_headers, + orig_headers + ); + apply_option!( + set_if_some, + builder, + request.default_headers, + default_headers + ); + apply_option!( + set_if_some_iter_inner_with_key, + builder, + request.cookies, + header, + COOKIE + ); + + // Authentication options. + apply_option!( + set_if_some_map_ref, + builder, + request.auth, + auth, + AsRef::::as_ref + ); + apply_option!(set_if_some, builder, request.bearer_auth, bearer_auth); + apply_option!(set_if_some_tuple, builder, request.basic_auth, basic_auth); + + // Query options. + apply_option!(set_if_some_ref, builder, request.query, query); + } + + // Send the WebSocket request. + builder + .send() + .and_then(|response| WebSocket::new(response, client.runtime.clone())) + .await + .map_err(Error::Library) + .map_err(Into::into) + }; + + tokio::select! { + biased; + _ = client.cancel.cancelled() => Err(CancelledError::new_err("Operation was cancelled: client has been closed")), + result = future => result, } - - // Send the WebSocket request. - builder - .send() - .and_then(WebSocket::new) - .await - .map_err(Error::Library) - .map_err(Into::into) } diff --git a/src/client/resp/http.rs b/src/client/resp/http.rs index 1831b822..e183ba75 100644 --- a/src/client/resp/http.rs +++ b/src/client/resp/http.rs @@ -24,6 +24,7 @@ use crate::{ header::HeaderMap, http::{StatusCode, Version}, redirect::History, + runtime::Runtime, tls::TlsInfo, }; @@ -33,6 +34,7 @@ pub struct Response { uri: Uri, parts: Parts, body: Arc>, + runtime: Runtime, } /// Represents the state of the HTTP response body. @@ -51,14 +53,19 @@ pub struct BlockingResponse(Response); impl Response { /// Create a new [`Response`] instance. - pub fn new(response: wreq::Response) -> Self { + pub fn new(response: wreq::Response, runtime: Runtime) -> Self { let uri = response.uri().clone(); let response = HttpResponse::from(response) .map(Body::Streamable) .map(ArcSwapOption::from_pointee) .map(Arc::new); let (parts, body) = response.into_parts(); - Response { uri, parts, body } + Response { + uri, + parts, + body, + runtime, + } } /// Builds a [`wreq::Response`] from the current response metadata and the given body. @@ -212,7 +219,7 @@ impl Response { /// Stream read-only memoryviews and any trailing headers from the body. pub fn stream(&self) -> PyResult { self.stream_response() - .map(Streamer::new) + .map(|response| Streamer::new(response, self.runtime.clone())) .map_err(Into::into) } @@ -227,7 +234,7 @@ impl Response { .cache_response() .and_then(|resp| ResponseExt::text(resp, encoding)) .map_err(Into::into); - NoGIL::new(fut, cancel).await + NoGIL::with_cancel(&self.runtime, fut, cancel).await } /// Get the JSON content of the response. @@ -236,7 +243,7 @@ impl Response { .cache_response() .and_then(ResponseExt::json::) .map_err(Into::into); - NoGIL::new(fut, cancel).await + NoGIL::with_cancel(&self.runtime, fut, cancel).await } /// Read the body as a read-only memoryview, retaining its data after the response closes. @@ -246,7 +253,7 @@ impl Response { .and_then(ResponseExt::bytes) .map_ok(PyBuffer::from) .map_err(Into::into); - NoGIL::new(fut, cancel).await + NoGIL::with_cancel(&self.runtime, fut, cancel).await } /// Close the response. @@ -381,7 +388,7 @@ impl BlockingResponse { .cache_response() .and_then(|resp| ResponseExt::text(resp, encoding)) .map_err(Into::into); - pyo3_async_runtimes::tokio::get_runtime().block_on(fut) + crate::client::nogil::block_on(&self.0.runtime, fut) }) } @@ -393,7 +400,7 @@ impl BlockingResponse { .cache_response() .and_then(ResponseExt::json::) .map_err(Into::into); - pyo3_async_runtimes::tokio::get_runtime().block_on(fut) + crate::client::nogil::block_on(&self.0.runtime, fut) }) } @@ -406,7 +413,7 @@ impl BlockingResponse { .and_then(ResponseExt::bytes) .map_ok(PyBuffer::from) .map_err(Into::into); - pyo3_async_runtimes::tokio::get_runtime().block_on(fut) + crate::client::nogil::block_on(&self.0.runtime, fut) }) } diff --git a/src/client/resp/ws.rs b/src/client/resp/ws.rs index 03e5a8c3..e273928c 100644 --- a/src/client/resp/ws.rs +++ b/src/client/resp/ws.rs @@ -18,6 +18,7 @@ use crate::{ extractor::StrInput, header::HeaderMap, http::{StatusCode, Version}, + runtime::Runtime, }; /// A WebSocket response. @@ -44,6 +45,7 @@ pub struct WebSocket { headers: HeaderMap, protocol: Option, cmd: mpsc::UnboundedSender, + runtime: Runtime, } /// A blocking WebSocket response. @@ -54,7 +56,7 @@ pub struct BlockingWebSocket(WebSocket); impl WebSocket { /// Creates a new [`WebSocket`] instance. - pub async fn new(response: WebSocketResponse) -> wreq::Result { + pub async fn new(response: WebSocketResponse, runtime: Runtime) -> wreq::Result { let (version, status, remote_addr, local_addr, headers) = ( Version::from_ffi(response.version()), StatusCode(response.status()), @@ -68,6 +70,7 @@ impl WebSocket { tokio::spawn(cmd::task(websocket, rx)); Ok(WebSocket { + runtime, version, status, remote_addr, @@ -106,7 +109,7 @@ impl WebSocket { timeout: Option, ) -> PyResult> { let tx = self.cmd.clone(); - NoGIL::new(cmd::recv(tx, timeout), cancel).await + NoGIL::with_cancel(&self.runtime, cmd::recv(tx, timeout), cancel).await } /// Send a message to the WebSocket. @@ -117,7 +120,7 @@ impl WebSocket { message: Message, ) -> PyResult<()> { let tx = self.cmd.clone(); - NoGIL::new(cmd::send(tx, message), cancel).await + NoGIL::with_cancel(&self.runtime, cmd::send(tx, message), cancel).await } /// Send multiple messages to the WebSocket. @@ -128,7 +131,7 @@ impl WebSocket { messages: Vec, ) -> PyResult<()> { let tx = self.cmd.clone(); - NoGIL::new(cmd::send_all(tx, messages), cancel).await + NoGIL::with_cancel(&self.runtime, cmd::send_all(tx, messages), cancel).await } /// Close the WebSocket connection. @@ -140,7 +143,7 @@ impl WebSocket { reason: Option, ) -> PyResult<()> { let tx = self.cmd.clone(); - NoGIL::new(cmd::close(tx, code, reason), cancel).await + NoGIL::with_cancel(&self.runtime, cmd::close(tx, code, reason), cancel).await } } @@ -159,7 +162,7 @@ impl WebSocket { _traceback: Py, ) -> PyResult<()> { let tx = self.cmd.clone(); - NoGIL::new(cmd::close(tx, None, None), CancelHandle::new()).await + NoGIL::new(&self.runtime, cmd::close(tx, None, None)).await } } @@ -219,8 +222,7 @@ impl BlockingWebSocket { #[pyo3(signature = (timeout=None))] pub fn recv(&self, py: Python, timeout: Option) -> PyResult> { py.detach(|| { - pyo3_async_runtimes::tokio::get_runtime() - .block_on(cmd::recv(self.0.cmd.clone(), timeout)) + crate::client::nogil::block_on(&self.0.runtime, cmd::recv(self.0.cmd.clone(), timeout)) }) } @@ -228,8 +230,7 @@ impl BlockingWebSocket { #[pyo3(signature = (message))] pub fn send(&self, py: Python, message: Message) -> PyResult<()> { py.detach(|| { - pyo3_async_runtimes::tokio::get_runtime() - .block_on(cmd::send(self.0.cmd.clone(), message)) + crate::client::nogil::block_on(&self.0.runtime, cmd::send(self.0.cmd.clone(), message)) }) } @@ -237,8 +238,10 @@ impl BlockingWebSocket { #[pyo3(signature = (messages))] pub fn send_all(&self, py: Python, messages: Vec) -> PyResult<()> { py.detach(|| { - pyo3_async_runtimes::tokio::get_runtime() - .block_on(cmd::send_all(self.0.cmd.clone(), messages)) + crate::client::nogil::block_on( + &self.0.runtime, + cmd::send_all(self.0.cmd.clone(), messages), + ) }) } @@ -246,11 +249,10 @@ impl BlockingWebSocket { #[pyo3(signature = (code=None, reason=None))] 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(), - code, - reason, - )) + crate::client::nogil::block_on( + &self.0.runtime, + cmd::close(self.0.cmd.clone(), code, reason), + ) }) } } diff --git a/src/dns.rs b/src/dns.rs index 8ce0123d..841ea2ae 100644 --- a/src/dns.rs +++ b/src/dns.rs @@ -2,7 +2,7 @@ use std::{ net::{IpAddr, SocketAddr}, - sync::{Arc, OnceLock}, + sync::Arc, }; use hickory_resolver::{ @@ -71,52 +71,31 @@ impl DnsOptions { } } -// Static resolvers for each IP strategy, lazily initialized -static RESOLVER_IPV4_ONLY: OnceLock = OnceLock::new(); -static RESOLVER_IPV6_ONLY: OnceLock = OnceLock::new(); -static RESOLVER_IPV4_AND_IPV6: OnceLock = OnceLock::new(); -static RESOLVER_IPV6_THEN_IPV4: OnceLock = OnceLock::new(); -static RESOLVER_IPV4_THEN_IPV6: OnceLock = OnceLock::new(); - /// Wrapper around an [`TokioResolver`], which implements the `Resolve` trait. #[derive(Clone)] pub struct HickoryResolver { - /// Shared, lazily-initialized Tokio-based DNS resolver. - resolver: &'static TokioResolver, + // DNS connections must not outlive or cross the client's selected runtime. + resolver: TokioResolver, } impl HickoryResolver { /// Use the system DNS configuration, falling back to Cloudflare if unreadable. - /// Only successfully built resolvers are cached for each IP strategy. pub fn new(strategy: LookupIpStrategy) -> Result { - let cell = match strategy { - LookupIpStrategy::IPV4_ONLY => &RESOLVER_IPV4_ONLY, - LookupIpStrategy::IPV6_ONLY => &RESOLVER_IPV6_ONLY, - LookupIpStrategy::IPV4_AND_IPV6 => &RESOLVER_IPV4_AND_IPV6, - LookupIpStrategy::IPV6_THEN_IPV4 => &RESOLVER_IPV6_THEN_IPV4, - LookupIpStrategy::IPV4_THEN_IPV6 => &RESOLVER_IPV4_THEN_IPV6, - }; - - let resolver = if let Some(resolver) = cell.get() { - resolver - } else { - let mut builder = match TokioResolver::builder_tokio() { - Ok(resolver) => resolver, - Err(err) => { - eprintln!( - "error reading DNS system conf: {}, using Cloudflare DNS", - err - ); - TokioResolver::builder_with_config( - ResolverConfig::udp_and_tcp(&CLOUDFLARE), - TokioRuntimeProvider::default(), - ) - } - }; - builder.options_mut().ip_strategy = strategy.into_ffi(); - let resolver = builder.build().map_err(Error::Dns)?; - cell.get_or_init(|| resolver) + let mut builder = match TokioResolver::builder_tokio() { + Ok(resolver) => resolver, + Err(err) => { + eprintln!( + "error reading DNS system conf: {}, using Cloudflare DNS", + err + ); + TokioResolver::builder_with_config( + ResolverConfig::udp_and_tcp(&CLOUDFLARE), + TokioRuntimeProvider::default(), + ) + } }; + builder.options_mut().ip_strategy = strategy.into_ffi(); + let resolver = builder.build().map_err(Error::Dns)?; Ok(Self { resolver }) } } diff --git a/src/error.rs b/src/error.rs index 9899b00f..7a78a8f3 100644 --- a/src/error.rs +++ b/src/error.rs @@ -1,11 +1,5 @@ -use std::{ - any::Any, - fmt, - panic::{self, AssertUnwindSafe}, -}; - use pyo3::{ - PyErr, Python, create_exception, + PyErr, create_exception, exceptions::{PyException, PyRuntimeError, PyStopAsyncIteration, PyStopIteration}, }; use wreq::header; @@ -63,9 +57,6 @@ macro_rules! wrap_error { }; } -const INTERPRETER_UNAVAILABLE_MSG: &str = - "The Python interpreter is not available (not initialized or shutting down)"; - /// Unified error enum #[derive(Debug)] pub enum Error { @@ -83,8 +74,6 @@ pub enum Error { Json(serde_json::Error), Form(serde_urlencoded::ser::Error), Library(wreq::Error), - InterpreterUnavailable, - Panic(Box), } impl From for PyErr { @@ -111,8 +100,6 @@ impl From for PyErr { Error::Dns(err) => BuilderError::new_err(format!("DNS resolver error: {err:?}")), Error::Json(err) => PyRuntimeError::new_err(format!("JSON error: {err:?}")), Error::Form(err) => PyRuntimeError::new_err(format!("Form error: {err:?}")), - Error::InterpreterUnavailable => PyRuntimeError::new_err(INTERPRETER_UNAVAILABLE_MSG), - Error::Panic(msg) => RustPanic::new_err(msg.into_string()), Error::Library(err) => wrap_error!(err, is_body => BodyError, is_tls => TlsError, @@ -131,41 +118,6 @@ impl From for PyErr { } } -impl fmt::Display for Error { - fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { - match self { - Error::InterpreterUnavailable => f.write_str(INTERPRETER_UNAVAILABLE_MSG), - Error::Panic(msg) => f.write_str(msg), - err => fmt::Debug::fmt(err, f), - } - } -} - -impl Error { - fn panic(payload: Box) -> Self { - let msg = payload - .downcast_ref::<&str>() - .copied() - .or_else(|| payload.downcast_ref::().map(String::as_str)) - .unwrap_or("unknown panic"); - Error::Panic(format!("Rust panic: {msg}").into_boxed_str()) - } -} - -/// Attaches to Python from a non-Python thread. -/// -/// `Python::attach` panics once the interpreter has shut down (see discussions/305), so -/// this uses `Python::try_attach` and returns [`Error::InterpreterUnavailable`] instead. -/// Any other panic in `f` is caught as [`Error::Panic`] to keep Tokio threads alive. -pub fn attach(f: F) -> Result -where - F: for<'py> FnOnce(Python<'py>) -> R, -{ - panic::catch_unwind(AssertUnwindSafe(|| Python::try_attach(f))) - .map_err(Error::panic)? - .ok_or(Error::InterpreterUnavailable) -} - impl From for Error { fn from(err: header::InvalidHeaderName) -> Self { Error::InvalidHeaderName(err) diff --git a/src/lib.rs b/src/lib.rs index 8f3b3cfc..183a9a1e 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -18,6 +18,7 @@ mod http1; mod http2; mod proxy; mod redirect; +mod runtime; mod tls; use client::{ @@ -47,6 +48,7 @@ use pyo3::{ coroutine::CancelHandle, intern, prelude::*, pybacked::PyBackedStr, types::PyDict, wrap_pymodule, }; +use runtime::Runtime; #[cfg(feature = "jemalloc")] use tikv_jemallocator as _; use tls::{ @@ -339,6 +341,7 @@ fn wreq(py: Python, m: &Bound<'_, PyModule>) -> PyResult<()> { m.add_class::()?; m.add_class::()?; m.add_class::()?; + m.add_class::()?; m.add_class::()?; m.add_class::()?; m.add_class::()?; @@ -357,6 +360,7 @@ fn wreq(py: Python, m: &Bound<'_, PyModule>) -> PyResult<()> { m.add_function(wrap_pyfunction!(websocket, m)?)?; m.add_wrapped(wrap_pymodule!(proxy_module))?; + m.add_wrapped(wrap_pymodule!(runtime_module))?; m.add_wrapped(wrap_pymodule!(dns_module))?; m.add_wrapped(wrap_pymodule!(http1_module))?; m.add_wrapped(wrap_pymodule!(http2_module))?; @@ -371,6 +375,10 @@ fn wreq(py: Python, m: &Bound<'_, PyModule>) -> PyResult<()> { let sys = PyModule::import(py, intern!(py, "sys"))?; let sys_modules: Bound<'_, PyDict> = sys.getattr(intern!(py, "modules"))?.cast_into()?; sys_modules.set_item(intern!(py, "wreq.proxy"), m.getattr(intern!(py, "proxy"))?)?; + sys_modules.set_item( + intern!(py, "wreq.runtime"), + m.getattr(intern!(py, "runtime"))?, + )?; sys_modules.set_item(intern!(py, "wreq.dns"), m.getattr(intern!(py, "dns"))?)?; sys_modules.set_item(intern!(py, "wreq.http1"), m.getattr(intern!(py, "http1"))?)?; sys_modules.set_item(intern!(py, "wreq.http2"), m.getattr(intern!(py, "http2"))?)?; @@ -408,6 +416,12 @@ fn proxy_module(m: &Bound<'_, PyModule>) -> PyResult<()> { Ok(()) } +#[pymodule(gil_used = false, name = "runtime")] +fn runtime_module(m: &Bound<'_, PyModule>) -> PyResult<()> { + m.add_class::()?; + Ok(()) +} + #[pymodule(gil_used = false, name = "dns")] fn dns_module(m: &Bound<'_, PyModule>) -> PyResult<()> { m.add_class::()?; diff --git a/src/redirect.rs b/src/redirect.rs index 2ee1bebe..97b0191a 100644 --- a/src/redirect.rs +++ b/src/redirect.rs @@ -2,7 +2,7 @@ use std::{fmt::Display, sync::Arc}; use pyo3::prelude::*; -use crate::{error, header::HeaderMap, http::StatusCode}; +use crate::{header::HeaderMap, http::StatusCode}; /// Represents the redirect policy for HTTP requests. #[derive(Clone)] @@ -103,14 +103,16 @@ impl Policy { attempt.pending(|attempt| async move { let args = Attempt::from(&attempt); let kind = tokio::task::spawn_blocking(move || { - error::attach(|py| { + Python::try_attach(|py| { callback .call1(py, (args,)) .and_then(|result| result.extract::(py).map_err(PyErr::from)) .map(|action| action.kind) .unwrap_or_else(|err| ActionKind::Error(err.to_string())) }) - .unwrap_or_else(|err| ActionKind::Error(err.to_string())) + .unwrap_or_else(|| { + ActionKind::Error("The Python interpreter is not available".into()) + }) }) .await; diff --git a/src/runtime.rs b/src/runtime.rs new file mode 100644 index 00000000..59db0a35 --- /dev/null +++ b/src/runtime.rs @@ -0,0 +1,220 @@ +use std::{ + sync::{Arc, OnceLock}, + time::Duration, +}; + +use pingora_runtime::{BlockingPoolOpts, Runtime as PingoraRuntime, RuntimeBuilder}; +use pyo3::{ + exceptions::{PyRuntimeError, PyValueError}, + prelude::*, +}; +use tokio::runtime::Handle; + +/// Shared Tokio runtime, released after its clients and active work are dropped. +#[derive(Clone)] +#[pyclass(frozen, skip_from_py_object)] +pub struct Runtime { + inner: Option>, + handle: Handle, +} + +impl Runtime { + /// Borrow the selected worker's handle without changing the selection. + pub fn handle(&self) -> &Handle { + &self.handle + } + + /// Share the runtime and select a worker for a new client. + pub fn select(&self) -> PyResult { + let inner = self + .inner + .as_ref() + .ok_or_else(|| PyRuntimeError::new_err("Runtime is unavailable"))?; + Ok(Self { + inner: Some(inner.clone()), + handle: inner.get_handle().clone(), + }) + } +} + +impl From for Runtime { + fn from(runtime: PingoraRuntime) -> Self { + // No-steal workers are lazy upstream. Start them before sharing the runtime. + let handle = runtime.get_handle().clone(); + Self { + inner: Some(Arc::new(runtime)), + handle, + } + } +} + +impl FromPyObject<'_, '_> for Runtime { + type Error = PyErr; + + fn extract(value: Borrowed<'_, '_, PyAny>) -> PyResult { + Ok(value.extract::>()?.clone()) + } +} + +#[pymethods] +impl Runtime { + /// Create and start the runtime's workers. + /// Without work stealing, each client stays on one worker (not CPU-pinned). + /// Workers default to CPU parallelism and names to the package name. + /// Thread counts must be positive; thread_keep_alive is a nonnegative timedelta. + #[new] + #[pyo3(signature = ( + *, + workers = None, + work_steal = true, + thread_name = None, + max_blocking_threads = None, + thread_keep_alive = None, + ))] + fn new( + py: Python<'_>, + workers: Option, + work_steal: bool, + thread_name: Option<&str>, + max_blocking_threads: Option, + thread_keep_alive: Option, + ) -> PyResult { + let workers = + workers.unwrap_or_else(|| std::thread::available_parallelism().map_or(1, usize::from)); + if workers == 0 + || max_blocking_threads == Some(0) + || workers + .checked_add(max_blocking_threads.unwrap_or(512)) + .is_none() + { + return Err(PyValueError::new_err("Invalid runtime thread counts")); + } + + let thread_name = thread_name.unwrap_or(env!("CARGO_PKG_NAME")); + if thread_name.contains('\0') { + return Err(PyValueError::new_err("thread_name must not contain NUL")); + } + + Ok(py.detach(|| { + RuntimeBuilder::new(workers, thread_name) + .work_steal(work_steal) + .blocking_pool_opts(BlockingPoolOpts { + max_threads: max_blocking_threads, + thread_keep_alive, + }) + .build() + .into() + })) + } +} + +impl Drop for Runtime { + fn drop(&mut self) { + // The final owner may be released on a worker or while holding the GIL. + if let Some(runtime) = self.inner.take().and_then(Arc::into_inner) { + match runtime { + PingoraRuntime::Steal { runtime, .. } => runtime.shutdown_background(), + PingoraRuntime::NoSteal(runtime) => drop(runtime), + } + } + } +} + +/// Create the shared runtime on first use and retain it for the process lifetime. +pub fn get() -> &'static Runtime { + static RUNTIME: OnceLock = OnceLock::new(); + + fn create() -> Runtime { + let workers = std::thread::available_parallelism().map_or(1, usize::from); + let runtime = RuntimeBuilder::new(workers, env!("CARGO_PKG_NAME")).build(); + runtime.into() + } + + if let Some(runtime) = RUNTIME.get() { + return runtime; + } + + // Never wait for another initializer while holding the interpreter. + Python::try_attach(|py| py.detach(|| RUNTIME.get_or_init(create))) + .unwrap_or_else(|| RUNTIME.get_or_init(create)) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn fixed_worker() { + let runtime = Runtime::from( + RuntimeBuilder::new(2, "wreq-affinity") + .work_steal(false) + .build(), + ); + let first = runtime.select().unwrap(); + let second = runtime.select().unwrap(); + let caller = std::thread::current().id(); + let mut ids = Vec::new(); + for worker in [&first, &second, &first.clone(), &second.clone()] { + let id = crate::client::nogil::block_on(worker, async { + let thread = std::thread::current().id(); + for _ in 0..8 { + tokio::task::yield_now().await; + assert_eq!(thread, std::thread::current().id()); + } + let child = tokio::spawn(async { std::thread::current().id() }) + .await + .unwrap(); + assert_eq!(thread, child); + Ok(thread) + }) + .unwrap(); + assert_ne!(id, caller); + ids.push(id); + } + assert_eq!(ids[0], ids[2]); + assert_eq!(ids[1], ids[3]); + // Different clients may select the same worker; each selection stays fixed. + } + + #[test] + fn last_owner_can_be_released_on_a_worker() { + struct NotifyOnDrop(std::sync::mpsc::Sender<()>); + + impl Drop for NotifyOnDrop { + fn drop(&mut self) { + let _ = self.0.send(()); + } + } + + for steal in [false, true] { + let runtime = Runtime::from( + RuntimeBuilder::new(1, "wreq-drop") + .work_steal(steal) + .build(), + ); + let weak = Arc::downgrade(runtime.inner.as_ref().unwrap()); + let handle = runtime.handle().clone(); + let (dropped, released) = std::sync::mpsc::channel(); + let guard = NotifyOnDrop(dropped); + let background = handle.spawn(async move { + let _guard = guard; + std::future::pending::<()>().await; + }); + let (tx, rx) = tokio::sync::oneshot::channel(); + let (done, wait) = std::sync::mpsc::channel(); + let owner = runtime.clone(); + let task = handle.spawn(async move { + let _owner = owner; + let _ = rx.await; + done.send(()).unwrap(); + }); + drop(runtime); + tx.send(()).unwrap(); + wait.recv_timeout(Duration::from_secs(5)).unwrap(); + released.recv_timeout(Duration::from_secs(5)).unwrap(); + assert!(weak.upgrade().is_none()); + handle.block_on(task).unwrap(); + drop(background); + } + } +} diff --git a/tests/cancellation_test.py b/tests/cancellation_test.py index 2bc55669..f409317f 100644 --- a/tests/cancellation_test.py +++ b/tests/cancellation_test.py @@ -1,5 +1,6 @@ import asyncio import gc +import sys import weakref from contextlib import asynccontextmanager @@ -12,6 +13,116 @@ class Cancellation(asyncio.CancelledError): pass +@pytest.mark.skipif( + sys.implementation.name != "pypy", reason="PyPy legacy throw protocol" +) +def test_legacy_coroutine_throw(): + from wreq._compat import _install + + try: + raise RuntimeError("traceback origin") + except RuntimeError as error: + origin = error.__traceback__ + + def contains(traceback): + while traceback is not None: + if traceback is origin: + return True + traceback = traceback.tb_next + return False + + def invoke(*args, **kwargs): + coroutine = wreq.get("") + try: + coroutine.throw(*args, **kwargs) + except BaseException as error: + return error + finally: + coroutine.close() + pytest.fail("throw did not raise") + + class ExceptionValue(ValueError): + def with_traceback(self, *_): + raise AssertionError("overridden method must not run") + + for instance_first in (False, True): + for traceback in (None, origin): + error = ExceptionValue("identity") + BaseException.with_traceback(error, origin) + args = ( + (error, None, traceback) + if instance_first + else (ExceptionValue, error, traceback) + ) + caught = invoke(*args) + assert caught is error + assert caught.args == ("identity",) + assert contains(caught.__traceback__) is ( + instance_first or traceback is not None + ) + + assert invoke(ValueError, ("tuple", 7)).args == ("tuple", 7) + error = ValueError("keyword") + assert invoke(exc=error) is error + + class PretendException: + @property + def __class__(self): + return ValueError + + value = PretendException() + assert invoke(ValueError, value).args == (value,) + + class ExceptionMeta(type): + def __subclasscheck__(cls, subclass): + return True + + class CustomException(Exception, metaclass=ExceptionMeta): + pass + + assert type(invoke(CustomException, value)) is CustomException + assert invoke(CustomException, value).args == (value,) + + class HiddenTraceback(RuntimeError): + def __getattribute__(self, name): + if name == "__traceback__": + return None + return super().__getattribute__(name) + + class BadConstructor(Exception): + def __new__(cls): + raise HiddenTraceback("constructor failure") + + class NotAnException(Exception): + def __new__(cls): + return PretendException() + + caught = invoke(BadConstructor, None, origin) + assert type(caught) is HiddenTraceback + assert caught.args == ("constructor failure",) + assert not contains(BaseException.__traceback__.__get__(caught)) + assert type(invoke(NotAnException)) is TypeError + + coroutine = wreq.get("") + try: + installed = vars(type(coroutine))["throw"] + _install(type(coroutine)) + assert vars(type(coroutine))["throw"] is installed + for args in ( + (object(),), + (ValueError, None, object()), + (error, "value"), + (ValueError,) * 4, + ): + with pytest.raises(TypeError): + coroutine.throw(*args) + with pytest.raises(ValueError) as caught: + coroutine.throw(error) + assert caught.value is error + finally: + coroutine.close() + + @asynccontextmanager async def local_server(): connections = asyncio.Queue() @@ -35,7 +146,7 @@ async def accept(reader, writer): @pytest.mark.asyncio -@pytest.mark.parametrize("operation", ["request", "request_error", "bytes", "text", "json"]) +@pytest.mark.parametrize("operation", ["request", "request_error", "stream"]) async def test_cancellation_after_rust_completion(operation): async with local_server() as (url, connections), wreq.Client(proxies=[]) as client: response = None @@ -49,7 +160,7 @@ async def test_cancellation_after_rust_completion(operation): writer.write(b"HTTP/1.1 200 OK\r\nContent-Length: 2\r\n\r\n") await writer.drain() response = await asyncio.wait_for(task, 5) - coroutine = getattr(response, operation)() + coroutine = anext(response.stream()) waiter = coroutine.send(None) try: @@ -81,7 +192,7 @@ async def test_cancellation_after_rust_completion(operation): @pytest.mark.asyncio -@pytest.mark.parametrize("action", ["cancel", "close_coroutine", "close_client"]) +@pytest.mark.parametrize("action", ["cancel", "close_coroutine"]) async def test_pending_request_cancellation(action): async with local_server() as (url, connections), wreq.Client(proxies=[]) as client: coroutine = client.get(url) @@ -94,20 +205,89 @@ async def test_pending_request_cancellation(action): if action == "close_coroutine": coroutine.close() else: - if action == "cancel": - task.cancel("caller cancellation message") - else: - client.close() + task.cancel("caller cancellation message") done, _ = await asyncio.wait({task}, timeout=5) assert task in done, "Cancellation did not finish" with pytest.raises(asyncio.CancelledError) as caught: await task - expected = ( - "caller cancellation message" - if action == "cancel" - else "Operation was cancelled: client has been closed" - ) - assert caught.value.args == (expected,) + assert caught.value.args == ("caller cancellation message",) # The cancelled operation must release its pending network request. assert await asyncio.wait_for(reader.read(), 5) == b"" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("action", ["cancel", "close_coroutine"]) +async def test_pending_stream_cancellation(action): + async with local_server() as (url, connections), wreq.Client(proxies=[]) as client: + task = asyncio.create_task(client.get(url)) + _, writer = await asyncio.wait_for(connections.get(), 5) + writer.write(b"HTTP/1.1 200 OK\r\nContent-Length: 2\r\n\r\n") + await writer.drain() + response = await asyncio.wait_for(task, 5) + stream = response.stream() + coroutine = anext(stream) + + if action == "close_coroutine": + assert isinstance(coroutine.send(None), asyncio.Future) + coroutine.close() + else: + started = asyncio.Event() + + async def read(): + started.set() + return await coroutine + + task = asyncio.create_task(read()) + await started.wait() + task.cancel("cancel stream read") + with pytest.raises(asyncio.CancelledError, match="cancel stream read"): + await asyncio.wait_for(task, 5) + + # Closing the stream must acquire the lock held by the pending read. + # Do not send a body: that would let a leaked read release it naturally. + await asyncio.wait_for(stream.__aexit__(None, None, None), 5) + with pytest.raises(StopAsyncIteration): + await anext(stream) + await response.close() + + +@pytest.mark.asyncio +async def test_stream_coroutine_iteration(): + async with local_server() as (url, connections), wreq.Client(proxies=[]) as client: + task = asyncio.create_task(client.get(url)) + _, writer = await asyncio.wait_for(connections.get(), 5) + writer.write( + b"HTTP/1.1 200 OK\r\nTransfer-Encoding: chunked\r\n" + b"Trailer: x-check\r\n\r\n1\r\na\r\n" + ) + await writer.drain() + response = await asyncio.wait_for(task, 5) + async with response.stream() as stream: + deferred = stream.__anext__() + assert asyncio.iscoroutine(deferred) + assert not isinstance(deferred, asyncio.Future) + assert deferred.__qualname__ == "Streamer.__anext__" + assert not hasattr(stream, "_anext") + try: + # An unawaited __anext__ must not consume the first frame. + assert await asyncio.wait_for(anext(stream), 5) == b"a" + writer.write(b"1\r\nb\r\n0\r\nx-check: done\r\n\r\n") + await writer.drain() + assert await asyncio.wait_for(deferred, 5) == b"b" + with pytest.raises( + RuntimeError, match="cannot reuse already awaited coroutine" + ): + await deferred + finally: + deferred.close() + + async with asyncio.timeout(5): + frames = [frame async for frame in stream] + assert len(frames) == 1 + assert isinstance(frames[0], wreq.HeaderMap) + assert frames[0]["x-check"] == b"done" + with pytest.raises(StopAsyncIteration): + await stream.__anext__() + assert await anext(stream, None) is None + await response.close() diff --git a/tests/dns_test.py b/tests/dns_test.py index 8d586948..c95c34f9 100644 --- a/tests/dns_test.py +++ b/tests/dns_test.py @@ -45,6 +45,7 @@ async def fetch(host, options): kwargs["dns_options"] = options if blocking_api: + def request(): with blocking.Client(**kwargs) as client: with client.get(url) as response: diff --git a/tests/redirect_test.py b/tests/redirect_test.py index 0f5f9703..a7afe7d5 100644 --- a/tests/redirect_test.py +++ b/tests/redirect_test.py @@ -1,10 +1,44 @@ +import asyncio + import pytest import wreq from wreq import redirect +from cancellation_test import local_server + client = wreq.Client(redirect=redirect.Policy.limited(10)) +@pytest.mark.asyncio +@pytest.mark.parametrize("fail", [False, True]) +async def test_custom_redirect_callback(fail): + def callback(attempt): + if fail: + raise ValueError("redirect callback failed") + return attempt.stop() + + async with ( + local_server() as (url, connections), + wreq.Client(proxies=[], redirect=redirect.Policy.custom(callback)) as client, + ): + task = asyncio.create_task(client.get(url)) + _, writer = await asyncio.wait_for(connections.get(), 5) + writer.write( + b"HTTP/1.1 302 Found\r\nLocation: /next\r\nContent-Length: 0\r\n\r\n" + ) + await writer.drain() + if fail: + with pytest.raises( + wreq.exceptions.RequestError, + match="ValueError: redirect callback failed", + ): + await asyncio.wait_for(task, 5) + else: + response = await asyncio.wait_for(task, 5) + assert response.status.is_redirection() + await response.close() + + @pytest.mark.asyncio @pytest.mark.flaky(reruns=3, reruns_delay=2) async def test_request_disable_redirect(): diff --git a/tests/runtime_test.py b/tests/runtime_test.py new file mode 100644 index 00000000..178e685a --- /dev/null +++ b/tests/runtime_test.py @@ -0,0 +1,332 @@ +import asyncio +import base64 +import hashlib +from datetime import timedelta + +import pytest +import wreq +from wreq.runtime import Runtime + +from cancellation_test import local_server +from upload_test import read_chunked + + +def test_runtime_configuration(): + assert Runtime is wreq.Runtime + assert Runtime is wreq.runtime.Runtime + assert "Runtime" in wreq.__all__ + for kwargs in ( + {"workers": 0}, + {"max_blocking_threads": 0}, + {"thread_name": "bad\0name"}, + {"thread_keep_alive": timedelta(microseconds=-1)}, + ): + with pytest.raises(ValueError): + Runtime(**kwargs) + with pytest.raises(TypeError): + Runtime(thread_keep_alive=0.25) + with pytest.raises(TypeError): + wreq.Client(runtime=object()) + for duration in (None, timedelta(), timedelta(microseconds=250001)): + runtime = Runtime( + workers=1, + work_steal=False, + thread_name=None if duration is None else "isolated", + max_blocking_threads=3, + thread_keep_alive=duration, + ) + for factory in (wreq.Client, wreq.blocking.Client): + client = factory(runtime=runtime) + alias = client.runtime + with pytest.raises(AttributeError): + client.runtime = runtime + client.close() + del client + # Releasing a client does not close a shared runtime. + other = factory(runtime=alias) + other.close() + client = factory(runtime=None) + assert isinstance(client.runtime, Runtime) + client.close() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("steal", [False, True]) +async def test_response_and_stream_keep_runtime_alive(steal): + runtime = wreq.Runtime(workers=1, work_steal=steal) + async with local_server() as (url, connections): + client = wreq.Client(runtime=runtime, proxies=[]) + task = asyncio.create_task(client.get(url)) + _, writer = await asyncio.wait_for(connections.get(), 5) + writer.write(b"HTTP/1.1 200 OK\r\nContent-Length: 4\r\n\r\n") + await writer.drain() + response = await asyncio.wait_for(task, 5) + del task + client.close() + del client, runtime + stream = response.stream() + del response + writer.write(b"body") + await writer.drain() + assert await asyncio.wait_for(anext(stream), 5) == b"body" + with pytest.raises(StopAsyncIteration): + await anext(stream) + del stream + + +@pytest.mark.asyncio +async def test_shared_runtime_cancellation_and_upload(): + runtime = wreq.Runtime(workers=2, work_steal=False, max_blocking_threads=2) + async with local_server() as (url, connections): + first = wreq.Client(runtime=runtime, proxies=[]) + second = wreq.Client(runtime=runtime, proxies=[]) + del runtime + pending = asyncio.create_task(first.get(url)) + await asyncio.wait_for(connections.get(), 5) + first.close() + with pytest.raises(asyncio.CancelledError): + await asyncio.wait_for(pending, 5) + with pytest.raises(asyncio.CancelledError): + await first.get("invalid URL") + del pending, first + finalized = asyncio.Event() + + async def chunks(): + try: + yield b"custom " + await asyncio.sleep(0) + yield b"runtime" + finally: + finalized.set() + + task = asyncio.create_task(second.post(url, body=chunks())) + reader, writer = await asyncio.wait_for(connections.get(), 5) + assert await asyncio.wait_for(read_chunked(reader), 5) == b"custom runtime" + writer.write(b"HTTP/1.1 200 OK\r\nContent-Length: 2\r\n\r\n{}") + await writer.drain() + response = await asyncio.wait_for(task, 5) + assert await response.json() == {} + await asyncio.wait_for(finalized.wait(), 5) + await response.close() + second.close() + del response, task, second + + +@pytest.mark.asyncio +@pytest.mark.parametrize("steal", [None, False, True]) +async def test_blocking_client_runtime(steal): + runtime = ( + None + if steal is None + else wreq.Runtime(workers=1, work_steal=steal, max_blocking_threads=2) + ) + + def request(url): + with wreq.blocking.Client(runtime=runtime, proxies=[]) as client: + assert isinstance(client.runtime, Runtime) + with client.post(url, body=iter((b"sync", b" upload"))) as response: + assert response.bytes() == b"{}" + assert response.json() == {} + with client.post(url, body=iter((b"sync", b" upload"))) as response: + with response.stream() as stream: + return b"".join(stream) + + async with local_server() as (url, connections): + task = asyncio.create_task(asyncio.to_thread(request, url)) + for _ in range(2): + reader, writer = await asyncio.wait_for(connections.get(), 5) + assert await asyncio.wait_for(read_chunked(reader), 5) == b"sync upload" + writer.write(b"HTTP/1.1 200 OK\r\nContent-Length: 2\r\n\r\n{}") + await writer.drain() + assert await asyncio.wait_for(task, 5) == b"{}" + del task + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("blocking", "operation"), + [(True, "get"), (False, "websocket"), (True, "websocket")], +) +async def test_client_close_cancels_requests(blocking, operation): + factory = wreq.blocking.Client if blocking else wreq.Client + client = factory( + runtime=Runtime(workers=1, work_steal=False), + proxies=[], + timeout=timedelta(seconds=10), + ) + task = None + try: + async with local_server() as (url, connections): + if operation == "websocket": + url = url.replace("http://", "ws://", 1) + request = getattr(client, operation) + + def start(): + return asyncio.create_task( + asyncio.to_thread(request, url) if blocking else request(url) + ) + + task = start() + reader, _ = await asyncio.wait_for(connections.get(), 5) + client.close() + done, _ = await asyncio.wait({task}, timeout=5) + assert task in done, "close() did not cancel the pending request" + with pytest.raises(asyncio.CancelledError) as caught: + await task + assert caught.value.args == ( + "Operation was cancelled: client has been closed", + ) + assert await asyncio.wait_for(reader.read(), 5) == b"" + + task = start() + done, _ = await asyncio.wait({task}, timeout=5) + assert task in done, "a closed client started another request" + with pytest.raises(asyncio.CancelledError): + await task + assert connections.empty() + finally: + client.close() + if task is not None: + await asyncio.gather(task, return_exceptions=True) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("blocking", [False, True]) +async def test_websocket_outlives_client(blocking): + connections = asyncio.Queue() + + async def accept(reader, writer): + header = await reader.readuntil(b"\r\n\r\n") + key = next( + line.split(b":", 1)[1].strip() + for line in header.split(b"\r\n") + if line.lower().startswith(b"sec-websocket-key:") + ) + digest = base64.b64encode( + hashlib.sha1(key + b"258EAFA5-E914-47DA-95CA-C5AB0DC85B11").digest() + ) + writer.write( + b"HTTP/1.1 101 Switching Protocols\r\nUpgrade: websocket\r\n" + b"Connection: Upgrade\r\nSec-WebSocket-Accept: " + digest + b"\r\n\r\n" + ) + await writer.drain() + connections.put_nowait((reader, writer)) + + server = await asyncio.start_server(accept, "127.0.0.1", 0) + runtime = wreq.Runtime(workers=1, work_steal=False) + writer = None + try: + url = f"ws://127.0.0.1:{server.sockets[0].getsockname()[1]}/" + client = (wreq.blocking.Client if blocking else wreq.Client)( + runtime=runtime, proxies=[] + ) + ws = ( + await asyncio.to_thread(client.websocket, url) + if blocking + else await client.websocket(url) + ) + reader, writer = await asyncio.wait_for(connections.get(), 5) + client.close() + del client, runtime + writer.write(b"\x81\x04pong") + await writer.drain() + message = await asyncio.to_thread(ws.recv) if blocking else await ws.recv() + assert message.text == "pong" + outgoing = wreq.Message.from_text("ping") + if blocking: + await asyncio.to_thread(ws.send, outgoing) + else: + await ws.send(outgoing) + frame = await asyncio.wait_for(reader.readexactly(10), 5) + assert frame[:2] == b"\x81\x84" + assert ( + bytes(byte ^ frame[2 + i % 4] for i, byte in enumerate(frame[6:])) + == b"ping" + ) + if blocking: + await asyncio.to_thread(ws.close) + else: + await ws.close() + del ws + finally: + if writer is not None: + writer.close() + await writer.wait_closed() + server.close() + await server.wait_closed() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("steal", [False, True]) +async def test_http2_multiplexing_on_custom_runtime(steal): + # Minimal h2c responder: indexed :status=200, then a two-byte DATA frame. + # No HPACK decoder is needed because the test does not inspect request headers. + def frame(kind, flags, stream, payload=b""): + return ( + len(payload).to_bytes(3, "big") + + bytes((kind, flags)) + + stream.to_bytes(4, "big") + + payload + ) + + connections = [] + handlers = set() + errors = [] + + async def accept(reader, writer): + handlers.add(asyncio.current_task()) + connections.append(writer) + try: + assert await reader.readexactly(24) == b"PRI * HTTP/2.0\r\n\r\nSM\r\n\r\n" + writer.write(frame(4, 0, 0)) + await writer.drain() + while True: + header = await reader.readexactly(9) + payload = await reader.readexactly(int.from_bytes(header[:3], "big")) + kind, flags = header[3:5] + stream = int.from_bytes(header[5:], "big") & 0x7FFFFFFF + if kind == 4 and not flags & 1: + writer.write(frame(4, 1, 0)) + elif kind == 6 and not flags & 1: + writer.write(frame(6, 1, 0, payload)) + elif kind == 1: + assert flags & 4 # These small requests fit in one header block. + writer.write( + frame(1, 4, stream, b"\x88") + frame(0, 1, stream, b"ok") + ) + await writer.drain() + except (asyncio.IncompleteReadError, ConnectionError): + pass + except Exception as error: + errors.append(error) + finally: + handlers.discard(asyncio.current_task()) + + runtime = wreq.Runtime(workers=2, work_steal=steal) + server = await asyncio.start_server(accept, "127.0.0.1", 0) + try: + url = f"http://127.0.0.1:{server.sockets[0].getsockname()[1]}/" + client = wreq.Client(runtime=runtime, http2_only=True, proxies=[]) + + async def request(): + response = await client.get(url) + assert response.version == wreq.Version.HTTP_2 + # Consume without close(), which intentionally forbids connection reuse. + return await response.bytes() + + assert await asyncio.wait_for(request(), 5) == b"ok" + assert ( + await asyncio.wait_for(asyncio.gather(*(request() for _ in range(16))), 5) + == [b"ok"] * 16 + ) + assert len(connections) == 1 + assert not errors + client.close() + del client, runtime + finally: + server.close() + for writer in connections: + writer.close() + await asyncio.gather(*(writer.wait_closed() for writer in connections)) + await asyncio.gather(*handlers) + await server.wait_closed() diff --git a/tests/shutdown_test.py b/tests/shutdown_test.py index 8aba518c..6a15f5e0 100644 --- a/tests/shutdown_test.py +++ b/tests/shutdown_test.py @@ -4,63 +4,140 @@ import pytest -# https://github.com/0x676e67/wreq-python/discussions/305: a Tokio worker that attaches -# to Python during interpreter shutdown used to hit PyO3's not-initialized panic. wreq -# now detects the unavailable interpreter and reports a wreq error without panicking. - -# A request is left pending on a listener that never accepts, then an uncaught error -# shuts the interpreter down. Module teardown runs after Py_IsInitialized() drops to 0; -# `HoldTeardown.__del__` closes the listener there, failing the request so the coroutine -# waker attaches from a Tokio worker, and stalls teardown until that happens. +# Keep a pending request alive into CPython module teardown, then fail its I/O. +# Upload mode also wakes the Python producer waiting for channel capacity. SCRIPT = """ import asyncio +import os import socket -import time +import sys +import threading +from types import FunctionType import wreq listener = socket.socket() listener.bind(("127.0.0.1", 0)) listener.listen(8) +listener.settimeout(5) url = f"http://127.0.0.1:{listener.getsockname()[1]}/" +wait_lock = threading.Lock() +wait_lock.acquire() + + class HoldTeardown: - def __init__(self, listener, task): - self.listener = listener + def __init__(self, peer, task): + self.peer = peer self.task = task - def __del__(self): - self.listener.close() - time.sleep(1) + def __del__(self, write=os.write, wait=wait_lock.acquire): + self.peer.close() + write(2, b"teardown: connection closed\\n") + wait(timeout=1) + write(2, b"teardown: wait complete\\n") +async def chunks(): + chunk = b"x" * (1024 * 1024) + while True: + yield chunk + + +# Do not let the suspended generator retain this module's teardown sentinel. +chunks = FunctionType(chunks.__code__, {}) loop = asyncio.new_event_loop() -task = loop.create_task(wreq.Client().get(url)) +client = wreq.Client(proxies=[]) +upload = sys.argv[1] == "upload" +task = loop.create_task( + client.post(url, body=chunks()) if upload else client.get(url) +) loop.run_until_complete(asyncio.sleep(0.05)) -hold = HoldTeardown(listener, task) -del listener, task +peer, _ = listener.accept() +listener.close() +peer.settimeout(5) +assert peer.recv(4096), "request did not reach the server" +loop.run_until_complete(asyncio.sleep(0.1)) +assert not task.done(), "request must remain pending" +if upload: + assert any( + getattr(getattr(t.get_coro(), "cr_await", None), "__name__", None) == "send" + for t in asyncio.all_tasks(loop) + ), "upload producer must be waiting for channel capacity" +hold = HoldTeardown(peer, task) +del peer, task raise RuntimeError("uncaught error while a request is in flight") """ -# PyPy doesn't guarantee `__del__` runs during interpreter exit, so the trigger may not fire. -pytestmark = pytest.mark.skipif( + +@pytest.mark.skipif( platform.python_implementation() != "CPython", - reason="relies on CPython running __del__ during module teardown", + reason="requires CPython module teardown to run __del__", ) - - -def test_shutdown_wake_reports_error_without_panic(): +@pytest.mark.parametrize("operation", ["request", "upload"]) +def test_shutdown_wake_without_panic(operation): proc = subprocess.run( - [sys.executable, "-c", SCRIPT], + [sys.executable, "-c", SCRIPT, operation], capture_output=True, text=True, - timeout=60, + timeout=15, ) - # Exits through the uncaught RuntimeError, not an abort. assert proc.returncode == 1, proc.stderr + assert "uncaught error while a request is in flight" in proc.stderr + assert "teardown: connection closed" in proc.stderr + assert "teardown: wait complete" in proc.stderr assert "panicked" not in proc.stderr - assert ( - "wreq: failed to wake a Python coroutine: " - "The Python interpreter is not available" in proc.stderr + assert "Exception ignored" not in proc.stderr + + +def test_unconsumed_upload_cleanup(): + script = """ +import asyncio +import gc +import wreq + +async def main(retain): + produced = [] + full = asyncio.Event() + closed = asyncio.Event() + + async def chunks(): + try: + for index in range(100): + produced.append(index) + if index == 1: + full.set() + yield b"chunk" + finally: + closed.set() + print("generator closed", flush=True) + + # A part retains the body without polling its Rust stream. + part = wreq.Part(name="file", value=chunks()) + try: + await asyncio.wait_for(full.wait(), 5) + await asyncio.sleep(0.05) + assert produced == [0, 1] + if retain: + return part + finally: + if not retain: + del part + # PyPy does not destroy native objects immediately after del. + for _ in range(3): + gc.collect() + await asyncio.wait_for(closed.wait(), 5) + +for retain in (False, True): + part = asyncio.run(main(retain)) + print("runner closed", flush=True) +""" + proc = subprocess.run( + [sys.executable, "-c", script], + capture_output=True, + text=True, + timeout=15, ) + assert proc.returncode == 0, proc.stderr + assert proc.stdout.splitlines() == ["generator closed", "runner closed"] * 2 diff --git a/tests/upload_test.py b/tests/upload_test.py new file mode 100644 index 00000000..4288132c --- /dev/null +++ b/tests/upload_test.py @@ -0,0 +1,115 @@ +import asyncio +import contextvars +import gc +import threading + +import pytest +import wreq + +from cancellation_test import local_server + + +@pytest.fixture +def no_automatic_gc(): + enabled = gc.isenabled() + gc.disable() + try: + yield + finally: + if enabled: + gc.enable() + + +async def read_chunked(reader): + body = bytearray() + while True: + size = int(await reader.readline(), 16) + if not size: + assert await reader.readline() == b"\r\n" + return bytes(body) + body.extend(await reader.readexactly(size)) + assert await reader.readexactly(2) == b"\r\n" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("multipart", [False, True]) +async def test_async_upload(multipart, no_automatic_gc): + context = contextvars.ContextVar("upload_context", default="missing") + context.set("caller") + thread = threading.get_ident() + closed = asyncio.Event() + + async def chunks(): + try: + for item in (b"hello ", "world"): + await asyncio.sleep(0) + assert context.get() == "caller" + assert threading.get_ident() == thread + yield item + finally: + closed.set() + + async with local_server() as (url, connections), wreq.Client(proxies=[]) as client: + kwds = ( + {"multipart": wreq.Multipart(wreq.Part(name="file", value=chunks()))} + if multipart + else {"body": chunks()} + ) + task = asyncio.create_task(client.post(url, **kwds)) + reader, writer = await asyncio.wait_for(connections.get(), 5) + body = await asyncio.wait_for(read_chunked(reader), 5) + assert (b"hello world" in body) if multipart else (body == b"hello world") + writer.write(b"HTTP/1.1 200 OK\r\nContent-Length: 0\r\n\r\n") + await writer.drain() + response = await asyncio.wait_for(task, 5) + await response.close() + assert closed.is_set() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("failure", ["exception", "type", "cancelled"]) +async def test_upload_errors(failure): + closed = asyncio.Event() + + async def chunks(): + try: + yield b"first" + if failure == "exception": + raise ValueError("upload exploded") + if failure == "cancelled": + raise asyncio.CancelledError("generator cancelled") + yield object() + finally: + closed.set() + + async with local_server() as (url, _), wreq.Client(proxies=[]) as client: + with pytest.raises(wreq.exceptions.RequestError): + await asyncio.wait_for(client.post(url, body=chunks()), 5) + await asyncio.wait_for(closed.wait(), 5) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("action", ["cancel", "close_client"]) +async def test_upload_cancellation(action): + started = asyncio.Event() + closed = asyncio.Event() + + async def chunks(): + try: + yield b"first" + started.set() + await asyncio.Event().wait() + finally: + await asyncio.sleep(0) + closed.set() + + async with local_server() as (url, _), wreq.Client(proxies=[]) as client: + task = asyncio.create_task(client.post(url, body=chunks())) + await asyncio.wait_for(started.wait(), 5) + if action == "cancel": + task.cancel("stop upload") + else: + client.close() + with pytest.raises(asyncio.CancelledError): + await asyncio.wait_for(task, 5) + await asyncio.wait_for(closed.wait(), 5)