mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
parent
5325ee4382
commit
a1964eab18
5 changed files with 42 additions and 65 deletions
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue