mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-19 00:01:29 +00:00
Revert "fix(http_handler): keep a handler alive while its responses are still streaming"
This reverts commit 250816eaeb.
That approach kept both `__del__` finalizers and made them unreachable while a
response was live, by stamping the owning handler onto `response.extensions`
from a client event hook. It only covers the shapes where a response is in
flight. A caller that keeps just `handler.client`, and a throwaway handler
wrapped around a client someone else owns, have no response to pin the handler
to, so the client is still closed underneath them.
The following commit deletes the finalizers instead, which covers those shapes
and removes code rather than adding it.
This commit is contained in:
parent
250816eaeb
commit
99ae409fd3
2 changed files with 1 additions and 175 deletions
|
|
@ -6,7 +6,6 @@ import socket
|
|||
import ssl
|
||||
import sys
|
||||
import time
|
||||
import weakref
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
|
|
@ -507,41 +506,6 @@ class MaskedHTTPStatusError(httpx.HTTPStatusError):
|
|||
self.status_code = original_error.response.status_code
|
||||
|
||||
|
||||
_CLIENT_OWNER_EXTENSION = "litellm_client_owner"
|
||||
|
||||
|
||||
def _pin_owner_on_async_client(owner: "AsyncHTTPHandler", client: httpx.AsyncClient) -> None:
|
||||
"""Record the handler that owns `client` on every response that client returns."""
|
||||
owner_ref = weakref.ref(owner)
|
||||
existing_hooks = client.event_hooks
|
||||
|
||||
async def _pin_owner(response: httpx.Response) -> None:
|
||||
handler = owner_ref()
|
||||
if handler is not None:
|
||||
response.extensions[_CLIENT_OWNER_EXTENSION] = handler
|
||||
|
||||
client.event_hooks = {
|
||||
"request": list(existing_hooks.get("request", [])),
|
||||
"response": [_pin_owner, *existing_hooks.get("response", [])],
|
||||
}
|
||||
|
||||
|
||||
def _pin_owner_on_sync_client(owner: "HTTPHandler", client: httpx.Client) -> None:
|
||||
"""Record the handler that owns `client` on every response that client returns."""
|
||||
owner_ref = weakref.ref(owner)
|
||||
existing_hooks = client.event_hooks
|
||||
|
||||
def _pin_owner(response: httpx.Response) -> None:
|
||||
handler = owner_ref()
|
||||
if handler is not None:
|
||||
response.extensions[_CLIENT_OWNER_EXTENSION] = handler
|
||||
|
||||
client.event_hooks = {
|
||||
"request": list(existing_hooks.get("request", [])),
|
||||
"response": [_pin_owner, *existing_hooks.get("response", [])],
|
||||
}
|
||||
|
||||
|
||||
class AsyncHTTPHandler:
|
||||
def __init__(
|
||||
self,
|
||||
|
|
@ -589,7 +553,7 @@ class AsyncHTTPHandler:
|
|||
# Get default headers (User-Agent, overridable via LITELLM_USER_AGENT)
|
||||
default_headers = get_default_headers()
|
||||
|
||||
client = httpx.AsyncClient(
|
||||
return httpx.AsyncClient(
|
||||
transport=transport,
|
||||
event_hooks=event_hooks,
|
||||
timeout=timeout,
|
||||
|
|
@ -598,8 +562,6 @@ class AsyncHTTPHandler:
|
|||
headers=default_headers,
|
||||
follow_redirects=True,
|
||||
)
|
||||
_pin_owner_on_async_client(self, client)
|
||||
return client
|
||||
|
||||
async def close(self):
|
||||
# Close the client when you're done with it
|
||||
|
|
@ -1137,7 +1099,6 @@ class HTTPHandler:
|
|||
headers=default_headers,
|
||||
follow_redirects=True,
|
||||
)
|
||||
_pin_owner_on_sync_client(self, self.client)
|
||||
else:
|
||||
self.client = client
|
||||
|
||||
|
|
|
|||
|
|
@ -1,14 +1,10 @@
|
|||
import asyncio
|
||||
import gc
|
||||
import io
|
||||
import os
|
||||
import pathlib
|
||||
import ssl
|
||||
import sys
|
||||
import threading
|
||||
import time
|
||||
import weakref
|
||||
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import certifi
|
||||
|
|
@ -797,134 +793,3 @@ class TestDefaultCachedClientTimeoutHonorsRequestTimeout:
|
|||
litellm.in_memory_llm_clients_cache = LLMClientCache()
|
||||
client = get_async_httpx_client(llm_provider=LlmProviders.BEDROCK)
|
||||
assert client.timeout.read == 300.0
|
||||
|
||||
|
||||
class _ChunkedSSEServer:
|
||||
"""In-process HTTP server that answers every request with chunked SSE frames."""
|
||||
|
||||
def __init__(self, frame_count: int = 6, frame_delay: float = 0.05) -> None:
|
||||
self.frame_count = frame_count
|
||||
self.frame_delay = frame_delay
|
||||
parent = self
|
||||
|
||||
class _Handler(BaseHTTPRequestHandler):
|
||||
protocol_version = "HTTP/1.1"
|
||||
|
||||
def _stream(self):
|
||||
self.send_response(200)
|
||||
self.send_header("Content-Type", "text/event-stream")
|
||||
self.send_header("Transfer-Encoding", "chunked")
|
||||
self.end_headers()
|
||||
try:
|
||||
for index in range(parent.frame_count):
|
||||
frame = f"data: frame-{index}\n\n".encode()
|
||||
self.wfile.write(b"%x\r\n" % len(frame) + frame + b"\r\n")
|
||||
self.wfile.flush()
|
||||
time.sleep(parent.frame_delay)
|
||||
self.wfile.write(b"0\r\n\r\n")
|
||||
self.wfile.flush()
|
||||
except (BrokenPipeError, ConnectionResetError):
|
||||
pass
|
||||
|
||||
do_GET = _stream
|
||||
do_POST = _stream
|
||||
|
||||
def log_message(self, *args):
|
||||
pass
|
||||
|
||||
self._server = ThreadingHTTPServer(("127.0.0.1", 0), _Handler)
|
||||
self.url = f"http://127.0.0.1:{self._server.server_address[1]}/stream"
|
||||
|
||||
def __enter__(self):
|
||||
threading.Thread(target=self._server.serve_forever, daemon=True).start()
|
||||
return self
|
||||
|
||||
def __exit__(self, *exc_info):
|
||||
self._server.shutdown()
|
||||
self._server.server_close()
|
||||
|
||||
|
||||
class TestHandlerCollectionDoesNotAbortInFlightStreams:
|
||||
"""Regression guard for https://github.com/BerriAI/litellm/issues/24929"""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("disable_aiohttp_transport", [False, True])
|
||||
async def test_async_stream_survives_handler_collection(self, monkeypatch, disable_aiohttp_transport):
|
||||
monkeypatch.delenv("DISABLE_AIOHTTP_TRANSPORT", raising=False)
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", disable_aiohttp_transport)
|
||||
monkeypatch.setattr(litellm, "force_ipv4", False)
|
||||
|
||||
with _ChunkedSSEServer() as server:
|
||||
handler = AsyncHTTPHandler(timeout=httpx.Timeout(10.0, connect=5.0))
|
||||
client = handler.client
|
||||
response = await client.send(client.build_request("GET", server.url), stream=True)
|
||||
del handler
|
||||
|
||||
async def read_frames(body: httpx.Response) -> int:
|
||||
total = 0
|
||||
async for chunk in body.aiter_bytes():
|
||||
total += chunk.count(b"data: frame-")
|
||||
gc.collect()
|
||||
return total
|
||||
|
||||
frames = await asyncio.wait_for(read_frames(response), timeout=20)
|
||||
|
||||
assert frames == 6
|
||||
assert client.is_closed is False
|
||||
|
||||
del response
|
||||
gc.collect()
|
||||
for _ in range(200):
|
||||
if client.is_closed:
|
||||
break
|
||||
await asyncio.sleep(0.01)
|
||||
assert client.is_closed is True
|
||||
|
||||
def test_sync_stream_survives_handler_collection(self, monkeypatch):
|
||||
monkeypatch.setattr(litellm, "force_ipv4", False)
|
||||
|
||||
with _ChunkedSSEServer() as server:
|
||||
handler = HTTPHandler(timeout=httpx.Timeout(10.0, connect=5.0))
|
||||
client = handler.client
|
||||
response = client.send(client.build_request("GET", server.url), stream=True)
|
||||
del handler
|
||||
|
||||
frames = 0
|
||||
for chunk in response.iter_bytes():
|
||||
frames += chunk.count(b"data: frame-")
|
||||
gc.collect()
|
||||
|
||||
assert frames == 6
|
||||
assert client.is_closed is False
|
||||
|
||||
del response
|
||||
gc.collect()
|
||||
assert client.is_closed is True
|
||||
|
||||
def test_client_does_not_keep_its_handler_alive(self, monkeypatch):
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
handler = AsyncHTTPHandler()
|
||||
client = handler.client
|
||||
handler_ref = weakref.ref(handler)
|
||||
|
||||
del handler
|
||||
gc.collect()
|
||||
|
||||
assert handler_ref() is None
|
||||
assert isinstance(client, httpx.AsyncClient)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_caller_supplied_event_hooks_still_fire(self):
|
||||
seen = []
|
||||
|
||||
async def user_response_hook(response: httpx.Response) -> None:
|
||||
seen.append(response.status_code)
|
||||
|
||||
with _ChunkedSSEServer(frame_count=1, frame_delay=0.0) as server:
|
||||
handler = AsyncHTTPHandler(event_hooks={"request": [], "response": [user_response_hook]})
|
||||
try:
|
||||
response = await handler.client.get(server.url)
|
||||
assert response.status_code == 200
|
||||
assert seen == [200]
|
||||
finally:
|
||||
await handler.close()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue