diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index 0ac4182ebd2..5a621163760 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -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. diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py index 1482937ab3b..85211f392ee 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py @@ -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(): """