Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 3 additions & 3 deletions Cargo.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

9 changes: 5 additions & 4 deletions docs/source/guide/advanced.md
Original file line number Diff line number Diff line change
Expand Up @@ -119,10 +119,11 @@ client = Client(runtime=runtime)

With `work_steal=False`, workers use independent single-thread Tokio runtimes.
Each client is assigned one worker for its lifetime; requests, response reads,
streams and WebSocket operations use that worker. With multiple workers, newly
created clients select a worker randomly and keep that selection. This is not
CPU pinning. Sharing the same `Runtime` between clients is supported, and
`client.runtime` returns the shared runtime object.
streams and WebSocket operations use that worker. Async reads of small HTTP/1
bodies that have already arrived finish on the event loop thread instead. With
multiple workers, newly created clients select a worker randomly and keep that
selection. This is not CPU pinning. Sharing the same `Runtime` between clients
is supported, and `client.runtime` returns the shared runtime object.

`workers=None` uses the available CPU parallelism, or 1 if it cannot be determined.
Custom runtimes start their threads during construction, before any client is
Expand Down
8 changes: 4 additions & 4 deletions python/wreq/wreq.py
Original file line number Diff line number Diff line change
Expand Up @@ -454,7 +454,7 @@ async def close(self) -> None:
"""

async def __aenter__(self) -> Any: ...
async def __aexit__(self, _exc_type: Any, _exc_value: Any, _traceback: Any) -> Any:
async def __aexit__(self, _exc_type: Any, _exc_value: Any, _traceback: Any) -> None:
r"""
Release the body without forbidding reuse: a fully read connection returns
to the pool, while an unread HTTP/1 body drains or closes its connection.
Expand Down Expand Up @@ -528,7 +528,7 @@ async def close(
"""

async def __aenter__(self) -> Any: ...
async def __aexit__(self, _exc_type: Any, _exc_value: Any, _traceback: Any) -> Any:
async def __aexit__(self, _exc_type: Any, _exc_value: Any, _traceback: Any) -> None:
r"""
Close the WebSocket connection without a close code or reason, unless already closed.
"""
Expand Down Expand Up @@ -1134,7 +1134,7 @@ class _RequestCoroutine(Coroutine[Any, Any, _T]):
"""

async def __aenter__(self) -> _T: ...
async def __aexit__(self, _exc_type: Any, _exc_value: Any, _traceback: Any) -> Any: ...
async def __aexit__(self, _exc_type: Any, _exc_value: Any, _traceback: Any) -> None: ...


class Client:
Expand Down Expand Up @@ -1458,7 +1458,7 @@ async def main():
...

async def __aenter__(self) -> Any: ...
async def __aexit__(self, _exc_type: Any, _exc_value: Any, _traceback: Any) -> Any:
async def __aexit__(self, _exc_type: Any, _exc_value: Any, _traceback: Any) -> None:
r"""
Close the client like `close()`: cancel pending requests and reject new ones.
"""
Expand Down
86 changes: 84 additions & 2 deletions src/client/body/json.rs
Original file line number Diff line number Diff line change
@@ -1,10 +1,15 @@
use std::fmt;

use indexmap::IndexMap;
use pyo3::{FromPyObject, prelude::*, pybacked::PyBackedStr};
use serde::{Deserialize, Deserializer, Serialize, Serializer};
use serde::{
Deserialize, Deserializer, Serialize, Serializer,
de::{Error, MapAccess, SeqAccess, Visitor},
};

/// Represents a JSON value for HTTP requests.
/// Supports objects, arrays, numbers, strings, booleans, and null.
#[derive(FromPyObject, IntoPyObject, Serialize, Deserialize)]
#[derive(FromPyObject, IntoPyObject, Serialize)]
#[serde(untagged)]
pub enum Json {
Object(IndexMap<JsonString, Json>),
Expand All @@ -24,6 +29,12 @@ pub enum JsonString {
RustString(String),
}

/// Builds a [`Json`] in one pass; an untagged derive would buffer every value and then try
/// each variant in turn.
struct JsonVisitor;

// ===== impl JsonString =====

impl FromPyObject<'_, '_> for JsonString {
type Error = PyErr;

Expand Down Expand Up @@ -55,3 +66,74 @@ impl<'de> Deserialize<'de> for JsonString {
String::deserialize(deserializer).map(JsonString::RustString)
}
}

// ===== impl Json =====

impl<'de> Deserialize<'de> for Json {
#[inline]
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: Deserializer<'de>,
{
deserializer.deserialize_any(JsonVisitor)
}
}

// ===== impl JsonVisitor =====

impl<'de> Visitor<'de> for JsonVisitor {
type Value = Json;

fn expecting(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("a JSON value")
}

fn visit_bool<E: Error>(self, v: bool) -> Result<Json, E> {
Ok(Json::Boolean(v))
}

// Integers outside `isize` become floats, as the untagged variant order did.
fn visit_i64<E: Error>(self, v: i64) -> Result<Json, E> {
Ok(isize::try_from(v).map_or(Json::Float(v as f64), Json::Number))
}

fn visit_u64<E: Error>(self, v: u64) -> Result<Json, E> {
Ok(isize::try_from(v).map_or(Json::Float(v as f64), Json::Number))
}

fn visit_f64<E: Error>(self, v: f64) -> Result<Json, E> {
Ok(Json::Float(v))
}

fn visit_str<E: Error>(self, v: &str) -> Result<Json, E> {
Ok(Json::String(JsonString::RustString(v.to_owned())))
}

fn visit_string<E: Error>(self, v: String) -> Result<Json, E> {
Ok(Json::String(JsonString::RustString(v)))
}

fn visit_unit<E: Error>(self) -> Result<Json, E> {
Ok(Json::Null(None))
}

fn visit_none<E: Error>(self) -> Result<Json, E> {
Ok(Json::Null(None))
}

fn visit_seq<A: SeqAccess<'de>>(self, mut seq: A) -> Result<Json, A::Error> {
let mut values = Vec::new();
while let Some(value) = seq.next_element()? {
values.push(value);
}
Ok(Json::Array(values))
}

fn visit_map<A: MapAccess<'de>>(self, mut map: A) -> Result<Json, A::Error> {
let mut values = IndexMap::new();
while let Some((key, value)) = map.next_entry()? {
values.insert(key, value);
}
Ok(Json::Object(values))
}
}
86 changes: 44 additions & 42 deletions src/client/body/stream.rs
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,10 @@

use std::{
pin::Pin,
sync::{Arc, Mutex, MutexGuard, PoisonError},
sync::{
Arc, Mutex, MutexGuard, PoisonError,
atomic::{AtomicBool, Ordering},
},
task::{Context, Poll},
time::Duration,
};
Expand Down Expand Up @@ -69,14 +72,19 @@ struct Pump {
/// A request body from a Python async generator, forwarded with one chunk of buffering
/// by a task on the loop that was running at extraction. Dropping it cancels that task.
struct PyAsyncStream {
rx: mpsc::Receiver<Option<Item>>,
rx: mpsc::Receiver<Item>,
/// Set by [`Sender::finish`], so the channel closing reads as the end of the body.
finished: Arc<AtomicBool>,
task: Option<(Py<PyAny>, Py<PyAny>)>,
}

/// The channel end given to the forwarding coroutine; awaiting `send` applies upload
/// backpressure. Closing it without `finish` fails the body.
#[pyclass(frozen)]
struct Sender(Mutex<Option<mpsc::Sender<Option<Item>>>>);
struct Sender {
tx: Mutex<Option<mpsc::Sender<Item>>>,
finished: Arc<AtomicBool>,
}

// ===== impl PyBytesLike =====

Expand Down Expand Up @@ -226,7 +234,7 @@ impl PyAsyncStream {
fn new(generator: Bound<'_, PyAny>) -> PyResult<Self> {
static FORWARD: PyOnceLock<Py<PyAny>> = PyOnceLock::new();
let py = generator.py();
let event_loop = py.import("asyncio")?.call_method0("get_running_loop")?;
let event_loop = coroutine::running_loop(py)?;
let forward = FORWARD.get_or_try_init(py, || {
PyModule::from_code(
py,
Expand All @@ -253,7 +261,7 @@ async def forward(gen, sender):
except BaseException as error:
await sender.send(error, True)
else:
await sender.finish()
sender.finish()
",
c"wreq/_async_stream.py",
c"wreq._async_stream",
Expand All @@ -262,19 +270,21 @@ async def forward(gen, sender):
.map(Bound::unbind)
})?;
let (tx, rx) = mpsc::channel(1);
let coroutine = forward
.bind(py)
.call1((generator, Sender(Mutex::new(Some(tx)))))?;
// create_task captures the caller's contextvars on the running loop.
let task = match event_loop.call_method1("create_task", (&coroutine,)) {
Ok(task) => task,
Err(err) => {
let _ = coroutine.call_method0("close");
return Err(err);
}
let finished = Arc::new(AtomicBool::new(false));
let sender = Sender {
tx: Mutex::new(Some(tx)),
finished: finished.clone(),
};
let coroutine = forward.bind(py).call1((generator, sender))?;
// create_task captures the caller's contextvars on the running loop.
let task = event_loop
.call_method1(intern!(py, "create_task"), (&coroutine,))
.inspect_err(|_| {
let _ = coroutine.call_method0(intern!(py, "close"));
})?;
Ok(Self {
rx,
finished,
task: Some((task.unbind(), event_loop.unbind())),
})
}
Expand All @@ -286,17 +296,16 @@ impl Stream for PyAsyncStream {
fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
let this = self.get_mut();
match this.rx.poll_recv(cx) {
Poll::Ready(Some(Some(item))) => Poll::Ready(Some(item)),
Poll::Ready(Some(None)) => {
this.rx.close();
this.task.take();
Poll::Ready(None)
Poll::Ready(Some(item)) => Poll::Ready(Some(item)),
// Every sender is gone: the end after `finish`. Otherwise the forwarding task was
// cancelled or destroyed, so the body is incomplete.
Poll::Ready(None)
if this.task.take().is_some() && !this.finished.load(Ordering::Acquire) =>
{
Poll::Ready(Some(Err(PyRuntimeError::new_err(
"async body generator stopped before it finished",
))))
}
// Every sender is gone without `finish`: the forwarding task was cancelled or
// destroyed, so the body is incomplete.
Poll::Ready(None) if this.task.take().is_some() => Poll::Ready(Some(Err(
PyRuntimeError::new_err("async body generator stopped before it finished"),
))),
Poll::Ready(None) => Poll::Ready(None),
Poll::Pending => Poll::Pending,
}
Expand Down Expand Up @@ -342,41 +351,34 @@ impl Sender {
} else {
Ok(item.extract()?)
};
let tx = self.sender();
let tx = self.lock().clone();
// Channel readiness is runtime-independent, so this waits on the Python loop.
coroutine::local(py, "Sender.send", async move {
Ok(match tx {
Some(tx) => tx.send(Some(item)).await.is_ok(),
Some(tx) => tx.send(item).await.is_ok(),
None => false,
})
})
}

/// Mark the normal end of the body.
fn finish<'py>(&self, py: Python<'py>) -> PyResult<Bound<'py, Coroutine>> {
/// Mark the normal end of the body. The last chunk may still be queued: it is read
/// before the closed channel, so finishing never waits for room.
fn finish(&self) {
self.finished.store(true, Ordering::Release);
// Python may retain the sender after completion, especially on PyPy.
let tx = self.sender();
coroutine::local(py, "Sender.finish", async move {
Ok(match tx {
Some(tx) => tx.send(None).await.is_ok(),
None => false,
})
})
self.close();
}

/// Drop the channel end at once, so the body fails instead of waiting for more.
fn close(&self) {
let tx = self.0.lock().unwrap_or_else(PoisonError::into_inner).take();
let tx = self.lock().take();
drop(tx);
}
}

impl Sender {
#[inline]
fn sender(&self) -> Option<mpsc::Sender<Option<Item>>> {
self.0
.lock()
.unwrap_or_else(PoisonError::into_inner)
.clone()
fn lock(&self) -> MutexGuard<'_, Option<mpsc::Sender<Item>>> {
self.tx.lock().unwrap_or_else(PoisonError::into_inner)
}
}
Loading
Loading