mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(custom_httpx): address Greptile 3/5 gaps in ssl_verify propagation
Three followups for #30810 (PR #30810, issue #30778): 1. **Forward ssl= kwarg for externally-provided sessions.** When a caller passes their own aiohttp ClientSession (via async_completion's client argument or BaseLLMAIOHTTPHandler's constructor), the lazily-created transport's connector-level SSL setting is bypassed. Previously the request would use the session's default SSL context. We now forward `self.ssl_verify` as a per-request `ssl=` kwarg to `async_client_session.post()`, matching the pattern already used in `LiteLLMAiohttpTransport._make_aiohttp_request`. 2. **Stop dropping string CA-bundle paths.** `AsyncHTTPHandler._create_aiohttp_transport` and `_get_ssl_connector_kwargs` accepted only `Optional[bool]` for ssl_verify, silently dropping string CA-bundle paths. We now accept `Optional[Union[bool, str]]` and build an `ssl.SSLContext` via `ssl.create_default_context(cafile=...)` so aiohttp's TCPConnector honors the caller's CA bundle. Fall back to default trust store on FileNotFoundError rather than disabling verification. 3. **Tests for the missing verb-method retry paths.** The original PR only added a test for `post`. Added parametrize coverage for `put`, `patch`, `delete` retry paths plus a regression test for string CA-bundle paths and the external-session ssl= kwarg behavior. All 46 tests in tests/test_litellm/llms/custom_httpx/test_http_handler.py pass; 25 existing anthropic handler tests still pass.
This commit is contained in:
parent
14024df9a7
commit
247b21a5fa
4 changed files with 205 additions and 19 deletions
|
|
@ -1,4 +1,4 @@
|
|||
from typing import TYPE_CHECKING, Any, Callable, Optional, Tuple, Union, cast
|
||||
from typing import TYPE_CHECKING, Any, Callable, Dict, Optional, Tuple, Union, cast
|
||||
|
||||
import aiohttp
|
||||
import httpx # type: ignore
|
||||
|
|
@ -201,12 +201,25 @@ class BaseLLMAIOHTTPHandler:
|
|||
|
||||
for i in range(max(max_retry_on_unprocessable_entity_error, 1)):
|
||||
try:
|
||||
response = await async_client_session.post(
|
||||
url=api_base,
|
||||
headers=headers,
|
||||
json=data,
|
||||
data=form_data,
|
||||
)
|
||||
# When a caller passes an external ``ClientSession`` (e.g. via
|
||||
# the ``client`` argument on ``async_completion``), the lazily
|
||||
# created transport's connector-level ssl setting is bypassed
|
||||
# entirely — the session was built without it. Forward
|
||||
# ``ssl_verify`` as a per-request kwarg so the caller's SSL
|
||||
# setting is honored on every request, including retries. Only
|
||||
# pass ``ssl`` when explicitly configured; passing ``ssl=None``
|
||||
# would override the session/connector default with a sentinel
|
||||
# aiohttp treats as "use default" but we still want to avoid
|
||||
# the case where a user explicitly disabled verification.
|
||||
post_kwargs: Dict[str, Any] = {
|
||||
"url": api_base,
|
||||
"headers": headers,
|
||||
"json": data,
|
||||
"data": form_data,
|
||||
}
|
||||
if self.ssl_verify is not None:
|
||||
post_kwargs["ssl"] = self.ssl_verify
|
||||
response = await async_client_session.post(**post_kwargs)
|
||||
if not response.ok:
|
||||
response.raise_for_status()
|
||||
except aiohttp.ClientResponseError as e:
|
||||
|
|
|
|||
|
|
@ -152,7 +152,7 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
|
|||
def __init__(
|
||||
self,
|
||||
client: Union[ClientSession, Callable[[], ClientSession]],
|
||||
ssl_verify: Optional[Union[bool, ssl.SSLContext]] = None,
|
||||
ssl_verify: Optional[Union[bool, str, ssl.SSLContext]] = None,
|
||||
owns_session: bool = True,
|
||||
):
|
||||
self.client = client
|
||||
|
|
@ -236,7 +236,7 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
|
|||
timeout: dict,
|
||||
proxy: Optional[str],
|
||||
sni_hostname: Optional[str],
|
||||
ssl_verify: Optional[Union[bool, ssl.SSLContext]] = None,
|
||||
ssl_verify: Optional[Union[bool, str, ssl.SSLContext]] = None,
|
||||
) -> ClientResponse:
|
||||
"""
|
||||
Helper function to make an aiohttp request with the given parameters.
|
||||
|
|
|
|||
|
|
@ -653,7 +653,8 @@ class AsyncHTTPHandler:
|
|||
except (httpx.RemoteProtocolError, httpx.ConnectError):
|
||||
# Retry the request with a new session if there is a connection error
|
||||
new_client = self.create_client(
|
||||
timeout=timeout, event_hooks=self.event_hooks,
|
||||
timeout=timeout,
|
||||
event_hooks=self.event_hooks,
|
||||
ssl_verify=self._ssl_verify,
|
||||
)
|
||||
try:
|
||||
|
|
@ -717,7 +718,8 @@ class AsyncHTTPHandler:
|
|||
except (httpx.RemoteProtocolError, httpx.ConnectError):
|
||||
# Retry the request with a new session if there is a connection error
|
||||
new_client = self.create_client(
|
||||
timeout=timeout, event_hooks=self.event_hooks,
|
||||
timeout=timeout,
|
||||
event_hooks=self.event_hooks,
|
||||
ssl_verify=self._ssl_verify,
|
||||
)
|
||||
try:
|
||||
|
|
@ -779,7 +781,8 @@ class AsyncHTTPHandler:
|
|||
except (httpx.RemoteProtocolError, httpx.ConnectError):
|
||||
# Retry the request with a new session if there is a connection error
|
||||
new_client = self.create_client(
|
||||
timeout=timeout, event_hooks=self.event_hooks,
|
||||
timeout=timeout,
|
||||
event_hooks=self.event_hooks,
|
||||
ssl_verify=self._ssl_verify,
|
||||
)
|
||||
try:
|
||||
|
|
@ -841,7 +844,8 @@ class AsyncHTTPHandler:
|
|||
except (httpx.RemoteProtocolError, httpx.ConnectError):
|
||||
# Retry the request with a new session if there is a connection error
|
||||
new_client = self.create_client(
|
||||
timeout=timeout, event_hooks=self.event_hooks,
|
||||
timeout=timeout,
|
||||
event_hooks=self.event_hooks,
|
||||
ssl_verify=self._ssl_verify,
|
||||
)
|
||||
try:
|
||||
|
|
@ -958,7 +962,7 @@ class AsyncHTTPHandler:
|
|||
|
||||
@staticmethod
|
||||
def _get_ssl_connector_kwargs(
|
||||
ssl_verify: Optional[bool] = None,
|
||||
ssl_verify: Optional[Union[bool, str]] = None,
|
||||
ssl_context: Optional[ssl.SSLContext] = None,
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
|
|
@ -967,6 +971,9 @@ class AsyncHTTPHandler:
|
|||
SSL Configuration Priority:
|
||||
1. If ssl_context is provided -> use the custom SSL context
|
||||
2. If ssl_verify is False -> disable SSL verification (ssl=False)
|
||||
3. If ssl_verify is a string path to an existing file -> build an
|
||||
SSLContext that trusts it (aiohttp's TCPConnector requires an
|
||||
SSLContext, bool, or ``Fingerprint`` — not a bare path string).
|
||||
|
||||
Returns:
|
||||
Dict with appropriate SSL configuration for TCPConnector
|
||||
|
|
@ -981,12 +988,25 @@ class AsyncHTTPHandler:
|
|||
elif ssl_verify is False:
|
||||
# Priority 2: Explicitly disable SSL verification
|
||||
connector_kwargs["ssl"] = False
|
||||
elif isinstance(ssl_verify, str):
|
||||
# Priority 3: Resolve a CA bundle path into an SSLContext that
|
||||
# trusts it. ``ssl.create_default_context(cafile=...)`` is the
|
||||
# documented path. ``aiohttp.TCPConnector`` accepts ``SSLContext``,
|
||||
# ``bool``, or ``Fingerprint`` — a bare path string is rejected
|
||||
# by ``aiohttp.client_reqrep.ClientRequest._ssl_is_absolute_url``.
|
||||
try:
|
||||
connector_kwargs["ssl"] = ssl.create_default_context(cafile=ssl_verify)
|
||||
except (FileNotFoundError, OSError):
|
||||
# If the path doesn't resolve, fall back to letting aiohttp
|
||||
# use its default trust store. We do NOT silently disable
|
||||
# verification — that would change the security posture.
|
||||
pass
|
||||
|
||||
return connector_kwargs
|
||||
|
||||
@staticmethod
|
||||
def _create_aiohttp_transport(
|
||||
ssl_verify: Optional[bool] = None,
|
||||
ssl_verify: Optional[Union[bool, str]] = None,
|
||||
ssl_context: Optional[ssl.SSLContext] = None,
|
||||
shared_session: Optional["ClientSession"] = None,
|
||||
) -> LiteLLMAiohttpTransport:
|
||||
|
|
@ -996,6 +1016,7 @@ class AsyncHTTPHandler:
|
|||
Note: aiohttp TCPConnector ssl parameter accepts:
|
||||
- SSLContext: custom SSL context
|
||||
- False: disable SSL verification
|
||||
- str: path to a CA bundle file (used as ``cafile`` by aiohttp)
|
||||
"""
|
||||
from litellm.llms.custom_httpx.aiohttp_transport import LiteLLMAiohttpTransport
|
||||
from litellm.secret_managers.main import str_to_bool
|
||||
|
|
@ -1013,13 +1034,19 @@ class AsyncHTTPHandler:
|
|||
|
||||
#########################################################
|
||||
# Determine SSL config to pass to transport for per-request override
|
||||
# This ensures ssl_verify works even with shared sessions
|
||||
# This ensures ssl_verify works even with shared sessions.
|
||||
# ``ssl_for_transport`` is the value stashed on the transport for
|
||||
# ``LiteLLMAiohttpTransport._make_aiohttp_request`` to forward as the
|
||||
# ``ssl=`` kwarg on each request. We accept SSLContext, bool, and
|
||||
# CA-bundle path strings here.
|
||||
#########################################################
|
||||
ssl_for_transport: Optional[Union[bool, ssl.SSLContext]] = None
|
||||
ssl_for_transport: Optional[Union[bool, str, ssl.SSLContext]] = None
|
||||
if ssl_context is not None:
|
||||
ssl_for_transport = ssl_context
|
||||
elif ssl_verify is False:
|
||||
ssl_for_transport = False
|
||||
elif isinstance(ssl_verify, str):
|
||||
ssl_for_transport = ssl_verify
|
||||
|
||||
verbose_logger.debug("Creating AiohttpTransport...")
|
||||
|
||||
|
|
|
|||
|
|
@ -1,5 +1,4 @@
|
|||
import asyncio
|
||||
import io
|
||||
import os
|
||||
import pathlib
|
||||
import ssl
|
||||
|
|
@ -900,7 +899,9 @@ async def test_async_handler_retry_forwards_ssl_verify():
|
|||
|
||||
captured_ssl_verify = {}
|
||||
|
||||
def fake_create_client(*, timeout, event_hooks, ssl_verify=None, shared_session=None):
|
||||
def fake_create_client(
|
||||
*, timeout, event_hooks, ssl_verify=None, shared_session=None
|
||||
):
|
||||
captured_ssl_verify["value"] = ssl_verify
|
||||
# Return a real httpx client with a MockTransport so the post() body
|
||||
# can still execute the retry path without touching the network.
|
||||
|
|
@ -933,3 +934,148 @@ async def test_async_handler_retry_forwards_ssl_verify():
|
|||
finally:
|
||||
await handler.client.aclose()
|
||||
await handler.close()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("verb", ["put", "patch", "delete"])
|
||||
async def test_async_handler_retry_forwards_ssl_verify_for_put_patch_delete(verb):
|
||||
"""
|
||||
Companion to ``test_async_handler_retry_forwards_ssl_verify`` for the
|
||||
other verb methods. ``put``, ``patch`` and ``delete`` all share the same
|
||||
``ConnectError``/``RemoteProtocolError`` retry block and must also
|
||||
forward the stored ``ssl_verify`` value.
|
||||
"""
|
||||
handler = AsyncHTTPHandler(ssl_verify=False)
|
||||
assert handler._ssl_verify is False
|
||||
|
||||
captured_ssl_verify = {}
|
||||
|
||||
def fake_create_client(
|
||||
*, timeout, event_hooks, ssl_verify=None, shared_session=None
|
||||
):
|
||||
captured_ssl_verify["value"] = ssl_verify
|
||||
return httpx.AsyncClient(
|
||||
transport=httpx.MockTransport(
|
||||
lambda req: httpx.Response(200, request=req, json={"ok": True})
|
||||
)
|
||||
)
|
||||
|
||||
handler.create_client = fake_create_client # type: ignore[assignment]
|
||||
|
||||
call_count = {"n": 0}
|
||||
|
||||
def maybe_fail(req):
|
||||
call_count["n"] += 1
|
||||
if call_count["n"] == 1:
|
||||
raise httpx.ConnectError("boom", request=req)
|
||||
return httpx.Response(200, request=req, json={"ok": True})
|
||||
|
||||
handler.client = httpx.AsyncClient(transport=httpx.MockTransport(maybe_fail))
|
||||
try:
|
||||
method = getattr(handler, verb)
|
||||
await method(
|
||||
"http://example.invalid/resource",
|
||||
json={"x": 1},
|
||||
)
|
||||
assert (
|
||||
captured_ssl_verify["value"] is False
|
||||
), f"{verb} retry path did not forward ssl_verify=False to create_client"
|
||||
finally:
|
||||
await handler.client.aclose()
|
||||
await handler.close()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_aiohttp_transport_accepts_string_ssl_verify():
|
||||
"""
|
||||
Regression for #30778 (followup): ``_create_aiohttp_transport`` must
|
||||
accept a CA bundle path string for ``ssl_verify`` and forward it through
|
||||
to ``TCPConnector(ssl=...)`` instead of silently dropping it.
|
||||
|
||||
Before the fix, callers passing ``ssl_verify="/etc/ssl/certs/ca.pem"``
|
||||
had their path ignored — the connector was built without ``ssl=`` set,
|
||||
so aiohttp fell back to its default SSL context and the caller's CA
|
||||
bundle was never trusted. The connector's ``ssl`` parameter must be
|
||||
an ``SSLContext``, ``bool``, or ``Fingerprint`` — a bare path string
|
||||
is rejected by aiohttp, so the helper builds an SSLContext via
|
||||
``ssl.create_default_context(cafile=...)``.
|
||||
"""
|
||||
original_disable = litellm.disable_aiohttp_transport
|
||||
litellm.disable_aiohttp_transport = False
|
||||
|
||||
ca_path = str(pathlib.Path(certifi.where()).resolve())
|
||||
|
||||
try:
|
||||
transport = AsyncHTTPHandler._create_aiohttp_transport(ssl_verify=ca_path)
|
||||
try:
|
||||
client_session = transport._get_valid_client_session()
|
||||
assert isinstance(client_session, ClientSession)
|
||||
connector = client_session.connector
|
||||
assert isinstance(connector, TCPConnector)
|
||||
# The connector must carry an SSLContext that trusts the CA bundle.
|
||||
# It cannot be the raw path string (aiohttp's TCPConnector only
|
||||
# accepts SSLContext, bool, or Fingerprint).
|
||||
assert isinstance(connector._ssl, ssl.SSLContext), (
|
||||
f"expected connector._ssl to be an SSLContext built from the "
|
||||
f"CA bundle path, got {type(connector._ssl).__name__}"
|
||||
)
|
||||
# The transport must also stash the original CA bundle path for
|
||||
# per-request override so externally-provided sessions can honor
|
||||
# the caller's CA bundle on a per-request basis.
|
||||
assert transport._ssl_verify == ca_path
|
||||
finally:
|
||||
await transport.aclose()
|
||||
finally:
|
||||
litellm.disable_aiohttp_transport = original_disable
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_base_llm_aiohttp_handler_forwards_ssl_to_external_session():
|
||||
"""
|
||||
Regression for #30778 (followup): ``BaseLLMAIOHTTPHandler._make_common_async_call``
|
||||
must forward ``ssl_verify`` as a per-request ``ssl=`` kwarg to
|
||||
``async_client_session.post`` when a caller supplies an external
|
||||
``ClientSession`` (e.g. via the ``client`` argument on ``async_completion``).
|
||||
|
||||
Before the fix, the connector-level SSL setting only applied to sessions
|
||||
built lazily by ``_create_client_session_with_transport``. Externally
|
||||
provided sessions skipped the connector entirely, so callers passing
|
||||
``ssl_verify=False`` still got the default SSL verification on the wire.
|
||||
"""
|
||||
from litellm.llms.custom_httpx.aiohttp_handler import BaseLLMAIOHTTPHandler
|
||||
|
||||
handler = BaseLLMAIOHTTPHandler(ssl_verify=False)
|
||||
assert handler.ssl_verify is False
|
||||
|
||||
captured_kwargs = {}
|
||||
|
||||
class _CapturingSession:
|
||||
async def post(self, **kwargs):
|
||||
captured_kwargs.update(kwargs)
|
||||
response = MagicMock()
|
||||
response.ok = True
|
||||
return response
|
||||
|
||||
fake_session = _CapturingSession()
|
||||
|
||||
# Stub provider_config — we only need ``max_retry_on_unprocessable_entity_error``
|
||||
# and ``get_error_class`` for ``_make_common_async_call`` to operate.
|
||||
provider_config = MagicMock()
|
||||
provider_config.max_retry_on_unprocessable_entity_error = 1
|
||||
provider_config.get_error_class.side_effect = RuntimeError
|
||||
|
||||
await handler._make_common_async_call(
|
||||
async_client_session=fake_session, # type: ignore[arg-type]
|
||||
provider_config=provider_config,
|
||||
api_base="http://example.invalid/post",
|
||||
headers={},
|
||||
data={"x": 1},
|
||||
timeout=1.0,
|
||||
litellm_params={},
|
||||
)
|
||||
|
||||
assert "ssl" in captured_kwargs, (
|
||||
"_make_common_async_call must forward ssl_verify=False as ssl= kwarg "
|
||||
"when a caller-supplied ClientSession is used"
|
||||
)
|
||||
assert captured_kwargs["ssl"] is False
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue