mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
Address PR review: URL-encode all Azure WS query params; forward query_params through provider_config branch
This commit is contained in:
parent
63d754d27b
commit
312f51a10b
4 changed files with 77 additions and 6 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue