Skip to content

Commit 25fcc9f

Browse files
fix(transport): handle unread streaming bodies and close 402 responses on retry (#6)
- Safely materialize request body via try/except RequestNotRead before retrying 402 responses in both sync and async transport paths\n- Close original 402 response before issuing retry to prevent connection leaks\n- Extend _clone_request_with_headers with explicit content parameter for deterministic body replay\n- Add comprehensive transport tests: passthrough, signing-failure fallback, close delegation, and 402 response cleanup verification
1 parent 66db9c8 commit 25fcc9f

2 files changed

Lines changed: 324 additions & 5 deletions

File tree

src/x402_openai/_transport.py

Lines changed: 30 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -23,15 +23,25 @@
2323
def _clone_request_with_headers(
2424
original: httpx.Request,
2525
extra_headers: dict[str, str],
26+
*,
27+
content: bytes | None = None,
2628
) -> httpx.Request:
27-
"""Clone *original* and merge *extra_headers* into the copy."""
29+
"""Clone *original* and merge *extra_headers* into the copy.
30+
31+
Parameters
32+
----------
33+
content:
34+
Optional explicit body bytes to use for the cloned request.
35+
When omitted, ``original.content`` is used and must already be
36+
materialized by the caller.
37+
"""
2838
headers = dict(original.headers)
2939
headers.update(extra_headers)
3040
return httpx.Request(
3141
method=original.method,
3242
url=original.url,
3343
headers=headers,
34-
content=original.content,
44+
content=original.content if content is None else content,
3545
extensions=dict(original.extensions),
3646
)
3747

@@ -78,7 +88,15 @@ def handle_request(self, request: httpx.Request) -> httpx.Response:
7888
logger.exception("x402: payment signing failed")
7989
return response
8090

81-
retry = _clone_request_with_headers(request, payment_headers)
91+
try:
92+
body = request.content
93+
except httpx.RequestNotRead:
94+
# Some transports/proxies can short-circuit with 402 before consuming
95+
# the request body. Ensure we materialize it so the retry is replayable.
96+
body = request.read()
97+
98+
retry = _clone_request_with_headers(request, payment_headers, content=body)
99+
response.close()
82100
return self._inner.handle_request(retry)
83101

84102
def close(self) -> None:
@@ -128,7 +146,15 @@ async def handle_async_request(self, request: httpx.Request) -> httpx.Response:
128146
logger.exception("x402: payment signing failed")
129147
return response
130148

131-
retry = _clone_request_with_headers(request, payment_headers)
149+
try:
150+
body = request.content
151+
except httpx.RequestNotRead:
152+
# Some transports/proxies can short-circuit with 402 before consuming
153+
# the request body. Ensure we materialize it so the retry is replayable.
154+
body = await request.aread()
155+
156+
retry = _clone_request_with_headers(request, payment_headers, content=body)
157+
await response.aclose()
132158
return await self._inner.handle_async_request(retry)
133159

134160
async def aclose(self) -> None:

tests/test_transport.py

Lines changed: 294 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,8 +1,18 @@
11
from __future__ import annotations
22

3+
from typing import TYPE_CHECKING
4+
35
import httpx
6+
import pytest
7+
8+
if TYPE_CHECKING:
9+
from collections.abc import AsyncIterator, Iterator
410

5-
from x402_openai._transport import _clone_request_with_headers
11+
from x402_openai._transport import (
12+
AsyncX402Transport,
13+
X402Transport,
14+
_clone_request_with_headers,
15+
)
616

717

818
def test_clone_request_merges_headers() -> None:
@@ -37,3 +47,286 @@ def test_clone_request_preserves_extensions() -> None:
3747
cloned = _clone_request_with_headers(original, {"x-payment": "signed"})
3848

3949
assert cloned.extensions == original.extensions
50+
51+
52+
class _FakeX402ClientSync:
53+
def handle_402_response(
54+
self,
55+
headers: dict[str, str],
56+
body: bytes,
57+
) -> tuple[dict[str, str], dict[str, str]]:
58+
assert headers["x-402"] == "required"
59+
assert body == b"challenge"
60+
return {"x-payment": "signed"}, {}
61+
62+
63+
class _FailingX402ClientSync:
64+
def handle_402_response(
65+
self,
66+
headers: dict[str, str],
67+
body: bytes,
68+
) -> tuple[dict[str, str], dict[str, str]]:
69+
raise RuntimeError("signing failed")
70+
71+
72+
class _ShortCircuit402Transport(httpx.BaseTransport):
73+
def __init__(self) -> None:
74+
self.calls = 0
75+
self.retry_body = b""
76+
self.retry_headers: dict[str, str] = {}
77+
78+
def handle_request(self, request: httpx.Request) -> httpx.Response:
79+
self.calls += 1
80+
if self.calls == 1:
81+
# Intentionally do not consume request body.
82+
return httpx.Response(402, headers={"x-402": "required"}, content=b"challenge")
83+
84+
self.retry_body = request.read()
85+
self.retry_headers = dict(request.headers)
86+
return httpx.Response(200, content=b"ok")
87+
88+
89+
class _PassthroughTransport(httpx.BaseTransport):
90+
def __init__(self) -> None:
91+
self.calls = 0
92+
93+
def handle_request(self, request: httpx.Request) -> httpx.Response:
94+
self.calls += 1
95+
return httpx.Response(200, content=b"ok")
96+
97+
98+
class _CloseTracking402Transport(httpx.BaseTransport):
99+
def __init__(self) -> None:
100+
self.calls = 0
101+
self.first_response_closed = False
102+
103+
def handle_request(self, request: httpx.Request) -> httpx.Response:
104+
self.calls += 1
105+
if self.calls == 1:
106+
response = httpx.Response(402, headers={"x-402": "required"}, content=b"challenge")
107+
original_close = response.close
108+
109+
def tracked_close() -> None:
110+
self.first_response_closed = True
111+
original_close()
112+
113+
response.close = tracked_close # type: ignore[method-assign]
114+
return response
115+
return httpx.Response(200, content=b"ok")
116+
117+
118+
class _CloseDelegatingTransport(httpx.BaseTransport):
119+
def __init__(self) -> None:
120+
self.closed = False
121+
122+
def handle_request(self, request: httpx.Request) -> httpx.Response:
123+
return httpx.Response(200, content=b"ok")
124+
125+
def close(self) -> None:
126+
self.closed = True
127+
128+
129+
def _iter_json() -> Iterator[bytes]:
130+
yield b'{"prompt":'
131+
yield b'"hi"}'
132+
133+
134+
def test_sync_transport_retries_even_if_402_response_short_circuits_body() -> None:
135+
inner = _ShortCircuit402Transport()
136+
transport = X402Transport(_FakeX402ClientSync(), inner=inner)
137+
138+
request = httpx.Request("POST", "https://example.com/v1/chat", content=_iter_json())
139+
response = transport.handle_request(request)
140+
141+
assert response.status_code == 200
142+
assert inner.calls == 2
143+
assert inner.retry_headers["x-payment"] == "signed"
144+
assert inner.retry_body == b'{"prompt":"hi"}'
145+
146+
147+
def test_sync_transport_passes_through_non_402_response() -> None:
148+
inner = _PassthroughTransport()
149+
transport = X402Transport(_FakeX402ClientSync(), inner=inner)
150+
151+
response = transport.handle_request(httpx.Request("GET", "https://example.com/v1/models"))
152+
153+
assert response.status_code == 200
154+
assert inner.calls == 1
155+
156+
157+
def test_sync_transport_returns_original_402_when_signing_fails() -> None:
158+
inner = _ShortCircuit402Transport()
159+
transport = X402Transport(_FailingX402ClientSync(), inner=inner)
160+
161+
response = transport.handle_request(
162+
httpx.Request("POST", "https://example.com/v1/chat", content=_iter_json())
163+
)
164+
165+
assert response.status_code == 402
166+
assert inner.calls == 1
167+
168+
169+
def test_sync_transport_closes_original_402_before_retry() -> None:
170+
inner = _CloseTracking402Transport()
171+
transport = X402Transport(_FakeX402ClientSync(), inner=inner)
172+
173+
response = transport.handle_request(
174+
httpx.Request("POST", "https://example.com/v1/chat", content=_iter_json())
175+
)
176+
177+
assert response.status_code == 200
178+
assert inner.calls == 2
179+
assert inner.first_response_closed is True
180+
181+
182+
def test_sync_transport_close_delegates_to_inner_transport() -> None:
183+
inner = _CloseDelegatingTransport()
184+
transport = X402Transport(_FakeX402ClientSync(), inner=inner)
185+
186+
transport.close()
187+
188+
assert inner.closed is True
189+
190+
191+
class _FakeX402ClientAsync:
192+
async def handle_402_response(
193+
self,
194+
headers: dict[str, str],
195+
body: bytes,
196+
) -> tuple[dict[str, str], dict[str, str]]:
197+
assert headers["x-402"] == "required"
198+
assert body == b"challenge"
199+
return {"x-payment": "signed"}, {}
200+
201+
202+
class _FailingX402ClientAsync:
203+
async def handle_402_response(
204+
self,
205+
headers: dict[str, str],
206+
body: bytes,
207+
) -> tuple[dict[str, str], dict[str, str]]:
208+
raise RuntimeError("signing failed")
209+
210+
211+
class _ShortCircuit402AsyncTransport(httpx.AsyncBaseTransport):
212+
def __init__(self) -> None:
213+
self.calls = 0
214+
self.retry_body = b""
215+
self.retry_headers: dict[str, str] = {}
216+
217+
async def handle_async_request(self, request: httpx.Request) -> httpx.Response:
218+
self.calls += 1
219+
if self.calls == 1:
220+
# Intentionally do not consume request body.
221+
return httpx.Response(402, headers={"x-402": "required"}, content=b"challenge")
222+
223+
self.retry_body = await request.aread()
224+
self.retry_headers = dict(request.headers)
225+
return httpx.Response(200, content=b"ok")
226+
227+
228+
class _PassthroughAsyncTransport(httpx.AsyncBaseTransport):
229+
def __init__(self) -> None:
230+
self.calls = 0
231+
232+
async def handle_async_request(self, request: httpx.Request) -> httpx.Response:
233+
self.calls += 1
234+
return httpx.Response(200, content=b"ok")
235+
236+
237+
class _CloseTracking402AsyncTransport(httpx.AsyncBaseTransport):
238+
def __init__(self) -> None:
239+
self.calls = 0
240+
self.first_response_closed = False
241+
242+
async def handle_async_request(self, request: httpx.Request) -> httpx.Response:
243+
self.calls += 1
244+
if self.calls == 1:
245+
response = httpx.Response(402, headers={"x-402": "required"}, content=b"challenge")
246+
original_aclose = response.aclose
247+
248+
async def tracked_aclose() -> None:
249+
self.first_response_closed = True
250+
await original_aclose()
251+
252+
response.aclose = tracked_aclose # type: ignore[method-assign]
253+
return response
254+
return httpx.Response(200, content=b"ok")
255+
256+
257+
class _CloseDelegatingAsyncTransport(httpx.AsyncBaseTransport):
258+
def __init__(self) -> None:
259+
self.closed = False
260+
261+
async def handle_async_request(self, request: httpx.Request) -> httpx.Response:
262+
return httpx.Response(200, content=b"ok")
263+
264+
async def aclose(self) -> None:
265+
self.closed = True
266+
267+
268+
async def _aiter_json() -> AsyncIterator[bytes]:
269+
yield b'{"prompt":'
270+
yield b'"hi"}'
271+
272+
273+
@pytest.mark.asyncio
274+
async def test_async_transport_retries_even_if_402_response_short_circuits_body() -> None:
275+
inner = _ShortCircuit402AsyncTransport()
276+
transport = AsyncX402Transport(_FakeX402ClientAsync(), inner=inner)
277+
278+
request = httpx.Request("POST", "https://example.com/v1/chat", content=_aiter_json())
279+
response = await transport.handle_async_request(request)
280+
281+
assert response.status_code == 200
282+
assert inner.calls == 2
283+
assert inner.retry_headers["x-payment"] == "signed"
284+
assert inner.retry_body == b'{"prompt":"hi"}'
285+
286+
287+
@pytest.mark.asyncio
288+
async def test_async_transport_passes_through_non_402_response() -> None:
289+
inner = _PassthroughAsyncTransport()
290+
transport = AsyncX402Transport(_FakeX402ClientAsync(), inner=inner)
291+
292+
response = await transport.handle_async_request(httpx.Request("GET", "https://example.com/v1/models"))
293+
294+
assert response.status_code == 200
295+
assert inner.calls == 1
296+
297+
298+
@pytest.mark.asyncio
299+
async def test_async_transport_returns_original_402_when_signing_fails() -> None:
300+
inner = _ShortCircuit402AsyncTransport()
301+
transport = AsyncX402Transport(_FailingX402ClientAsync(), inner=inner)
302+
303+
response = await transport.handle_async_request(
304+
httpx.Request("POST", "https://example.com/v1/chat", content=_aiter_json())
305+
)
306+
307+
assert response.status_code == 402
308+
assert inner.calls == 1
309+
310+
311+
@pytest.mark.asyncio
312+
async def test_async_transport_closes_original_402_before_retry() -> None:
313+
inner = _CloseTracking402AsyncTransport()
314+
transport = AsyncX402Transport(_FakeX402ClientAsync(), inner=inner)
315+
316+
response = await transport.handle_async_request(
317+
httpx.Request("POST", "https://example.com/v1/chat", content=_aiter_json())
318+
)
319+
320+
assert response.status_code == 200
321+
assert inner.calls == 2
322+
assert inner.first_response_closed is True
323+
324+
325+
@pytest.mark.asyncio
326+
async def test_async_transport_aclose_delegates_to_inner_transport() -> None:
327+
inner = _CloseDelegatingAsyncTransport()
328+
transport = AsyncX402Transport(_FakeX402ClientAsync(), inner=inner)
329+
330+
await transport.aclose()
331+
332+
assert inner.closed is True

0 commit comments

Comments
 (0)