mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-11 22:51:28 +00:00
fix(ssl): trust the OS CA store and SSL_CERT_DIR alongside certifi
This commit is contained in:
parent
168a0055a2
commit
5ad31913ab
3 changed files with 213 additions and 23 deletions
|
|
@ -582,7 +582,7 @@ class MCPClient:
|
|||
1. Check ssl_verify parameter (can be SSLContext, bool, or path to CA bundle)
|
||||
2. Check SSL_VERIFY environment variable
|
||||
3. Check SSL_CERT_FILE environment variable
|
||||
4. Fall back to certifi CA bundle
|
||||
4. Fall back to the OS trust store with certifi layered on top
|
||||
"""
|
||||
|
||||
def factory(
|
||||
|
|
|
|||
|
|
@ -235,9 +235,9 @@ def _prepare_request_data_and_content(
|
|||
|
||||
|
||||
# Cache for SSL contexts to avoid creating duplicate contexts with the same configuration
|
||||
# Key: tuple of (cafile, ssl_security_level, ssl_ecdh_curve)
|
||||
# Key: tuple of (cafile, ssl_cert_dir, ssl_security_level, ssl_ecdh_curve)
|
||||
# Value: ssl.SSLContext
|
||||
_ssl_context_cache: Final[dict[tuple[str | None, str | None, str | None], ssl.SSLContext]] = {}
|
||||
_ssl_context_cache: Final[dict[tuple[str | None, str | None, str | None, str | None], ssl.SSLContext]] = {}
|
||||
|
||||
|
||||
def _create_ssl_context(
|
||||
|
|
@ -248,8 +248,14 @@ def _create_ssl_context(
|
|||
"""
|
||||
Create an SSL context with the given configuration.
|
||||
This is separated from get_ssl_configuration to enable caching.
|
||||
|
||||
With no CA bundle configured, create_default_context loads the OS trust store, which is also the
|
||||
only place OpenSSL reads SSL_CERT_DIR from; certifi is layered on top so a host with a thin or
|
||||
absent system store still reaches public providers.
|
||||
"""
|
||||
custom_ssl_context: Final = ssl.create_default_context(cafile=cafile)
|
||||
if cafile is None:
|
||||
custom_ssl_context.load_verify_locations(cafile=certifi.where())
|
||||
|
||||
# Optimize SSL handshake performance
|
||||
# Set minimum TLS version to 1.2 for better performance
|
||||
|
|
@ -328,6 +334,18 @@ def get_ssl_verify(
|
|||
return ssl_verify if ssl_verify is not None else True
|
||||
|
||||
|
||||
def _resolve_cafile(ssl_verify: bool | str) -> str | None:
|
||||
"""
|
||||
Resolve the explicitly configured CA bundle path, or None when the operator configured none.
|
||||
"""
|
||||
if isinstance(ssl_verify, str) and os.path.exists(ssl_verify):
|
||||
return ssl_verify
|
||||
ssl_cert_file: Final = os.getenv("SSL_CERT_FILE")
|
||||
if ssl_cert_file and os.path.exists(ssl_cert_file):
|
||||
return ssl_cert_file
|
||||
return None
|
||||
|
||||
|
||||
def get_ssl_configuration(
|
||||
ssl_verify: VerifyTypes | None = None,
|
||||
) -> bool | str | ssl.SSLContext:
|
||||
|
|
@ -339,7 +357,7 @@ def get_ssl_configuration(
|
|||
2. If ssl_verify is False -> disable SSL verification (ssl=False)
|
||||
3. If ssl_verify is a string -> use it as a path to CA bundle file
|
||||
4. If SSL_CERT_FILE environment variable is set and exists -> use it as CA bundle file
|
||||
5. Else will use default SSL context with certifi CA bundle
|
||||
5. Else use the OS trust store (where OpenSSL reads SSL_CERT_DIR) with certifi layered on top
|
||||
|
||||
If ssl_security_level is set, it will apply the security level to the SSL context.
|
||||
|
||||
|
|
@ -366,19 +384,12 @@ def get_ssl_configuration(
|
|||
ssl_security_level: Final = os.getenv("SSL_SECURITY_LEVEL", litellm.ssl_security_level)
|
||||
ssl_ecdh_curve: Final = os.getenv("SSL_ECDH_CURVE", litellm.ssl_ecdh_curve)
|
||||
|
||||
cafile = None
|
||||
if isinstance(ssl_verify, str) and os.path.exists(ssl_verify):
|
||||
cafile = ssl_verify
|
||||
if not cafile:
|
||||
ssl_cert_file: Final = os.getenv("SSL_CERT_FILE")
|
||||
if ssl_cert_file and os.path.exists(ssl_cert_file):
|
||||
cafile = ssl_cert_file
|
||||
else:
|
||||
cafile = certifi.where()
|
||||
cafile: Final = _resolve_cafile(ssl_verify)
|
||||
ssl_cert_dir: Final = os.getenv("SSL_CERT_DIR")
|
||||
|
||||
if ssl_verify is not False:
|
||||
# Create cache key from configuration parameters
|
||||
cache_key: Final = (cafile, ssl_security_level, ssl_ecdh_curve)
|
||||
cache_key: Final = (cafile, ssl_cert_dir, ssl_security_level, ssl_ecdh_curve)
|
||||
|
||||
# Check if we have a cached SSL context for this configuration
|
||||
if cache_key not in _ssl_context_cache:
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@ import gc
|
|||
import io
|
||||
import os
|
||||
import pathlib
|
||||
import socket
|
||||
import ssl
|
||||
import threading
|
||||
import weakref
|
||||
|
|
@ -279,8 +280,7 @@ async def test_ssl_context_with_shared_session(monkeypatch: pytest.MonkeyPatch):
|
|||
|
||||
|
||||
def test_get_ssl_configuration():
|
||||
"""Test that get_ssl_configuration() returns a proper SSL context with certifi CA bundle
|
||||
when no environment variables are set."""
|
||||
"""With nothing configured, the context must load the OS trust store and add certifi on top."""
|
||||
from litellm.llms.custom_httpx.http_handler import _ssl_context_cache
|
||||
|
||||
# Clear cache to ensure ssl.create_default_context is called
|
||||
|
|
@ -288,20 +288,15 @@ def test_get_ssl_configuration():
|
|||
|
||||
with patch.dict(os.environ, clear=True):
|
||||
with patch("ssl.create_default_context") as mock_create_context:
|
||||
# Mock the return value
|
||||
mock_ssl_context = MagicMock(spec=ssl.SSLContext)
|
||||
mock_ssl_context.set_ciphers = MagicMock()
|
||||
mock_ssl_context.minimum_version = ssl.TLSVersion.TLSv1_2
|
||||
mock_create_context.return_value = mock_ssl_context
|
||||
|
||||
# Call the static method
|
||||
result = get_ssl_configuration()
|
||||
|
||||
# Verify ssl.create_default_context was called with certifi's CA file
|
||||
expected_ca_file = certifi.where()
|
||||
mock_create_context.assert_called_once_with(cafile=expected_ca_file)
|
||||
|
||||
# Verify it returns the mocked SSL context
|
||||
mock_create_context.assert_called_once_with(cafile=None)
|
||||
mock_ssl_context.load_verify_locations.assert_called_once_with(cafile=certifi.where())
|
||||
assert result == mock_ssl_context
|
||||
|
||||
|
||||
|
|
@ -318,6 +313,190 @@ def test_get_ssl_configuration_integration():
|
|||
assert ssl_context.verify_mode is not None
|
||||
|
||||
|
||||
def _handshake(context: ssl.SSLContext, port: int) -> None:
|
||||
with socket.create_connection(("127.0.0.1", port), timeout=5) as sock:
|
||||
with context.wrap_socket(sock, server_hostname="upstream.invalid") as tls:
|
||||
tls.send(b"hi")
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def private_ca_hash_dir(tmp_path: pathlib.Path):
|
||||
"""A throwaway root CA reachable only through SSL_CERT_DIR, serving TLS on localhost."""
|
||||
import datetime
|
||||
import hashlib
|
||||
|
||||
from cryptography import x509
|
||||
from cryptography.hazmat.primitives import hashes, serialization
|
||||
from cryptography.hazmat.primitives.asymmetric import rsa
|
||||
from cryptography.x509.oid import NameOID
|
||||
|
||||
key = rsa.generate_private_key(public_exponent=65537, key_size=2048)
|
||||
name = x509.Name([x509.NameAttribute(NameOID.COMMON_NAME, "litellm-test-ca")])
|
||||
now = datetime.datetime.now(datetime.timezone.utc)
|
||||
cert = (
|
||||
x509.CertificateBuilder()
|
||||
.subject_name(name)
|
||||
.issuer_name(name)
|
||||
.public_key(key.public_key())
|
||||
.serial_number(x509.random_serial_number())
|
||||
.not_valid_before(now - datetime.timedelta(minutes=1))
|
||||
.not_valid_after(now + datetime.timedelta(hours=1))
|
||||
.add_extension(x509.SubjectAlternativeName([x509.DNSName("upstream.invalid")]), critical=False)
|
||||
.add_extension(x509.BasicConstraints(ca=True, path_length=None), critical=True)
|
||||
.sign(key, hashes.SHA256())
|
||||
)
|
||||
pem = cert.public_bytes(serialization.Encoding.PEM)
|
||||
ca_pem = tmp_path / "ca.pem"
|
||||
ca_pem.write_bytes(pem)
|
||||
key_pem = tmp_path / "key.pem"
|
||||
key_pem.write_bytes(
|
||||
key.private_bytes(
|
||||
serialization.Encoding.PEM, serialization.PrivateFormat.PKCS8, serialization.NoEncryption()
|
||||
)
|
||||
)
|
||||
|
||||
# OpenSSL finds a CA in a hash dir by <X509_NAME_hash>.0, the sha1 of the canonical name
|
||||
# encoding: the subject DER with the Name's outer SEQUENCE header stripped, RDN SET tag kept.
|
||||
subject_der = cert.subject.public_bytes()
|
||||
header = 2 if subject_der[1] < 0x80 else 2 + (subject_der[1] & 0x7F)
|
||||
subject_hash = int.from_bytes(hashlib.sha1(subject_der[header:]).digest()[:4], "little")
|
||||
cadir = tmp_path / "cadir"
|
||||
cadir.mkdir()
|
||||
(cadir / f"{subject_hash:08x}.0").write_bytes(pem)
|
||||
|
||||
server_ctx = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER)
|
||||
server_ctx.load_cert_chain(str(ca_pem), str(key_pem))
|
||||
listener = socket.socket()
|
||||
listener.bind(("127.0.0.1", 0))
|
||||
listener.listen(8)
|
||||
port = listener.getsockname()[1]
|
||||
|
||||
def handshake(connection):
|
||||
try:
|
||||
with server_ctx.wrap_socket(connection, server_side=True) as tls:
|
||||
tls.recv(16)
|
||||
except OSError:
|
||||
pass
|
||||
|
||||
def serve():
|
||||
while True:
|
||||
try:
|
||||
connection, _ = listener.accept()
|
||||
except OSError:
|
||||
return
|
||||
threading.Thread(target=handshake, args=(connection,), daemon=True).start()
|
||||
|
||||
thread = threading.Thread(target=serve, daemon=True)
|
||||
thread.start()
|
||||
try:
|
||||
yield cadir, ca_pem, port
|
||||
finally:
|
||||
listener.close()
|
||||
thread.join(timeout=5)
|
||||
|
||||
|
||||
def test_ssl_cert_dir_ca_is_trusted(private_ca_hash_dir, monkeypatch: pytest.MonkeyPatch):
|
||||
"""A CA installed only in the SSL_CERT_DIR hash dir must verify (GH #12451)."""
|
||||
from litellm.llms.custom_httpx.http_handler import _ssl_context_cache
|
||||
|
||||
cadir, _, port = private_ca_hash_dir
|
||||
_ssl_context_cache.clear()
|
||||
try:
|
||||
with patch.dict(os.environ, clear=True):
|
||||
monkeypatch.setenv("SSL_CERT_DIR", str(cadir))
|
||||
context = get_ssl_configuration()
|
||||
|
||||
assert isinstance(context, ssl.SSLContext)
|
||||
_handshake(context, port)
|
||||
finally:
|
||||
_ssl_context_cache.clear()
|
||||
|
||||
|
||||
def test_a_second_ssl_cert_dir_is_not_served_from_the_first_ones_cached_context(
|
||||
private_ca_hash_dir, tmp_path: pathlib.Path, monkeypatch: pytest.MonkeyPatch
|
||||
):
|
||||
"""Changing SSL_CERT_DIR must build a new context, not reuse the one keyed on cafile=None."""
|
||||
from litellm.llms.custom_httpx.http_handler import _ssl_context_cache
|
||||
|
||||
cadir, _, port = private_ca_hash_dir
|
||||
empty_dir = tmp_path / "empty-cadir"
|
||||
empty_dir.mkdir()
|
||||
_ssl_context_cache.clear()
|
||||
try:
|
||||
with patch.dict(os.environ, clear=True):
|
||||
monkeypatch.setenv("SSL_CERT_DIR", str(cadir))
|
||||
trusting = get_ssl_configuration()
|
||||
monkeypatch.setenv("SSL_CERT_DIR", str(empty_dir))
|
||||
untrusting = get_ssl_configuration()
|
||||
|
||||
assert isinstance(trusting, ssl.SSLContext)
|
||||
assert isinstance(untrusting, ssl.SSLContext)
|
||||
_handshake(trusting, port)
|
||||
with pytest.raises(ssl.SSLCertVerificationError):
|
||||
_handshake(untrusting, port)
|
||||
finally:
|
||||
_ssl_context_cache.clear()
|
||||
|
||||
|
||||
def test_default_context_stays_a_superset_of_the_certifi_bundle():
|
||||
"""Widening to the OS trust store must not drop a CA that certifi carries and the OS lacks."""
|
||||
from litellm.llms.custom_httpx.http_handler import _ssl_context_cache
|
||||
|
||||
_ssl_context_cache.clear()
|
||||
try:
|
||||
with patch.dict(os.environ, clear=True):
|
||||
context = get_ssl_configuration()
|
||||
|
||||
assert isinstance(context, ssl.SSLContext)
|
||||
certifi_certs = set(ssl.create_default_context(cafile=certifi.where()).get_ca_certs(binary_form=True))
|
||||
assert certifi_certs <= set(context.get_ca_certs(binary_form=True))
|
||||
finally:
|
||||
_ssl_context_cache.clear()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("configured_via", ["ssl_verify", "SSL_CERT_FILE"])
|
||||
def test_an_explicit_ca_bundle_still_replaces_the_default_trust(
|
||||
configured_via, private_ca_hash_dir, monkeypatch: pytest.MonkeyPatch
|
||||
):
|
||||
"""SSL_CERT_FILE and ssl_verify=<path> keep OpenSSL replace semantics, unchanged by this fix."""
|
||||
from litellm.llms.custom_httpx.http_handler import _ssl_context_cache
|
||||
|
||||
_, ca_pem, port = private_ca_hash_dir
|
||||
_ssl_context_cache.clear()
|
||||
try:
|
||||
with patch.dict(os.environ, clear=True):
|
||||
if configured_via == "SSL_CERT_FILE":
|
||||
monkeypatch.setenv("SSL_CERT_FILE", str(ca_pem))
|
||||
context = get_ssl_configuration()
|
||||
else:
|
||||
context = get_ssl_configuration(ssl_verify=str(ca_pem))
|
||||
|
||||
assert isinstance(context, ssl.SSLContext)
|
||||
_handshake(context, port)
|
||||
certifi_certs = set(ssl.create_default_context(cafile=certifi.where()).get_ca_certs(binary_form=True))
|
||||
assert not certifi_certs & set(context.get_ca_certs(binary_form=True))
|
||||
finally:
|
||||
_ssl_context_cache.clear()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("env_var", ["SSL_CERT_FILE", "SSL_CERT_DIR"])
|
||||
def test_an_unreadable_ssl_env_path_degrades_instead_of_failing(env_var, monkeypatch: pytest.MonkeyPatch):
|
||||
"""A typo in SSL_CERT_FILE / SSL_CERT_DIR must not brick outbound TLS."""
|
||||
from litellm.llms.custom_httpx.http_handler import _ssl_context_cache
|
||||
|
||||
_ssl_context_cache.clear()
|
||||
try:
|
||||
with patch.dict(os.environ, clear=True):
|
||||
monkeypatch.setenv(env_var, "/nonexistent/litellm-test-path")
|
||||
context = get_ssl_configuration()
|
||||
|
||||
assert isinstance(context, ssl.SSLContext)
|
||||
certifi_certs = set(ssl.create_default_context(cafile=certifi.where()).get_ca_certs(binary_form=True))
|
||||
assert certifi_certs <= set(context.get_ca_certs(binary_form=True))
|
||||
finally:
|
||||
_ssl_context_cache.clear()
|
||||
|
||||
|
||||
# Session Reuse Tests
|
||||
class MockClientSession:
|
||||
"""Mock ClientSession that is not callable"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue