diff --git a/routstr/proxy.py b/routstr/proxy.py index d9e7bc8d..ce50f528 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -444,7 +444,7 @@ async def _read_bounded_body(request: Request) -> bytes | Response: return bytes(body) try: - return await asyncio.wait_for(read(), timeout) + body = await asyncio.wait_for(read(), timeout) except _BodyLimitExceeded: error_type, message, status = ( "invalid_request", @@ -457,6 +457,12 @@ async def _read_bounded_body(request: Request) -> bytes | Response: f"Request body not received within {timeout} seconds", 408, ) + else: + # Draining the stream leaves Starlette unable to serve a second read. + # Cache the body so later readers (EHBP forwarding, upstream stream + # passthrough) get it instead of "Stream consumed". + request._body = body + return body return create_error_response(error_type, message, status, request=request) diff --git a/tests/unit/test_bounded_request_body.py b/tests/unit/test_bounded_request_body.py index e6c1e363..8b6cd06b 100644 --- a/tests/unit/test_bounded_request_body.py +++ b/tests/unit/test_bounded_request_body.py @@ -7,6 +7,7 @@ from unittest.mock import ANY, AsyncMock, MagicMock, patch import pytest from fastapi.responses import Response +from starlette.requests import Request from routstr import proxy as proxy_module from routstr.core.settings import settings @@ -101,3 +102,39 @@ async def test_normal_request_reaches_proxy_with_body() -> None: assert response.status_code == 200 session_factory.assert_called_once() inner.assert_awaited_once_with(request, "v1/chat/completions", ANY, body) + + +def _starlette_request(body: bytes) -> Request: + messages: list[dict[str, Any]] = [ + {"type": "http.request", "body": body, "more_body": False} + ] + + async def receive() -> dict[str, Any]: + return messages.pop(0) if messages else {"type": "http.disconnect"} + + return Request( + { + "type": "http", + "method": "POST", + "headers": [(b"content-length", str(len(body)).encode())], + "path": "/v1/chat/completions", + "query_string": b"", + "state": {}, + }, + receive, + ) + + +@pytest.mark.asyncio +async def test_body_stays_readable_after_bounded_read() -> None: + """EHBP forwarding and upstream passthrough re-read the same request.""" + body = b'{"model": "test-model"}' + request = _starlette_request(body) + + assert await proxy_module._read_bounded_body(request) == body + + assert await request.body() == body + streamed = bytearray() + async for chunk in request.stream(): + streamed += chunk + assert bytes(streamed) == body