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:
mateo 2026-08-08 20:05:45 +00:00
parent cbefb1ce5f
commit bb0bb48da8
6 changed files with 57 additions and 23 deletions

View file

@ -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)

View file

@ -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:

View file

@ -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):

View file

@ -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:

View file

@ -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()

View file

@ -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