fix(passthrough): flush spend tracking on interrupted Bedrock streams

When a client disconnects mid-stream from a Bedrock pass-through endpoint,
Starlette calls aclose() on the async generator, raising GeneratorExit
(a BaseException, not Exception) at the suspended yield. The previous
`except Exception` blocks in _async_streaming/_sync_streaming
(litellm/passthrough/main.py) and PassThroughStreamingHandler.chunk_processor
did not catch GeneratorExit, so the post-loop flush that hands collected
raw bytes to async_flush_passthrough_collected_chunks /
_route_streaming_logging_to_handler never ran. All per-chunk usage data
was silently dropped, undercounting spend for interrupted Bedrock invoke
and converse streams.

Move the flush into a finally block in all three sites and guard with a
`flush_scheduled` flag so the success path still flushes exactly once.
Also pull raise_for_status() out of the chunk-collection try block in
_async_streaming so 4xx/5xx responses still raise and don't enter the
flush path with zero bytes (preserving the behavior tested by
test_async_streaming_error_propagation.py).

Add regression coverage:
- test_async_streaming_flushes_on_client_disconnect
- test_async_streaming_flushes_on_upstream_exception_with_partial_data
- test_sync_streaming_flushes_on_early_close
- test_chunk_processor_logs_on_client_disconnect
plus baseline tests for normal completion and the 4xx no-flush path.

Fixes LIT-2642.

Co-authored-by: Mateo Wang <mateo-berri@users.noreply.github.com>
This commit is contained in:
Cursor Agent 2026-04-28 21:07:17 +00:00
parent 2e561bd04e
commit 1ef034bff6
No known key found for this signature in database
4 changed files with 534 additions and 46 deletions

View file

@ -390,19 +390,29 @@ def _sync_streaming(
):
from litellm.utils import executor
raw_bytes: List[bytes] = []
flush_scheduled = False
try:
raw_bytes: List[bytes] = []
for chunk in response.iter_bytes(): # type: ignore
raw_bytes.append(chunk)
yield chunk
executor.submit(
litellm_logging_obj.flush_passthrough_collected_chunks,
raw_bytes=raw_bytes,
provider_config=provider_config,
)
except Exception as e:
raise e
finally:
# Always flush collected chunks for spend tracking, even if the
# consumer terminates the generator early (GeneratorExit). Without
# this, an interrupted stream loses all per-chunk usage data
# because the post-loop flush never runs. See LIT-2642.
if not flush_scheduled and raw_bytes:
flush_scheduled = True
try:
executor.submit(
litellm_logging_obj.flush_passthrough_collected_chunks,
raw_bytes=raw_bytes,
provider_config=provider_config,
)
except Exception:
# Don't mask the original exception (incl. GeneratorExit)
# if scheduling the flush itself fails.
pass
async def _async_streaming(
@ -411,23 +421,47 @@ async def _async_streaming(
provider_config: "BasePassthroughConfig",
):
iter_response = await response
# Validate response status before consuming the body so 4xx/5xx
# responses raise without entering the chunk-collection path.
try:
iter_response.raise_for_status()
raw_bytes: List[bytes] = []
async for chunk in iter_response.aiter_bytes(): # type: ignore
raw_bytes.append(chunk)
yield chunk
asyncio.create_task(
litellm_logging_obj.async_flush_passthrough_collected_chunks(
raw_bytes=raw_bytes,
provider_config=provider_config,
)
)
except Exception:
try:
await iter_response.aclose()
except Exception:
pass
raise
raw_bytes: List[bytes] = []
flush_scheduled = False
try:
async for chunk in iter_response.aiter_bytes(): # type: ignore
raw_bytes.append(chunk)
yield chunk
except Exception:
try:
await iter_response.aclose()
except Exception:
pass
raise
finally:
# Always flush collected chunks for spend tracking, even if the
# client disconnects mid-stream. On disconnect, Starlette calls
# aclose() on this generator, which raises GeneratorExit at the
# suspended `yield` — `except Exception` does not catch it, so
# the post-loop flush would otherwise be skipped and all
# captured per-chunk usage data lost. See LIT-2642.
if not flush_scheduled and raw_bytes:
flush_scheduled = True
try:
asyncio.create_task(
litellm_logging_obj.async_flush_passthrough_collected_chunks(
raw_bytes=raw_bytes,
provider_config=provider_config,
)
)
except Exception:
# Don't mask the original exception (incl. GeneratorExit)
# if scheduling the flush itself fails.
pass

View file

@ -41,16 +41,17 @@ class PassThroughStreamingHandler:
- Collect non-empty chunks for post-processing (logging)
- Inject cost into chunks if include_cost_in_streaming_usage is enabled
"""
try:
raw_bytes: List[bytes] = []
# Extract model name for cost injection
model_name = PassThroughStreamingHandler._extract_model_for_cost_injection(
request_body=request_body,
url_route=url_route,
endpoint_type=endpoint_type,
litellm_logging_obj=litellm_logging_obj,
)
raw_bytes: List[bytes] = []
logging_scheduled = False
# Extract model name for cost injection
model_name = PassThroughStreamingHandler._extract_model_for_cost_injection(
request_body=request_body,
url_route=url_route,
endpoint_type=endpoint_type,
litellm_logging_obj=litellm_logging_obj,
)
try:
async for chunk in response.aiter_bytes():
raw_bytes.append(chunk)
if (
@ -73,25 +74,37 @@ class PassThroughStreamingHandler:
chunk = modified_chunk
yield chunk
# After all chunks are processed, handle post-processing
end_time = datetime.now()
asyncio.create_task(
PassThroughStreamingHandler._route_streaming_logging_to_handler(
litellm_logging_obj=litellm_logging_obj,
passthrough_success_handler_obj=passthrough_success_handler_obj,
url_route=url_route,
request_body=request_body or {},
endpoint_type=endpoint_type,
start_time=start_time,
raw_bytes=raw_bytes,
end_time=end_time,
)
)
except Exception as e:
verbose_proxy_logger.error(f"Error in chunk_processor: {str(e)}")
raise
finally:
# Always log collected chunks for spend tracking, even if the
# client disconnects mid-stream. On disconnect, Starlette calls
# aclose() on this async generator, which raises GeneratorExit
# at the suspended `yield` — `except Exception` does not catch
# it, so post-loop logging would otherwise be skipped and all
# captured per-chunk usage data lost (e.g. for interrupted
# Bedrock streams). See LIT-2642.
if not logging_scheduled and raw_bytes:
logging_scheduled = True
try:
end_time = datetime.now()
asyncio.create_task(
PassThroughStreamingHandler._route_streaming_logging_to_handler(
litellm_logging_obj=litellm_logging_obj,
passthrough_success_handler_obj=passthrough_success_handler_obj,
url_route=url_route,
request_body=request_body or {},
endpoint_type=endpoint_type,
start_time=start_time,
raw_bytes=raw_bytes,
end_time=end_time,
)
)
except Exception as e:
verbose_proxy_logger.error(
f"Error scheduling chunk_processor logging: {str(e)}"
)
@staticmethod
async def _route_streaming_logging_to_handler(

View file

@ -0,0 +1,297 @@
"""
Regression tests for LIT-2642 — interrupted streaming responses must still
flush collected chunks so spend is tracked even when the client disconnects
mid-stream.
Bedrock invoke streaming was the reported reproducer: the proxy passes the
upstream stream through `_async_streaming` in `litellm/passthrough/main.py`,
which collects bytes and triggers `async_flush_passthrough_collected_chunks`
once the loop completes. When a FastAPI client disconnects mid-stream,
Starlette calls `aclose()` on the async generator and raises `GeneratorExit`
at the suspended `yield`. The previous `except Exception` branch did not
catch `GeneratorExit`, so the post-loop flush was skipped and all per-chunk
usage data was dropped.
"""
from typing import List
from unittest.mock import AsyncMock, MagicMock
import httpx
import pytest
def _make_streaming_response(chunks: List[bytes]):
"""Build a mock httpx.Response that streams the given chunks via aiter_bytes."""
mock = MagicMock(spec=httpx.Response)
mock.status_code = 200
mock.headers = httpx.Headers({"content-type": "application/vnd.amazon.eventstream"})
mock.raise_for_status = MagicMock(return_value=None)
async def _aiter_bytes():
for chunk in chunks:
yield chunk
mock.aiter_bytes = _aiter_bytes
mock.aclose = AsyncMock()
return mock
def _make_logging_obj():
mock = MagicMock()
mock.async_flush_passthrough_collected_chunks = AsyncMock()
return mock
@pytest.mark.asyncio
async def test_async_streaming_flushes_on_normal_completion():
"""Baseline: full stream consumption flushes collected chunks once."""
from litellm.passthrough.main import _async_streaming
chunks = [b"chunk-1", b"chunk-2", b"chunk-3"]
mock_response = _make_streaming_response(chunks)
async def response_coro():
return mock_response
mock_logging_obj = _make_logging_obj()
provider_config = MagicMock()
received = []
async for chunk in _async_streaming(
response=response_coro(),
litellm_logging_obj=mock_logging_obj,
provider_config=provider_config,
):
received.append(chunk)
assert received == chunks
# Allow the scheduled task to run.
import asyncio
await asyncio.sleep(0)
mock_logging_obj.async_flush_passthrough_collected_chunks.assert_called_once()
call_kwargs = (
mock_logging_obj.async_flush_passthrough_collected_chunks.call_args.kwargs
)
assert call_kwargs["raw_bytes"] == chunks
assert call_kwargs["provider_config"] is provider_config
@pytest.mark.asyncio
async def test_async_streaming_flushes_on_client_disconnect():
"""
LIT-2642 regression: GeneratorExit (raised when the consumer disconnects
mid-stream) must still flush whatever chunks we already collected so
spend tracking captures the partial usage.
"""
from litellm.passthrough.main import _async_streaming
chunks = [
b'{"chunk": 1, "outputTokens": 10}',
b'{"chunk": 2, "outputTokens": 12}',
b'{"chunk": 3, "outputTokens": 8}',
]
mock_response = _make_streaming_response(chunks)
async def response_coro():
return mock_response
mock_logging_obj = _make_logging_obj()
provider_config = MagicMock()
gen = _async_streaming(
response=response_coro(),
litellm_logging_obj=mock_logging_obj,
provider_config=provider_config,
)
# Pull one chunk, then close the generator early — mirrors what
# Starlette does when the HTTP client disconnects mid-stream.
received = [await gen.__anext__()]
await gen.aclose()
assert received == [chunks[0]]
# Allow the scheduled flush task to run.
import asyncio
await asyncio.sleep(0)
mock_logging_obj.async_flush_passthrough_collected_chunks.assert_called_once()
call_kwargs = (
mock_logging_obj.async_flush_passthrough_collected_chunks.call_args.kwargs
)
# Only the first chunk was consumed before disconnect; that's what we
# must hand off to the cost-tracking flush so partial usage isn't
# silently dropped.
assert call_kwargs["raw_bytes"] == [chunks[0]]
@pytest.mark.asyncio
async def test_async_streaming_does_not_flush_on_4xx():
"""Error responses must still raise without entering the flush path."""
from litellm.passthrough.main import _async_streaming
err_response = MagicMock(spec=httpx.Response)
err_response.status_code = 429
def _raise():
raise httpx.HTTPStatusError(
"429",
request=httpx.Request("POST", "https://example.com"),
response=httpx.Response(
429, request=httpx.Request("POST", "https://example.com")
),
)
err_response.raise_for_status = _raise
err_response.aclose = AsyncMock()
async def response_coro():
return err_response
mock_logging_obj = _make_logging_obj()
with pytest.raises(httpx.HTTPStatusError):
async for _ in _async_streaming(
response=response_coro(),
litellm_logging_obj=mock_logging_obj,
provider_config=MagicMock(),
):
pass
# No bytes were collected, so no flush should have been scheduled.
mock_logging_obj.async_flush_passthrough_collected_chunks.assert_not_called()
@pytest.mark.asyncio
async def test_async_streaming_flushes_on_upstream_exception_with_partial_data():
"""
If the upstream connection drops mid-stream and aiter_bytes raises,
we still surface the exception, but partial chunks already collected
are flushed so spend tracking isn't fully lost.
"""
from litellm.passthrough.main import _async_streaming
partial_chunks = [b"partial-chunk-1", b"partial-chunk-2"]
mock_response = MagicMock(spec=httpx.Response)
mock_response.status_code = 200
mock_response.raise_for_status = MagicMock(return_value=None)
mock_response.aclose = AsyncMock()
async def _aiter_bytes_then_raise():
for c in partial_chunks:
yield c
raise httpx.ReadError("upstream disconnected")
mock_response.aiter_bytes = _aiter_bytes_then_raise
async def response_coro():
return mock_response
mock_logging_obj = _make_logging_obj()
provider_config = MagicMock()
received = []
with pytest.raises(httpx.ReadError):
async for chunk in _async_streaming(
response=response_coro(),
litellm_logging_obj=mock_logging_obj,
provider_config=provider_config,
):
received.append(chunk)
assert received == partial_chunks
import asyncio
await asyncio.sleep(0)
mock_logging_obj.async_flush_passthrough_collected_chunks.assert_called_once()
call_kwargs = (
mock_logging_obj.async_flush_passthrough_collected_chunks.call_args.kwargs
)
assert call_kwargs["raw_bytes"] == partial_chunks
def test_sync_streaming_flushes_on_normal_completion():
"""Baseline for the sync codepath."""
from litellm.passthrough.main import _sync_streaming
chunks = [b"a", b"b", b"c"]
mock_response = MagicMock(spec=httpx.Response)
def _iter_bytes():
yield from chunks
mock_response.iter_bytes = _iter_bytes
mock_logging_obj = MagicMock()
mock_logging_obj.flush_passthrough_collected_chunks = MagicMock()
provider_config = MagicMock()
# Use a synchronous in-process executor so we can assert immediately.
class _ImmediateExecutor:
def submit(self, fn, *args, **kwargs):
fn(*args, **kwargs)
from unittest.mock import patch
with patch("litellm.utils.executor", _ImmediateExecutor()):
received = list(
_sync_streaming(
response=mock_response,
litellm_logging_obj=mock_logging_obj,
provider_config=provider_config,
)
)
assert received == chunks
mock_logging_obj.flush_passthrough_collected_chunks.assert_called_once()
def test_sync_streaming_flushes_on_early_close():
"""
Sync analog of LIT-2642: closing the generator early must still flush
so per-chunk usage data is not silently dropped.
"""
from litellm.passthrough.main import _sync_streaming
chunks = [b"first", b"second", b"third"]
mock_response = MagicMock(spec=httpx.Response)
def _iter_bytes():
yield from chunks
mock_response.iter_bytes = _iter_bytes
mock_logging_obj = MagicMock()
mock_logging_obj.flush_passthrough_collected_chunks = MagicMock()
provider_config = MagicMock()
class _ImmediateExecutor:
def submit(self, fn, *args, **kwargs):
fn(*args, **kwargs)
from unittest.mock import patch
with patch("litellm.utils.executor", _ImmediateExecutor()):
gen = _sync_streaming(
response=mock_response,
litellm_logging_obj=mock_logging_obj,
provider_config=provider_config,
)
# Consume one chunk, then close — analog of a client disconnect.
first = next(gen)
gen.close()
assert first == chunks[0]
mock_logging_obj.flush_passthrough_collected_chunks.assert_called_once()
call_kwargs = mock_logging_obj.flush_passthrough_collected_chunks.call_args.kwargs
assert call_kwargs["raw_bytes"] == [chunks[0]]

View file

@ -0,0 +1,144 @@
"""
Regression tests for LIT-2642 — interrupted pass-through streams must still
trigger logging so spend is tracked.
`PassThroughStreamingHandler.chunk_processor` collects bytes from the
upstream response and schedules `_route_streaming_logging_to_handler` once
the chunk loop completes. When a FastAPI client disconnects mid-stream,
Starlette calls `aclose()` on the async generator and raises `GeneratorExit`
at the suspended `yield`. The previous `except Exception` branch did not
catch `GeneratorExit`, so the post-loop logging task was never scheduled.
"""
from datetime import datetime
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
import pytest
from litellm.proxy.pass_through_endpoints.streaming_handler import (
PassThroughStreamingHandler,
)
from litellm.types.passthrough_endpoints.pass_through_endpoints import EndpointType
def _make_streaming_response(chunks):
mock = MagicMock(spec=httpx.Response)
mock.status_code = 200
async def _aiter_bytes():
for c in chunks:
yield c
mock.aiter_bytes = _aiter_bytes
return mock
@pytest.mark.asyncio
async def test_chunk_processor_logs_on_normal_completion():
"""Baseline: full consumption schedules logging exactly once."""
chunks = [b"chunk-1", b"chunk-2", b"chunk-3"]
response = _make_streaming_response(chunks)
mock_logging_obj = MagicMock()
mock_passthrough_handler = MagicMock()
with patch.object(
PassThroughStreamingHandler,
"_route_streaming_logging_to_handler",
new=AsyncMock(),
) as mock_route:
received = []
async for chunk in PassThroughStreamingHandler.chunk_processor(
response=response,
request_body={"model": "claude-3-haiku"},
litellm_logging_obj=mock_logging_obj,
endpoint_type=EndpointType.GENERIC,
start_time=datetime.now(),
passthrough_success_handler_obj=mock_passthrough_handler,
url_route="/bedrock/model/claude/invoke-with-response-stream",
):
received.append(chunk)
import asyncio
await asyncio.sleep(0)
assert received == chunks
mock_route.assert_called_once()
call_kwargs = mock_route.call_args.kwargs
assert call_kwargs["raw_bytes"] == chunks
@pytest.mark.asyncio
async def test_chunk_processor_logs_on_client_disconnect():
"""
LIT-2642 regression: closing the generator early (e.g. client
disconnect) must still schedule logging so per-chunk spend data
isn't dropped.
"""
chunks = [b"event-1", b"event-2", b"event-3"]
response = _make_streaming_response(chunks)
mock_logging_obj = MagicMock()
mock_passthrough_handler = MagicMock()
with patch.object(
PassThroughStreamingHandler,
"_route_streaming_logging_to_handler",
new=AsyncMock(),
) as mock_route:
gen = PassThroughStreamingHandler.chunk_processor(
response=response,
request_body={"model": "claude-3-haiku"},
litellm_logging_obj=mock_logging_obj,
endpoint_type=EndpointType.GENERIC,
start_time=datetime.now(),
passthrough_success_handler_obj=mock_passthrough_handler,
url_route="/bedrock/model/claude/invoke-with-response-stream",
)
# Consume one chunk, then close the generator — same path Starlette
# takes when the HTTP client disconnects mid-stream.
first = await gen.__anext__()
await gen.aclose()
import asyncio
await asyncio.sleep(0)
assert first == chunks[0]
mock_route.assert_called_once()
call_kwargs = mock_route.call_args.kwargs
# Only one chunk made it through before disconnect — that is what
# the logging handler must be given so partial usage is captured.
assert call_kwargs["raw_bytes"] == [chunks[0]]
@pytest.mark.asyncio
async def test_chunk_processor_does_not_schedule_logging_when_no_chunks():
"""If no chunks were ever received, don't schedule a no-op logging task."""
response = _make_streaming_response([])
mock_logging_obj = MagicMock()
mock_passthrough_handler = MagicMock()
with patch.object(
PassThroughStreamingHandler,
"_route_streaming_logging_to_handler",
new=AsyncMock(),
) as mock_route:
received = []
async for chunk in PassThroughStreamingHandler.chunk_processor(
response=response,
request_body={"model": "claude-3-haiku"},
litellm_logging_obj=mock_logging_obj,
endpoint_type=EndpointType.GENERIC,
start_time=datetime.now(),
passthrough_success_handler_obj=mock_passthrough_handler,
url_route="/bedrock/model/claude/invoke-with-response-stream",
):
received.append(chunk)
assert received == []
mock_route.assert_not_called()