Realtime API - Set 'headers' in scope for websocket auth requests + reliability fix infinite loop when model_name not found for realtime models (#10679)

* fix(user_api_key_auth.py): add 'headers' to constructed request for websocket

Fix issue on some datastructure versions which require a headers field in scope

* test(test_user_api_key_auth.py): add unit testing for headers in scope change

* fix(router.py): migrate `_arealtime` to generic router endpoint

Fix infinite loop on model name missing for realtime api calls

* test(test_router_helper_utils.py): cleanup test post refactor
This commit is contained in:
Krish Dholakia 2025-05-08 22:50:09 -07:00 • committed by GitHub
parent 5325ee4382
commit a1964eab18
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 42 additions and 65 deletions

View file

@ -51,3 +51,11 @@ model_list:
model_info:
input_cost_per_token: 0.75
output_cost_per_token: 3
- model_name: "gpt-3.5-turbo"
litellm_params:
model: gpt-3.5-turbo
- model_name: gpt-4o-realtime-preview
litellm_params:
model: azure/gpt-4o-realtime-preview
api_key: os.environ/AZURE_SWEDEN_API_KEY
api_base: os.environ/AZURE_SWEDEN_API_BASE

View file

@ -106,7 +106,15 @@ def _get_bearer_token(
async def user_api_key_auth_websocket(websocket: WebSocket):
# Accept the WebSocket connection
request = Request(scope={"type": "http"})
request = Request(
scope={
"type": "http",
"headers": [
(k.lower().encode(), v.encode()) for k, v in websocket.headers.items()
],
}
)
request._url = websocket.url
query_params = websocket.query_params
@ -120,9 +128,7 @@ async def user_api_key_auth_websocket(websocket: WebSocket):
request.body = return_body # type: ignore
# Extract the Authorization header
authorization = websocket.headers.get("authorization")
# If no Authorization header, try the api-key header
if not authorization:
api_key = websocket.headers.get("api-key")
@ -521,23 +527,23 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
proxy_logging_obj=proxy_logging_obj,
)
if _end_user_object is not None:
end_user_params["allowed_model_region"] = (
_end_user_object.allowed_model_region
)
end_user_params[
"allowed_model_region"
] = _end_user_object.allowed_model_region
if _end_user_object.litellm_budget_table is not None:
budget_info = _end_user_object.litellm_budget_table
if budget_info.tpm_limit is not None:
end_user_params["end_user_tpm_limit"] = (
budget_info.tpm_limit
)
end_user_params[
"end_user_tpm_limit"
] = budget_info.tpm_limit
if budget_info.rpm_limit is not None:
end_user_params["end_user_rpm_limit"] = (
budget_info.rpm_limit
)
end_user_params[
"end_user_rpm_limit"
] = budget_info.rpm_limit
if budget_info.max_budget is not None:
end_user_params["end_user_max_budget"] = (
budget_info.max_budget
)
end_user_params[
"end_user_max_budget"
] = budget_info.max_budget
except Exception as e:
if isinstance(e, litellm.BudgetExceededError):
raise e

View file

@ -745,6 +745,9 @@ class Router:
self.adelete_responses = self.factory_function(
litellm.adelete_responses, call_type="adelete_responses"
)
self._arealtime = self.factory_function(
litellm._arealtime, call_type="_arealtime"
)
def validate_fallbacks(self, fallback_param: Optional[List]):
"""
@ -2151,40 +2154,6 @@ class Router:
self.fail_calls[model_name] += 1
raise e
async def _arealtime(self, model: str, **kwargs):
messages = [{"role": "user", "content": "dummy-text"}]
try:
kwargs["num_retries"] = kwargs.get("num_retries", self.num_retries)
self._update_kwargs_before_fallbacks(model=model, kwargs=kwargs)
# pick the one that is available (lowest TPM/RPM)
deployment = await self.async_get_available_deployment(
model=model,
messages=messages,
specific_deployment=kwargs.pop("specific_deployment", None),
request_kwargs=kwargs,
)
self._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs)
data = deployment["litellm_params"].copy()
for k, v in self.default_litellm_params.items():
if (
k not in kwargs
): # prioritize model-specific params > default router params
kwargs[k] = v
elif k == "metadata":
kwargs[k].update(v)
return await litellm._arealtime(**{**data, "caching": self.cache_responses, **kwargs}) # type: ignore
except Exception as e:
if self.num_retries > 0:
kwargs["model"] = model
kwargs["messages"] = messages
kwargs["original_function"] = self._arealtime
return await self.async_function_with_retries(**kwargs)
else:
raise e
def text_completion(
self,
model: str,
@ -3096,6 +3065,7 @@ class Router:
"adelete_responses",
"afile_delete",
"afile_content",
"_arealtime",
] = "assistants",
):
"""
@ -3143,6 +3113,7 @@ class Router:
elif call_type in (
"anthropic_messages",
"aresponses",
"_arealtime",
):
return await self._ageneric_api_call_with_fallbacks(
original_function=original_function,

View file

@ -793,6 +793,14 @@ async def test_user_api_key_auth_websocket():
# Assert that `user_api_key_auth` was called with the correct parameters
mock_user_api_key_auth.assert_called_once()
# Get the request object that was passed to user_api_key_auth
request_arg = mock_user_api_key_auth.call_args.kwargs["request"]
# Verify that the request has headers set
assert hasattr(request_arg, "headers"), "Request object should have headers attribute"
assert "authorization" in request_arg.headers, "Request headers should contain authorization"
assert request_arg.headers["authorization"] == "Bearer some_api_key"
assert (
mock_user_api_key_auth.call_args.kwargs["api_key"] == "Bearer some_api_key"
)

View file

@ -209,22 +209,6 @@ async def test_router_schedule_factory(model_list):
assert "priority" not in mock_atext_completion.call_args.kwargs
@pytest.mark.asyncio
async def test_router_arealtime(model_list):
"""Test if the '_arealtime' function is working correctly"""
import litellm
router = Router(model_list=model_list)
with patch.object(litellm, "_arealtime", AsyncMock()) as mock_arealtime:
mock_arealtime.return_value = "I'm fine, thank you!"
await router._arealtime(
model="gpt-3.5-turbo",
messages=[{"role": "user", "content": "Hello, how are you?"}],
)
mock_arealtime.assert_awaited_once()
@pytest.mark.parametrize("sync_mode", [True, False])
@pytest.mark.asyncio
async def test_router_function_with_fallbacks(model_list, sync_mode):