diff --git a/litellm/proxy/_new_secret_config.yaml b/litellm/proxy/_new_secret_config.yaml index ff9e0c8f367..da1a17094f1 100644 --- a/litellm/proxy/_new_secret_config.yaml +++ b/litellm/proxy/_new_secret_config.yaml @@ -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 \ No newline at end of file diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 60b7474bd6c..cbabcf7f768 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -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 diff --git a/litellm/router.py b/litellm/router.py index ffb589a8fdc..3e9087a100a 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -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, diff --git a/tests/proxy_unit_tests/test_user_api_key_auth.py b/tests/proxy_unit_tests/test_user_api_key_auth.py index dd23a6a9cf7..e4330d96475 100644 --- a/tests/proxy_unit_tests/test_user_api_key_auth.py +++ b/tests/proxy_unit_tests/test_user_api_key_auth.py @@ -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" ) diff --git a/tests/router_unit_tests/test_router_helper_utils.py b/tests/router_unit_tests/test_router_helper_utils.py index b17f0c0a5e1..8a27f2147ce 100644 --- a/tests/router_unit_tests/test_router_helper_utils.py +++ b/tests/router_unit_tests/test_router_helper_utils.py @@ -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):