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
14 changes: 7 additions & 7 deletions src/client.rs
Original file line number Diff line number Diff line change
Expand Up @@ -21,8 +21,8 @@ use self::{
resp::{BlockingResponse, BlockingWebSocket, Response, WebSocket},
};
use crate::{
aio::{self, Coroutine},
cookie::Jar,
coroutine::{self, Coroutine},
dns::{DnsOptions, HickoryResolver, LookupIpStrategy},
emulate::EmulationLike,
error::Error,
Expand Down Expand Up @@ -265,7 +265,7 @@ impl Client {
kwds: Option<Bound<'py, PyDict>>,
) -> PyResult<Bound<'py, Coroutine>> {
let kwds = kwds.map(Bound::unbind);
aio::managed(py, qualname, self.clone().execute(method, url, kwds))
coroutine::managed(py, qualname, self.clone().execute(method, url, kwds))
}

/// Send a request on the client's runtime, extracting options on first await so an
Expand All @@ -277,7 +277,7 @@ impl Client {
kwds: Option<Py<PyDict>>,
) -> PyResult<Response> {
let kwds = Python::attach(|py| kwds.map(|kwds| kwds.bind(py).extract()).transpose())?;
aio::run(
coroutine::run(
self.runtime.clone(),
execute_request(self, method, url, kwds),
)
Expand All @@ -291,7 +291,7 @@ impl Client {
kwds: Option<Py<PyDict>>,
) -> PyResult<WebSocket> {
let kwds = Python::attach(|py| kwds.map(|kwds| kwds.bind(py).extract()).transpose())?;
aio::run(
coroutine::run(
self.runtime.clone(),
execute_websocket_request(self, url, kwds),
)
Expand Down Expand Up @@ -663,14 +663,14 @@ impl Client {
kwds: Option<Bound<'py, PyDict>>,
) -> PyResult<Bound<'py, Coroutine>> {
let kwds = kwds.map(Bound::unbind);
aio::managed(py, "Client.websocket", self.clone().connect(url, kwds))
coroutine::managed(py, "Client.websocket", self.clone().connect(url, kwds))
}
}

#[pymethods]
impl Client {
fn __aenter__(slf: Bound<'_, Self>) -> PyResult<Bound<'_, Coroutine>> {
aio::ready("Client.__aenter__", slf)
coroutine::ready("Client.__aenter__", slf)
}

/// Close the client like `close()`: cancel pending requests and reject new ones.
Expand All @@ -682,7 +682,7 @@ impl Client {
_traceback: Py<PyAny>,
) -> PyResult<Bound<'py, Coroutine>> {
let cancel = self.cancel.clone();
aio::local(py, "Client.__aexit__", async move {
coroutine::local(py, "Client.__aexit__", async move {
cancel.cancel();
Ok(())
})
Expand Down
65 changes: 47 additions & 18 deletions src/client/body/stream.rs
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,7 @@ use std::{
use bytes::Bytes;
use futures_util::Stream;
use pyo3::{
exceptions::PyStopIteration,
exceptions::{PyRuntimeError, PyStopIteration},
intern,
prelude::*,
sync::PyOnceLock,
Expand All @@ -24,7 +24,7 @@ use tokio::{
};

use crate::{
aio::{self, Coroutine},
coroutine::{self, Coroutine},
extractor::{Binary, Text},
runtime,
};
Expand Down Expand Up @@ -74,9 +74,9 @@ struct PyAsyncStream {
}

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

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

Expand Down Expand Up @@ -243,8 +243,11 @@ async def forward(gen, sender):
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():
# Task cancellation must not wait for space in a retained body, and must not leave
# the body waiting while a traceback keeps this frame and its sender alive.
if asyncio.current_task().cancelling():
sender.close()
else:
await sender.send(error, True)
raise
except BaseException as error:
Expand All @@ -259,7 +262,9 @@ async def forward(gen, sender):
.map(Bound::unbind)
})?;
let (tx, rx) = mpsc::channel(1);
let coroutine = forward.bind(py).call1((generator, Sender(tx)))?;
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,
Expand Down Expand Up @@ -287,10 +292,12 @@ impl Stream for PyAsyncStream {
this.task.take();
Poll::Ready(None)
}
Poll::Ready(_) => {
this.rx.close();
Poll::Ready(None)
}
// 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 @@ -323,7 +330,7 @@ impl Drop for PyAsyncStream {
#[pymethods]
impl Sender {
/// Queue a chunk, or `item` as the error that ends the body. Resolves to `False`
/// once the body is dropped, which stops forwarding.
/// once the body is dropped or the sender closed, which stops forwarding.
fn send<'py>(
&self,
py: Python<'py>,
Expand All @@ -335,19 +342,41 @@ impl Sender {
} else {
Ok(item.extract()?)
};
let tx = self.0.clone();
let tx = self.sender();
// Channel readiness is runtime-independent, so this waits on the Python loop.
aio::local(py, "Sender.send", async move {
Ok(tx.send(Some(item)).await.is_ok())
coroutine::local(py, "Sender.send", async move {
Ok(match tx {
Some(tx) => tx.send(Some(item)).await.is_ok(),
None => false,
})
})
}

/// Mark the normal end of the body.
fn finish<'py>(&self, py: Python<'py>) -> PyResult<Bound<'py, Coroutine>> {
// Python may retain the sender after completion, especially on PyPy.
let tx = self.0.clone();
aio::local(py, "Sender.finish", async move {
Ok(tx.send(None).await.is_ok())
let tx = self.sender();
coroutine::local(py, "Sender.finish", async move {
Ok(match tx {
Some(tx) => tx.send(None).await.is_ok(),
None => false,
})
})
}

/// 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();
drop(tx);
}
}

impl Sender {
#[inline]
fn sender(&self) -> Option<mpsc::Sender<Option<Item>>> {
self.0
.lock()
.unwrap_or_else(PoisonError::into_inner)
.clone()
}
}
12 changes: 6 additions & 6 deletions src/client/resp/http.rs
Original file line number Diff line number Diff line change
Expand Up @@ -18,10 +18,10 @@ use wreq::Uri;

use super::{ext::ResponseExt, stream::Streamer};
use crate::{
aio::{self, Coroutine},
buffer::PyBuffer,
client::{SocketAddr, body::Json, nogil},
cookie::Cookie,
coroutine::{self, Coroutine},
error::Error,
header::HeaderMap,
http::{StatusCode, Version},
Expand Down Expand Up @@ -185,9 +185,9 @@ impl Response {
{
let py = slf.py();
let slf = slf.unbind();
aio::local(py, qualname, async move {
coroutine::local(py, qualname, async move {
let this = slf.get();
aio::run(this.runtime.clone(), this.read_body(read)).await
coroutine::run(this.runtime.clone(), this.read_body(read)).await
})
}

Expand Down Expand Up @@ -353,7 +353,7 @@ impl Response {
pub fn close(slf: Bound<'_, Self>) -> PyResult<Bound<'_, Coroutine>> {
let py = slf.py();
let slf = slf.unbind();
aio::local(py, "Response.close", async move {
coroutine::local(py, "Response.close", async move {
slf.get().discard();
Ok(())
})
Expand All @@ -363,7 +363,7 @@ impl Response {
#[pymethods]
impl Response {
fn __aenter__(slf: Bound<'_, Self>) -> PyResult<Bound<'_, Coroutine>> {
aio::ready("Response.__aenter__", slf)
coroutine::ready("Response.__aenter__", slf)
}

/// Release the body without forbidding reuse: a fully read connection returns
Expand All @@ -376,7 +376,7 @@ impl Response {
) -> PyResult<Bound<'py, Coroutine>> {
let py = slf.py();
let slf = slf.unbind();
aio::local(py, "Response.__aexit__", async move {
coroutine::local(py, "Response.__aexit__", async move {
slf.get().destroy();
Ok(())
})
Expand Down
10 changes: 5 additions & 5 deletions src/client/resp/stream.rs
Original file line number Diff line number Diff line change
Expand Up @@ -22,9 +22,9 @@ use tokio::sync::{
use tokio_util::task::AbortOnDropHandle;

use crate::{
aio::{self, Coroutine},
buffer::PyBuffer,
client::nogil,
coroutine::{self, Coroutine},
error::Error,
header::HeaderMap,
runtime::Runtime,
Expand Down Expand Up @@ -295,20 +295,20 @@ impl Streamer {
fn __anext__(slf: Bound<'_, Self>) -> PyResult<Bound<'_, Coroutine>> {
let py = slf.py();
let slf = slf.unbind();
aio::local(py, "Streamer.__anext__", async move {
coroutine::local(py, "Streamer.__anext__", async move {
let this = slf.get();
// Buffered frames complete without suspending; yield to the event loop
// periodically so timeouts, cancellation and other tasks run.
if this.reader.since_yield.fetch_add(1, Ordering::Relaxed) >= Self::YIELD_EVERY {
this.reader.since_yield.store(0, Ordering::Relaxed);
aio::yield_now().await;
coroutine::yield_now().await;
}
this.next(|| Error::StopAsyncIteration).await
})
}

fn __aenter__(slf: Bound<'_, Self>) -> PyResult<Bound<'_, Coroutine>> {
aio::ready("Streamer.__aenter__", slf)
coroutine::ready("Streamer.__aenter__", slf)
}

/// Release the body and end any pending read; returned views stay valid.
Expand All @@ -320,7 +320,7 @@ impl Streamer {
_traceback: Py<PyAny>,
) -> PyResult<Bound<'py, Coroutine>> {
let reader = self.reader.clone();
aio::local(py, "Streamer.__aexit__", async move {
coroutine::local(py, "Streamer.__aexit__", async move {
reader.close();
Ok(())
})
Expand Down
20 changes: 9 additions & 11 deletions src/client/resp/ws.rs
Original file line number Diff line number Diff line change
Expand Up @@ -8,13 +8,12 @@ use std::{

use msg::Message;
use pyo3::prelude::*;
use tokio::sync::mpsc;
use wreq::{header::HeaderValue, ws::WebSocketResponse};

use crate::{
aio::{self, Coroutine},
client::{SocketAddr, nogil},
cookie::Cookie,
coroutine::{self, Coroutine},
extractor::Text,
header::HeaderMap,
http::{StatusCode, Version},
Expand Down Expand Up @@ -44,7 +43,7 @@ pub struct WebSocket {
#[pyo3(get)]
headers: HeaderMap,
protocol: Option<HeaderValue>,
cmd: mpsc::UnboundedSender<cmd::Command>,
cmd: cmd::Handle,
runtime: Runtime,
}

Expand All @@ -66,8 +65,7 @@ impl WebSocket {
);
let websocket = response.into_websocket().await?;
let protocol = websocket.protocol().cloned();
let (cmd, rx) = mpsc::unbounded_channel();
tokio::spawn(cmd::task(websocket, rx));
let cmd = cmd::spawn(websocket);

Ok(WebSocket {
runtime,
Expand Down Expand Up @@ -108,7 +106,7 @@ impl WebSocket {
py: Python<'py>,
timeout: Option<Duration>,
) -> PyResult<Bound<'py, Coroutine>> {
aio::spawn(
coroutine::spawn(
py,
"WebSocket.recv",
&self.runtime,
Expand All @@ -119,7 +117,7 @@ impl WebSocket {
/// Send a message to the WebSocket.
#[pyo3(signature = (message))]
pub fn send<'py>(&self, py: Python<'py>, message: Message) -> PyResult<Bound<'py, Coroutine>> {
aio::spawn(
coroutine::spawn(
py,
"WebSocket.send",
&self.runtime,
Expand All @@ -134,7 +132,7 @@ impl WebSocket {
py: Python<'py>,
messages: Vec<Message>,
) -> PyResult<Bound<'py, Coroutine>> {
aio::spawn(
coroutine::spawn(
py,
"WebSocket.send_all",
&self.runtime,
Expand All @@ -150,7 +148,7 @@ impl WebSocket {
code: Option<u16>,
reason: Option<Text>,
) -> PyResult<Bound<'py, Coroutine>> {
aio::spawn(
coroutine::spawn(
py,
"WebSocket.close",
&self.runtime,
Expand All @@ -162,7 +160,7 @@ impl WebSocket {
#[pymethods]
impl WebSocket {
fn __aenter__(slf: Bound<'_, Self>) -> PyResult<Bound<'_, Coroutine>> {
aio::ready("WebSocket.__aenter__", slf)
coroutine::ready("WebSocket.__aenter__", slf)
}

/// Close the WebSocket connection without a close code or reason, unless already closed.
Expand All @@ -173,7 +171,7 @@ impl WebSocket {
_exc_val: Py<PyAny>,
_traceback: Py<PyAny>,
) -> PyResult<Bound<'py, Coroutine>> {
aio::spawn(
coroutine::spawn(
py,
"WebSocket.__aexit__",
&self.runtime,
Expand Down
Loading
Loading