mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge 617630bde9 into 6e569ee0c7
This commit is contained in:
commit
ce6ea2676e
6 changed files with 159 additions and 7 deletions
|
|
@ -268,6 +268,7 @@ class BaseOpenAILLM:
|
|||
"max_retries",
|
||||
"organization",
|
||||
"api_base",
|
||||
"ssl_verify",
|
||||
)
|
||||
openai_client_fields: Final = (
|
||||
BaseOpenAILLM.get_openai_client_initialization_param_fields(client_type=client_type)
|
||||
|
|
@ -293,6 +294,7 @@ class BaseOpenAILLM:
|
|||
@staticmethod
|
||||
def _get_async_http_client(
|
||||
shared_session: Optional["ClientSession"] = None,
|
||||
ssl_verify: bool | str | None = None,
|
||||
) -> httpx.AsyncClient | None:
|
||||
if litellm.aclient_session is not None:
|
||||
return litellm.aclient_session
|
||||
|
|
@ -303,7 +305,7 @@ class BaseOpenAILLM:
|
|||
return httpx.AsyncClient(transport=MockOpenAITransport())
|
||||
|
||||
# Get unified SSL configuration
|
||||
ssl_config: Final = get_ssl_configuration()
|
||||
ssl_config: Final = ssl_verify if ssl_verify is not None else get_ssl_configuration()
|
||||
|
||||
return httpx.AsyncClient(
|
||||
verify=ssl_config,
|
||||
|
|
@ -316,7 +318,9 @@ class BaseOpenAILLM:
|
|||
)
|
||||
|
||||
@staticmethod
|
||||
def _get_sync_http_client() -> httpx.Client | None:
|
||||
def _get_sync_http_client(
|
||||
ssl_verify: bool | str | None = None,
|
||||
) -> httpx.Client | None:
|
||||
if litellm.client_session is not None:
|
||||
return litellm.client_session
|
||||
|
||||
|
|
@ -326,7 +330,7 @@ class BaseOpenAILLM:
|
|||
return httpx.Client(transport=MockOpenAITransport())
|
||||
|
||||
# Get unified SSL configuration
|
||||
ssl_config: Final = get_ssl_configuration()
|
||||
ssl_config: Final = ssl_verify if ssl_verify is not None else get_ssl_configuration()
|
||||
|
||||
return httpx.Client(
|
||||
verify=ssl_config,
|
||||
|
|
|
|||
|
|
@ -347,6 +347,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
|
|||
organization: str | None = None,
|
||||
client: OpenAI | AsyncOpenAI | None = None,
|
||||
shared_session: Optional["ClientSession"] = None,
|
||||
ssl_verify: Any | None = None,
|
||||
) -> OpenAI | AsyncOpenAI | None:
|
||||
client_initialization_params: Final[dict] = locals()
|
||||
if client is None:
|
||||
|
|
@ -364,9 +365,9 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
|
|||
if isinstance(cached_client, OpenAI) or isinstance(cached_client, AsyncOpenAI):
|
||||
return cached_client
|
||||
http_client: Final[httpx.Client | httpx.AsyncClient | None] = (
|
||||
OpenAIChatCompletion._get_async_http_client(shared_session=shared_session)
|
||||
OpenAIChatCompletion._get_async_http_client(shared_session=shared_session, ssl_verify=ssl_verify)
|
||||
if is_async
|
||||
else OpenAIChatCompletion._get_sync_http_client()
|
||||
else OpenAIChatCompletion._get_sync_http_client(ssl_verify=ssl_verify)
|
||||
)
|
||||
if is_async:
|
||||
_new_client: OpenAI | AsyncOpenAI = AsyncOpenAI(
|
||||
|
|
@ -716,10 +717,12 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
|
|||
max_retries=max_retries,
|
||||
organization=organization,
|
||||
stream_options=stream_options,
|
||||
ssl_verify=litellm_params.get("ssl_verify", None),
|
||||
)
|
||||
else:
|
||||
if not isinstance(max_retries, int):
|
||||
raise OpenAIError(status_code=422, message="max retries must be an int")
|
||||
ssl_verify: Final = litellm_params.get("ssl_verify", None) if litellm_params else None
|
||||
openai_client: OpenAI = self._get_openai_client(
|
||||
is_async=False,
|
||||
api_key=api_key,
|
||||
|
|
@ -729,6 +732,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
|
|||
max_retries=max_retries,
|
||||
organization=organization,
|
||||
client=client,
|
||||
ssl_verify=ssl_verify,
|
||||
)
|
||||
|
||||
## LOGGING
|
||||
|
|
@ -860,6 +864,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
|
|||
)
|
||||
for _ in range(2): # if call fails due to alternating messages, retry with reformatted message
|
||||
try:
|
||||
ssl_verify: Final = litellm_params.get("ssl_verify", None) if litellm_params else None
|
||||
openai_aclient: AsyncOpenAI = self._get_openai_client(
|
||||
is_async=True,
|
||||
api_key=api_key,
|
||||
|
|
@ -870,6 +875,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
|
|||
organization=organization,
|
||||
client=client,
|
||||
shared_session=shared_session,
|
||||
ssl_verify=ssl_verify,
|
||||
)
|
||||
|
||||
## LOGGING
|
||||
|
|
@ -978,6 +984,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
|
|||
max_retries=max_retries,
|
||||
organization=organization,
|
||||
client=client,
|
||||
ssl_verify=ssl_verify,
|
||||
)
|
||||
## LOGGING
|
||||
logging_obj.pre_call(
|
||||
|
|
@ -1040,6 +1047,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
|
|||
data.update(self.get_stream_options(stream_options=stream_options, api_base=api_base))
|
||||
for _ in range(2):
|
||||
try:
|
||||
ssl_verify: Final = litellm_params.get("ssl_verify", None) if litellm_params else None
|
||||
openai_aclient: AsyncOpenAI = self._get_openai_client(
|
||||
is_async=True,
|
||||
api_key=api_key,
|
||||
|
|
@ -1050,6 +1058,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
|
|||
organization=organization,
|
||||
client=client,
|
||||
shared_session=shared_session,
|
||||
ssl_verify=ssl_verify,
|
||||
)
|
||||
## LOGGING
|
||||
logging_obj.pre_call(
|
||||
|
|
@ -2269,6 +2278,7 @@ class OpenAIAssistantsAPI(BaseLLM):
|
|||
max_retries=max_retries,
|
||||
organization=organization,
|
||||
client=client,
|
||||
ssl_verify=ssl_verify,
|
||||
)
|
||||
request_params: Final = {
|
||||
"order": order,
|
||||
|
|
@ -2343,6 +2353,7 @@ class OpenAIAssistantsAPI(BaseLLM):
|
|||
max_retries=max_retries,
|
||||
organization=organization,
|
||||
client=client,
|
||||
ssl_verify=ssl_verify,
|
||||
)
|
||||
|
||||
request_params: Final = {
|
||||
|
|
@ -2377,6 +2388,7 @@ class OpenAIAssistantsAPI(BaseLLM):
|
|||
max_retries=max_retries,
|
||||
organization=organization,
|
||||
client=client,
|
||||
ssl_verify=ssl_verify,
|
||||
)
|
||||
|
||||
response: Final = await openai_client.beta.assistants.create(**create_assistant_data)
|
||||
|
|
@ -2411,6 +2423,7 @@ class OpenAIAssistantsAPI(BaseLLM):
|
|||
max_retries=max_retries,
|
||||
organization=organization,
|
||||
client=client,
|
||||
ssl_verify=ssl_verify,
|
||||
)
|
||||
|
||||
response: Final = openai_client.beta.assistants.create(**create_assistant_data)
|
||||
|
|
@ -2434,6 +2447,7 @@ class OpenAIAssistantsAPI(BaseLLM):
|
|||
max_retries=max_retries,
|
||||
organization=organization,
|
||||
client=client,
|
||||
ssl_verify=ssl_verify,
|
||||
)
|
||||
|
||||
response: Final = await openai_client.beta.assistants.delete(assistant_id=assistant_id)
|
||||
|
|
@ -2468,6 +2482,7 @@ class OpenAIAssistantsAPI(BaseLLM):
|
|||
max_retries=max_retries,
|
||||
organization=organization,
|
||||
client=client,
|
||||
ssl_verify=ssl_verify,
|
||||
)
|
||||
|
||||
response: Final = openai_client.beta.assistants.delete(assistant_id=assistant_id)
|
||||
|
|
@ -2493,6 +2508,7 @@ class OpenAIAssistantsAPI(BaseLLM):
|
|||
max_retries=max_retries,
|
||||
organization=organization,
|
||||
client=client,
|
||||
ssl_verify=ssl_verify,
|
||||
)
|
||||
|
||||
thread_message: Final[OpenAIMessage] = await openai_client.beta.threads.messages.create(
|
||||
|
|
@ -2572,6 +2588,7 @@ class OpenAIAssistantsAPI(BaseLLM):
|
|||
max_retries=max_retries,
|
||||
organization=organization,
|
||||
client=client,
|
||||
ssl_verify=ssl_verify,
|
||||
)
|
||||
|
||||
thread_message: Final[OpenAIMessage] = openai_client.beta.threads.messages.create(
|
||||
|
|
@ -2604,6 +2621,7 @@ class OpenAIAssistantsAPI(BaseLLM):
|
|||
max_retries=max_retries,
|
||||
organization=organization,
|
||||
client=client,
|
||||
ssl_verify=ssl_verify,
|
||||
)
|
||||
|
||||
response: Final = await openai_client.beta.threads.messages.list(thread_id=thread_id)
|
||||
|
|
@ -2670,6 +2688,7 @@ class OpenAIAssistantsAPI(BaseLLM):
|
|||
max_retries=max_retries,
|
||||
organization=organization,
|
||||
client=client,
|
||||
ssl_verify=ssl_verify,
|
||||
)
|
||||
|
||||
response: Final = openai_client.beta.threads.messages.list(thread_id=thread_id)
|
||||
|
|
@ -2696,6 +2715,7 @@ class OpenAIAssistantsAPI(BaseLLM):
|
|||
max_retries=max_retries,
|
||||
organization=organization,
|
||||
client=client,
|
||||
ssl_verify=ssl_verify,
|
||||
)
|
||||
|
||||
data: Final = {}
|
||||
|
|
@ -2782,6 +2802,7 @@ class OpenAIAssistantsAPI(BaseLLM):
|
|||
max_retries=max_retries,
|
||||
organization=organization,
|
||||
client=client,
|
||||
ssl_verify=ssl_verify,
|
||||
)
|
||||
|
||||
data: Final = {}
|
||||
|
|
@ -2811,6 +2832,7 @@ class OpenAIAssistantsAPI(BaseLLM):
|
|||
max_retries=max_retries,
|
||||
organization=organization,
|
||||
client=client,
|
||||
ssl_verify=ssl_verify,
|
||||
)
|
||||
|
||||
response: Final = await openai_client.beta.threads.retrieve(thread_id=thread_id)
|
||||
|
|
@ -2877,6 +2899,7 @@ class OpenAIAssistantsAPI(BaseLLM):
|
|||
max_retries=max_retries,
|
||||
organization=organization,
|
||||
client=client,
|
||||
ssl_verify=ssl_verify,
|
||||
)
|
||||
|
||||
response: Final = openai_client.beta.threads.retrieve(thread_id=thread_id)
|
||||
|
|
@ -2912,6 +2935,7 @@ class OpenAIAssistantsAPI(BaseLLM):
|
|||
max_retries=max_retries,
|
||||
organization=organization,
|
||||
client=client,
|
||||
ssl_verify=ssl_verify,
|
||||
)
|
||||
|
||||
response: Final = await openai_client.beta.threads.runs.create_and_poll(
|
||||
|
|
@ -3096,6 +3120,7 @@ class OpenAIAssistantsAPI(BaseLLM):
|
|||
max_retries=max_retries,
|
||||
organization=organization,
|
||||
client=client,
|
||||
ssl_verify=ssl_verify,
|
||||
)
|
||||
|
||||
if stream is not None and stream is True:
|
||||
|
|
|
|||
|
|
@ -5169,7 +5169,7 @@ def completion(
|
|||
context_window_fallback_dict: Final = kwargs.get("context_window_fallback_dict", None)
|
||||
organization: Final = kwargs.get("organization", None)
|
||||
### VERIFY SSL ###
|
||||
ssl_verify: Final = kwargs.get("ssl_verify", None)
|
||||
ssl_verify: Final = kwargs.pop("ssl_verify", None)
|
||||
### CUSTOM MODEL COST ###
|
||||
input_cost_per_token: Final = kwargs.get("input_cost_per_token", None)
|
||||
output_cost_per_token: Final = kwargs.get("output_cost_per_token", None)
|
||||
|
|
@ -7071,7 +7071,7 @@ def embedding(
|
|||
optional_params=optional_params,
|
||||
client=client,
|
||||
aembedding=aembedding,
|
||||
litellm_params={"ssl_verify": kwargs.get("ssl_verify", None)},
|
||||
litellm_params={"ssl_verify": kwargs.pop("ssl_verify", None)},
|
||||
)
|
||||
elif custom_llm_provider == "perplexity":
|
||||
response = base_llm_http_handler.embedding(
|
||||
|
|
|
|||
|
|
@ -4666,6 +4666,8 @@ def add_provider_specific_params_to_optional_params(
|
|||
extra_body: Final = dict(passed_params.pop("extra_body", None) or {})
|
||||
for k in passed_params:
|
||||
if k not in openai_params and passed_params[k] is not None:
|
||||
if k in ["ssl_verify"]:
|
||||
continue
|
||||
extra_body[k] = passed_params[k]
|
||||
if not isinstance(optional_params.get("extra_body"), dict):
|
||||
optional_params["extra_body"] = {}
|
||||
|
|
|
|||
71
tests/test_litellm/llms/openai/test_openai_ssl_verify.py
Normal file
71
tests/test_litellm/llms/openai/test_openai_ssl_verify.py
Normal file
|
|
@ -0,0 +1,71 @@
|
|||
from unittest.mock import MagicMock, patch
|
||||
import httpx
|
||||
import litellm
|
||||
from litellm.llms.openai.common_utils import BaseOpenAILLM, OpenAIChatCompletion
|
||||
from litellm.utils import get_optional_params, add_provider_specific_params_to_optional_params
|
||||
|
||||
|
||||
def test_ssl_verify_not_in_extra_body():
|
||||
"""
|
||||
Ensure ssl_verify is NOT dumped into extra_body for openai and openai-compatible providers.
|
||||
Issue #38178: ssl_verify was leaking into extra_body payload sent to OpenAI-compatible endpoints.
|
||||
"""
|
||||
optional_params = {}
|
||||
passed_params = {
|
||||
"ssl_verify": "/custom/path/ca.pem",
|
||||
"temperature": 0.7,
|
||||
"custom_param": "value",
|
||||
}
|
||||
|
||||
result = add_provider_specific_params_to_optional_params(
|
||||
optional_params=optional_params,
|
||||
passed_params=passed_params,
|
||||
custom_llm_provider="openai",
|
||||
openai_params=["temperature"],
|
||||
)
|
||||
|
||||
extra_body = result.get("extra_body", {})
|
||||
assert "ssl_verify" not in extra_body
|
||||
assert extra_body.get("custom_param") == "value"
|
||||
|
||||
|
||||
def test_get_sync_http_client_with_ssl_verify():
|
||||
"""
|
||||
Verify _get_sync_http_client applies per-call ssl_verify to httpx.Client(verify=...).
|
||||
"""
|
||||
custom_ca = "/path/to/my_custom_ca.pem"
|
||||
client = OpenAIChatCompletion._get_sync_http_client(ssl_verify=custom_ca)
|
||||
assert client is not None
|
||||
client_false = OpenAIChatCompletion._get_sync_http_client(ssl_verify=False)
|
||||
assert client_false is not None
|
||||
|
||||
|
||||
def test_get_async_http_client_with_ssl_verify():
|
||||
"""
|
||||
Verify _get_async_http_client applies per-call ssl_verify to httpx.AsyncClient(verify=...).
|
||||
"""
|
||||
client_false = OpenAIChatCompletion._get_async_http_client(ssl_verify=False)
|
||||
assert client_false is not None
|
||||
|
||||
|
||||
def test_cache_key_differs_by_ssl_verify():
|
||||
"""
|
||||
Verify cache keys differ when ssl_verify differs to prevent client poisoning across CAs.
|
||||
"""
|
||||
params_ca1 = {
|
||||
"api_key": "sk-1234",
|
||||
"is_async": True,
|
||||
"ssl_verify": "/path/ca1.pem",
|
||||
}
|
||||
params_ca2 = {
|
||||
"api_key": "sk-1234",
|
||||
"is_async": True,
|
||||
"ssl_verify": "/path/ca2.pem",
|
||||
}
|
||||
|
||||
key1 = BaseOpenAILLM.get_openai_client_cache_key(params_ca1, "openai")
|
||||
key2 = BaseOpenAILLM.get_openai_client_cache_key(params_ca2, "openai")
|
||||
|
||||
assert key1 != key2
|
||||
assert "ssl_verify=/path/ca1.pem" in key1
|
||||
assert "ssl_verify=/path/ca2.pem" in key2
|
||||
50
tests/test_litellm_ssl_verify.py
Normal file
50
tests/test_litellm_ssl_verify.py
Normal file
|
|
@ -0,0 +1,50 @@
|
|||
import pytest
|
||||
from unittest.mock import patch, MagicMock
|
||||
import litellm
|
||||
import httpx
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ssl_verify_false():
|
||||
with patch("httpx.AsyncClient") as mock_client:
|
||||
mock_client.return_value.post.return_value = MagicMock(status_code=200, json=lambda: {"choices": [{"message": {"content": "hello"}}]})
|
||||
|
||||
response = await litellm.acompletion(
|
||||
model="openai/gpt-3.5-turbo",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
ssl_verify=False,
|
||||
api_key="sk-123"
|
||||
)
|
||||
# Check that AsyncClient was initialized with verify=False
|
||||
called_kwargs = mock_client.call_args[1]
|
||||
assert called_kwargs.get("verify") is False
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ssl_verify_custom_ca():
|
||||
with patch("httpx.AsyncClient") as mock_client:
|
||||
mock_client.return_value.post.return_value = MagicMock(status_code=200, json=lambda: {"choices": [{"message": {"content": "hello"}}]})
|
||||
|
||||
custom_ca_path = "/path/to/custom-ca.pem"
|
||||
response = await litellm.acompletion(
|
||||
model="openai/gpt-3.5-turbo",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
ssl_verify=custom_ca_path,
|
||||
api_key="sk-123"
|
||||
)
|
||||
# Check that AsyncClient was initialized with verify=custom_ca_path
|
||||
called_kwargs = mock_client.call_args[1]
|
||||
assert called_kwargs.get("verify") == custom_ca_path
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ssl_verify_default():
|
||||
with patch("httpx.AsyncClient") as mock_client:
|
||||
mock_client.return_value.post.return_value = MagicMock(status_code=200, json=lambda: {"choices": [{"message": {"content": "hello"}}]})
|
||||
|
||||
response = await litellm.acompletion(
|
||||
model="openai/gpt-3.5-turbo",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
api_key="sk-123"
|
||||
)
|
||||
# By default, should not pass verify=False (usually defaults to True or SSLContext depending on get_ssl_configuration)
|
||||
called_kwargs = mock_client.call_args[1]
|
||||
verify_arg = called_kwargs.get("verify")
|
||||
assert verify_arg is not False and verify_arg is not None
|
||||
Loading…
Add table
Reference in a new issue