From 984534525686d72dc399ef752caa05fa93ca6f67 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Wed, 22 Jan 2025 10:47:41 -0800 Subject: [PATCH] fix(http_handler.py): support passing ssl verify dynamically and using the correct httpx client based on passed ssl verify param Fixes https://github.com/BerriAI/litellm/issues/6499 --- litellm/llms/custom_httpx/http_handler.py | 17 ++++++++-- litellm/llms/custom_httpx/llm_http_handler.py | 4 ++- litellm/main.py | 3 ++ litellm/utils.py | 2 ++ tests/local_testing/test_ollama.py | 31 +++++++++++++++++++ 5 files changed, 54 insertions(+), 3 deletions(-) diff --git a/litellm/llms/custom_httpx/http_handler.py b/litellm/llms/custom_httpx/http_handler.py index 469bb693fb3..ee0cae44e23 100644 --- a/litellm/llms/custom_httpx/http_handler.py +++ b/litellm/llms/custom_httpx/http_handler.py @@ -440,13 +440,17 @@ class HTTPHandler: timeout: Optional[Union[float, httpx.Timeout]] = None, concurrent_limit=1000, client: Optional[httpx.Client] = None, + ssl_verify: Optional[Union[bool, str]] = None, ): if timeout is None: timeout = _DEFAULT_TIMEOUT # SSL certificates (a.k.a CA bundle) used to verify the identity of requested hosts. # /path/to/certificate.pem - ssl_verify = os.getenv("SSL_VERIFY", litellm.ssl_verify) + + if ssl_verify is None: + ssl_verify = os.getenv("SSL_VERIFY", litellm.ssl_verify) + # An SSL certificate used by the requested host to authenticate the client. # /path/to/client.pem cert = os.getenv("SSL_CERTIFICATE", litellm.ssl_certificate) @@ -506,7 +510,15 @@ class HTTPHandler: try: if timeout is not None: req = self.client.build_request( - "POST", url, data=data, json=json, params=params, headers=headers, timeout=timeout, files=files, content=content # type: ignore + "POST", + url, + data=data, # type: ignore + json=json, + params=params, + headers=headers, + timeout=timeout, + files=files, + content=content, # type: ignore ) else: req = self.client.build_request( @@ -684,6 +696,7 @@ def _get_httpx_client(params: Optional[dict] = None) -> HTTPHandler: pass _cache_key_name = "httpx_client" + _params_key_name + _cached_client = litellm.in_memory_llm_clients_cache.get_cache(_cache_key_name) if _cached_client: return _cached_client diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index c7ba9cd0961..5041c8252ee 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -318,7 +318,9 @@ class BaseLLMHTTPHandler: ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client() + sync_httpx_client = _get_httpx_client( + params={"ssl_verify": litellm_params.get("ssl_verify", None)} + ) else: sync_httpx_client = client diff --git a/litellm/main.py b/litellm/main.py index 8042fb1cc80..db0c6ac54b9 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -851,6 +851,8 @@ def completion( # type: ignore # noqa: PLR0915 cooldown_time = kwargs.get("cooldown_time", None) context_window_fallback_dict = kwargs.get("context_window_fallback_dict", None) organization = kwargs.get("organization", None) + ### VERIFY SSL ### + ssl_verify = kwargs.get("ssl_verify", None) ### CUSTOM MODEL COST ### input_cost_per_token = kwargs.get("input_cost_per_token", None) output_cost_per_token = kwargs.get("output_cost_per_token", None) @@ -1091,6 +1093,7 @@ def completion( # type: ignore # noqa: PLR0915 drop_params=kwargs.get("drop_params"), prompt_id=prompt_id, prompt_variables=prompt_variables, + ssl_verify=ssl_verify, ) logging.update_environment_variables( model=model, diff --git a/litellm/utils.py b/litellm/utils.py index 37d721fc29c..ffe1586f879 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -2074,6 +2074,7 @@ def get_litellm_params( prompt_id: Optional[str] = None, prompt_variables: Optional[dict] = None, async_call: Optional[bool] = None, + ssl_verify: Optional[bool] = None, **kwargs, ) -> dict: litellm_params = { @@ -2112,6 +2113,7 @@ def get_litellm_params( "prompt_id": prompt_id, "prompt_variables": prompt_variables, "async_call": async_call, + "ssl_verify": ssl_verify, } return litellm_params diff --git a/tests/local_testing/test_ollama.py b/tests/local_testing/test_ollama.py index 2066859091e..dab099b65e9 100644 --- a/tests/local_testing/test_ollama.py +++ b/tests/local_testing/test_ollama.py @@ -174,3 +174,34 @@ def test_ollama_chat_function_calling(): print(json.loads(tool_calls[0].function.arguments)) print(response) + + +def test_ollama_ssl_verify(): + from litellm.llms.custom_httpx.http_handler import HTTPHandler + import ssl + import httpx + + try: + response = litellm.completion( + model="ollama/llama3.1", + messages=[ + { + "role": "user", + "content": "What's the weather like in San Francisco?", + } + ], + ssl_verify=False, + ) + except Exception as e: + print(e) + + client: HTTPHandler = litellm.in_memory_llm_clients_cache.get_cache( + "httpx_clientssl_verify_False" + ) + + test_client = httpx.Client(verify=False) + print(client) + assert ( + client.client._transport._pool._ssl_context.verify_mode + == test_client._transport._pool._ssl_context.verify_mode + )