fix(proxy): cancel upstream gemini request and release httpx connection on client disconnect

- add _check_request_disconnection to common_request_processing; wrap llm_call
  as asyncio.Task so it can be cancelled; catch CancelledError and raise
  HTTPException(499) when client disconnects before LLM responds (non-streaming path)

- pass raw httpx.Response into ModelResponseIterator in make_call/make_sync_call
  so the iterator holds a reference to the underlying connection

- implement ModelResponseIterator.aclose() and .close(): close the line iterator
  then explicitly call response.aclose()/response.close() to release the httpx
  connection when the client drops mid-stream; errors are debug-logged, not raised

- add tests for _check_request_disconnection (cancels task, graceful on exception,
  does not cancel when client stays connected) and base_process_llm_request 499
  behavior; add TestModelResponseIteratorCleanup verifying aclose/close propagation
  through CustomStreamWrapper
This commit is contained in:
S0ngRu1 2026-06-10 11:54:09 +08:00
parent 51ba6e39cd
commit 4664963003
4 changed files with 356 additions and 2 deletions

View file

@ -2846,6 +2846,7 @@ async def make_call(
sync_stream=False,
logging_obj=logging_obj,
response_headers=response.headers,
response=response,
)
# LOGGING
logging_obj.post_call(
@ -2889,6 +2890,7 @@ def make_sync_call(
sync_stream=True,
logging_obj=logging_obj,
response_headers=response.headers,
response=response,
)
# LOGGING
@ -3350,12 +3352,14 @@ class ModelResponseIterator:
sync_stream: bool,
logging_obj: LoggingClass,
response_headers: Optional[Dict[str, str]] = None,
response: Optional[httpx.Response] = None,
):
from litellm.litellm_core_utils.prompt_templates.common_utils import (
check_is_function_call,
)
self.streaming_response = streaming_response
self.response = response
self.chunk_type: Literal["valid_json", "accumulated_json"] = "valid_json"
self.accumulated_json = ""
self.sent_first_chunk = False
@ -3653,3 +3657,37 @@ class ModelResponseIterator:
raise StopAsyncIteration
except ValueError as e:
raise RuntimeError(f"Error parsing chunk: {e},\nReceived chunk: {chunk}")
async def aclose(self) -> None:
iterator = getattr(self, "async_response_iterator", self.streaming_response)
if iterator is not None and hasattr(iterator, "aclose"):
try:
await iterator.aclose()
except Exception as e:
verbose_logger.debug(
"ModelResponseIterator.aclose: error closing iterator: %s", e
)
if self.response is not None:
try:
await self.response.aclose()
except Exception as e:
verbose_logger.debug(
"ModelResponseIterator.aclose: error closing response: %s", e
)
def close(self) -> None:
iterator = getattr(self, "response_iterator", self.streaming_response)
if iterator is not None and hasattr(iterator, "close"):
try:
iterator.close()
except Exception as e:
verbose_logger.debug(
"ModelResponseIterator.close: error closing iterator: %s", e
)
if self.response is not None:
try:
self.response.close()
except Exception as e:
verbose_logger.debug(
"ModelResponseIterator.close: error closing response: %s", e
)

View file

@ -73,6 +73,20 @@ from litellm.types.utils import ModelResponse, ModelResponseStream, Usage
_DD_STREAMING_TRACE_ENABLED = not isinstance(tracer, NullTracer)
async def _check_request_disconnection(
request: Request, llm_api_call_task: "asyncio.Task[Any]"
) -> None:
start_time = time.time()
while time.time() - start_time < 600:
await asyncio.sleep(1)
try:
if await request.is_disconnected():
llm_api_call_task.cancel()
return
except Exception:
return
def _serialize_http_exception_detail(
detail: Any,
) -> Tuple[str, Optional[dict]]:
@ -1215,14 +1229,31 @@ class ProxyBaseLLMRequestProcessing:
user_model=user_model,
user_api_key_dict=user_api_key_dict,
)
tasks.append(llm_call)
llm_call_task = asyncio.create_task(llm_call)
tasks.append(llm_call_task)
disconnect_task = asyncio.create_task(
_check_request_disconnection(request, llm_call_task)
)
# wait for call to end
llm_responses = asyncio.gather(
*tasks
) # run the moderation check in parallel to the actual llm api call
responses = await llm_responses
try:
responses = await llm_responses
except asyncio.CancelledError:
raise HTTPException(
status_code=499,
detail="Client disconnected the request",
)
finally:
disconnect_task.cancel()
try:
await disconnect_task
except asyncio.CancelledError:
pass
response = responses[1]

View file

@ -4996,3 +4996,146 @@ def test_mid_stream_429_error_raises_during_iteration():
# Verify: 429 error is properly raised
assert exc_info.value.status_code == 429
assert "RESOURCE_EXHAUSTED" in str(exc_info.value.message)
class TestModelResponseIteratorCleanup:
def _make_logging_obj(self):
from unittest.mock import Mock
obj = Mock()
obj.optional_params = {}
return obj
def test_aclose_closes_iterator_and_response(self):
import asyncio
from unittest.mock import AsyncMock, MagicMock
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
ModelResponseIterator,
)
mock_response = MagicMock()
mock_response.aclose = AsyncMock()
mock_iterator = MagicMock()
mock_iterator.aclose = AsyncMock()
iterator = ModelResponseIterator(
streaming_response=MagicMock(),
sync_stream=False,
logging_obj=self._make_logging_obj(),
response=mock_response,
)
iterator.async_response_iterator = mock_iterator
asyncio.run(iterator.aclose())
mock_iterator.aclose.assert_awaited_once()
mock_response.aclose.assert_awaited_once()
def test_close_closes_iterator_and_response(self):
from unittest.mock import MagicMock
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
ModelResponseIterator,
)
mock_response = MagicMock()
mock_iterator = MagicMock()
iterator = ModelResponseIterator(
streaming_response=MagicMock(),
sync_stream=True,
logging_obj=self._make_logging_obj(),
response=mock_response,
)
iterator.response_iterator = mock_iterator
iterator.close()
mock_iterator.close.assert_called_once()
mock_response.close.assert_called_once()
def test_aclose_without_response_does_not_raise(self):
import asyncio
from unittest.mock import AsyncMock, MagicMock
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
ModelResponseIterator,
)
mock_iterator = MagicMock()
mock_iterator.aclose = AsyncMock()
iterator = ModelResponseIterator(
streaming_response=MagicMock(),
sync_stream=False,
logging_obj=self._make_logging_obj(),
)
iterator.async_response_iterator = mock_iterator
asyncio.run(iterator.aclose())
mock_iterator.aclose.assert_awaited_once()
def test_aclose_tolerates_iterator_error(self):
import asyncio
from unittest.mock import AsyncMock, MagicMock
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
ModelResponseIterator,
)
mock_response = MagicMock()
mock_response.aclose = AsyncMock()
mock_iterator = MagicMock()
mock_iterator.aclose = AsyncMock(side_effect=RuntimeError("transport error"))
iterator = ModelResponseIterator(
streaming_response=MagicMock(),
sync_stream=False,
logging_obj=self._make_logging_obj(),
response=mock_response,
)
iterator.async_response_iterator = mock_iterator
asyncio.run(iterator.aclose())
mock_response.aclose.assert_awaited_once()
def test_custom_stream_wrapper_aclose_triggers_model_response_iterator_aclose(self):
"""CustomStreamWrapper.aclose() must propagate to ModelResponseIterator.aclose()."""
import asyncio
from unittest.mock import AsyncMock, MagicMock
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
ModelResponseIterator,
)
mock_response = MagicMock()
mock_response.aclose = AsyncMock()
mock_iterator = MagicMock()
mock_iterator.aclose = AsyncMock()
model_response_iter = ModelResponseIterator(
streaming_response=MagicMock(),
sync_stream=False,
logging_obj=self._make_logging_obj(),
response=mock_response,
)
model_response_iter.async_response_iterator = mock_iterator
wrapper = CustomStreamWrapper(
completion_stream=model_response_iter,
model="gemini-2.0-flash",
custom_llm_provider="vertex_ai",
logging_obj=MagicMock(),
)
asyncio.run(wrapper.aclose())
mock_iterator.aclose.assert_awaited_once()
mock_response.aclose.assert_awaited_once()

View file

@ -2342,3 +2342,145 @@ class TestAsyncStreamingDataGeneratorFastPath:
hook_spy.assert_awaited_once()
ProxyLogging._callback_capabilities_cache.clear()
class TestCheckRequestDisconnection:
@pytest.mark.asyncio
async def test_cancels_task_when_client_disconnects(self, monkeypatch):
import asyncio
import litellm.proxy.common_request_processing as cpr
from litellm.proxy.common_request_processing import _check_request_disconnection
task_cancelled = False
async def never_ending():
nonlocal task_cancelled
try:
await asyncio.sleep(9999)
except asyncio.CancelledError:
task_cancelled = True
raise
task = asyncio.create_task(never_ending())
mock_request = MagicMock(spec=Request)
mock_request.is_disconnected = AsyncMock(return_value=True)
monkeypatch.setattr(cpr.asyncio, "sleep", AsyncMock())
await _check_request_disconnection(mock_request, task)
await asyncio.sleep(0)
assert task.cancelled() or task_cancelled
@pytest.mark.asyncio
async def test_does_not_cancel_task_when_client_stays_connected(self, monkeypatch):
import asyncio
import litellm.proxy.common_request_processing as cpr
from litellm.proxy.common_request_processing import _check_request_disconnection
call_count = 0
async def fake_sleep(_):
nonlocal call_count
call_count += 1
if call_count >= 3:
raise asyncio.CancelledError
task = asyncio.create_task(asyncio.sleep(9999))
mock_request = MagicMock(spec=Request)
mock_request.is_disconnected = AsyncMock(return_value=False)
monkeypatch.setattr(cpr.asyncio, "sleep", fake_sleep)
with pytest.raises(asyncio.CancelledError):
await _check_request_disconnection(mock_request, task)
assert not task.cancelled()
task.cancel()
try:
await task
except asyncio.CancelledError:
pass
@pytest.mark.asyncio
async def test_exits_gracefully_when_is_disconnected_raises(self, monkeypatch):
import asyncio
import litellm.proxy.common_request_processing as cpr
from litellm.proxy.common_request_processing import _check_request_disconnection
task = asyncio.create_task(asyncio.sleep(9999))
mock_request = MagicMock(spec=Request)
mock_request.is_disconnected = AsyncMock(
side_effect=RuntimeError("transport closed")
)
monkeypatch.setattr(cpr.asyncio, "sleep", AsyncMock())
await _check_request_disconnection(mock_request, task)
assert not task.cancelled()
task.cancel()
try:
await task
except asyncio.CancelledError:
pass
@pytest.mark.asyncio
async def test_base_process_llm_request_raises_499_on_client_disconnect(
self, monkeypatch
):
"""When _check_request_disconnection cancels the LLM task, base_process_llm_request
must raise HTTPException(499) instead of propagating CancelledError."""
import asyncio
import litellm.proxy.common_request_processing as cpr
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
async def slow_llm():
await asyncio.sleep(9999)
async def fake_route_request(**_kwargs):
return slow_llm()
async def instant_disconnect(request, task):
task.cancel()
mock_logging_obj = MagicMock()
mock_logging_obj.litellm_call_id = "test-call-id"
mock_logging_obj._defer_async_logging = False
mock_proxy_logging = MagicMock(spec=ProxyLogging)
mock_proxy_logging.during_call_hook = AsyncMock(return_value=None)
mock_proxy_logging._callback_capabilities_cache = {}
monkeypatch.setattr(cpr, "route_request", fake_route_request)
monkeypatch.setattr(cpr, "_check_request_disconnection", instant_disconnect)
processing_obj = ProxyBaseLLMRequestProcessing(data={"model": "gemini-2.0-flash"})
monkeypatch.setattr(
processing_obj,
"common_processing_pre_call_logic",
AsyncMock(return_value=({"model": "gemini-2.0-flash"}, mock_logging_obj)),
)
monkeypatch.setattr(
processing_obj, "_has_post_call_guardrails", MagicMock(return_value=False)
)
mock_request = MagicMock(spec=Request)
mock_request.is_disconnected = AsyncMock(return_value=True)
mock_request.headers = {}
with pytest.raises(HTTPException) as exc_info:
await processing_obj.base_process_llm_request(
request=mock_request,
fastapi_response=MagicMock(),
user_api_key_dict=MagicMock(spec=UserAPIKeyAuth),
proxy_logging_obj=mock_proxy_logging,
general_settings={},
proxy_config=MagicMock(spec=ProxyConfig),
route_type="acompletion",
version=None,
)
assert exc_info.value.status_code == 499
assert "disconnected" in exc_info.value.detail.lower()