mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge 03b05fe8c8 into b370996b9d
This commit is contained in:
commit
f2f7e9de5a
7 changed files with 559 additions and 15 deletions
|
|
@ -341,6 +341,7 @@ BEDROCK_REALTIME_SDK_DISTRIBUTION: Final = "aws-sdk-bedrock-runtime"
|
|||
BEDROCK_REALTIME_SDK_SUPPORTED_RANGE: Final = ">=0.10.0,<0.12.0"
|
||||
CLIENT_REQUESTED_MODEL_SCOPE_KEY: Final = "litellm.client_requested_model"
|
||||
MODEL_GROUP_ALIAS_RESOLVED_SCOPE_KEY: Final = "litellm.model_group_alias_resolved"
|
||||
NON_OBJECT_JSON_BODY_SCOPE_KEY: Final = "litellm.non_object_json_body"
|
||||
REALTIME_SESSION_SUCCESS_LOGGED_KEY: Final = "realtime_session_success_logged"
|
||||
REALTIME_SESSION_FAILURE_LOGGED_KEY: Final = "realtime_session_failure_logged"
|
||||
|
||||
|
|
|
|||
|
|
@ -13,6 +13,7 @@ from litellm.constants import (
|
|||
AZURE_SPEECH_PASS_THROUGH_ROUTE_PREFIX,
|
||||
CLIENT_REQUESTED_MODEL_SCOPE_KEY,
|
||||
MAX_REQUEST_BODY_SIZE_TO_REPAIR_MB,
|
||||
NON_OBJECT_JSON_BODY_SCOPE_KEY,
|
||||
)
|
||||
from litellm.proxy._types import ProxyException
|
||||
from litellm.proxy.common_utils.callback_utils import (
|
||||
|
|
@ -127,7 +128,14 @@ async def _read_request_body(request: Request | None) -> dict:
|
|||
- request: The request object to read the body from
|
||||
|
||||
Returns:
|
||||
- dict: Parsed request data as a dictionary or an empty dictionary if parsing fails
|
||||
- dict: Parsed request data as a dictionary. A body that is valid JSON but not an
|
||||
object reads as ``{}``, the same as an empty body: every field a caller can look
|
||||
for is a key, so a non-object body carries none of them. Honouring the annotation
|
||||
matters because auth reads the body before any route does, and both read it with
|
||||
``.get(...)``, so returning a list here surfaced as AttributeError -> 500 (#43711)
|
||||
|
||||
Raises:
|
||||
- ProxyException: 400, when the body is present but malformed
|
||||
"""
|
||||
try:
|
||||
if request is None:
|
||||
|
|
@ -208,9 +216,15 @@ async def _read_request_body(request: Request | None) -> dict:
|
|||
code=status.HTTP_400_BAD_REQUEST,
|
||||
)
|
||||
|
||||
# Cache the parsed result
|
||||
_safe_set_request_parsed_body(request=request, parsed_body=parsed_body)
|
||||
return parsed_body
|
||||
if isinstance(parsed_body, dict):
|
||||
_safe_set_request_parsed_body(request=request, parsed_body=parsed_body)
|
||||
return parsed_body
|
||||
|
||||
# Anything that forwards the parsed body upstream would now send an empty object in
|
||||
# place of the caller's payload, so mark the request for ``non_object_raw_body``.
|
||||
_mark_non_object_body(request=request)
|
||||
_safe_set_request_parsed_body(request=request, parsed_body={})
|
||||
return {}
|
||||
|
||||
except (json.JSONDecodeError, orjson.JSONDecodeError, ProxyException) as e:
|
||||
# Re-raise ProxyException as-is
|
||||
|
|
@ -230,6 +244,32 @@ def is_opaque_audio_pass_through_request(route: str, content_type: str) -> bool:
|
|||
)
|
||||
|
||||
|
||||
def _mark_non_object_body(request: Request | None) -> None:
|
||||
try:
|
||||
if request is not None:
|
||||
request.scope[NON_OBJECT_JSON_BODY_SCOPE_KEY] = True
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.debug("Unexpected error marking non-object request body - %s", e)
|
||||
|
||||
|
||||
async def non_object_raw_body(request: Request | None) -> bytes | None:
|
||||
"""The bytes of a body that is valid JSON but not an object, else ``None``.
|
||||
|
||||
``_read_request_body`` reads such a body as ``{}`` so every caller that treats it as a
|
||||
mapping is safe, which means forwarding the parsed body would send an empty object in
|
||||
place of the caller's payload. Passthrough sends these bytes instead.
|
||||
"""
|
||||
if request is None:
|
||||
return None
|
||||
scope: Final[object] = getattr(request, "scope", None)
|
||||
if not isinstance(scope, Mapping) or scope.get(NON_OBJECT_JSON_BODY_SCOPE_KEY) is not True:
|
||||
return None
|
||||
try:
|
||||
return await request.body()
|
||||
except RuntimeError:
|
||||
return None
|
||||
|
||||
|
||||
async def read_raw_json_body(request: Request | None) -> bytes | None:
|
||||
if request is None or _safe_get_request_parsed_body(request=request) is None:
|
||||
return None
|
||||
|
|
|
|||
|
|
@ -86,6 +86,7 @@ from litellm.proxy.common_utils.error_body_call_id import JSON_OBJECT, error_bod
|
|||
from litellm.proxy.common_utils.http_parsing_utils import (
|
||||
_read_request_body,
|
||||
_safe_get_request_headers,
|
||||
non_object_raw_body,
|
||||
)
|
||||
from litellm.proxy.common_utils.openai_error_payload import (
|
||||
LITELLM_CALL_ID_HEADER,
|
||||
|
|
@ -1073,6 +1074,8 @@ async def pass_through_request(
|
|||
|
||||
# parsed request body
|
||||
_parsed_body: dict | None = None
|
||||
# bytes of a non-object JSON body: forwarded verbatim, but unreadable to guardrails
|
||||
uninspectable_body: bytes | None = None
|
||||
# kwargs for pass through endpoint, contains metadata, litellm_params, call_type, litellm_call_id, passthrough_logging_payload
|
||||
kwargs: dict | None = None
|
||||
logging_obj: Logging | None = None
|
||||
|
|
@ -1120,6 +1123,11 @@ async def pass_through_request(
|
|||
_parsed_body = {}
|
||||
else:
|
||||
_parsed_body = await _read_request_body(request)
|
||||
if state_raw_body is None:
|
||||
# A non-object JSON body reads as ``{}``, so forwarding the parsed body would
|
||||
# send an empty object in place of the caller's own provider payload.
|
||||
uninspectable_body = await non_object_raw_body(request)
|
||||
state_raw_body = uninspectable_body
|
||||
verbose_proxy_logger.debug(
|
||||
"Pass through endpoint sending request to \nURL %s\nheaders: %s\nbody: %s\n",
|
||||
url,
|
||||
|
|
@ -1135,6 +1143,25 @@ async def pass_through_request(
|
|||
passthrough_guardrails_config=guardrails_config,
|
||||
)
|
||||
|
||||
if (
|
||||
guardrails_to_run
|
||||
and uninspectable_body is not None
|
||||
and PassthroughGuardrailHandler.any_inspects_the_request_body(guardrails_to_run)
|
||||
):
|
||||
# A request-inspecting guardrail reads the parsed body, which for a non-object
|
||||
# payload carries none of the caller's content, while the bytes forwarded upstream
|
||||
# carry all of it. Running it would report "inspected" on content nobody looked at,
|
||||
# so refuse instead. A post_call guardrail reads the response, so it does not gate.
|
||||
raise ProxyException(
|
||||
message=(
|
||||
"Guardrails are configured for this route and cannot inspect a JSON body "
|
||||
"that is not an object. Send the payload as a JSON object."
|
||||
),
|
||||
type="invalid_request_error",
|
||||
param="request_body",
|
||||
code=status.HTTP_400_BAD_REQUEST,
|
||||
)
|
||||
|
||||
# Add guardrails to metadata if any should run
|
||||
if guardrails_to_run and len(guardrails_to_run) > 0:
|
||||
if _parsed_body is None:
|
||||
|
|
@ -1914,6 +1941,13 @@ class _PassThroughRequestEnvelope(TypedDict, total=False):
|
|||
stream: bool | None
|
||||
|
||||
|
||||
def _passthrough_envelope(body: object) -> _PassThroughRequestEnvelope | None:
|
||||
"""``request.json()`` returns any JSON kind, and a passthrough body is the caller's own
|
||||
provider payload, so a list or scalar is legitimate here and simply carries no envelope
|
||||
fields. Reading it as one raised AttributeError -> 500 (#43711)."""
|
||||
return body if isinstance(body, dict) else None
|
||||
|
||||
|
||||
async def _parse_request_data_by_content_type(
|
||||
request: Request,
|
||||
) -> tuple[object, object, None, bool | None]:
|
||||
|
|
@ -1935,10 +1969,11 @@ async def _parse_request_data_by_content_type(
|
|||
if "application/json" in content_type:
|
||||
# ✅ Handle JSON
|
||||
try:
|
||||
body: _PassThroughRequestEnvelope = await request.json()
|
||||
query_params_data = body.get("query_params")
|
||||
custom_body_data = body.get("custom_body")
|
||||
stream = body.get("stream")
|
||||
body: _PassThroughRequestEnvelope | None = _passthrough_envelope(await request.json())
|
||||
if body is not None:
|
||||
query_params_data = body.get("query_params")
|
||||
custom_body_data = body.get("custom_body")
|
||||
stream = body.get("stream")
|
||||
except json.JSONDecodeError:
|
||||
# Handle requests with no body (e.g., DELETE requests)
|
||||
pass
|
||||
|
|
@ -1946,14 +1981,15 @@ async def _parse_request_data_by_content_type(
|
|||
# ✅ Try to parse as JSON first (handles misconfigured clients sending JSON with multipart content-type)
|
||||
# If that fails, skip parsing - pass_through_request will handle actual multipart
|
||||
try:
|
||||
body = await request.json()
|
||||
body = _passthrough_envelope(await request.json())
|
||||
# Successfully parsed as JSON - treat as JSON body
|
||||
query_params_data = body.get("query_params")
|
||||
custom_body_data = body.get("custom_body")
|
||||
stream = body.get("stream")
|
||||
# If custom_body is not set, use the entire body
|
||||
if custom_body_data is None and body:
|
||||
custom_body_data = body
|
||||
if body is not None:
|
||||
query_params_data = body.get("query_params")
|
||||
custom_body_data = body.get("custom_body")
|
||||
stream = body.get("stream")
|
||||
# If custom_body is not set, use the entire body
|
||||
if custom_body_data is None and body:
|
||||
custom_body_data = body
|
||||
except (json.JSONDecodeError, Exception):
|
||||
# Not JSON - this is actual multipart data
|
||||
# Skip parsing here to avoid consuming the request body stream
|
||||
|
|
|
|||
|
|
@ -7,15 +7,37 @@ Handles guardrail execution for passthrough endpoints with:
|
|||
- Automatic inheritance from org/team/key levels when enabled
|
||||
"""
|
||||
|
||||
from collections.abc import Collection
|
||||
from typing import Any, Final
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.proxy._types import (
|
||||
PassThroughGuardrailsConfig,
|
||||
PassThroughGuardrailSettings,
|
||||
UserAPIKeyAuth,
|
||||
)
|
||||
from litellm.proxy.pass_through_endpoints.jsonpath_extractor import JsonPathExtractor
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
|
||||
# The hooks that read the request body. ``post_call`` and ``logging_only`` read the response.
|
||||
_REQUEST_INSPECTING_EVENT_HOOKS: Final = (
|
||||
GuardrailEventHooks.pre_call,
|
||||
GuardrailEventHooks.during_call,
|
||||
)
|
||||
|
||||
|
||||
def _reads_the_request_body(callback: object, guardrail_names: Collection[str]) -> bool:
|
||||
if not isinstance(callback, CustomGuardrail) or callback.guardrail_name not in guardrail_names:
|
||||
return False
|
||||
# Same private predicate common_request_processing.py uses to ask a callback which
|
||||
# lifecycle hooks it is configured for, rather than re-deriving the matching rules.
|
||||
return any(
|
||||
callback._event_hook_is_event_type(hook) # pyright: ignore[reportPrivateUsage] # no public equivalent
|
||||
for hook in _REQUEST_INSPECTING_EVENT_HOOKS
|
||||
)
|
||||
|
||||
|
||||
# Type for raw guardrails config input (before normalization)
|
||||
# Can be a list of names or a dict with settings
|
||||
|
|
@ -273,6 +295,16 @@ class PassthroughGuardrailHandler:
|
|||
|
||||
return guardrails_to_run if guardrails_to_run else None
|
||||
|
||||
@staticmethod
|
||||
def any_inspects_the_request_body(guardrail_names: Collection[str]) -> bool:
|
||||
"""Whether any of these guardrails reads the request body.
|
||||
|
||||
A ``post_call`` or ``logging_only`` guardrail inspects the response, so a request
|
||||
body it never reads is no reason to turn the request away. A name with no
|
||||
initialized callback inspects nothing either, since there is nothing to run.
|
||||
"""
|
||||
return any(_reads_the_request_body(callback, guardrail_names) for callback in litellm.callbacks)
|
||||
|
||||
@staticmethod
|
||||
def get_field_targeted_text(
|
||||
data: dict,
|
||||
|
|
|
|||
|
|
@ -53,6 +53,7 @@ from litellm.proxy.auth.user_api_key_auth import (
|
|||
_ensure_parent_otel_span_on_request_state,
|
||||
_PendingAutoRegister,
|
||||
_matches_routing_override,
|
||||
_read_request_body_deferring_parse_failure,
|
||||
_reserve_budget_after_common_checks,
|
||||
_route_requires_auth_despite_public,
|
||||
_routing_selector_matches_claim,
|
||||
|
|
@ -6419,6 +6420,43 @@ async def test_user_api_key_auth_authenticates_before_raising_malformed_body_err
|
|||
setattr(_proxy_server_mod, k, v)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"body",
|
||||
[
|
||||
pytest.param(b"[]", id="array"),
|
||||
pytest.param(b'[{"role": "user"}]', id="array-of-objects"),
|
||||
pytest.param(b"123", id="number"),
|
||||
pytest.param(b'"gpt-4o"', id="string"),
|
||||
pytest.param(b"true", id="boolean"),
|
||||
pytest.param(b"null", id="null"),
|
||||
],
|
||||
)
|
||||
async def test_auth_reads_a_non_object_body_as_no_fields(body: bytes):
|
||||
"""Auth reads the body before any route does and then treats it as a mapping, in
|
||||
``pre_db_read_auth_checks`` and again in ``populate_request_with_path_params``. A body
|
||||
that parsed to a list/int/str raised AttributeError in both, and in the handler meant
|
||||
to turn the first one into a response, so the caller got a bare 500 (#43711)."""
|
||||
from fastapi import Request
|
||||
from starlette.datastructures import URL
|
||||
|
||||
request = Request(
|
||||
scope={
|
||||
"type": "http",
|
||||
"headers": [(b"content-type", b"application/json")],
|
||||
"method": "POST",
|
||||
"path_params": {"vector_store_id": "vs-1"},
|
||||
}
|
||||
)
|
||||
request._url = URL(url="/v1/vector_stores/vs-1/search")
|
||||
request._body = body
|
||||
|
||||
request_data, parse_exception = await _read_request_body_deferring_parse_failure(request=request)
|
||||
|
||||
assert parse_exception is None
|
||||
assert request_data == {"vector_store_id": "vs-1", "vector_store_ids": ["vs-1"]}
|
||||
|
||||
|
||||
async def _run_auth_with_malformed_body(post_call_failure_hook):
|
||||
"""Drive ``user_api_key_auth`` for an authenticated caller whose body never parses,
|
||||
with ``proxy_logging_obj.post_call_failure_hook`` swapped for the passed double.
|
||||
|
|
|
|||
|
|
@ -26,6 +26,7 @@ from litellm.proxy.common_utils.http_parsing_utils import (
|
|||
get_tags_from_request_body,
|
||||
numeric_form_fields,
|
||||
populate_request_with_path_params,
|
||||
non_object_raw_body,
|
||||
read_raw_json_body,
|
||||
)
|
||||
|
||||
|
|
@ -504,6 +505,103 @@ def _make_json_request(body: bytes) -> MagicMock:
|
|||
return mock_request
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"body",
|
||||
[
|
||||
pytest.param(b"[]", id="empty-array"),
|
||||
pytest.param(b'[{"model": "gpt-4o"}]', id="array-of-objects"),
|
||||
pytest.param(b"123", id="integer"),
|
||||
pytest.param(b"1.5", id="float"),
|
||||
pytest.param(b'"gpt-4o"', id="string"),
|
||||
pytest.param(b"true", id="boolean"),
|
||||
pytest.param(b"null", id="null"),
|
||||
],
|
||||
)
|
||||
async def test_non_object_json_body_reads_as_no_fields(body: bytes):
|
||||
"""
|
||||
``orjson.loads`` returns whatever JSON kind it found, so this ``-> dict`` used to hand
|
||||
back a list/int/str. Auth reads the body before any route does, and both read it with
|
||||
``.get(...)``, so that surfaced as AttributeError -> 500 (#43711). Every field a caller
|
||||
looks for is a key, so a non-object body carries none of them and reads as ``{}``.
|
||||
"""
|
||||
assert await _read_request_body(_make_json_request(body)) == {}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_non_object_body_is_not_cached_for_later_readers():
|
||||
"""The coercion only helps if a second read, or a route reading the cache after auth,
|
||||
cannot pull the list back out and call ``.get()`` on it."""
|
||||
request = _starlette_request(b"[1, 2, 3]", "application/json")
|
||||
|
||||
assert await _read_request_body(request) == {}
|
||||
assert _safe_get_request_parsed_body(request=request) == {}
|
||||
assert await _read_request_body(request) == {}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"body",
|
||||
[
|
||||
pytest.param(b'[{"role": "user", "content": "hi"}]', id="array"),
|
||||
pytest.param(b"123", id="number"),
|
||||
pytest.param(b'"gpt-4o"', id="string"),
|
||||
pytest.param(b"true", id="boolean"),
|
||||
pytest.param(b"null", id="null"),
|
||||
],
|
||||
)
|
||||
async def test_non_object_body_is_recoverable_verbatim_for_forwarding(body: bytes):
|
||||
"""Passthrough forwards ``_parsed_body`` as JSON unless it is handed exact bytes, so the
|
||||
coerced ``{}`` would reach the provider in place of the caller's payload. The original
|
||||
bytes must stay recoverable."""
|
||||
request = _starlette_request(body, "application/json")
|
||||
|
||||
assert await _read_request_body(request) == {}
|
||||
assert await non_object_raw_body(request) == body
|
||||
assert await read_raw_json_body(request) == body
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"body",
|
||||
[
|
||||
pytest.param(b'{"model": "gpt-4o"}', id="object"),
|
||||
pytest.param(b"{}", id="empty-object"),
|
||||
pytest.param(b"", id="empty-body"),
|
||||
],
|
||||
)
|
||||
async def test_object_body_is_not_offered_for_raw_forwarding(body: bytes):
|
||||
"""Only a coerced body may bypass ``json=_parsed_body``: an object body must keep going
|
||||
through the parsed path, so hooks that mutate it are still what gets sent."""
|
||||
request = _starlette_request(body, "application/json")
|
||||
|
||||
assert isinstance(await _read_request_body(request), dict)
|
||||
assert await non_object_raw_body(request) is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_no_raw_forwarding_before_the_body_has_been_read():
|
||||
"""Nothing may be marked for raw forwarding until a read has established the body is
|
||||
actually a non-object."""
|
||||
assert await non_object_raw_body(_starlette_request(b"[]", "application/json")) is None
|
||||
assert await non_object_raw_body(None) is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"body, expected",
|
||||
[
|
||||
pytest.param(b"{}", {}, id="empty-object"),
|
||||
pytest.param(b"", {}, id="empty-body"),
|
||||
pytest.param(b'{"model": "gpt-4o", "n": [1, 2]}', {"model": "gpt-4o", "n": [1, 2]}, id="object"),
|
||||
],
|
||||
)
|
||||
async def test_object_bodies_are_untouched(body: bytes, expected: dict):
|
||||
"""The coercion must only ever fire on a non-object body: an object's own list values
|
||||
stay exactly as the caller sent them."""
|
||||
assert await _read_request_body(_make_json_request(body)) == expected
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_surrogate_repair_skipped_above_size_limit(monkeypatch):
|
||||
"""
|
||||
|
|
@ -1210,3 +1308,37 @@ class TestCoerceNumericFormFields:
|
|||
numeric_fields=self.numeric_fields,
|
||||
)
|
||||
assert result == {"n": 3, "temperature": None, "image": buffer}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_request_that_cannot_be_marked_still_parses():
|
||||
"""Marking is best-effort bookkeeping for passthrough forwarding. A request object that
|
||||
rejects the write (a test double, a frozen scope) must still get its parsed body rather
|
||||
than turning a client's bad body into a 500."""
|
||||
|
||||
class _UnwritableScope(dict):
|
||||
def __setitem__(self, key, value):
|
||||
raise TypeError("scope is read-only")
|
||||
|
||||
request = MagicMock()
|
||||
request.body = AsyncMock(return_value=b"[1, 2, 3]")
|
||||
request.headers = {"content-type": "application/json"}
|
||||
request.scope = _UnwritableScope()
|
||||
|
||||
assert await _read_request_body(request) == {}
|
||||
assert await non_object_raw_body(request) is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_consumed_body_stream_yields_no_bytes_to_forward():
|
||||
"""Starlette raises RuntimeError once a body stream has been consumed. Forwarding must
|
||||
fall back to the parsed view rather than propagating that as a 500."""
|
||||
request = MagicMock()
|
||||
request.body = AsyncMock(return_value=b"[1, 2, 3]")
|
||||
request.headers = {"content-type": "application/json"}
|
||||
request.scope = {}
|
||||
|
||||
assert await _read_request_body(request) == {}
|
||||
|
||||
request.body = AsyncMock(side_effect=RuntimeError("Stream consumed"))
|
||||
assert await non_object_raw_body(request) is None
|
||||
|
|
|
|||
|
|
@ -24,6 +24,7 @@ from starlette.datastructures import UploadFile as StarletteUploadFile
|
|||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.constants import DEFAULT_REQUEST_TIMEOUT_SECONDS
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.proxy._types import ProxyException, UserAPIKeyAuth
|
||||
|
|
@ -32,6 +33,7 @@ from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
|
|||
LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY,
|
||||
HttpPassThroughEndpointHelpers,
|
||||
InitPassThroughEndpointHelpers,
|
||||
_parse_request_data_by_content_type,
|
||||
_registered_pass_through_routes,
|
||||
_truncate_upstream_error_body,
|
||||
_with_trace_context,
|
||||
|
|
@ -48,6 +50,7 @@ from litellm.proxy.pass_through_endpoints.success_handler import (
|
|||
)
|
||||
from litellm.proxy.route_llm_request import ProxyModelNotFoundError
|
||||
from litellm.types import utils as types_utils
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
from litellm.types.passthrough_endpoints.pass_through_endpoints import (
|
||||
LITELLM_PASS_THROUGH_DEPLOYMENT_MODEL_INFO_STATE_KEY,
|
||||
LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY,
|
||||
|
|
@ -7622,3 +7625,265 @@ def test_passthrough_attributes_a_cli_session_to_its_alias_not_the_login_token()
|
|||
metadata = kwargs["litellm_params"]["metadata"]
|
||||
assert metadata["user_api_key"] == "cli-session-alice"
|
||||
assert _get_spend_logs_metadata(metadata)["user_api_key"] == "cli-session-alice"
|
||||
|
||||
|
||||
def _json_request(body: bytes, content_type: str = "application/json") -> Request:
|
||||
async def receive():
|
||||
return {"type": "http.request", "body": body, "more_body": False}
|
||||
|
||||
return Request(
|
||||
{
|
||||
"type": "http",
|
||||
"method": "POST",
|
||||
"path": "/anthropic/v1/messages",
|
||||
"headers": [(b"content-type", content_type.encode())],
|
||||
"query_string": b"",
|
||||
},
|
||||
receive,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"body",
|
||||
[
|
||||
pytest.param(b"[]", id="empty-array"),
|
||||
pytest.param(b'[{"role": "user", "content": "hi"}]', id="array-of-objects"),
|
||||
pytest.param(b"123", id="number"),
|
||||
pytest.param(b'"claude-sonnet-4-5"', id="string"),
|
||||
pytest.param(b"true", id="boolean"),
|
||||
pytest.param(b"null", id="null"),
|
||||
],
|
||||
)
|
||||
async def test_non_object_passthrough_body_carries_no_envelope_fields(body: bytes):
|
||||
"""A passthrough body is the caller's own provider payload, so it need not be a JSON
|
||||
object. Reading a list or scalar as the query_params/custom_body envelope raised
|
||||
AttributeError and the caller got a bare 500 (#43711)."""
|
||||
query_params_data, custom_body_data, file_data, stream = await _parse_request_data_by_content_type(
|
||||
_json_request(body)
|
||||
)
|
||||
|
||||
assert (query_params_data, custom_body_data, file_data, stream) == (None, None, None, None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_object_passthrough_body_still_yields_its_envelope_fields():
|
||||
"""The guard must only fire on a non-object body: a real envelope is unaffected."""
|
||||
body = json.dumps({"query_params": {"alt": "sse"}, "custom_body": {"model": "x"}, "stream": True}).encode()
|
||||
|
||||
query_params_data, custom_body_data, _, stream = await _parse_request_data_by_content_type(_json_request(body))
|
||||
|
||||
assert query_params_data == {"alt": "sse"}
|
||||
assert custom_body_data == {"model": "x"}
|
||||
assert stream is True
|
||||
|
||||
|
||||
async def _capture_upstream_request(body: bytes, guardrails: list[str] | None = None) -> httpx.Request:
|
||||
"""Drive ``pass_through_request`` for ``body`` and return the request it built for the
|
||||
provider, by letting a real httpx client encode it and failing the send."""
|
||||
captured: list[httpx.Request] = [] # mutable-ok: the send double records what it was given
|
||||
|
||||
async def _send(req, **kwargs):
|
||||
captured.append(req)
|
||||
raise httpx.HTTPError("stop after capture")
|
||||
|
||||
real_client = httpx.AsyncClient()
|
||||
real_client.send = _send
|
||||
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging,
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client"
|
||||
) as mock_get_client,
|
||||
):
|
||||
mock_proxy_logging.pre_call_hook = AsyncMock(side_effect=lambda **kwargs: kwargs["data"])
|
||||
mock_get_client.return_value = SimpleNamespace(client=real_client)
|
||||
|
||||
# The failure handler runs on a MagicMock proxy_logging_obj and dies with TypeError,
|
||||
# the same way test_pass_through_request_uses_resolved_timeout relies on.
|
||||
with pytest.raises(TypeError):
|
||||
await pass_through_request(
|
||||
request=_json_request(body),
|
||||
target="http://upstream.test/v1/messages",
|
||||
custom_headers={},
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="sk-test"),
|
||||
guardrails_config=guardrails,
|
||||
)
|
||||
|
||||
assert captured, "pass_through_request never built an upstream request"
|
||||
return captured[0]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"body",
|
||||
[
|
||||
pytest.param(b'[{"role": "user", "content": "hi"}]', id="array"),
|
||||
pytest.param(b"[1, 2, 3]", id="array-of-numbers"),
|
||||
pytest.param(b"123", id="number"),
|
||||
pytest.param(b'"claude-sonnet-4-5"', id="string"),
|
||||
pytest.param(b"true", id="boolean"),
|
||||
pytest.param(b"null", id="null"),
|
||||
],
|
||||
)
|
||||
async def test_non_object_body_is_forwarded_to_the_provider_verbatim(body: bytes):
|
||||
"""A passthrough body is the caller's own provider payload. The parsed view of a non-object
|
||||
body is ``{}``, and the forwarding path sends the parsed body as JSON, so without the raw
|
||||
bytes the provider silently received ``{}`` in place of what the caller sent."""
|
||||
upstream = await _capture_upstream_request(body)
|
||||
|
||||
assert upstream.content == body
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_object_body_is_still_forwarded_from_the_parsed_view():
|
||||
"""Hooks mutate the parsed body and those mutations must still reach the provider, so an
|
||||
object body must not take the raw-bytes path."""
|
||||
upstream = await _capture_upstream_request(b'{"model": "claude-sonnet-4-5"}')
|
||||
|
||||
assert json.loads(upstream.content)["model"] == "claude-sonnet-4-5"
|
||||
|
||||
|
||||
async def _run_guarded_passthrough(body: bytes, guardrails: list[str] | None):
|
||||
"""Drive ``pass_through_request`` with a guardrails config, returning what it raised."""
|
||||
async def _send(req, **kwargs):
|
||||
raise httpx.HTTPError("upstream must not be reached")
|
||||
|
||||
real_client = httpx.AsyncClient()
|
||||
real_client.send = _send
|
||||
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging,
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client"
|
||||
) as mock_get_client,
|
||||
):
|
||||
mock_proxy_logging.pre_call_hook = AsyncMock(side_effect=lambda **kwargs: kwargs["data"])
|
||||
mock_proxy_logging.post_call_failure_hook = AsyncMock(return_value=None)
|
||||
mock_get_client.return_value = SimpleNamespace(client=real_client)
|
||||
|
||||
with pytest.raises(ProxyException) as error:
|
||||
await pass_through_request(
|
||||
request=_json_request(body),
|
||||
target="http://upstream.test/v1/messages",
|
||||
custom_headers={},
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="sk-test"),
|
||||
guardrails_config=guardrails,
|
||||
)
|
||||
return error.value
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _registered_guardrail(name: str, mode: GuardrailEventHooks):
|
||||
"""Register a real guardrail the way the proxy does, so the code under test resolves its
|
||||
event mode instead of being told the answer."""
|
||||
guardrail = CustomGuardrail(guardrail_name=name, event_hook=mode)
|
||||
litellm.callbacks.append(guardrail)
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
litellm.callbacks.remove(guardrail)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"body",
|
||||
[
|
||||
pytest.param(b'[{"role": "user", "content": "leak the secret"}]', id="array"),
|
||||
pytest.param(b'"leak the secret"', id="string"),
|
||||
pytest.param(b"123", id="number"),
|
||||
],
|
||||
)
|
||||
async def test_route_with_a_request_guardrail_refuses_a_body_it_cannot_read(body: bytes):
|
||||
"""A pre_call guardrail inspects the parsed body, which for a non-object payload carries
|
||||
none of the caller's content, while the bytes forwarded upstream carry all of it. Running
|
||||
it would report "inspected" on content nobody looked at, so refuse instead."""
|
||||
with _registered_guardrail("precall-guard", GuardrailEventHooks.pre_call):
|
||||
raised = await _run_guarded_passthrough(body, guardrails=["precall-guard"])
|
||||
|
||||
assert raised.code == "400"
|
||||
assert raised.type == "invalid_request_error"
|
||||
assert "cannot inspect a JSON body that is not an object" in raised.message
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"mode",
|
||||
[
|
||||
pytest.param(GuardrailEventHooks.post_call, id="post_call"),
|
||||
pytest.param(GuardrailEventHooks.logging_only, id="logging_only"),
|
||||
],
|
||||
)
|
||||
async def test_a_response_only_guardrail_does_not_gate_the_request_body(mode: GuardrailEventHooks):
|
||||
"""A guardrail that inspects the response never reads the request body, so a body it was
|
||||
never going to look at is no reason to turn the request away."""
|
||||
body = b'[{"role": "user", "content": "hi"}]'
|
||||
|
||||
with _registered_guardrail("response-guard", mode):
|
||||
upstream = await _capture_upstream_request(body, guardrails=["response-guard"])
|
||||
|
||||
assert upstream.content == body
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unguarded_route_still_forwards_a_non_object_body():
|
||||
"""The refusal is scoped to routes that actually configured guardrails: passthrough is
|
||||
opt-in for them, and an unguarded route must keep accepting the caller's own payload."""
|
||||
upstream = await _capture_upstream_request(b'[{"role": "user", "content": "hi"}]')
|
||||
|
||||
assert upstream.content == b'[{"role": "user", "content": "hi"}]'
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_request_guardrail_still_accepts_an_object_body():
|
||||
"""An object body is fully readable, so a pre_call guardrail being configured must not turn
|
||||
it away: it still reaches the provider."""
|
||||
with _registered_guardrail("precall-guard", GuardrailEventHooks.pre_call):
|
||||
upstream = await _capture_upstream_request(
|
||||
b'{"model": "claude-sonnet-4-5"}', guardrails=["precall-guard"]
|
||||
)
|
||||
|
||||
assert json.loads(upstream.content)["model"] == "claude-sonnet-4-5"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_json_body_sent_with_a_multipart_content_type_still_yields_its_envelope():
|
||||
"""A misconfigured client can send JSON under a multipart content-type; that branch parses
|
||||
it as JSON, so it must read the envelope the same way the JSON branch does."""
|
||||
body = json.dumps({"query_params": {"alt": "sse"}, "stream": True}).encode()
|
||||
|
||||
query_params_data, custom_body_data, _, stream = await _parse_request_data_by_content_type(
|
||||
_json_request(body, content_type="multipart/form-data; boundary=x")
|
||||
)
|
||||
|
||||
assert query_params_data == {"alt": "sse"}
|
||||
assert stream is True
|
||||
assert custom_body_data == {"query_params": {"alt": "sse"}, "stream": True}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_non_object_json_under_a_multipart_content_type_is_left_to_the_multipart_handler():
|
||||
"""The same branch must not read a list as an envelope: it falls through so the real
|
||||
multipart handler deals with the body."""
|
||||
query_params_data, custom_body_data, file_data, stream = await _parse_request_data_by_content_type(
|
||||
_json_request(b"[1, 2, 3]", content_type="multipart/form-data; boundary=x")
|
||||
)
|
||||
|
||||
assert (query_params_data, custom_body_data, file_data, stream) == (None, None, None, None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_another_routes_request_guardrail_does_not_gate_this_route():
|
||||
"""A proxy registers every callback globally, so this route must look only at the
|
||||
guardrails it configured. Another route's pre_call guardrail, or a plain logger, must not
|
||||
make this route turn a body away."""
|
||||
body = b'[{"role": "user", "content": "hi"}]'
|
||||
unrelated_logger = CustomLogger()
|
||||
litellm.callbacks.append(unrelated_logger)
|
||||
try:
|
||||
with _registered_guardrail("someone-elses-guard", GuardrailEventHooks.pre_call):
|
||||
upstream = await _capture_upstream_request(body, guardrails=["postcall-guard"])
|
||||
finally:
|
||||
litellm.callbacks.remove(unrelated_logger)
|
||||
|
||||
assert upstream.content == body
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue