mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(proxy): do not let held-back keepalive pings block the budget reservation refund on client disconnect
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
cbefb1ce5f
commit
bb0bb48da8
6 changed files with 57 additions and 23 deletions
|
|
@ -434,6 +434,8 @@ CONNECTION_ERROR_PATTERNS: Final[list[str]] = [
|
|||
]
|
||||
STREAM_SSE_DONE_STRING: Final[str] = "[DONE]"
|
||||
STREAM_SSE_DATA_PREFIX: Final[str] = "data: "
|
||||
STREAM_SSE_KEEPALIVE_PING_CHUNK: Final[str] = 'event: ping\ndata: {"type": "ping"}\n\n'
|
||||
STREAM_SSE_KEEPALIVE_PING_BYTES: Final[bytes] = STREAM_SSE_KEEPALIVE_PING_CHUNK.encode("utf-8")
|
||||
### SPEND TRACKING ###
|
||||
DEFAULT_REPLICATE_GPU_PRICE_PER_SECOND: Final = float(
|
||||
os.getenv("DEFAULT_REPLICATE_GPU_PRICE_PER_SECOND", 0.001400)
|
||||
|
|
|
|||
|
|
@ -21,8 +21,8 @@ from collections.abc import AsyncIterator
|
|||
from typing import Any, Final, cast
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.constants import STREAM_SSE_KEEPALIVE_PING_BYTES
|
||||
|
||||
PING_SSE_BYTES: Final = b'event: ping\ndata: {"type": "ping"}\n\n'
|
||||
HOLD_BACK_PING_INTERVAL_SECONDS: Final = 15.0
|
||||
SERVER_FULFILLED_TOOL_LEAK_ERROR_SSE_BYTES: Final = (
|
||||
b"event: error\n"
|
||||
|
|
@ -249,17 +249,17 @@ class AgenticAnthropicStreamingIterator:
|
|||
async def _anext_held_back(self) -> bytes:
|
||||
if self._drain_task is None:
|
||||
self._drain_task = asyncio.create_task(self._drain_upstream())
|
||||
return PING_SSE_BYTES
|
||||
return STREAM_SSE_KEEPALIVE_PING_BYTES
|
||||
|
||||
if not self._stream_exhausted:
|
||||
if not await self._completed_within_ping_interval(self._drain_task):
|
||||
return PING_SSE_BYTES
|
||||
return STREAM_SSE_KEEPALIVE_PING_BYTES
|
||||
self._stream_exhausted = True
|
||||
|
||||
if self._hook_task is None:
|
||||
self._hook_task = asyncio.create_task(self._process_agentic_hooks())
|
||||
if not await self._completed_within_ping_interval(self._hook_task):
|
||||
return PING_SSE_BYTES
|
||||
return STREAM_SSE_KEEPALIVE_PING_BYTES
|
||||
|
||||
if self._follow_up_iterator is not None:
|
||||
return await self._next_follow_up_chunk(self._follow_up_iterator)
|
||||
|
|
@ -286,7 +286,7 @@ class AgenticAnthropicStreamingIterator:
|
|||
if self._follow_up_chunk_task is None:
|
||||
self._follow_up_chunk_task = asyncio.create_task(_anext_or_none(follow_up_iterator))
|
||||
if not await self._completed_within_ping_interval(self._follow_up_chunk_task):
|
||||
return PING_SSE_BYTES
|
||||
return STREAM_SSE_KEEPALIVE_PING_BYTES
|
||||
chunk: Final = self._follow_up_chunk_task.result()
|
||||
self._follow_up_chunk_task = None
|
||||
if chunk is None:
|
||||
|
|
|
|||
|
|
@ -27,6 +27,7 @@ from litellm.constants import (
|
|||
MAX_PAYLOAD_SIZE_FOR_DEBUG_LOG,
|
||||
RETURN_RAW_MODEL_NAME_METADATA_KEY,
|
||||
STREAM_SSE_DATA_PREFIX,
|
||||
STREAM_SSE_KEEPALIVE_PING_BYTES,
|
||||
UNSAFE_PROXY_RESPONSE_HEADERS,
|
||||
)
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
|
|
@ -2953,8 +2954,9 @@ class ProxyBaseLLMRequestProcessing:
|
|||
# so a GeneratorExit on client disconnect is raised there and any
|
||||
# statement after the yield never runs. The slow-path hook is
|
||||
# awaited above, so a cancellation during it still leaves this
|
||||
# False and refunds.
|
||||
delivered_chunk = True
|
||||
# False and refunds. A keepalive ping carries no provider output,
|
||||
# so it must not suppress that refund.
|
||||
delivered_chunk = delivered_chunk or chunk != STREAM_SSE_KEEPALIVE_PING_BYTES
|
||||
yield serialize_chunk(chunk)
|
||||
stream_completed = True
|
||||
except (asyncio.CancelledError, GeneratorExit):
|
||||
|
|
|
|||
|
|
@ -6,7 +6,9 @@ from typing import Final
|
|||
|
||||
import anyio
|
||||
|
||||
ANTHROPIC_PING_SSE_CHUNK: Final = 'event: ping\ndata: {"type": "ping"}\n\n'
|
||||
from litellm.constants import STREAM_SSE_KEEPALIVE_PING_CHUNK
|
||||
|
||||
ANTHROPIC_PING_SSE_CHUNK: Final = STREAM_SSE_KEEPALIVE_PING_CHUNK
|
||||
|
||||
|
||||
def _coerce_interval(ping_interval_seconds: float | str | None) -> float | None:
|
||||
|
|
|
|||
|
|
@ -13,8 +13,8 @@ import pytest
|
|||
|
||||
sys.path.insert(0, os.path.abspath("../../../../.."))
|
||||
|
||||
from litellm.constants import STREAM_SSE_KEEPALIVE_PING_BYTES
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.agentic_streaming_iterator import (
|
||||
PING_SSE_BYTES,
|
||||
SERVER_FULFILLED_TOOL_LEAK_ERROR_SSE_BYTES,
|
||||
AgenticAnthropicStreamingIterator,
|
||||
_handle_content_block_delta,
|
||||
|
|
@ -858,10 +858,10 @@ class TestAgenticStreamingIteratorHoldBack:
|
|||
async for chunk in iterator:
|
||||
collected.append(chunk)
|
||||
|
||||
non_ping = [c for c in collected if c != PING_SSE_BYTES]
|
||||
non_ping = [c for c in collected if c != STREAM_SSE_KEEPALIVE_PING_BYTES]
|
||||
assert non_ping == phase2_chunks
|
||||
assert b"litellm_content_retrieve" not in b"".join(collected)
|
||||
assert collected[0] == PING_SSE_BYTES
|
||||
assert collected[0] == STREAM_SSE_KEEPALIVE_PING_BYTES
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_should_replay_buffer_verbatim_when_no_hook_fires(self):
|
||||
|
|
@ -877,7 +877,7 @@ class TestAgenticStreamingIteratorHoldBack:
|
|||
async for chunk in iterator:
|
||||
collected.append(chunk)
|
||||
|
||||
assert [c for c in collected if c != PING_SSE_BYTES] == chunks
|
||||
assert [c for c in collected if c != STREAM_SSE_KEEPALIVE_PING_BYTES] == chunks
|
||||
mock_handler._call_agentic_completion_hooks.assert_awaited_once()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -898,8 +898,8 @@ class TestAgenticStreamingIteratorHoldBack:
|
|||
async for chunk in iterator:
|
||||
collected.append(chunk)
|
||||
|
||||
assert collected.count(PING_SSE_BYTES) >= 2
|
||||
assert [c for c in collected if c != PING_SSE_BYTES] == chunks
|
||||
assert collected.count(STREAM_SSE_KEEPALIVE_PING_BYTES) >= 2
|
||||
assert [c for c in collected if c != STREAM_SSE_KEEPALIVE_PING_BYTES] == chunks
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_should_propagate_upstream_error_instead_of_partial_message(self):
|
||||
|
|
@ -919,7 +919,7 @@ class TestAgenticStreamingIteratorHoldBack:
|
|||
async for chunk in iterator:
|
||||
collected.append(chunk)
|
||||
|
||||
assert all(c == PING_SSE_BYTES for c in collected)
|
||||
assert all(c == STREAM_SSE_KEEPALIVE_PING_BYTES for c in collected)
|
||||
mock_handler._call_agentic_completion_hooks.assert_not_awaited()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -945,8 +945,8 @@ class TestAgenticStreamingIteratorHoldBack:
|
|||
async for chunk in iterator:
|
||||
collected.append(chunk)
|
||||
|
||||
assert collected.count(PING_SSE_BYTES) >= 4
|
||||
assert [c for c in collected if c != PING_SSE_BYTES] == phase2_chunks
|
||||
assert collected.count(STREAM_SSE_KEEPALIVE_PING_BYTES) >= 4
|
||||
assert [c for c in collected if c != STREAM_SSE_KEEPALIVE_PING_BYTES] == phase2_chunks
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_should_error_instead_of_replaying_server_fulfilled_tool_use_when_hook_crashes(self):
|
||||
|
|
@ -962,7 +962,7 @@ class TestAgenticStreamingIteratorHoldBack:
|
|||
async for chunk in iterator:
|
||||
collected.append(chunk)
|
||||
|
||||
assert [c for c in collected if c != PING_SSE_BYTES] == [SERVER_FULFILLED_TOOL_LEAK_ERROR_SSE_BYTES]
|
||||
assert [c for c in collected if c != STREAM_SSE_KEEPALIVE_PING_BYTES] == [SERVER_FULFILLED_TOOL_LEAK_ERROR_SSE_BYTES]
|
||||
assert b"litellm_content_retrieve" not in b"".join(collected)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -979,7 +979,7 @@ class TestAgenticStreamingIteratorHoldBack:
|
|||
async for chunk in iterator:
|
||||
collected.append(chunk)
|
||||
|
||||
assert [c for c in collected if c != PING_SSE_BYTES] == [SERVER_FULFILLED_TOOL_LEAK_ERROR_SSE_BYTES]
|
||||
assert [c for c in collected if c != STREAM_SSE_KEEPALIVE_PING_BYTES] == [SERVER_FULFILLED_TOOL_LEAK_ERROR_SSE_BYTES]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_should_replay_client_owned_tool_use_verbatim(self):
|
||||
|
|
@ -999,7 +999,7 @@ class TestAgenticStreamingIteratorHoldBack:
|
|||
async for chunk in iterator:
|
||||
collected.append(chunk)
|
||||
|
||||
assert [c for c in collected if c != PING_SSE_BYTES] == chunks
|
||||
assert [c for c in collected if c != STREAM_SSE_KEEPALIVE_PING_BYTES] == chunks
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_should_emit_pings_while_the_follow_up_stream_is_slow(self):
|
||||
|
|
@ -1023,8 +1023,8 @@ class TestAgenticStreamingIteratorHoldBack:
|
|||
collected.append(chunk)
|
||||
|
||||
first_follow_up_index = collected.index(phase2_chunks[0])
|
||||
assert collected[first_follow_up_index + 1] == PING_SSE_BYTES
|
||||
assert [c for c in collected if c != PING_SSE_BYTES] == phase2_chunks
|
||||
assert collected[first_follow_up_index + 1] == STREAM_SSE_KEEPALIVE_PING_BYTES
|
||||
assert [c for c in collected if c != STREAM_SSE_KEEPALIVE_PING_BYTES] == phase2_chunks
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_should_propagate_follow_up_stream_error(self):
|
||||
|
|
@ -1074,7 +1074,7 @@ class TestAgenticStreamingIteratorHoldBack:
|
|||
)
|
||||
|
||||
first = await iterator.__anext__()
|
||||
assert first == PING_SSE_BYTES
|
||||
assert first == STREAM_SSE_KEEPALIVE_PING_BYTES
|
||||
assert iterator._drain_task is not None
|
||||
|
||||
await iterator.aclose()
|
||||
|
|
|
|||
|
|
@ -7,6 +7,7 @@ from fastapi import HTTPException
|
|||
|
||||
import litellm
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.constants import STREAM_SSE_KEEPALIVE_PING_BYTES
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_BudgetTable,
|
||||
LiteLLM_EndUserTable,
|
||||
|
|
@ -2453,6 +2454,33 @@ async def test_streaming_cancel_after_chunk_keeps_reservation(
|
|||
streaming_logging_obj._arelease_max_parallel_requests_on_disconnect.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_streaming_cancel_after_only_keepalive_pings_reconciles_to_input_cost(
|
||||
spend_counter_state,
|
||||
):
|
||||
counter_cache, key_cache = spend_counter_state
|
||||
proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache)
|
||||
valid_token, reservation = await _reserve_for_stream(
|
||||
counter_cache, key_cache, proxy_logging_obj, "key-cancel-after-ping"
|
||||
)
|
||||
|
||||
async def cancel_after_ping(user_api_key_dict, response, request_data):
|
||||
yield STREAM_SSE_KEEPALIVE_PING_BYTES
|
||||
raise asyncio.CancelledError()
|
||||
|
||||
generator, streaming_logging_obj = _drive_streaming_cancel(valid_token, cancel_after_ping)
|
||||
received = []
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
async for chunk in generator:
|
||||
received.append(chunk)
|
||||
|
||||
assert received == [STREAM_SSE_KEEPALIVE_PING_BYTES]
|
||||
assert counter_cache.in_memory_cache.get_cache(
|
||||
key="spend:key:key-cancel-after-ping"
|
||||
) == pytest.approx(0.5)
|
||||
assert reservation["finalized"] is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_release_budget_reservation_on_cancel_swallows_release_errors():
|
||||
# If the release itself fails (e.g. Redis unavailable) it must not escape
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue