mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
Merge pull request #23338 from BerriAI/litellm_fix_multipart_passthrough
fix(proxy): preserve multipart/form-data boundary in passthrough endpoints
This commit is contained in:
commit
d99ddc67d4
2 changed files with 132 additions and 31 deletions
|
|
@ -404,7 +404,9 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils):
|
|||
headers=headers,
|
||||
params=requested_query_params,
|
||||
)
|
||||
elif HttpPassThroughEndpointHelpers.is_multipart(request) is True:
|
||||
elif HttpPassThroughEndpointHelpers.is_multipart(request) is True and not _parsed_body:
|
||||
# Only use multipart handler if we don't have a parsed body
|
||||
# (parsed body means it was JSON despite multipart content-type header)
|
||||
return await HttpPassThroughEndpointHelpers.make_multipart_http_request(
|
||||
request=request,
|
||||
async_client=async_client,
|
||||
|
|
@ -677,8 +679,15 @@ async def pass_through_request( # noqa: PLR0915
|
|||
str(url)
|
||||
)
|
||||
|
||||
# Skip body parsing for multipart requests - make_multipart_http_request will handle it
|
||||
# But if custom_body is provided (e.g., JSON parsed despite multipart content-type), use it
|
||||
is_multipart = HttpPassThroughEndpointHelpers.is_multipart(request) and not custom_body
|
||||
|
||||
if custom_body:
|
||||
_parsed_body = custom_body
|
||||
elif is_multipart:
|
||||
# Don't parse multipart body here - it will be handled by make_multipart_http_request
|
||||
_parsed_body = {}
|
||||
else:
|
||||
_parsed_body = await _read_request_body(request)
|
||||
verbose_proxy_logger.debug(
|
||||
|
|
@ -1043,30 +1052,22 @@ async def _parse_request_data_by_content_type(
|
|||
# Handle requests with no body (e.g., DELETE requests)
|
||||
pass
|
||||
elif "multipart/form-data" in content_type:
|
||||
# ✅ Handle multipart form-data
|
||||
form = await request.form()
|
||||
if "query_params" in form:
|
||||
form_value = form["query_params"]
|
||||
if isinstance(form_value, str):
|
||||
try:
|
||||
query_params_data = json.loads(form_value)
|
||||
except Exception:
|
||||
query_params_data = form_value
|
||||
else:
|
||||
query_params_data = form_value
|
||||
|
||||
if "custom_body" in form:
|
||||
form_value = form["custom_body"]
|
||||
if isinstance(form_value, str):
|
||||
try:
|
||||
custom_body_data = json.loads(form_value)
|
||||
except Exception:
|
||||
custom_body_data = form_value
|
||||
else:
|
||||
custom_body_data = form_value
|
||||
|
||||
if "file" in form:
|
||||
file_data = form["file"] # this is a Starlette UploadFile object
|
||||
# ✅ 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()
|
||||
# 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
|
||||
except (json.JSONDecodeError, Exception):
|
||||
# Not JSON - this is actual multipart data
|
||||
# Skip parsing here to avoid consuming the request body stream
|
||||
# make_multipart_http_request will handle it
|
||||
pass
|
||||
|
||||
elif "application/x-www-form-urlencoded" in content_type:
|
||||
# ✅ Handle URL-encoded form data
|
||||
|
|
@ -1132,7 +1133,6 @@ def create_pass_through_route(
|
|||
fastapi_response: Response,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
subpath: str = "", # captures sub-paths when include_subpath=True
|
||||
custom_body: Optional[dict] = None,
|
||||
):
|
||||
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
|
||||
InitPassThroughEndpointHelpers,
|
||||
|
|
@ -1208,12 +1208,9 @@ def create_pass_through_route(
|
|||
)
|
||||
if query_params:
|
||||
final_query_params.update(query_params)
|
||||
# When a caller (e.g. bedrock_proxy_route) supplies a pre-built
|
||||
# body, use it instead of the body parsed from the raw request.
|
||||
# Use the body parsed from the raw request
|
||||
final_custom_body: Optional[dict] = None
|
||||
if custom_body is not None:
|
||||
final_custom_body = custom_body
|
||||
elif isinstance(custom_body_data, dict):
|
||||
if isinstance(custom_body_data, dict):
|
||||
final_custom_body = custom_body_data
|
||||
|
||||
return await pass_through_request( # type: ignore
|
||||
|
|
|
|||
|
|
@ -2369,3 +2369,107 @@ def test_get_registered_pass_through_route_with_custom_root():
|
|||
|
||||
# Clean up
|
||||
_registered_pass_through_routes.clear()
|
||||
|
||||
|
||||
def test_mapped_pass_through_routes_with_server_root_path():
|
||||
"""
|
||||
Mapped passthrough routes (vertex_ai, bedrock, etc) should match
|
||||
even when SERVER_ROOT_PATH is set and the incoming route is prefixed.
|
||||
|
||||
Regression test for https://github.com/BerriAI/litellm/issues/22272
|
||||
"""
|
||||
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
|
||||
InitPassThroughEndpointHelpers,
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_server_root_path"
|
||||
) as mock_get_root:
|
||||
mock_get_root.return_value = "/litellm"
|
||||
|
||||
# prefixed route should match mapped routes like /vertex_ai
|
||||
assert (
|
||||
InitPassThroughEndpointHelpers.is_registered_pass_through_route(
|
||||
"/litellm/vertex_ai/v1/projects/foo"
|
||||
)
|
||||
is True
|
||||
)
|
||||
assert (
|
||||
InitPassThroughEndpointHelpers.is_registered_pass_through_route(
|
||||
"/litellm/bedrock/model/invoke"
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
||||
# bare route without prefix should not match when root is set
|
||||
assert (
|
||||
InitPassThroughEndpointHelpers.is_registered_pass_through_route(
|
||||
"/vertex_ai/v1/projects/foo"
|
||||
)
|
||||
is False
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_multipart_passthrough_preserves_boundary():
|
||||
"""
|
||||
Test that multipart/form-data requests through passthrough preserve the boundary
|
||||
and can be correctly parsed by the upstream server.
|
||||
|
||||
Regression test for multipart boundary stripping issue.
|
||||
"""
|
||||
from io import BytesIO
|
||||
|
||||
# Mock the httpx request to verify files are passed correctly
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.headers = httpx.Headers({"content-type": "application/json"})
|
||||
mock_response.aread = AsyncMock(return_value=b'{"filename": "test.txt", "size": 17}')
|
||||
mock_response.text = '{"filename": "test.txt", "size": 17}'
|
||||
|
||||
async def mock_httpx_request(method, url, **kwargs):
|
||||
# Verify that files parameter is passed (not json)
|
||||
assert "files" in kwargs, "Files should be passed for multipart requests"
|
||||
assert "file" in kwargs["files"], "File field should be in files dict"
|
||||
|
||||
# Verify content-type is NOT in headers (httpx will set it with correct boundary)
|
||||
headers = kwargs.get("headers", {})
|
||||
assert "content-type" not in headers, "content-type should be removed for multipart"
|
||||
|
||||
filename, content, content_type = kwargs["files"]["file"]
|
||||
assert filename == "test.txt"
|
||||
assert content == b"test file content"
|
||||
assert content_type == "text/plain"
|
||||
|
||||
return mock_response
|
||||
|
||||
async_client = MagicMock()
|
||||
async_client.request = AsyncMock(side_effect=mock_httpx_request)
|
||||
|
||||
# Create mock request
|
||||
request = MagicMock(spec=Request)
|
||||
request.method = "POST"
|
||||
request.headers = Headers({"content-type": "multipart/form-data; boundary=test123"})
|
||||
|
||||
# Mock form data
|
||||
file_content = b"test file content"
|
||||
file = BytesIO(file_content)
|
||||
headers = Headers({"content-type": "text/plain"})
|
||||
upload_file = UploadFile(file=file, filename="test.txt", headers=headers)
|
||||
upload_file.read = AsyncMock(return_value=file_content)
|
||||
|
||||
form_data = {"file": upload_file}
|
||||
request.form = AsyncMock(return_value=form_data)
|
||||
|
||||
# Test the multipart handler directly
|
||||
response = await HttpPassThroughEndpointHelpers.make_multipart_http_request(
|
||||
request=request,
|
||||
async_client=async_client,
|
||||
url=httpx.URL("http://test.com/upload"),
|
||||
headers={},
|
||||
requested_query_params=None,
|
||||
)
|
||||
|
||||
# Verify the response
|
||||
assert response.status_code == 200
|
||||
async_client.request.assert_called_once()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue