fix(openai): honor per-call SSL verification

This commit is contained in:
King Star 2026-08-25 15:43:51 +08:00
parent 2488f84b02
commit e14f99dd38
No known key found for this signature in database
3 changed files with 109 additions and 7 deletions

View file

@ -33,6 +33,7 @@ from litellm.llms.custom_httpx.http_handler import (
AsyncHTTPHandler,
get_ssl_configuration,
)
from litellm.types.llms.custom_http import VerifyTypes
def _get_client_init_params(cls: type) -> tuple[str, ...]:
@ -278,6 +279,7 @@ class BaseOpenAILLM:
"organization",
"api_base",
"workload_identity_config",
"ssl_verify",
)
openai_client_fields: Final = (
BaseOpenAILLM.get_openai_client_initialization_param_fields(client_type=client_type)
@ -303,6 +305,7 @@ class BaseOpenAILLM:
@staticmethod
def _get_async_http_client(
shared_session: Optional["ClientSession"] = None,
ssl_verify: VerifyTypes | None = None,
) -> httpx.AsyncClient | None:
if litellm.aclient_session is not None:
return litellm.aclient_session
@ -313,7 +316,7 @@ class BaseOpenAILLM:
return httpx.AsyncClient(transport=MockOpenAITransport())
# Get unified SSL configuration
ssl_config: Final = get_ssl_configuration()
ssl_config: Final = get_ssl_configuration(ssl_verify=ssl_verify)
transport: Final = AsyncHTTPHandler._create_async_transport(
ssl_context=(ssl_config if isinstance(ssl_config, ssl.SSLContext) else None),
ssl_verify=ssl_config if isinstance(ssl_config, bool) else None,
@ -328,7 +331,7 @@ class BaseOpenAILLM:
)
@staticmethod
def _get_sync_http_client() -> httpx.Client | None:
def _get_sync_http_client(ssl_verify: VerifyTypes | None = None) -> httpx.Client | None:
if litellm.client_session is not None:
return litellm.client_session
@ -338,7 +341,7 @@ class BaseOpenAILLM:
return httpx.Client(transport=MockOpenAITransport())
# Get unified SSL configuration
ssl_config: Final = get_ssl_configuration()
ssl_config: Final = get_ssl_configuration(ssl_verify=ssl_verify)
return httpx.Client(
verify=ssl_config,

View file

@ -31,6 +31,7 @@ from litellm.litellm_core_utils.logging_utils import speech_request_body, track_
from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator
from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException
from litellm.llms.bedrock.chat.invoke_handler import MockResponseIterator
from litellm.types.llms.custom_http import VerifyTypes
from litellm.types.utils import (
EmbeddingResponse,
ImageResponse,
@ -382,6 +383,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
organization: str | None = None,
client: OpenAI | AsyncOpenAI | None = None,
shared_session: Optional["ClientSession"] = None,
ssl_verify: VerifyTypes | None = None,
) -> OpenAI | AsyncOpenAI | None:
workload_identity_config: Final = resolve_openai_workload_identity_config(api_key=api_key, api_base=api_base)
client_initialization_params: Final[dict] = locals()
@ -400,7 +402,9 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
if isinstance(cached_client, OpenAI) or isinstance(cached_client, AsyncOpenAI):
return cached_client
if is_async:
async_http_client: Final = OpenAIChatCompletion._get_async_http_client(shared_session=shared_session)
async_http_client: Final = OpenAIChatCompletion._get_async_http_client(
shared_session=shared_session, ssl_verify=ssl_verify
)
http_client: httpx.Client | httpx.AsyncClient | None = async_http_client
_new_client: OpenAI | AsyncOpenAI = (
AsyncOpenAI(
@ -422,7 +426,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
)
)
else:
sync_http_client: Final = OpenAIChatCompletion._get_sync_http_client()
sync_http_client: Final = OpenAIChatCompletion._get_sync_http_client(ssl_verify=ssl_verify)
http_client = sync_http_client
_new_client = (
OpenAI(
@ -750,6 +754,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
drop_params=drop_params,
fake_stream=fake_stream,
shared_session=shared_session,
ssl_verify=litellm_params.get("ssl_verify"),
)
data = provider_config.transform_request(
@ -773,6 +778,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
max_retries=max_retries,
organization=organization,
stream_options=stream_options,
ssl_verify=litellm_params.get("ssl_verify"),
)
else:
if not isinstance(max_retries, int):
@ -786,6 +792,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
max_retries=max_retries,
organization=organization,
client=client,
ssl_verify=litellm_params.get("ssl_verify"),
)
## LOGGING
@ -906,6 +913,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
stream_options: dict | None = None,
fake_stream: bool = False,
shared_session: Optional["ClientSession"] = None,
ssl_verify: VerifyTypes | None = None,
):
response = None
data = await provider_config.async_transform_request(
@ -927,6 +935,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
organization=organization,
client=client,
shared_session=shared_session,
ssl_verify=ssl_verify,
)
## LOGGING
@ -1022,6 +1031,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
max_retries=None,
headers=None,
stream_options: dict | None = None,
ssl_verify: VerifyTypes | None = None,
):
data["stream"] = True
data.update(self.get_stream_options(stream_options=stream_options, api_base=api_base))
@ -1035,6 +1045,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
max_retries=max_retries,
organization=organization,
client=client,
ssl_verify=ssl_verify,
)
## LOGGING
logging_obj.pre_call(
@ -1107,6 +1118,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
organization=organization,
client=client,
shared_session=shared_session,
ssl_verify=litellm_params.get("ssl_verify"),
)
## LOGGING
logging_obj.pre_call(

View file

@ -1,10 +1,9 @@
from unittest.mock import MagicMock, call, patch
from unittest.mock import MagicMock, patch
import httpx
import openai
import pytest
import litellm
from litellm.litellm_core_utils.token_counter import token_counter
from litellm.llms.openai.common_utils import BaseOpenAILLM, is_openai_backed_api_base
@ -174,6 +173,94 @@ def test_get_openai_client_cache_key(client_type):
assert "api_key=sk-test" in key
def test_get_openai_client_cache_key_includes_ssl_verify():
first_key = BaseOpenAILLM.get_openai_client_cache_key(
client_initialization_params={"api_key": "sk-test", "ssl_verify": "/tmp/first-ca.pem"},
client_type="openai",
)
second_key = BaseOpenAILLM.get_openai_client_cache_key(
client_initialization_params={"api_key": "sk-test", "ssl_verify": "/tmp/second-ca.pem"},
client_type="openai",
)
assert first_key != second_key
def test_get_sync_http_client_uses_per_call_ssl_verify(monkeypatch):
from litellm.llms.openai import common_utils
monkeypatch.setattr(litellm, "client_session", None)
monkeypatch.setattr(litellm, "network_mock", False)
with (
patch.object( # test-quality-ok: verify the per-call SSL setting reaches the resolver
common_utils, "get_ssl_configuration", return_value=False
) as get_ssl_configuration,
patch.object( # test-quality-ok: capture the constructed HTTP client options
common_utils.httpx, "Client"
) as http_client,
):
result = BaseOpenAILLM._get_sync_http_client(ssl_verify="/tmp/custom-ca.pem")
assert result is http_client.return_value
get_ssl_configuration.assert_called_once_with(ssl_verify="/tmp/custom-ca.pem")
http_client.assert_called_once_with(verify=False, follow_redirects=True)
@pytest.mark.asyncio
async def test_get_async_http_client_uses_per_call_ssl_verify(monkeypatch):
from litellm.llms.openai import common_utils
monkeypatch.setattr(litellm, "aclient_session", None)
monkeypatch.setattr(litellm, "network_mock", False)
with (
patch.object( # test-quality-ok: verify the per-call SSL setting reaches the resolver
common_utils, "get_ssl_configuration", return_value=False
) as get_ssl_configuration,
patch.object( # test-quality-ok: capture async transport options
common_utils.AsyncHTTPHandler, "_create_async_transport", return_value=None
) as transport,
patch.object( # test-quality-ok: capture the constructed HTTP client options
common_utils.httpx, "AsyncClient"
) as http_client,
):
result = BaseOpenAILLM._get_async_http_client(ssl_verify="/tmp/custom-ca.pem")
assert result is http_client.return_value
get_ssl_configuration.assert_called_once_with(ssl_verify="/tmp/custom-ca.pem")
transport.assert_called_once_with(ssl_context=None, ssl_verify=False, shared_session=None)
http_client.assert_called_once_with(verify=False, transport=None, follow_redirects=True)
def test_openai_client_ssl_verify(monkeypatch): # test-quality-ok: verifies per-call SSL reaches the client boundary
from litellm.llms.openai.openai import OpenAIChatCompletion
monkeypatch.setattr(litellm, "in_memory_llm_clients_cache", MagicMock())
with (
patch.object( # test-quality-ok: force the fresh-client path
BaseOpenAILLM, "get_cached_openai_client", return_value=None
),
patch.object( # test-quality-ok: avoid mutating the shared client cache
BaseOpenAILLM, "set_cached_openai_client"
),
patch.object( # test-quality-ok: observe the HTTP client boundary
OpenAIChatCompletion, "_get_sync_http_client"
) as get_http_client,
patch( # test-quality-ok: avoid constructing a provider client
"litellm.llms.openai.openai.OpenAI"
) as openai_client,
):
OpenAIChatCompletion()._get_openai_client(
is_async=False,
api_key="sk-test",
api_base="https://example.test/v1",
max_retries=2,
ssl_verify="/tmp/custom-ca.pem",
)
get_http_client.assert_called_once_with(ssl_verify="/tmp/custom-ca.pem")
assert openai_client.call_count == 1
def test_evicting_a_client_built_on_the_callers_session_leaves_that_session_open(monkeypatch):
"""`litellm.aclient_session` belongs to the caller, who goes on using it.