fix(passthrough): stop request params from clobbering merged target query params (#32404)

* fix(passthrough): stop request params from clobbering merged target query params

* fix(passthrough): rewrite managed ids in query params before folding them into the URL
This commit is contained in:
Mateo Wang 2026-07-07 18:56:56 -07:00 • committed by GitHub
parent ae0d84116a
commit 07aeaa17a0
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 176 additions and 18 deletions

View file

@ -820,21 +820,7 @@ async def pass_through_request(
forward_headers=forward_headers,
)
# Apply default query parameters if provided, regardless of merge_query_params setting
if default_query_params or merge_query_params:
# Determine what to merge based on settings
request_params = dict(request.query_params) if merge_query_params else {}
# Create a new URL with the merged query params
url = url.copy_with(
query=urlencode(
HttpPassThroughEndpointHelpers.get_merged_query_parameters(
existing_url=url,
request_query_params=request_params,
default_query_params=default_query_params,
)
).encode("ascii")
)
requested_query_params: Optional[dict] = query_params or dict(request.query_params)
endpoint_type: EndpointType = HttpPassThroughEndpointHelpers.get_endpoint_type(str(url))
@ -952,9 +938,6 @@ async def pass_through_request(
)
logging_obj.model_call_details["litellm_call_id"] = litellm_call_id
# combine url with query params for logging
requested_query_params: Optional[dict] = query_params or dict(request.query_params)
## PASSTHROUGH MANAGED ID RESOLUTION (INPUT) ##
# Resolve managed IDs in path, query params, and body back to raw
# provider IDs before forwarding upstream. Gated by feature flag and
@ -1024,6 +1007,20 @@ async def pass_through_request(
request.method,
)
# Apply default query parameters if provided, regardless of merge_query_params setting
if default_query_params or merge_query_params:
# Create a new URL with the merged query params
url = url.copy_with(
query=urlencode(
HttpPassThroughEndpointHelpers.get_merged_query_parameters(
existing_url=url,
request_query_params=requested_query_params,
default_query_params=default_query_params,
)
).encode("ascii")
)
requested_query_params = None
## PASSTHROUGH MANAGED LIST (DB-only response) ##
# For GET /v1/files and GET /v1/batches passthrough routes, serve the
# listing entirely from our DB so each caller only sees their own IDs.

View file

@ -5,6 +5,7 @@ import sys
from contextlib import ExitStack
from io import BytesIO
from types import SimpleNamespace
from typing import Optional
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
@ -29,6 +30,7 @@ from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
resolve_pass_through_request_timeout,
resolve_llm_passthrough_timeout,
)
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.passthrough_endpoints.pass_through_endpoints import (
LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY,
)
@ -2251,6 +2253,165 @@ async def test_pass_through_request_query_params_forwarding():
assert call_kwargs["_parsed_body"] == test_body
class _FakeManagedFilesHook:
def __init__(self, file_row: SimpleNamespace):
self._file_row = file_row
async def get_unified_file_id(self, file_id: str, litellm_parent_otel_span=None) -> SimpleNamespace:
return self._file_row
async def _run_pass_through_and_capture_wire_url(
target: str,
incoming_query: str,
merge_query_params: bool = False,
default_query_params: Optional[dict] = None,
custom_llm_provider: Optional[str] = None,
managed_files_hook: Optional[_FakeManagedFilesHook] = None,
user_api_key_dict: Optional[UserAPIKeyAuth] = None,
) -> httpx.URL:
import litellm
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.types.llms.custom_http import httpxSpecialProvider
recorded_requests = []
def transport_handler(upstream_request: httpx.Request) -> httpx.Response:
recorded_requests.append(upstream_request)
return httpx.Response(200, json={"ok": True})
real_handler = get_async_httpx_client(
llm_provider=httpxSpecialProvider.PassThroughEndpoint,
params={"timeout": resolve_pass_through_request_timeout(None)},
)
cache_dict = litellm.in_memory_llm_clients_cache.cache_dict
cache_key = next((key for key, cached in cache_dict.items() if cached is real_handler), None)
assert cache_key is not None, (
"PassThroughEndpoint client not found in in_memory_llm_clients_cache; "
"get_async_httpx_client may not be caching this provider."
)
cache_dict[cache_key] = SimpleNamespace(
client=httpx.AsyncClient(transport=httpx.MockTransport(transport_handler))
)
mock_request = MagicMock(spec=Request)
mock_request.method = "GET"
mock_request.headers = Headers({})
mock_request.query_params = QueryParams(incoming_query)
mock_request.body = AsyncMock(return_value=b"")
mock_proxy_logging = MagicMock()
mock_proxy_logging.pre_call_hook = AsyncMock(
side_effect=lambda user_api_key_dict, data, call_type: data
)
mock_proxy_logging.post_call_failure_hook = AsyncMock()
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value={})
mock_proxy_logging.get_proxy_hook = MagicMock(return_value=managed_files_hook)
try:
with ExitStack() as stack:
stack.enter_context(
patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging)
)
if managed_files_hook is not None:
stack.enter_context(
patch(
"litellm.proxy.proxy_server.general_settings",
{"passthrough_managed_object_ids": True},
)
)
stack.enter_context(patch("litellm.proxy.proxy_server.prisma_client", None))
response = await pass_through_request(
request=mock_request,
target=target,
custom_headers={},
user_api_key_dict=user_api_key_dict if user_api_key_dict is not None else MagicMock(),
merge_query_params=merge_query_params,
default_query_params=default_query_params,
custom_llm_provider=custom_llm_provider,
)
finally:
cache_dict[cache_key] = real_handler
assert response.status_code == 200
assert len(recorded_requests) == 1
return recorded_requests[0].url
@pytest.mark.asyncio
async def test_pass_through_request_merge_query_params_preserves_target_query_on_wire():
"""
Regression test: with merge_query_params=True, the target URL's own query
params must survive on the final outgoing request. Passing the incoming
params via httpx's params= replaces the URL's entire query string, which
used to silently drop the merged target params.
"""
wire_url = await _run_pass_through_and_capture_wire_url(
target="https://www.bing.com/search?setLang=en-US&mkt=en-US",
incoming_query="q=litellm",
merge_query_params=True,
)
assert dict(wire_url.params) == {
"setLang": "en-US",
"mkt": "en-US",
"q": "litellm",
}
@pytest.mark.asyncio
async def test_pass_through_request_default_query_params_reach_the_wire():
"""
default_query_params are sent with every request and can be overridden
per-key by client-provided query params; params the client does not
override must not be dropped from the outgoing request.
"""
wire_url = await _run_pass_through_and_capture_wire_url(
target="https://example.com/api",
incoming_query="limit=5&api-version=client-version",
default_query_params={"api-version": "2024-01-01", "setLang": "en-US"},
)
assert dict(wire_url.params) == {
"api-version": "client-version",
"setLang": "en-US",
"limit": "5",
}
@pytest.mark.asyncio
async def test_pass_through_request_without_merge_replaces_target_query():
wire_url = await _run_pass_through_and_capture_wire_url(
target="https://www.bing.com/search?setLang=en-US",
incoming_query="q=litellm",
)
assert dict(wire_url.params) == {"q": "litellm"}
@pytest.mark.asyncio
async def test_pass_through_request_merge_query_params_rewrites_managed_ids_on_the_wire():
"""
Regression test: on merge-enabled endpoints the managed-ID rewrite must see
the incoming query params before they are folded into the URL. Folding
first bakes the un-rewritten managed ID into the URL and hands the rewriter
None, leaking the managed ID upstream.
"""
from litellm.proxy.pass_through_endpoints.managed_id_codec import new_managed_id
managed_id = new_managed_id("openai", "file-raw-123")
hook = _FakeManagedFilesHook(SimpleNamespace(created_by="user-1", team_id=None))
wire_url = await _run_pass_through_and_capture_wire_url(
target="https://api.openai.com/v1/files/content?api-version=preview",
incoming_query=f"file_id={managed_id}",
merge_query_params=True,
custom_llm_provider="openai",
managed_files_hook=hook,
user_api_key_dict=UserAPIKeyAuth(user_id="user-1"),
)
assert dict(wire_url.params) == {
"api-version": "preview",
"file_id": "file-raw-123",
}
@pytest.mark.asyncio
async def test_pass_through_with_httpbin_redirect():
"""