diff --git a/litellm/llms/custom_httpx/http_handler.py b/litellm/llms/custom_httpx/http_handler.py index ee0cae44e23..517cad25b0e 100644 --- a/litellm/llms/custom_httpx/http_handler.py +++ b/litellm/llms/custom_httpx/http_handler.py @@ -93,11 +93,15 @@ class AsyncHTTPHandler: event_hooks: Optional[Mapping[str, List[Callable[..., Any]]]] = None, concurrent_limit=1000, client_alias: Optional[str] = None, # name for client in logs + ssl_verify: Optional[Union[bool, str]] = None, ): self.timeout = timeout self.event_hooks = event_hooks self.client = self.create_client( - timeout=timeout, concurrent_limit=concurrent_limit, event_hooks=event_hooks + timeout=timeout, + concurrent_limit=concurrent_limit, + event_hooks=event_hooks, + ssl_verify=ssl_verify, ) self.client_alias = client_alias @@ -106,11 +110,13 @@ class AsyncHTTPHandler: timeout: Optional[Union[float, httpx.Timeout]], concurrent_limit: int, event_hooks: Optional[Mapping[str, List[Callable[..., Any]]]], + ssl_verify: Optional[Union[bool, str]] = None, ) -> httpx.AsyncClient: # 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) @@ -672,6 +678,7 @@ def get_async_httpx_client( _new_client = AsyncHTTPHandler( timeout=httpx.Timeout(timeout=600.0, connect=5.0) ) + litellm.in_memory_llm_clients_cache.set_cache( key=_cache_key_name, value=_new_client, diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 5041c8252ee..71a8a8168b5 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -158,7 +158,8 @@ class BaseLLMHTTPHandler: ): if client is None: async_httpx_client = get_async_httpx_client( - llm_provider=litellm.LlmProviders(custom_llm_provider) + llm_provider=litellm.LlmProviders(custom_llm_provider), + params={"ssl_verify": litellm_params.get("ssl_verify", None)}, ) else: async_httpx_client = client @@ -361,7 +362,11 @@ class BaseLLMHTTPHandler: client: Optional[HTTPHandler] = None, ) -> Tuple[Any, dict]: if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client() + sync_httpx_client = _get_httpx_client( + { + "ssl_verify": litellm_params.get("ssl_verify", None), + } + ) else: sync_httpx_client = client stream = True @@ -413,7 +418,7 @@ class BaseLLMHTTPHandler: fake_stream: bool = False, client: Optional[AsyncHTTPHandler] = None, ): - completion_stream, _response_headers = await self.make_async_call( + completion_stream, _response_headers = await self.make_async_call_stream_helper( custom_llm_provider=custom_llm_provider, provider_config=provider_config, api_base=api_base, @@ -434,7 +439,7 @@ class BaseLLMHTTPHandler: ) return streamwrapper - async def make_async_call( + async def make_async_call_stream_helper( self, custom_llm_provider: str, provider_config: BaseConfig, @@ -448,9 +453,15 @@ class BaseLLMHTTPHandler: fake_stream: bool = False, client: Optional[AsyncHTTPHandler] = None, ) -> Tuple[Any, httpx.Headers]: + """ + Helper function for making an async call with stream. + + Handles fake stream as well. + """ if client is None: async_httpx_client = get_async_httpx_client( - llm_provider=litellm.LlmProviders(custom_llm_provider) + llm_provider=litellm.LlmProviders(custom_llm_provider), + params={"ssl_verify": litellm_params.get("ssl_verify", None)}, ) else: async_httpx_client = client diff --git a/tests/local_testing/test_ollama.py b/tests/local_testing/test_ollama.py index dab099b65e9..81cd331263f 100644 --- a/tests/local_testing/test_ollama.py +++ b/tests/local_testing/test_ollama.py @@ -205,3 +205,36 @@ def test_ollama_ssl_verify(): client.client._transport._pool._ssl_context.verify_mode == test_client._transport._pool._ssl_context.verify_mode ) + + +@pytest.mark.parametrize("stream", [True, False]) +@pytest.mark.asyncio +async def test_async_ollama_ssl_verify(stream): + from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler + import httpx + + try: + response = await litellm.acompletion( + model="ollama/llama3.1", + messages=[ + { + "role": "user", + "content": "What's the weather like in San Francisco?", + } + ], + ssl_verify=False, + stream=stream, + ) + except Exception as e: + print(e) + + client: AsyncHTTPHandler = litellm.in_memory_llm_clients_cache.get_cache( + "async_httpx_clientssl_verify_Falseollama" + ) + + test_client = httpx.AsyncClient(verify=False) + print(client) + assert ( + client.client._transport._pool._ssl_context.verify_mode + == test_client._transport._pool._ssl_context.verify_mode + )