mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
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:
parent
51ba6e39cd
commit
4664963003
4 changed files with 356 additions and 2 deletions
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue