diff --git a/litellm/llms/custom_httpx/async_client_cleanup.py b/litellm/llms/custom_httpx/async_client_cleanup.py index d5984e457eb..6ef61b6f58f 100644 --- a/litellm/llms/custom_httpx/async_client_cleanup.py +++ b/litellm/llms/custom_httpx/async_client_cleanup.py @@ -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: diff --git a/litellm/llms/custom_httpx/http_handler.py b/litellm/llms/custom_httpx/http_handler.py index 619341be62b..726392577e5 100644 --- a/litellm/llms/custom_httpx/http_handler.py +++ b/litellm/llms/custom_httpx/http_handler.py @@ -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, diff --git a/tests/test_litellm/llms/custom_httpx/test_async_client_cleanup.py b/tests/test_litellm/llms/custom_httpx/test_async_client_cleanup.py new file mode 100644 index 00000000000..e8fb0808019 --- /dev/null +++ b/tests/test_litellm/llms/custom_httpx/test_async_client_cleanup.py @@ -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 diff --git a/tests/test_litellm/llms/custom_httpx/test_http_handler.py b/tests/test_litellm/llms/custom_httpx/test_http_handler.py index 87d67e0e8b7..35db698ad76 100644 --- a/tests/test_litellm/llms/custom_httpx/test_http_handler.py +++ b/tests/test_litellm/llms/custom_httpx/test_http_handler.py @@ -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()