mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 20:28:23 +00:00
fix: bound proxy request body reads by size and time
This commit is contained in:
@@ -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
@@ -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
|
||||
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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()
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user