mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
parent
2e561bd04e
commit
1ef034bff6
4 changed files with 534 additions and 46 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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]]
|
||||
|
|
@ -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()
|
||||
Loading…
Add table
Reference in a new issue