Merge pull request #35862 from BerriAI/litellm_self_heal_evicted_httpx_clients

fix(http_handler): self-heal handler clients closed after cache eviction
This commit is contained in:
Mateo Wang 2026-08-04 23:55:53 -07:00 • committed by GitHub
commit 732bba00df
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 208 additions and 23 deletions

View file

@ -30,8 +30,8 @@ async def close_litellm_async_clients():
pass
# Handle AsyncHTTPHandler instances (used by Gemini and other providers)
elif hasattr(handler, "client"):
client = handler.client
elif hasattr(handler, "_client") or hasattr(handler, "client"):
client = handler._client if hasattr(handler, "_client") else handler.client
# Check if the httpx client has an aiohttp transport
if hasattr(client, "_transport") and hasattr(client._transport, "aclose"):
try:

View file

@ -5,6 +5,7 @@ import os
import socket
import ssl
import sys
import threading
import time
from collections.abc import Callable, Mapping
from typing import TYPE_CHECKING, Any, Final, Optional
@ -510,7 +511,10 @@ class AsyncHTTPHandler:
):
self.timeout = timeout
self.event_hooks = event_hooks
self.client = self.create_client(
self.ssl_verify = ssl_verify
self.shared_session = shared_session
self._owns_client = True
self._client = self.create_client(
timeout=timeout,
event_hooks=event_hooks,
ssl_verify=ssl_verify,
@ -518,6 +522,22 @@ class AsyncHTTPHandler:
)
self.client_alias = client_alias
@property
def client(self) -> httpx.AsyncClient:
if self._owns_client and self._client.is_closed:
self._client = self.create_client(
timeout=self.timeout,
event_hooks=self.event_hooks,
ssl_verify=self.ssl_verify,
shared_session=self.shared_session,
)
return self._client
@client.setter
def client(self, client: httpx.AsyncClient) -> None:
self._client = client
self._owns_client = False
def create_client(
self,
timeout: float | httpx.Timeout | None,
@ -557,14 +577,14 @@ class AsyncHTTPHandler:
async def close(self):
# Close the client when you're done with it
await self.client.aclose()
await self._client.aclose()
async def __aenter__(self):
return self.client
async def __aexit__(self):
# close the client when exiting
await self.client.aclose()
await self._client.aclose()
async def get(
self,
@ -1069,37 +1089,50 @@ class HTTPHandler:
disable_default_headers: bool
| None = False, # arize phoenix returns different API responses when user agent header in request
):
if timeout is None:
timeout = _DEFAULT_TIMEOUT
self.timeout = timeout
self.ssl_verify = ssl_verify
self.disable_default_headers = disable_default_headers
self._owns_client = client is None
self._heal_lock = threading.Lock()
self._client = self.create_client() if client is None else client
def create_client(self) -> httpx.Client:
# Get unified SSL configuration
ssl_config: Final = get_ssl_configuration(ssl_verify)
ssl_config: Final = get_ssl_configuration(self.ssl_verify)
# An SSL certificate used by the requested host to authenticate the client.
# /path/to/client.pem
cert: Final = os.getenv("SSL_CERTIFICATE", litellm.ssl_certificate)
# Get default headers (User-Agent, overridable via LITELLM_USER_AGENT)
default_headers: Final = get_default_headers() if not disable_default_headers else None
default_headers: Final = get_default_headers() if not self.disable_default_headers else None
if client is None:
transport: Final = self._create_sync_transport()
# Create a client with a connection pool
return httpx.Client(
transport=self._create_sync_transport(),
timeout=self.timeout if self.timeout is not None else _DEFAULT_TIMEOUT,
verify=ssl_config,
cert=cert,
headers=default_headers,
follow_redirects=True,
)
# Create a client with a connection pool
self.client = httpx.Client(
transport=transport,
timeout=timeout,
verify=ssl_config,
cert=cert,
headers=default_headers,
follow_redirects=True,
)
else:
self.client = client
@property
def client(self) -> httpx.Client:
if self._owns_client and self._client.is_closed:
with self._heal_lock:
if self._owns_client and self._client.is_closed:
self._client = self.create_client()
return self._client
@client.setter
def client(self, client: httpx.Client) -> None:
self._client = client
self._owns_client = False
def close(self):
# Close the client when you're done with it
self.client.close()
self._client.close()
def get(
self,

View file

@ -0,0 +1,21 @@
import pytest
import litellm
from litellm.llms.custom_httpx.async_client_cleanup import close_litellm_async_clients
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
@pytest.mark.asyncio
async def test_second_cleanup_pass_does_not_resurrect_owned_client():
handler = AsyncHTTPHandler()
original_client = handler._client
cache_key = "test-cleanup-no-resurrect"
litellm.in_memory_llm_clients_cache.cache_dict[cache_key] = handler
try:
await close_litellm_async_clients()
assert original_client.is_closed
await close_litellm_async_clients()
finally:
litellm.in_memory_llm_clients_cache.cache_dict.pop(cache_key, None)
assert handler._client is original_client

View file

@ -793,3 +793,134 @@ class TestDefaultCachedClientTimeoutHonorsRequestTimeout:
litellm.in_memory_llm_clients_cache = LLMClientCache()
client = get_async_httpx_client(llm_provider=LlmProviders.BEDROCK)
assert client.timeout.read == 300.0
async def _read_http_request(reader: asyncio.StreamReader) -> None:
raw = b""
while b"\r\n\r\n" not in raw:
chunk = await reader.read(1024)
if not chunk:
return
raw += chunk
head, _, body = raw.partition(b"\r\n\r\n")
content_length = next(
(int(line.split(b":", 1)[1]) for line in head.split(b"\r\n") if line.lower().startswith(b"content-length")),
0,
)
while len(body) < content_length:
body += await reader.read(content_length - len(body))
@pytest.mark.asyncio
async def test_init_held_async_handler_survives_external_client_close():
handler = AsyncHTTPHandler(timeout=42.5)
held_client = handler.client
await held_client.aclose()
assert held_client.is_closed
async def respond(reader: asyncio.StreamReader, writer: asyncio.StreamWriter) -> None:
await _read_http_request(reader)
writer.write(b"HTTP/1.1 200 OK\r\nContent-Length: 2\r\nConnection: close\r\n\r\nok")
await writer.drain()
writer.close()
server = await asyncio.start_server(respond, "127.0.0.1", 0)
port = server.sockets[0].getsockname()[1]
try:
response = await handler.post(f"http://127.0.0.1:{port}/v1/compress", json={"messages": []})
finally:
server.close()
await server.wait_closed()
assert response.status_code == 200
assert handler.client is not held_client
assert handler.client.timeout == httpx.Timeout(42.5)
await handler.close()
def test_init_held_sync_handler_recreates_closed_client():
from http.server import BaseHTTPRequestHandler, HTTPServer
class OkRequestHandler(BaseHTTPRequestHandler):
def do_GET(self):
self.send_response(200)
self.send_header("Content-Length", "2")
self.end_headers()
self.wfile.write(b"ok")
def log_message(self, format, *args):
pass
handler = HTTPHandler(timeout=7)
held_client = handler.client
held_client.close()
server = HTTPServer(("127.0.0.1", 0), OkRequestHandler)
thread = threading.Thread(target=server.serve_forever, daemon=True)
thread.start()
try:
response = handler.get(f"http://127.0.0.1:{server.server_port}/")
finally:
server.shutdown()
server.server_close()
thread.join(timeout=5)
assert response.status_code == 200
assert handler.client is not held_client
assert handler.client.timeout == httpx.Timeout(7)
handler.close()
def test_caller_supplied_sync_client_is_not_replaced_when_closed():
supplied = httpx.Client()
handler = HTTPHandler(client=supplied)
supplied.close()
assert handler.client is supplied
@pytest.mark.asyncio
async def test_assigned_async_client_is_not_replaced():
handler = AsyncHTTPHandler()
await handler.client.aclose()
replacement = MagicMock()
handler.client = replacement
assert handler.client is replacement
def test_concurrent_sync_heal_creates_exactly_one_replacement():
class GatedHealHandler(HTTPHandler):
def __init__(self):
self.heal_started = threading.Event()
self.release_heal = threading.Event()
self.heal_calls = 0
super().__init__(timeout=7)
def create_client(self) -> httpx.Client:
if hasattr(self, "_client"):
self.heal_calls += 1
self.heal_started.set()
assert self.release_heal.wait(timeout=5)
return super().create_client()
handler = GatedHealHandler()
handler.client.close()
seen = []
def grab_client():
seen.append(handler.client)
first = threading.Thread(target=grab_client)
second = threading.Thread(target=grab_client)
first.start()
assert handler.heal_started.wait(timeout=5)
second.start()
second.join(timeout=0.3)
handler.release_heal.set()
first.join(timeout=5)
second.join(timeout=5)
assert handler.heal_calls == 1
assert seen[0] is seen[1]
assert not seen[0].is_closed
handler.close()