Address PR review: URL-encode all Azure WS query params; forward query_params through provider_config branch

This commit is contained in:
Emerson Gomes 2026-06-05 07:34:59 -05:00
parent 63d754d27b
commit 312f51a10b
No known key found for this signature in database
GPG key ID: D3DF28AB5D1B5E17
4 changed files with 77 additions and 6 deletions

View file

@ -5,7 +5,6 @@ This requires websockets, and is currently only supported on LiteLLM Proxy.
"""
from typing import Any, Optional, cast
from urllib.parse import quote
from litellm._logging import _redact_string, verbose_proxy_logger
from litellm.constants import REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES
@ -57,6 +56,8 @@ class AzureOpenAIRealtime(AzureChatCompletion):
beta/default: "wss://.../openai/realtime?api-version=2024-10-01-preview&deployment=gpt-4o-realtime-preview"
GA/v1: "wss://.../openai/v1/realtime?model=gpt-realtime-deployment"
"""
from urllib.parse import urlencode
api_base = api_base.replace("https://", "wss://")
# Determine path based on realtime_protocol (case-insensitive)
@ -66,16 +67,17 @@ class AzureOpenAIRealtime(AzureChatCompletion):
)
if _is_ga:
path = "/openai/v1/realtime"
url = f"{api_base}{path}?model={model}"
qs = urlencode({"model": model})
else:
# Default to beta path for backwards compatibility
path = "/openai/realtime"
url = f"{api_base}{path}?api-version={api_version}&deployment={model}"
qs = urlencode({"api-version": api_version, "deployment": model})
intent = (query_params or {}).get("intent")
if intent:
url = f"{url}&intent={quote(str(intent), safe='')}"
return url
qs = f"{qs}&{urlencode({'intent': intent})}"
return f"{api_base}{path}?{qs}"
async def async_realtime(
self,

View file

@ -5260,6 +5260,21 @@ class BaseLLMHTTPHandler:
headers=error_headers,
)
@staticmethod
def _append_query_params(url: str, query_params: Optional[Dict[str, Any]]) -> str:
"""Append query_params to url, skipping keys already present in the URL."""
if not query_params:
return url
from urllib.parse import parse_qsl, urlencode, urlparse, urlunparse
parsed = urlparse(url)
existing = dict(parse_qsl(parsed.query))
extras = {k: v for k, v in query_params.items() if k not in existing}
if not extras:
return url
new_query = parsed.query + ("&" if parsed.query else "") + urlencode(extras)
return urlunparse(parsed._replace(query=new_query))
async def async_realtime(
self,
model: str,
@ -5273,11 +5288,14 @@ class BaseLLMHTTPHandler:
timeout: Optional[float] = None,
user_api_key_dict: Optional[Any] = None,
litellm_metadata: Optional[Dict[str, Any]] = None,
query_params: Optional[Dict[str, Any]] = None,
):
import websockets
from websockets.asyncio.client import ClientConnection
url = provider_config.get_complete_url(api_base, model, api_key)
url = self._append_query_params(
provider_config.get_complete_url(api_base, model, api_key), query_params
)
headers = provider_config.validate_environment(
headers=headers,
model=model,

View file

@ -345,6 +345,7 @@ async def _arealtime( # noqa: PLR0915
headers=headers,
user_api_key_dict=kwargs.get("user_api_key_dict"),
litellm_metadata=_build_litellm_metadata(kwargs),
query_params=query_params,
)
elif _custom_llm_provider == "azure":
api_base = (
@ -518,6 +519,7 @@ async def _arealtime( # noqa: PLR0915
headers=headers,
user_api_key_dict=kwargs.get("user_api_key_dict"),
litellm_metadata=_build_litellm_metadata(kwargs),
query_params=query_params,
)
else:
raise ValueError(f"Unsupported model: {model}")

View file

@ -172,3 +172,52 @@ async def test_sdk_fn_routes_openai_transcription_session(monkeypatch):
assert kwargs["json"]["input_audio_transcription"] == {
"model": "gpt-realtime-whisper"
}
def test_append_query_params_skips_existing_keys():
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
url = "wss://example.com/v1/realtime?model=gpt-4o"
result = BaseLLMHTTPHandler._append_query_params(
url, {"model": "ignored", "intent": "transcription"}
)
assert "model=ignored" not in result
assert "intent=transcription" in result
def test_append_query_params_no_params_returns_unchanged():
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
url = "wss://example.com/v1/realtime?model=gpt-4o"
assert BaseLLMHTTPHandler._append_query_params(url, None) == url
assert BaseLLMHTTPHandler._append_query_params(url, {}) == url
def test_append_query_params_encodes_special_chars():
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
url = "wss://example.com/v1/realtime"
result = BaseLLMHTTPHandler._append_query_params(url, {"intent": "a&b=c"})
assert "intent=a%26b%3Dc" in result
assert "&b=c" not in result
def test_azure_construct_url_encodes_model_and_api_version():
"""model and api-version must be URL-encoded to prevent query-string injection."""
from litellm.llms.azure.realtime.handler import AzureOpenAIRealtime
h = AzureOpenAIRealtime()
url = h._construct_url(
"https://x.openai.azure.com",
"deploy&evil=1",
"2024-10-01-preview",
)
assert "evil=1" not in url.split("?", 1)[1]
url_ga = h._construct_url(
"https://x.openai.azure.com",
"deploy&evil=1",
None,
realtime_protocol="GA",
)
assert "evil=1" not in url_ga.split("?", 1)[1]