From 57c6dec506a166973b35ac9caaa70accea405a5e Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sun, 27 Sep 2026 02:20:19 +0200 Subject: [PATCH] fix: bound proxy request body reads by size and time --- routstr/core/settings.py | 8 ++ routstr/proxy.py | 78 ++++++++++--- .../test_proxy_session_lifecycle.py | 6 +- tests/unit/proxy_test_utils.py | 27 +++++ tests/unit/test_bounded_request_body.py | 103 ++++++++++++++++++ tests/unit/test_model_path_routing.py | 10 +- tests/unit/test_proxy_session_lifecycle.py | 17 ++- tests/unit/test_stale_reservations.py | 7 +- tests/unit/test_tinfoil_integration.py | 7 +- tests/unit/test_upstream_rate_limit.py | 9 +- 10 files changed, 237 insertions(+), 35 deletions(-) create mode 100644 tests/unit/proxy_test_utils.py create mode 100644 tests/unit/test_bounded_request_body.py diff --git a/routstr/core/settings.py b/routstr/core/settings.py index da503a50..39448a58 100644 --- a/routstr/core/settings.py +++ b/routstr/core/settings.py @@ -125,6 +125,14 @@ class Settings(BaseSettings): # widens what the provider credential can be spent against, so wildcards # and prefixes are not supported. proxy_extra_allowed_paths: str = Field(default="", env="PROXY_EXTRA_ALLOWED_PATHS") + # Bound the client request body: a slow or oversized upload otherwise blocks + # the proxy before authentication and holds server resources for its duration. + request_body_timeout_seconds: float = Field( + default=30.0, gt=0, env="REQUEST_BODY_TIMEOUT_SECONDS" + ) + max_request_body_bytes: int = Field( + default=20 * 1024 * 1024, gt=0, env="MAX_REQUEST_BODY_BYTES" + ) tor_proxy_url: str = Field(default="socks5://127.0.0.1:9050", env="TOR_PROXY_URL") providers_refresh_interval_seconds: int = Field( default=0, env="PROVIDERS_REFRESH_INTERVAL_SECONDS" diff --git a/routstr/proxy.py b/routstr/proxy.py index ae36d1d4..d9e7bc8d 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -3,7 +3,7 @@ import inspect import json from typing import Any -from fastapi import APIRouter, Depends, HTTPException, Request +from fastapi import APIRouter, HTTPException, Request from fastapi.responses import Response, StreamingResponse from sqlmodel import select @@ -21,7 +21,6 @@ from .core.db import ( ModelRow, UpstreamProviderRow, create_session, - get_session, ) from .core.error_scope import ( ERROR_SCOPE_UPSTREAM, @@ -418,23 +417,71 @@ def _attribute_request( request.state.provider = upstream.provider_type -@proxy_router.api_route("/{path:path}", methods=["GET", "POST"], response_model=None) -async def proxy( - request: Request, path: str, session: AsyncSession = Depends(get_session) -) -> Response | StreamingResponse: - """Run proxy setup in a short request session, never across response streaming.""" +class _BodyLimitExceeded(Exception): + """The client body is larger than ``max_request_body_bytes``.""" + + +async def _read_bounded_body(request: Request) -> bytes | Response: + """Read the request body under a size and time bound. + + Returns the body, or the error response to send instead. Both bounds run + before any authentication or DB work, so an oversized or slowly uploaded + body cannot occupy the request for longer than the timeout. + """ + max_bytes = settings.max_request_body_bytes + timeout = settings.request_body_timeout_seconds + + async def read() -> bytes: + declared = request.headers.get("content-length", "") + if declared.isdigit() and int(declared) > max_bytes: + raise _BodyLimitExceeded + body = bytearray() + async for chunk in request.stream(): + body += chunk + # Chunked uploads declare no length, so the cap is enforced here. + if len(body) > max_bytes: + raise _BodyLimitExceeded + return bytes(body) + try: - return await _proxy(request, path, session) - finally: - # FastAPI yield dependencies normally close after the response body is - # sent. Close explicitly so a long stream cannot retain DB resources. - close_result = session.close() - if inspect.isawaitable(close_result): - await close_result + return await asyncio.wait_for(read(), timeout) + except _BodyLimitExceeded: + error_type, message, status = ( + "invalid_request", + f"Request body exceeds the {max_bytes} byte limit", + 413, + ) + except asyncio.TimeoutError: + error_type, message, status = ( + "timeout", + f"Request body not received within {timeout} seconds", + 408, + ) + return create_error_response(error_type, message, status, request=request) + + +@proxy_router.api_route("/{path:path}", methods=["GET", "POST"], response_model=None) +async def proxy(request: Request, path: str) -> Response | StreamingResponse: + """Run proxy setup in a short request session, never across response streaming.""" + # Read the body before opening a session: a slow uploader must not hold a + # DB connection while its request trickles in. + request_body = await _read_bounded_body(request) + if isinstance(request_body, Response): + return request_body + + async with create_session() as session: + try: + return await _proxy(request, path, session, request_body) + finally: + # Close explicitly so a long stream cannot retain DB resources + # while its response body is being sent. + close_result = session.close() + if inspect.isawaitable(close_result): + await close_result async def _proxy( - request: Request, path: str, session: AsyncSession + request: Request, path: str, session: AsyncSession, request_body: bytes ) -> Response | StreamingResponse: # Screen the path before any routing decision: reject ambiguous spellings, # then require a known API prefix so nothing unknown is forwarded with the @@ -449,7 +496,6 @@ async def _proxy( return build_not_found_response(request, path) is_responses_api = path.startswith("v1/responses") or path.startswith("responses") - request_body = await request.body() # EHBP (Encrypted HTTP Body Protocol) requests carry an Ehbp-Encapsulated-Key # header and a binary HPKE-sealed body. The proxy cannot parse the body to diff --git a/tests/integration/test_proxy_session_lifecycle.py b/tests/integration/test_proxy_session_lifecycle.py index 9a70f294..4233b653 100644 --- a/tests/integration/test_proxy_session_lifecycle.py +++ b/tests/integration/test_proxy_session_lifecycle.py @@ -33,7 +33,6 @@ async def test_authenticated_proxy_releases_db_connection_before_upstream_header request = MagicMock() request.method = "POST" request.headers = {"authorization": "Bearer test-key"} - request.body = AsyncMock(return_value=json.dumps({"model": "test-model"}).encode()) request.url.path = "/v1/chat/completions" request.state.request_id = "pool-hold-regression" @@ -59,7 +58,10 @@ async def test_authenticated_proxy_releases_db_connection_before_upstream_header patch("routstr.proxy.get_bearer_token_key", AsyncMock(return_value=key)), ): response = await proxy_module._proxy( - request, "v1/chat/completions", integration_session + request, + "v1/chat/completions", + integration_session, + json.dumps({"model": "test-model"}).encode(), ) assert response.status_code == 200 diff --git a/tests/unit/proxy_test_utils.py b/tests/unit/proxy_test_utils.py new file mode 100644 index 00000000..a4fe4d3b --- /dev/null +++ b/tests/unit/proxy_test_utils.py @@ -0,0 +1,27 @@ +"""Helpers for driving ``routstr.proxy.proxy`` with mocked request and session.""" + +from collections.abc import AsyncIterator +from contextlib import asynccontextmanager +from typing import Any +from unittest.mock import MagicMock, patch + +from routstr import proxy as proxy_module + + +def mock_request_stream(request: MagicMock, body: bytes) -> None: + """Give a mocked request a readable body stream (the proxy reads the stream).""" + + async def stream() -> AsyncIterator[bytes]: + yield body + + request.stream = stream + + +def patch_proxy_session(session: Any) -> Any: + """Make the proxy route use ``session`` instead of opening its own.""" + + @asynccontextmanager + async def factory() -> AsyncIterator[Any]: + yield session + + return patch.object(proxy_module, "create_session", factory) diff --git a/tests/unit/test_bounded_request_body.py b/tests/unit/test_bounded_request_body.py new file mode 100644 index 00000000..e6c1e363 --- /dev/null +++ b/tests/unit/test_bounded_request_body.py @@ -0,0 +1,103 @@ +"""Bounded request-body read: size cap, read timeout, and late DB session.""" + +import asyncio +from collections.abc import AsyncIterator +from typing import Any +from unittest.mock import ANY, AsyncMock, MagicMock, patch + +import pytest +from fastapi.responses import Response + +from routstr import proxy as proxy_module +from routstr.core.settings import settings + + +def _make_request(headers: dict[str, str], chunks: list[bytes]) -> MagicMock: + request = MagicMock() + request.method = "POST" + request.headers = headers + request.state.request_id = "req-bounded-body" + request.consumed = [] + + async def stream() -> AsyncIterator[bytes]: + for chunk in chunks: + request.consumed.append(chunk) + yield chunk + + request.stream = stream + return request + + +def _slow_request(delay: float) -> MagicMock: + request = MagicMock() + request.method = "POST" + request.headers = {} + request.state.request_id = "req-slow-body" + + async def stream() -> AsyncIterator[bytes]: + yield b"{" + await asyncio.sleep(delay) + yield b"}" + + request.stream = stream + return request + + +async def _run(request: MagicMock) -> tuple[Any, MagicMock, AsyncMock]: + """Run the proxy route with the session factory and _proxy stubbed out.""" + session_factory = MagicMock() + inner = AsyncMock(return_value=Response(status_code=200)) + with ( + patch.object(proxy_module, "create_session", session_factory), + patch.object(proxy_module, "_proxy", inner), + ): + response = await proxy_module.proxy(request, "v1/chat/completions") + return response, session_factory, inner + + +@pytest.mark.asyncio +async def test_oversize_content_length_rejected_without_reading() -> None: + request = _make_request({"content-length": "999999999"}, [b"x" * 16]) + + response, session_factory, inner = await _run(request) + + assert response.status_code == 413 + assert request.consumed == [] + inner.assert_not_awaited() + session_factory.assert_not_called() + + +@pytest.mark.asyncio +async def test_oversize_chunked_body_rejected_mid_stream() -> None: + with patch.object(settings, "max_request_body_bytes", 8): + request = _make_request({}, [b"1234", b"5678", b"9012", b"3456"]) + response, session_factory, inner = await _run(request) + + assert response.status_code == 413 + # Reading stops as soon as the cap is exceeded; the last chunk is never read. + assert request.consumed == [b"1234", b"5678", b"9012"] + inner.assert_not_awaited() + session_factory.assert_not_called() + + +@pytest.mark.asyncio +async def test_slow_body_times_out() -> None: + with patch.object(settings, "request_body_timeout_seconds", 0.05): + request = _slow_request(delay=5) + response, session_factory, inner = await _run(request) + + assert response.status_code == 408 + inner.assert_not_awaited() + session_factory.assert_not_called() + + +@pytest.mark.asyncio +async def test_normal_request_reaches_proxy_with_body() -> None: + body = b'{"model": "test-model"}' + request = _make_request({"content-length": str(len(body))}, [body]) + + response, session_factory, inner = await _run(request) + + assert response.status_code == 200 + session_factory.assert_called_once() + inner.assert_awaited_once_with(request, "v1/chat/completions", ANY, body) diff --git a/tests/unit/test_model_path_routing.py b/tests/unit/test_model_path_routing.py index 5b23720f..e2c3e52f 100644 --- a/tests/unit/test_model_path_routing.py +++ b/tests/unit/test_model_path_routing.py @@ -17,6 +17,8 @@ from routstr.core.error_scope import ( ) from routstr.upstream.model_paths import decode_model_path, encode_model_path +from .proxy_test_utils import mock_request_stream, patch_proxy_session + MODEL_ID = "test-model" @@ -38,7 +40,7 @@ def _make_request(headers: dict[str, str], body: bytes) -> MagicMock: request = MagicMock() request.method = "POST" request.headers = headers - request.body = AsyncMock(return_value=body) + mock_request_stream(request, body) request.state = MagicMock() request.state.request_id = "req-model-path" return request @@ -72,8 +74,9 @@ async def _run_proxy( proxy_module, "pay_for_request", AsyncMock(return_value=reservation) ), patch.object(proxy_module, "revert_pay_for_request", AsyncMock()), + patch_proxy_session(MagicMock()), ): - return await proxy_module.proxy(request, path, session=MagicMock()) + return await proxy_module.proxy(request, path) def test_decode_model_path_round_trips_encode() -> None: @@ -529,8 +532,9 @@ async def test_unsupported_endpoint_pins_fail_before_payment( patch.object( proxy_module, "get_candidates", return_value=[(MagicMock(), upstream)] ), + patch_proxy_session(MagicMock()), ): - response = await proxy_module.proxy(request, path, MagicMock()) + response = await proxy_module.proxy(request, path) assert response.status_code == 400 assert json.loads(response.body)["error"]["type"] == "unsupported_request" payment.assert_not_called() diff --git a/tests/unit/test_proxy_session_lifecycle.py b/tests/unit/test_proxy_session_lifecycle.py index 5d0416d5..7cc01b03 100644 --- a/tests/unit/test_proxy_session_lifecycle.py +++ b/tests/unit/test_proxy_session_lifecycle.py @@ -6,6 +6,8 @@ from fastapi.responses import StreamingResponse from routstr import proxy as proxy_module +from .proxy_test_utils import mock_request_stream, patch_proxy_session + @pytest.mark.asyncio async def test_proxy_closes_request_session_before_returning_response() -> None: @@ -15,9 +17,11 @@ async def test_proxy_closes_request_session_before_returning_response() -> None: request.headers = {"accept": "application/json"} request.url.path = "/not-an-api-route" request.state.request_id = "test-request" + mock_request_stream(request, b"") session = AsyncMock() - response = await proxy_module.proxy(request, "not-an-api-route", session=session) + with patch_proxy_session(session): + response = await proxy_module.proxy(request, "not-an-api-route") assert response.status_code == 404 session.close.assert_awaited_once() @@ -26,6 +30,8 @@ async def test_proxy_closes_request_session_before_returning_response() -> None: @pytest.mark.asyncio async def test_proxy_session_is_closed_before_first_stream_chunk() -> None: request = MagicMock() + request.headers = {} + mock_request_stream(request, b"") session = AsyncMock() async def stream() -> AsyncIterator[bytes]: @@ -33,10 +39,11 @@ async def test_proxy_session_is_closed_before_first_stream_chunk() -> None: yield b"chunk" upstream_response = StreamingResponse(stream()) - with patch("routstr.proxy._proxy", AsyncMock(return_value=upstream_response)): - response = await proxy_module.proxy( - request, "v1/chat/completions", session=session - ) + with ( + patch("routstr.proxy._proxy", AsyncMock(return_value=upstream_response)), + patch_proxy_session(session), + ): + response = await proxy_module.proxy(request, "v1/chat/completions") assert isinstance(response, StreamingResponse) chunks = [chunk async for chunk in response.body_iterator] diff --git a/tests/unit/test_stale_reservations.py b/tests/unit/test_stale_reservations.py index 2b5d8e31..558fdda5 100644 --- a/tests/unit/test_stale_reservations.py +++ b/tests/unit/test_stale_reservations.py @@ -29,6 +29,8 @@ from routstr.core.db import ( reset_all_reserved_balances, ) +from .proxy_test_utils import mock_request_stream, patch_proxy_session + def _make_engine() -> AsyncEngine: return create_async_engine( @@ -387,7 +389,7 @@ async def test_proxy_reverts_reservation_on_client_disconnect() -> None: request = MagicMock() request.method = "POST" request.headers = {"authorization": "Bearer sk-cancelkey"} - request.body = AsyncMock(return_value=b'{"model": "test-model"}') + mock_request_stream(request, b'{"model": "test-model"}') upstream = MagicMock() upstream.provider_type = "test" @@ -420,8 +422,9 @@ async def test_proxy_reverts_reservation_on_client_disconnect() -> None: AsyncMock(return_value=reservation_snapshot), ), patch.object(proxy_module, "revert_pay_for_request", revert_mock), + patch_proxy_session(session), ): with pytest.raises(asyncio.CancelledError): - await proxy_module.proxy(request, "v1/chat/completions", session=session) + await proxy_module.proxy(request, "v1/chat/completions") revert_mock.assert_awaited_once_with(key, session, 1000, reservation_snapshot) diff --git a/tests/unit/test_tinfoil_integration.py b/tests/unit/test_tinfoil_integration.py index ec5e8c11..14715f58 100644 --- a/tests/unit/test_tinfoil_integration.py +++ b/tests/unit/test_tinfoil_integration.py @@ -31,6 +31,8 @@ from routstr.upstream.tinfoil import ( ) from routstr.upstream.tinfoil_trailer import TrailerResponse +from .proxy_test_utils import patch_proxy_session + # --------------------------------------------------------------------------- # parse_tinfoil_usage_metrics # --------------------------------------------------------------------------- @@ -1303,10 +1305,9 @@ async def test_bearer_key_config_422_releases_reservation_and_passes_through() - "routstr.upstream.ehbp.forward_with_trailer", AsyncMock(return_value=upstream_resp), ), + patch_proxy_session(session), ): - response = await proxy_module.proxy( - request, "v1/chat/completions", session=session - ) + response = await proxy_module.proxy(request, "v1/chat/completions") # The reservation was released despite the early passthrough return. revert_mock.assert_awaited_once_with(key, session, 1_000, reservation_snapshot) diff --git a/tests/unit/test_upstream_rate_limit.py b/tests/unit/test_upstream_rate_limit.py index c111b74a..0b73199c 100644 --- a/tests/unit/test_upstream_rate_limit.py +++ b/tests/unit/test_upstream_rate_limit.py @@ -28,6 +28,8 @@ from routstr.upstream.rate_limit import ( classify_rate_limit, ) +from .proxy_test_utils import mock_request_stream, patch_proxy_session + # The exact scenario from the issue, with a realistic (fake) org identifier. RAW_ORG_ID = "org-abc123XYZ456def" RATE_LIMIT_MESSAGE = ( @@ -353,7 +355,7 @@ async def test_proxy_loop_surfaces_rate_limit_and_reverts_once() -> None: request = MagicMock() request.method = "POST" request.headers = {"authorization": "Bearer sk-rlkey"} - request.body = AsyncMock(return_value=b'{"model": "test-model"}') + mock_request_stream(request, b'{"model": "test-model"}') request.state = MagicMock() request.state.request_id = "req-rl" @@ -400,10 +402,9 @@ async def test_proxy_loop_surfaces_rate_limit_and_reverts_once() -> None: AsyncMock(return_value=reservation), ), patch.object(proxy_module, "revert_pay_for_request", revert_mock), + patch_proxy_session(session), ): - response = await proxy_module.proxy( - request, "v1/chat/completions", session=session - ) + response = await proxy_module.proxy(request, "v1/chat/completions") # Original 429 status and the stable code/details survive to the client. assert response.status_code == 429