fix: bound proxy request body reads by size and time

This commit is contained in:
9qeklajc
2026-09-28 23:59:36 +02:00
parent 8cc800ea67
commit 57c6dec506
10 changed files with 237 additions and 35 deletions
+8
View File
@@ -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"
+62 -16
View File
@@ -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
@@ -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
+27
View File
@@ -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)
+103
View File
@@ -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)
+7 -3
View File
@@ -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()
+12 -5
View File
@@ -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]
+5 -2
View File
@@ -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)
+4 -3
View File
@@ -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)
+5 -4
View File
@@ -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