Skip to content
Open
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
47 changes: 34 additions & 13 deletions httpcore/_async/http2.py
Original file line number Diff line number Diff line change
Expand Up @@ -60,6 +60,15 @@ def __init__(
self._state_lock = AsyncLock()
self._read_lock = AsyncLock()
self._write_lock = AsyncLock()
# Guards the shared `h2` state machine on the send path, plus
# stream-ID allocation and the `_events` mapping. The same
# `HTTP2Connection` is handed to multiple threads by the connection
# pool (HTTP/2 multiplexing), and `h2` is not thread-safe: concurrent
# `send_headers` calls corrupt the HPACK encoder table ("deque mutated
# during iteration") and the streams dict ("dictionary changed size
# during iteration"), and concurrent `get_next_available_stream_id`
# calls can hand out duplicate stream IDs. (encode/httpx#3566)
self._send_lock = AsyncLock()
self._sent_connection_init = False
self._used_all_stream_ids = False
self._connection_error = False
Expand Down Expand Up @@ -131,19 +140,27 @@ async def handle_async_request(self, request: Request) -> Response:
await self._max_streams_semaphore.acquire()

try:
stream_id = self._h2_state.get_next_available_stream_id()
self._events[stream_id] = []
except h2.exceptions.NoAvailableStreamIDError: # pragma: nocover
self._used_all_stream_ids = True
self._request_count -= 1
raise ConnectionNotAvailable()

try:
kwargs = {"request": request, "stream_id": stream_id}
async with Trace("send_request_headers", logger, request, kwargs):
await self._send_request_headers(request=request, stream_id=stream_id)
async with Trace("send_request_body", logger, request, kwargs):
await self._send_request_body(request=request, stream_id=stream_id)
# The send path mutates the shared `h2` state machine, which is
# not safe for concurrent use, so stream ID allocation and the
# send itself are serialized under a single lock. Note that h2
# requires `get_next_available_stream_id()` to be immediately
# followed by the matching `send_headers()` call, otherwise
# concurrent callers may be handed duplicate stream IDs. Without
# this, clients sharing one connection hit errors such as
# "deque mutated during iteration", "dictionary changed size
# during iteration", and `StreamIDTooLowError`.
# (encode/httpx#3566)
async with self._send_lock:
stream_id = self._h2_state.get_next_available_stream_id()
self._events[stream_id] = []
kwargs = {"request": request, "stream_id": stream_id}
async with Trace("send_request_headers", logger, request, kwargs):
await self._send_request_headers(
request=request,
stream_id=stream_id,
)
async with Trace("send_request_body", logger, request, kwargs):
await self._send_request_body(request=request, stream_id=stream_id)
async with Trace(
"receive_response_headers", logger, request, kwargs
) as trace:
Expand All @@ -162,6 +179,10 @@ async def handle_async_request(self, request: Request) -> Response:
"stream_id": stream_id,
},
)
except h2.exceptions.NoAvailableStreamIDError: # pragma: nocover
self._used_all_stream_ids = True
self._request_count -= 1
raise ConnectionNotAvailable()
except BaseException as exc: # noqa: PIE786
with AsyncShieldCancellation():
kwargs = {"stream_id": stream_id}
Expand Down
47 changes: 34 additions & 13 deletions httpcore/_sync/http2.py
Original file line number Diff line number Diff line change
Expand Up @@ -60,6 +60,15 @@ def __init__(
self._state_lock = Lock()
self._read_lock = Lock()
self._write_lock = Lock()
# Guards the shared `h2` state machine on the send path, plus
# stream-ID allocation and the `_events` mapping. The same
# `HTTP2Connection` is handed to multiple threads by the connection
# pool (HTTP/2 multiplexing), and `h2` is not thread-safe: concurrent
# `send_headers` calls corrupt the HPACK encoder table ("deque mutated
# during iteration") and the streams dict ("dictionary changed size
# during iteration"), and concurrent `get_next_available_stream_id`
# calls can hand out duplicate stream IDs. (encode/httpx#3566)
self._send_lock = Lock()
self._sent_connection_init = False
self._used_all_stream_ids = False
self._connection_error = False
Expand Down Expand Up @@ -131,19 +140,27 @@ def handle_request(self, request: Request) -> Response:
self._max_streams_semaphore.acquire()

try:
stream_id = self._h2_state.get_next_available_stream_id()
self._events[stream_id] = []
except h2.exceptions.NoAvailableStreamIDError: # pragma: nocover
self._used_all_stream_ids = True
self._request_count -= 1
raise ConnectionNotAvailable()

try:
kwargs = {"request": request, "stream_id": stream_id}
with Trace("send_request_headers", logger, request, kwargs):
self._send_request_headers(request=request, stream_id=stream_id)
with Trace("send_request_body", logger, request, kwargs):
self._send_request_body(request=request, stream_id=stream_id)
# The send path mutates the shared `h2` state machine, which is
# not safe for concurrent use, so stream ID allocation and the
# send itself are serialized under a single lock. Note that h2
# requires `get_next_available_stream_id()` to be immediately
# followed by the matching `send_headers()` call, otherwise
# concurrent callers may be handed duplicate stream IDs. Without
# this, clients sharing one connection hit errors such as
# "deque mutated during iteration", "dictionary changed size
# during iteration", and `StreamIDTooLowError`.
# (encode/httpx#3566)
with self._send_lock:
stream_id = self._h2_state.get_next_available_stream_id()
self._events[stream_id] = []
kwargs = {"request": request, "stream_id": stream_id}
with Trace("send_request_headers", logger, request, kwargs):
self._send_request_headers(
request=request,
stream_id=stream_id,
)
with Trace("send_request_body", logger, request, kwargs):
self._send_request_body(request=request, stream_id=stream_id)
with Trace(
"receive_response_headers", logger, request, kwargs
) as trace:
Expand All @@ -162,6 +179,10 @@ def handle_request(self, request: Request) -> Response:
"stream_id": stream_id,
},
)
except h2.exceptions.NoAvailableStreamIDError: # pragma: nocover
self._used_all_stream_ids = True
self._request_count -= 1
raise ConnectionNotAvailable()
except BaseException as exc: # noqa: PIE786
with ShieldCancellation():
kwargs = {"stream_id": stream_id}
Expand Down
2 changes: 1 addition & 1 deletion requirements.txt
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,7 @@ jinja2==3.1.6

# Packaging
build==1.2.2.post1
twine==6.1.0
twine==7.0.0; python_version >= "3.10" # twine 7 needs py3.10+ (Metadata-Version 2.5)

# Tests & Linting
coverage[toml]==7.5.4
Expand Down
6 changes: 5 additions & 1 deletion scripts/build
Original file line number Diff line number Diff line change
Expand Up @@ -11,5 +11,9 @@ fi
set -x

${PREFIX}python -m build
${PREFIX}twine check dist/*
# twine>=7 (needed for Metadata-Version 2.5) requires Python 3.10+,
# so only check the distribution where twine is installed.
if ${PREFIX}python -c "import twine" 2>/dev/null; then
${PREFIX}twine check dist/*
fi
${PREFIX}mkdocs build
150 changes: 150 additions & 0 deletions tests/_sync/test_http2_thread_safety.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,150 @@
"""Regression test for encode/httpx#3566.

`httpx.Client(http2=True)` shares a single connection across threads (HTTP/2
multiplexing). That connection is `httpcore.HTTP2Connection`, which wraps one
`h2.H2Connection` state machine. The send path — stream-ID allocation, the
`_events` mapping, and the HPACK encoding in `send_headers` — used to be
completely unserialized, so concurrent threads could corrupt the state
machine ("deque mutated during iteration", "dictionary changed size during
iteration", `StreamIDTooLowError`).

This test drives one shared `HTTP2Connection` from multiple threads over a
fully in-memory fake HTTP/2 server (a real server-side `h2` connection
guarded by its own lock, so only the *client-side* httpcore code is under
test) and asserts every request succeeds.
"""

import threading
import time
import typing

import h2.config
import h2.connection
import h2.events

import httpcore


class FakeStream(httpcore.NetworkStream):
"""In-memory full-duplex socket backed by a real server-side h2 connection."""

def __init__(self) -> None:
self._server = h2.connection.H2Connection(
config=h2.config.H2Configuration(client_side=False)
)
self._server.initiate_connection()
self._lock = threading.Lock()
self._cond = threading.Condition(self._lock)
self._closed = False
self._responded: typing.Set[int] = set()

def _respond(self, stream_id: int) -> None:
self._server.send_headers(
stream_id,
[(":status", "200"), ("content-length", "2")],
end_stream=False,
)
self._server.send_data(stream_id, b"ok", end_stream=True)

# -- NetworkStream interface --
def write(self, buffer: bytes, timeout: typing.Optional[float] = None) -> None:
with self._lock:
if buffer:
for event in self._server.receive_data(buffer):
if isinstance(event, h2.events.RequestReceived):
if event.stream_ended is not None:
self._responded.add(event.stream_id)
self._respond(event.stream_id)
elif isinstance(event, h2.events.DataReceived):
# Replenish the server's inbound flow-control window,
# like any real HTTP/2 server does.
self._server.acknowledge_received_data(
event.flow_controlled_length, event.stream_id
)
elif isinstance(event, h2.events.StreamEnded):
if event.stream_id not in self._responded:
self._responded.add(event.stream_id)
self._respond(event.stream_id)
self._cond.notify_all()

def read(self, max_bytes: int, timeout: typing.Optional[float] = None) -> bytes:
deadline = None if timeout is None else time.monotonic() + timeout
with self._lock:
while True:
data: bytes = self._server.data_to_send(max_bytes)
if data:
return data
if self._closed: # pragma: no cover
return b""
remaining = None if deadline is None else deadline - time.monotonic() # pragma: no cover
if remaining is not None and remaining <= 0: # pragma: no cover
return b""
self._cond.wait(timeout=1.0) # pragma: no cover

def close(self) -> None: # pragma: no cover
with self._lock:
self._closed = True
self._cond.notify_all()

def get_extra_info(self, info: str) -> typing.Any: # pragma: no cover
return None


def test_http2_connection_is_thread_safe():
"""
One shared HTTP2Connection must survive N threads racing the send path.

Without the `_send_lock` fix, this fails with errors such as
"deque mutated during iteration", "dictionary changed size during
iteration", `StreamIDTooLowError`, or `KeyError`.
"""
n_threads = 8
n_requests = 50

origin = httpcore.Origin(b"https", b"example.org", 443)
connection = httpcore.HTTP2Connection(origin=origin, stream=FakeStream())

errors = []
errors_lock = threading.Lock()

def worker(worker_id):
for i in range(n_requests):
body = b"x" * 64 if i % 2 else b""
headers = [
(b"host", b"example.org"),
(b"user-agent", b"thread-safety-test"),
(b"x-request", f"{worker_id}-{i}".encode()),
]
if body:
headers.append((b"content-length", str(len(body)).encode()))
request = httpcore.Request(
"POST" if body else "GET",
f"https://example.org/{worker_id}/{i}",
headers=headers,
content=body,
extensions={"timeout": {"read": 10, "write": 10}},
)
try:
response = connection.handle_request(request)
content = response.read()
response.close()
assert response.status == 200
assert content == b"ok"
except Exception as exc: # noqa: BLE001 # pragma: no cover
with errors_lock:
errors.append(exc)

threads = [
threading.Thread(target=worker, args=(w,), name=f"h2-worker-{w}")
for w in range(n_threads)
]
for thread in threads:
thread.start()
for thread in threads:
thread.join(timeout=120)

assert not any(thread.is_alive() for thread in threads), "worker threads hung"
assert errors == [], (
f"{len(errors)} request(s) failed out of {n_threads * n_requests}: "
f"{sorted({type(e).__name__ for e in errors})}"
)
Loading