mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
parent
ae0d84116a
commit
07aeaa17a0
2 changed files with 176 additions and 18 deletions
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue