mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(openai): pop ssl_verify in main and wire per-request TLS to httpx client (#38178)
This commit is contained in:
parent
bb992386bb
commit
4107ae68f2
4 changed files with 72 additions and 4 deletions
|
|
@ -294,7 +294,7 @@ class BaseOpenAILLM:
|
|||
@staticmethod
|
||||
def _get_async_http_client(
|
||||
shared_session: Optional["ClientSession"] = None,
|
||||
ssl_verify: Optional[Any] = None,
|
||||
ssl_verify: bool | str | None = None,
|
||||
) -> httpx.AsyncClient | None:
|
||||
if litellm.aclient_session is not None:
|
||||
return litellm.aclient_session
|
||||
|
|
@ -319,7 +319,7 @@ class BaseOpenAILLM:
|
|||
|
||||
@staticmethod
|
||||
def _get_sync_http_client(
|
||||
ssl_verify: Optional[Any] = None,
|
||||
ssl_verify: bool | str | None = None,
|
||||
) -> httpx.Client | None:
|
||||
if litellm.client_session is not None:
|
||||
return litellm.client_session
|
||||
|
|
|
|||
|
|
@ -717,6 +717,7 @@ 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):
|
||||
|
|
@ -983,6 +984,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
|
|||
max_retries=max_retries,
|
||||
organization=organization,
|
||||
client=client,
|
||||
ssl_verify=ssl_verify,
|
||||
)
|
||||
## LOGGING
|
||||
logging_obj.pre_call(
|
||||
|
|
@ -2276,6 +2278,7 @@ class OpenAIAssistantsAPI(BaseLLM):
|
|||
max_retries=max_retries,
|
||||
organization=organization,
|
||||
client=client,
|
||||
ssl_verify=ssl_verify,
|
||||
)
|
||||
request_params: Final = {
|
||||
"order": order,
|
||||
|
|
@ -2350,6 +2353,7 @@ class OpenAIAssistantsAPI(BaseLLM):
|
|||
max_retries=max_retries,
|
||||
organization=organization,
|
||||
client=client,
|
||||
ssl_verify=ssl_verify,
|
||||
)
|
||||
|
||||
request_params: Final = {
|
||||
|
|
@ -2384,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)
|
||||
|
|
@ -2418,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)
|
||||
|
|
@ -2441,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)
|
||||
|
|
@ -2475,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)
|
||||
|
|
@ -2500,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(
|
||||
|
|
@ -2579,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(
|
||||
|
|
@ -2611,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)
|
||||
|
|
@ -2677,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)
|
||||
|
|
@ -2703,6 +2715,7 @@ class OpenAIAssistantsAPI(BaseLLM):
|
|||
max_retries=max_retries,
|
||||
organization=organization,
|
||||
client=client,
|
||||
ssl_verify=ssl_verify,
|
||||
)
|
||||
|
||||
data: Final = {}
|
||||
|
|
@ -2789,6 +2802,7 @@ class OpenAIAssistantsAPI(BaseLLM):
|
|||
max_retries=max_retries,
|
||||
organization=organization,
|
||||
client=client,
|
||||
ssl_verify=ssl_verify,
|
||||
)
|
||||
|
||||
data: Final = {}
|
||||
|
|
@ -2818,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)
|
||||
|
|
@ -2884,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)
|
||||
|
|
@ -2919,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(
|
||||
|
|
@ -3103,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:
|
||||
|
|
|
|||
|
|
@ -5118,7 +5118,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)
|
||||
|
|
@ -7027,7 +7027,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(
|
||||
|
|
|
|||
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