Merge pull request #41448 from BerriAI/litellm_fix_passthrough_empty_query_params_drop_url_query

fix(passthrough): keep target URL query when client sends no query params
This commit is contained in:
Mateo Wang 2026-09-17 18:02:48 -07:00 committed by GitHub
commit 3424390101
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 29 additions and 22 deletions

View file

@ -9,6 +9,7 @@ from collections.abc import AsyncGenerator, Callable, Iterable, Mapping, Sequenc
from dataclasses import dataclass
from datetime import datetime
from itertools import groupby
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, TypedDict, cast
from urllib.parse import urlencode, urlparse
@ -991,7 +992,7 @@ async def pass_through_request(
)
upstream_headers: Final = _with_trace_context(headers, parent_span=user_api_key_dict.parent_otel_span)
requested_query_params: dict | None = query_params or dict(request.query_params)
requested_query_params: dict | None = query_params or dict(request.query_params) or None
endpoint_type: Final[EndpointType] = HttpPassThroughEndpointHelpers.get_endpoint_type(str(url))
@ -1193,7 +1194,7 @@ async def pass_through_request(
query=urlencode(
HttpPassThroughEndpointHelpers.get_merged_query_parameters(
existing_url=url,
request_query_params=requested_query_params,
request_query_params=requested_query_params or MappingProxyType({}),
default_query_params=default_query_params,
)
).encode("ascii")

View file

@ -325,7 +325,7 @@ async def test_pass_through_request_stream_param_override(
"POST",
httpx.URL("https://api.anthropic.com/v1/messages"),
json=request_body,
params={},
params=None,
headers={"Authorization": "Bearer test-key"},
)
@ -424,7 +424,7 @@ async def test_pass_through_request_stream_param_no_override(
"POST",
httpx.URL("https://api.anthropic.com/v1/messages"),
headers={"Authorization": "Bearer test-key"},
params={},
params=None,
json=request_body,
)
mock_async_client.send.assert_called_once()

View file

@ -7,7 +7,6 @@ from collections.abc import Callable
from contextlib import ExitStack, contextmanager
from io import BytesIO
from types import SimpleNamespace
from typing import Optional
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
@ -16,34 +15,32 @@ from fastapi import Request, Response, UploadFile
from starlette.datastructures import FormData, Headers, QueryParams
from starlette.datastructures import UploadFile as StarletteUploadFile
import litellm
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
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
DEFAULT_PASS_THROUGH_REQUEST_TIMEOUT_SECONDS,
LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY,
HttpPassThroughEndpointHelpers,
InitPassThroughEndpointHelpers,
LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY,
_registered_pass_through_routes,
chat_completion_pass_through_endpoint,
create_pass_through_route,
initialize_pass_through_endpoints,
pass_through_request,
resolve_pass_through_request_timeout,
resolve_llm_passthrough_timeout,
resolve_pass_through_request_timeout,
websocket_passthrough_request,
_with_trace_context,
)
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
from litellm.types.passthrough_endpoints.pass_through_endpoints import (
LITELLM_PASS_THROUGH_DEPLOYMENT_MODEL_INFO_STATE_KEY,
LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY,
)
from litellm.proxy.pass_through_endpoints.success_handler import (
PassThroughEndpointLogging,
)
import litellm
from litellm.types.passthrough_endpoints.pass_through_endpoints import (
LITELLM_PASS_THROUGH_DEPLOYMENT_MODEL_INFO_STATE_KEY,
LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY,
)
MESSAGE_START_SSE_FRAME = b'event: message_start\ndata: {"type": "message_start"}\n\n'
@ -2436,10 +2433,10 @@ 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,
default_query_params: dict | None = None,
custom_llm_provider: str | None = None,
managed_files_hook: _FakeManagedFilesHook | None = None,
user_api_key_dict: UserAPIKeyAuth | None = None,
) -> httpx.URL:
import litellm
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
@ -2551,6 +2548,15 @@ async def test_pass_through_request_without_merge_replaces_target_query():
assert dict(wire_url.params) == {"q": "litellm"}
@pytest.mark.asyncio
async def test_pass_through_request_preserves_target_query_without_client_query():
wire_url = await _run_pass_through_and_capture_wire_url(
target="https://example.com/v1/models/gemini:streamGenerateContent?alt=sse",
incoming_query="",
)
assert dict(wire_url.params) == {"alt": "sse"}
@pytest.mark.asyncio
async def test_pass_through_request_merge_query_params_rewrites_managed_ids_on_the_wire():
"""
@ -5361,7 +5367,7 @@ async def test_websocket_passthrough_does_not_close_twice_when_success_logging_f
def _passthrough_kwargs_for_reservation(
user_api_key_dict: UserAPIKeyAuth,
parsed_body: Optional[dict] = None,
parsed_body: dict | None = None,
user_defined_route: bool = False,
) -> dict:
mock_request = MagicMock(spec=Request)