diff --git a/litellm/llms/azure/realtime/handler.py b/litellm/llms/azure/realtime/handler.py index 8f0037570b0..a95f65be417 100644 --- a/litellm/llms/azure/realtime/handler.py +++ b/litellm/llms/azure/realtime/handler.py @@ -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, diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 377faf1400c..321bab82208 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -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, diff --git a/litellm/realtime_api/main.py b/litellm/realtime_api/main.py index 4ec00e1f65b..60f6fabd561 100644 --- a/litellm/realtime_api/main.py +++ b/litellm/realtime_api/main.py @@ -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}") diff --git a/tests/test_litellm/llms/openai/realtime/test_transcription_sessions.py b/tests/test_litellm/llms/openai/realtime/test_transcription_sessions.py index afcfcc75336..08ef9485677 100644 --- a/tests/test_litellm/llms/openai/realtime/test_transcription_sessions.py +++ b/tests/test_litellm/llms/openai/realtime/test_transcription_sessions.py @@ -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]