mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
commit
732bba00df
4 changed files with 208 additions and 23 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue