From 6202f9bbb0648d2a5f04787049cb18741acb7f9d Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Wed, 31 Jul 2024 14:51:16 -0700 Subject: [PATCH] fix(http_handler.py): correctly re-raise timeout exception --- litellm/exceptions.py | 7 ++++++- litellm/llms/custom_httpx/http_handler.py | 23 +++++++++++++---------- litellm/llms/predibase.py | 12 ++++++++++++ litellm/proxy/_new_secret_config.yaml | 5 ++--- litellm/proxy/proxy_server.py | 1 + litellm/tests/test_completion.py | 18 +++++++++--------- 6 files changed, 43 insertions(+), 23 deletions(-) diff --git a/litellm/exceptions.py b/litellm/exceptions.py index 04558e437a5..197e64c7589 100644 --- a/litellm/exceptions.py +++ b/litellm/exceptions.py @@ -199,8 +199,12 @@ class Timeout(openai.APITimeoutError): # type: ignore litellm_debug_info: Optional[str] = None, max_retries: Optional[int] = None, num_retries: Optional[int] = None, + headers: Optional[dict] = None, ): - request = httpx.Request(method="POST", url="https://api.openai.com/v1") + request = httpx.Request( + method="POST", + url="https://api.openai.com/v1", + ) super().__init__( request=request ) # Call the base class constructor with the parameters it needs @@ -211,6 +215,7 @@ class Timeout(openai.APITimeoutError): # type: ignore self.litellm_debug_info = litellm_debug_info self.max_retries = max_retries self.num_retries = num_retries + self.headers = headers # custom function to convert to str def __str__(self): diff --git a/litellm/llms/custom_httpx/http_handler.py b/litellm/llms/custom_httpx/http_handler.py index 9fee97bd788..e3a0a4f1c5a 100644 --- a/litellm/llms/custom_httpx/http_handler.py +++ b/litellm/llms/custom_httpx/http_handler.py @@ -84,20 +84,17 @@ class AsyncHTTPHandler: stream: bool = False, ): try: - if timeout is not None: - req = self.client.build_request( - "POST", url, data=data, json=json, params=params, headers=headers, timeout=timeout # type: ignore - ) - else: - req = self.client.build_request( - "POST", url, data=data, json=json, params=params, headers=headers # type: ignore - ) + if timeout is None: + timeout = self.timeout + req = self.client.build_request( + "POST", url, data=data, json=json, params=params, headers=headers, timeout=timeout # type: ignore + ) response = await self.client.send(req, stream=stream) response.raise_for_status() return response except (httpx.RemoteProtocolError, httpx.ConnectError): # Retry the request with a new session if there is a connection error - new_client = self.create_client(timeout=self.timeout, concurrent_limit=1) + new_client = self.create_client(timeout=timeout, concurrent_limit=1) try: return await self.single_connection_post_request( url=url, @@ -110,11 +107,17 @@ class AsyncHTTPHandler: ) finally: await new_client.aclose() - except httpx.TimeoutException: + except httpx.TimeoutException as e: + headers = {} + if hasattr(e, "response") and e.response is not None: + for key, value in e.response.headers.items(): + headers["response_headers-{}".format(key)] = value + raise litellm.Timeout( message=f"Connection timed out after {timeout} seconds.", model="default-model-name", llm_provider="litellm-httpx-handler", + headers=headers, ) except httpx.HTTPStatusError as e: setattr(e, "status_code", e.response.status_code) diff --git a/litellm/llms/predibase.py b/litellm/llms/predibase.py index 68c52ef2e36..8055e06945d 100644 --- a/litellm/llms/predibase.py +++ b/litellm/llms/predibase.py @@ -362,6 +362,15 @@ class PredibaseChatCompletion(BaseLLM): total_tokens=total_tokens, ) model_response.usage = usage # type: ignore + + ## RESPONSE HEADERS + predibase_headers = response.headers + response_headers = {} + for k, v in predibase_headers.items(): + if k.startswith("x-"): + response_headers["llm_provider-{}".format(k)] = v + + model_response._hidden_params["additional_headers"] = response_headers return model_response def completion( @@ -550,6 +559,9 @@ class PredibaseChatCompletion(BaseLLM): ), ) except Exception as e: + for exception in litellm.LITELLM_EXCEPTION_TYPES: + if isinstance(e, exception): + raise e raise PredibaseError( status_code=500, message="{}\n{}".format(str(e), traceback.format_exc()) ) diff --git a/litellm/proxy/_new_secret_config.yaml b/litellm/proxy/_new_secret_config.yaml index 0bd00067a28..eff98ae672e 100644 --- a/litellm/proxy/_new_secret_config.yaml +++ b/litellm/proxy/_new_secret_config.yaml @@ -1,5 +1,4 @@ model_list: - - model_name: claude-3-haiku-20240307 + - model_name: "*" litellm_params: - model: anthropic/claude-3-haiku-20240307 - max_tokens: 4096 \ No newline at end of file + model: "*" \ No newline at end of file diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 5a2970df51d..e12ae5bc285 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -3069,6 +3069,7 @@ async def chat_completion( type=getattr(e, "type", "None"), param=getattr(e, "param", "None"), code=getattr(e, "status_code", 500), + headers=getattr(e, "headers", {}), ) diff --git a/litellm/tests/test_completion.py b/litellm/tests/test_completion.py index aa37d0e11c8..6e26c28c8d8 100644 --- a/litellm/tests/test_completion.py +++ b/litellm/tests/test_completion.py @@ -261,16 +261,16 @@ async def test_completion_predibase(): try: litellm.set_verbose = True - with patch("requests.post", side_effect=predibase_mock_post): - response = completion( - model="predibase/llama-3-8b-instruct", - tenant_id="c4768f95", - api_key=os.getenv("PREDIBASE_API_KEY"), - messages=[{"role": "user", "content": "What is the meaning of life?"}], - max_tokens=10, - ) + # with patch("requests.post", side_effect=predibase_mock_post): + response = await litellm.acompletion( + model="predibase/llama-3-8b-instruct", + tenant_id="c4768f95", + api_key=os.getenv("PREDIBASE_API_KEY"), + messages=[{"role": "user", "content": "What is the meaning of life?"}], + max_tokens=10, + ) - print(response) + print(response) except litellm.Timeout as e: pass except Exception as e: