mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-09 22:31:41 +00:00
fix(http_handler): honor HTTP(S)_PROXY / NO_PROXY when force_ipv4 uses the httpx transport (#39443)
* fix(http_handler): honor HTTP(S)_PROXY / NO_PROXY when force_ipv4 uses the httpx transport Passing an explicit transport to httpx.AsyncClient / httpx.Client disables its automatic environment proxy mounts, so force_ipv4 on the httpx path sent every LLM request direct and silently bypassed HTTPS_PROXY. Mount the same env-derived proxy transports next to the IPv4-pinned direct transport in AsyncHTTPHandler, HTTPHandler and the OpenAI async client factory. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(http_handler): carry the client's TLS verify and cert settings onto env proxy mounts Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: yassin <yassin@berri.ai> Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
e1a7af7e58
commit
5ad330f620
3 changed files with 285 additions and 9 deletions
|
|
@ -9,13 +9,15 @@ import threading
|
|||
import time
|
||||
from collections.abc import AsyncIterable, Callable, Iterable, Mapping
|
||||
from http.cookiejar import CookieJar, DefaultCookiePolicy
|
||||
from typing import TYPE_CHECKING, Any, ClassVar, Final, NoReturn, Optional, TypeAlias, TypedDict
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, ClassVar, Final, NoReturn, Optional, TypeAlias, TypedDict, TypeVar
|
||||
|
||||
import certifi
|
||||
import httpx
|
||||
from aiohttp import ClientSession, DummyCookieJar, TCPConnector
|
||||
from httpx import USE_CLIENT_DEFAULT, AsyncHTTPTransport, HTTPTransport
|
||||
from httpx._types import RequestFiles
|
||||
from httpx._types import CertTypes, RequestFiles
|
||||
from httpx._utils import get_environment_proxies
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
|
|
@ -66,6 +68,22 @@ _AddrInfo: TypeAlias = tuple[int | socket.AddressFamily, int | socket.SocketKind
|
|||
|
||||
_RequestContent: TypeAlias = str | bytes | Iterable[bytes] | AsyncIterable[bytes]
|
||||
|
||||
_IPV4_LOCAL_ADDRESS: Final = "0.0.0.0"
|
||||
|
||||
_HttpxTransportT = TypeVar("_HttpxTransportT", HTTPTransport, AsyncHTTPTransport)
|
||||
|
||||
|
||||
def _environment_proxy_mounts(
|
||||
build_proxy_transport: Callable[[str], _HttpxTransportT],
|
||||
) -> Mapping[str, _HttpxTransportT | None]:
|
||||
"""httpx skips its own HTTP(S)_PROXY / NO_PROXY mounts whenever an explicit `transport=` is passed."""
|
||||
return MappingProxyType(
|
||||
{
|
||||
pattern: None if proxy_url is None else build_proxy_transport(proxy_url)
|
||||
for pattern, proxy_url in get_environment_proxies().items()
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
class _TCPConnectorKwargs(TypedDict, total=False):
|
||||
local_addr: tuple[str, int] | None
|
||||
|
|
@ -607,6 +625,7 @@ class AsyncHTTPHandler:
|
|||
|
||||
return httpx.AsyncClient(
|
||||
transport=transport,
|
||||
mounts=AsyncHTTPHandler._create_httpx_proxy_mounts(transport, verify=ssl_config, cert=cert),
|
||||
event_hooks=event_hooks,
|
||||
timeout=timeout,
|
||||
verify=ssl_config,
|
||||
|
|
@ -1191,10 +1210,22 @@ class AsyncHTTPHandler:
|
|||
- [Default] If force_ipv4 is False, it will return None
|
||||
"""
|
||||
if litellm.force_ipv4:
|
||||
return AsyncHTTPTransport(local_address="0.0.0.0")
|
||||
return AsyncHTTPTransport(local_address=_IPV4_LOCAL_ADDRESS)
|
||||
else:
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _create_httpx_proxy_mounts(
|
||||
transport: LiteLLMAiohttpTransport | AsyncHTTPTransport | None,
|
||||
verify: VerifyTypes,
|
||||
cert: CertTypes | None,
|
||||
) -> Mapping[str, AsyncHTTPTransport | None] | None:
|
||||
if not isinstance(transport, AsyncHTTPTransport):
|
||||
return None
|
||||
return _environment_proxy_mounts(
|
||||
lambda proxy_url: AsyncHTTPTransport(proxy=proxy_url, verify=verify, cert=cert)
|
||||
)
|
||||
|
||||
|
||||
class HTTPHandler:
|
||||
def __init__(
|
||||
|
|
@ -1227,6 +1258,7 @@ class HTTPHandler:
|
|||
# Create a client with a connection pool
|
||||
return httpx.Client(
|
||||
transport=self._create_sync_transport(),
|
||||
mounts=self._create_sync_proxy_mounts(verify=ssl_config, cert=cert),
|
||||
timeout=self.timeout if self.timeout is not None else _DEFAULT_TIMEOUT,
|
||||
verify=ssl_config,
|
||||
cert=cert,
|
||||
|
|
@ -1507,10 +1539,19 @@ class HTTPHandler:
|
|||
Some users have seen httpx ConnectionError when using ipv6 - forcing ipv4 resolves the issue for them
|
||||
"""
|
||||
if litellm.force_ipv4:
|
||||
return HTTPTransport(local_address="0.0.0.0")
|
||||
return HTTPTransport(local_address=_IPV4_LOCAL_ADDRESS)
|
||||
else:
|
||||
return getattr(litellm, "sync_transport", None)
|
||||
|
||||
@staticmethod
|
||||
def _create_sync_proxy_mounts(
|
||||
verify: VerifyTypes,
|
||||
cert: CertTypes | None,
|
||||
) -> Mapping[str, HTTPTransport | None] | None:
|
||||
if not litellm.force_ipv4:
|
||||
return None
|
||||
return _environment_proxy_mounts(lambda proxy_url: HTTPTransport(proxy=proxy_url, verify=verify, cert=cert))
|
||||
|
||||
|
||||
def get_async_httpx_client(
|
||||
llm_provider: LlmProviders | httpxSpecialProvider,
|
||||
|
|
|
|||
|
|
@ -305,14 +305,16 @@ class BaseOpenAILLM:
|
|||
|
||||
# Get unified SSL configuration
|
||||
ssl_config: Final = get_ssl_configuration()
|
||||
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,
|
||||
shared_session=shared_session,
|
||||
)
|
||||
|
||||
return httpx.AsyncClient(
|
||||
verify=ssl_config,
|
||||
transport=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,
|
||||
shared_session=shared_session,
|
||||
),
|
||||
transport=transport,
|
||||
mounts=AsyncHTTPHandler._create_httpx_proxy_mounts(transport, verify=ssl_config, cert=None),
|
||||
follow_redirects=True,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -1314,3 +1314,236 @@ async def test_finalizer_on_live_loop_disposes_foreign_loop_session_without_sche
|
|||
|
||||
assert AsyncHTTPHandler._finalizer_close_tasks == baseline_tasks
|
||||
assert session.closed
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def forward_proxy_server():
|
||||
"""Plain HTTP forward proxy that records the absolute URIs it is asked to fetch."""
|
||||
from http.server import BaseHTTPRequestHandler, HTTPServer
|
||||
from socketserver import ThreadingMixIn
|
||||
|
||||
seen_uris: list[str] = []
|
||||
|
||||
class RecordingProxyHandler(BaseHTTPRequestHandler):
|
||||
protocol_version = "HTTP/1.1"
|
||||
|
||||
def do_GET(self):
|
||||
seen_uris.append(self.path)
|
||||
self.send_response(200)
|
||||
self.send_header("Content-Length", "9")
|
||||
self.end_headers()
|
||||
self.wfile.write(b"via-proxy")
|
||||
|
||||
def log_message(self, format, *args):
|
||||
pass
|
||||
|
||||
class ThreadedServer(ThreadingMixIn, HTTPServer):
|
||||
daemon_threads = True
|
||||
|
||||
server = ThreadedServer(("127.0.0.1", 0), RecordingProxyHandler)
|
||||
thread = threading.Thread(target=server.serve_forever, daemon=True)
|
||||
thread.start()
|
||||
try:
|
||||
yield f"http://127.0.0.1:{server.server_port}", seen_uris
|
||||
finally:
|
||||
server.shutdown()
|
||||
server.server_close()
|
||||
thread.join(timeout=5)
|
||||
|
||||
|
||||
# `.invalid` never resolves (RFC 6761), so the only way this request can succeed is through the proxy
|
||||
_PROXY_ONLY_UPSTREAM_URL = "http://upstream.invalid/v1/models"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("disable_aiohttp_transport", [True, False])
|
||||
@pytest.mark.parametrize("force_ipv4", [True, False])
|
||||
async def test_async_handler_honours_proxy_env_for_every_transport(
|
||||
forward_proxy_server, monkeypatch: pytest.MonkeyPatch, disable_aiohttp_transport: bool, force_ipv4: bool
|
||||
):
|
||||
proxy_url, seen_uris = forward_proxy_server
|
||||
monkeypatch.setenv("HTTP_PROXY", proxy_url)
|
||||
monkeypatch.delenv("NO_PROXY", raising=False)
|
||||
monkeypatch.delenv("no_proxy", raising=False)
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", disable_aiohttp_transport)
|
||||
monkeypatch.setattr(litellm, "force_ipv4", force_ipv4)
|
||||
|
||||
handler = AsyncHTTPHandler()
|
||||
try:
|
||||
response = await handler.get(_PROXY_ONLY_UPSTREAM_URL)
|
||||
finally:
|
||||
await handler.close()
|
||||
|
||||
assert response.text == "via-proxy"
|
||||
assert seen_uris == [_PROXY_ONLY_UPSTREAM_URL]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("force_ipv4", [True, False])
|
||||
def test_sync_handler_honours_proxy_env(forward_proxy_server, monkeypatch: pytest.MonkeyPatch, force_ipv4: bool):
|
||||
proxy_url, seen_uris = forward_proxy_server
|
||||
monkeypatch.setenv("HTTP_PROXY", proxy_url)
|
||||
monkeypatch.delenv("NO_PROXY", raising=False)
|
||||
monkeypatch.delenv("no_proxy", raising=False)
|
||||
monkeypatch.setattr(litellm, "force_ipv4", force_ipv4)
|
||||
|
||||
handler = HTTPHandler()
|
||||
try:
|
||||
response = handler.get(_PROXY_ONLY_UPSTREAM_URL)
|
||||
finally:
|
||||
handler.close()
|
||||
|
||||
assert response.text == "via-proxy"
|
||||
assert seen_uris == [_PROXY_ONLY_UPSTREAM_URL]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_force_ipv4_httpx_transport_honours_no_proxy(keepalive_server, monkeypatch: pytest.MonkeyPatch):
|
||||
"""NO_PROXY hosts must still go direct when the proxy mounts are supplied by litellm instead of httpx."""
|
||||
monkeypatch.setenv("HTTP_PROXY", "http://proxy.invalid:3128")
|
||||
monkeypatch.setenv("NO_PROXY", "127.0.0.1")
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
monkeypatch.setattr(litellm, "force_ipv4", True)
|
||||
|
||||
handler = AsyncHTTPHandler()
|
||||
try:
|
||||
response = await handler.get(keepalive_server)
|
||||
finally:
|
||||
await handler.close()
|
||||
|
||||
assert response.text == "ok"
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def private_ca_tls_upstream(tmp_path: pathlib.Path):
|
||||
"""HTTPS server behind a CONNECT proxy, both on localhost; the server's cert is signed by a test-only CA."""
|
||||
import datetime
|
||||
import select
|
||||
import socket
|
||||
from http.server import BaseHTTPRequestHandler, HTTPServer
|
||||
from socketserver import ThreadingMixIn
|
||||
|
||||
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, "upstream.invalid")])
|
||||
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())
|
||||
)
|
||||
ca_pem = tmp_path / "ca.pem"
|
||||
ca_pem.write_bytes(cert.public_bytes(serialization.Encoding.PEM))
|
||||
key_pem = tmp_path / "key.pem"
|
||||
key_pem.write_bytes(
|
||||
key.private_bytes(
|
||||
serialization.Encoding.PEM, serialization.PrivateFormat.PKCS8, serialization.NoEncryption()
|
||||
)
|
||||
)
|
||||
|
||||
class OkTlsHandler(BaseHTTPRequestHandler):
|
||||
protocol_version = "HTTP/1.1"
|
||||
|
||||
def do_GET(self):
|
||||
self.send_response(200)
|
||||
self.send_header("Content-Length", "6")
|
||||
self.end_headers()
|
||||
self.wfile.write(b"ok-tls")
|
||||
|
||||
def log_message(self, format, *args):
|
||||
pass
|
||||
|
||||
class ThreadedServer(ThreadingMixIn, HTTPServer):
|
||||
daemon_threads = True
|
||||
|
||||
tls_server = ThreadedServer(("127.0.0.1", 0), OkTlsHandler)
|
||||
server_ctx = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER)
|
||||
server_ctx.load_cert_chain(str(ca_pem), str(key_pem))
|
||||
tls_server.socket = server_ctx.wrap_socket(tls_server.socket, server_side=True)
|
||||
tls_port = tls_server.server_port
|
||||
|
||||
class ConnectProxyHandler(BaseHTTPRequestHandler):
|
||||
protocol_version = "HTTP/1.1"
|
||||
|
||||
def do_CONNECT(self):
|
||||
upstream = socket.create_connection(("127.0.0.1", tls_port))
|
||||
self.send_response(200, "Connection established")
|
||||
self.end_headers()
|
||||
sockets = [self.connection, upstream]
|
||||
while True:
|
||||
readable, _, _ = select.select(sockets, [], [], 5)
|
||||
if not readable:
|
||||
break
|
||||
for src in readable:
|
||||
data = src.recv(65536)
|
||||
if not data:
|
||||
upstream.close()
|
||||
return
|
||||
(upstream if src is self.connection else self.connection).sendall(data)
|
||||
|
||||
def log_message(self, format, *args):
|
||||
pass
|
||||
|
||||
proxy_server = ThreadedServer(("127.0.0.1", 0), ConnectProxyHandler)
|
||||
threads = [
|
||||
threading.Thread(target=tls_server.serve_forever, daemon=True),
|
||||
threading.Thread(target=proxy_server.serve_forever, daemon=True),
|
||||
]
|
||||
for thread in threads:
|
||||
thread.start()
|
||||
try:
|
||||
yield f"http://127.0.0.1:{proxy_server.server_port}", str(ca_pem)
|
||||
finally:
|
||||
for server in (proxy_server, tls_server):
|
||||
server.shutdown()
|
||||
server.server_close()
|
||||
for thread in threads:
|
||||
thread.join(timeout=5)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_force_ipv4_https_proxy_mount_uses_handler_ca_bundle(
|
||||
private_ca_tls_upstream, monkeypatch: pytest.MonkeyPatch
|
||||
):
|
||||
proxy_url, ca_pem = private_ca_tls_upstream
|
||||
monkeypatch.setenv("HTTPS_PROXY", proxy_url)
|
||||
monkeypatch.delenv("NO_PROXY", raising=False)
|
||||
monkeypatch.delenv("no_proxy", raising=False)
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
monkeypatch.setattr(litellm, "force_ipv4", True)
|
||||
|
||||
handler = AsyncHTTPHandler(ssl_verify=ca_pem)
|
||||
try:
|
||||
response = await handler.get("https://upstream.invalid/v1/models")
|
||||
finally:
|
||||
await handler.close()
|
||||
|
||||
assert response.text == "ok-tls"
|
||||
|
||||
|
||||
def test_sync_force_ipv4_https_proxy_mount_uses_handler_ca_bundle(
|
||||
private_ca_tls_upstream, monkeypatch: pytest.MonkeyPatch
|
||||
):
|
||||
proxy_url, ca_pem = private_ca_tls_upstream
|
||||
monkeypatch.setenv("HTTPS_PROXY", proxy_url)
|
||||
monkeypatch.delenv("NO_PROXY", raising=False)
|
||||
monkeypatch.delenv("no_proxy", raising=False)
|
||||
monkeypatch.setattr(litellm, "force_ipv4", True)
|
||||
|
||||
handler = HTTPHandler(ssl_verify=ca_pem)
|
||||
try:
|
||||
response = handler.get("https://upstream.invalid/v1/models")
|
||||
finally:
|
||||
handler.close()
|
||||
|
||||
assert response.text == "ok-tls"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue