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
13 changes: 5 additions & 8 deletions src/client/resp.rs
Original file line number Diff line number Diff line change
Expand Up @@ -14,12 +14,9 @@ pub use self::{
const READ_ATTACHED: u64 = 64 * 1024;

/// The largest unread body the event loop polls for a response of `version`: `limit` up to
/// HTTP/1.1, else 0. An HTTP/2 stream shares its connection's state behind a lock that the
/// connection task holds while it handles frames, so polling one would stall the loop.
fn loop_limit(version: wreq::Version, limit: u64) -> u64 {
if version <= wreq::Version::HTTP_11 {
limit
} else {
0
}
/// HTTP/1.1, and none after it, not even an empty one. An HTTP/2 stream shares its
/// connection's state behind a lock that the connection task holds while it handles frames,
/// so polling one would stall the loop.
fn loop_limit(version: wreq::Version, limit: u64) -> Option<u64> {
(version <= wreq::Version::HTTP_11).then_some(limit)
}
21 changes: 13 additions & 8 deletions src/client/resp/http.rs
Original file line number Diff line number Diff line change
Expand Up @@ -112,9 +112,10 @@ impl Response {
}

/// Take the body for reading. Cached bytes are shared at once, and a body of known
/// length up to `limit` that is already buffered is read now on the calling thread;
/// overlapping reads and `stream()` fail with [`Error::Memory`] while a read runs.
fn take_bytes(&self, limit: u64) -> Result<BodyRead, Error> {
/// length up to `limit` that is already buffered is read now on the calling thread, none
/// without a `limit`; overlapping reads and `stream()` fail with [`Error::Memory`] while
/// a read runs.
fn take_bytes(&self, limit: Option<u64>) -> Result<BodyRead, Error> {
let mut slot = self.slot();
let body = match mem::replace(&mut *slot, Body::Taken) {
Body::Unread(body) => body,
Expand All @@ -128,7 +129,11 @@ impl Response {
}
};
drop(slot);
let attached = body.size_hint().exact().is_some_and(|len| len <= limit);
let attached = body
.size_hint()
.exact()
.zip(limit)
.is_some_and(|(len, limit)| len <= limit);
let mut collect = body.collect();
let ready = if attached {
// A read timeout starts a timer on its first poll, which needs the runtime.
Expand Down Expand Up @@ -176,7 +181,7 @@ impl Response {
/// the body is still arriving or too large to decode on the caller.
fn read_body<F, Fut, T>(
&self,
limit: u64,
limit: Option<u64>,
read: F,
) -> Result<(impl Future<Output = PyResult<T>> + Send + 'static, bool), Error>
where
Expand All @@ -185,7 +190,7 @@ impl Response {
{
let (decode, inline) = match self.take_bytes(limit)? {
BodyRead::Ready(bytes) => {
let inline = bytes.len() as u64 <= limit;
let inline = limit.is_some_and(|limit| bytes.len() as u64 <= limit);
(Either::Left(read(self.build_response(bytes))), inline)
}
BodyRead::Pending(collect) => (
Expand Down Expand Up @@ -484,7 +489,7 @@ impl BlockingResponse {
T: Send,
{
let runtime = &self.0.runtime;
let (read, inline) = self.0.read_body(READ_ATTACHED, read)?;
let (read, inline) = self.0.read_body(Some(READ_ATTACHED), read)?;
if inline {
nogil::run(py, runtime, read)
} else {
Expand Down Expand Up @@ -579,7 +584,7 @@ impl BlockingResponse {
/// Read the body as a read-only memoryview, retaining its data after the response closes.
pub fn bytes(&self, py: Python) -> PyResult<PyBuffer> {
let response = &self.0;
match response.take_bytes(READ_ATTACHED)? {
match response.take_bytes(Some(READ_ATTACHED))? {
BodyRead::Ready(bytes) => Ok(PyBuffer::from(bytes)),
BodyRead::Pending(collect) => {
let read = response
Expand Down
3 changes: 2 additions & 1 deletion src/client/resp/stream.rs
Original file line number Diff line number Diff line change
Expand Up @@ -171,7 +171,8 @@ impl Streamer {
&& !resp
.size_hint()
.exact()
.is_some_and(|len| len <= loop_limit(resp.version(), limit))
.zip(loop_limit(resp.version(), limit))
.is_some_and(|(len, limit)| len <= limit)
{
return None;
}
Expand Down
97 changes: 97 additions & 0 deletions tests/response_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
import socket
import threading
import time
from contextlib import asynccontextmanager
from datetime import timedelta

import pytest
Expand Down Expand Up @@ -251,6 +252,102 @@ async def test_pending_read_holds_the_body_until_it_ends():
await response.text()


async def steps(coroutine):
"""Drive `coroutine` by hand; return its result and how often it suspended."""
suspended = 0
while True:
try:
waiter = coroutine.send(None)
except StopIteration as stop:
return stop.value, suspended
except StopAsyncIteration:
return None, suspended
suspended += 1
if waiter is None:
await asyncio.sleep(0)
else:
await asyncio.wait({waiter}, timeout=5)


@asynccontextmanager
async def body_server(http2, body):
"""Answer every request with `body` and its Content-Length, over HTTP/1 or h2c."""
length = str(len(body)).encode()
handlers, writers = set(), []

def frame(kind, flags, stream, payload=b""):
size = len(payload).to_bytes(3, "big")
return size + bytes((kind, flags)) + stream.to_bytes(4, "big") + payload

async def accept(reader, writer):
handlers.add(asyncio.current_task())
writers.append(writer)
try:
if not http2:
await reader.readuntil(b"\r\n\r\n")
head = b"HTTP/1.1 200 OK\r\nConnection: close\r\nContent-Length: "
writer.write(head + length + b"\r\n\r\n" + body)
await writer.drain()
await reader.read()
return
assert await reader.readexactly(24) == b"PRI * HTTP/2.0\r\n\r\nSM\r\n\r\n"
writer.write(frame(4, 0, 0))
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 == 1:
# `:status: 200` and `content-length`, then the body ending the stream.
block = b"\x88\x0f\x0d" + bytes((len(length),)) + length
if body:
writer.write(
frame(1, 4, stream, block) + frame(0, 1, stream, body)
)
else:
writer.write(frame(1, 5, stream, block))
await writer.drain()
except (asyncio.IncompleteReadError, ConnectionError):
pass
finally:
handlers.discard(asyncio.current_task())
writer.close()

server = await asyncio.start_server(accept, "127.0.0.1", 0)
try:
yield f"http://127.0.0.1:{server.sockets[0].getsockname()[1]}/"
finally:
server.close()
for writer in writers:
writer.close()
await asyncio.gather(*handlers, return_exceptions=True)
await server.wait_closed()


@pytest.mark.asyncio
@pytest.mark.parametrize("read", ["bytes", "stream"])
@pytest.mark.parametrize("body", [b"ok", b""], ids=["body", "empty"])
@pytest.mark.parametrize("http2", [False, True], ids=["http1", "http2"])
async def test_only_http1_bodies_are_read_on_the_event_loop(http2, body, read):
# A buffered HTTP/1 body is read in the first step of the coroutine. An HTTP/2 body,
# even an empty one, is always read on the runtime, so the coroutine suspends first.
async with body_server(http2, body) as url:
async with wreq.Client(http2_only=http2, proxies=[]) as client:
response = await asyncio.wait_for(client.get(url), 5)
assert response.version == (Version.HTTP_2 if http2 else Version.HTTP_11)
# Let the whole body arrive before the read starts.
await asyncio.sleep(0.05)
if read == "bytes":
result, suspended = await steps(response.bytes())
assert bytes(result) == body
else:
result, suspended = await steps(response.stream().__anext__())
assert (bytes(result) if result is not None else b"") == body
assert (suspended > 0) == http2


@pytest.mark.asyncio
async def test_stream_read_ahead_yields_and_closes_waiting_readers():
stalled = threading.Event()
Expand Down
7 changes: 5 additions & 2 deletions tests/websocket_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -73,7 +73,7 @@ async def test_websocket_close_frame(code, reason, expected):
async with wreq.Client(proxies=[]) as client:
# Leaving the block after an explicit close must not fail.
async with client.websocket(url) as ws:
await ws.close(code, reason)
assert await ws.close(code, reason) is None
assert await asyncio.wait_for(frames.get(), 5) == expected
finally:
server.close()
Expand Down Expand Up @@ -137,7 +137,10 @@ async def send_later():
await asyncio.wait_for(ws.recv(), 0.2)
# A pending receive does not hold up a send.
pending = asyncio.ensure_future(ws.recv())
await asyncio.wait_for(ws.send(wreq.Message.from_text("ping")), 5)
assert (
await asyncio.wait_for(ws.send(wreq.Message.from_text("ping")), 5)
is None
)
assert await asyncio.wait_for(received.get(), 5) == (0x1, b"ping")
release.set()
assert (await asyncio.wait_for(pending, 5)).text == "first"
Expand Down
Loading