mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
Merge pull request #33222 from BerriAI/litellm_fix_stream_reset_empty_200
fix(streaming): surface upstream connection resets instead of empty 200 streams
This commit is contained in:
commit
fba7ac4428
6 changed files with 446 additions and 167 deletions
|
|
@ -2000,99 +2000,7 @@ class CustomStreamWrapper:
|
|||
self.chunks.append(processed_chunk)
|
||||
return processed_chunk
|
||||
except (StopAsyncIteration, StopIteration):
|
||||
if self.sent_last_chunk is True:
|
||||
# log the final chunk with accurate streaming values
|
||||
try:
|
||||
complete_streaming_response = litellm.stream_chunk_builder(
|
||||
chunks=self.chunks,
|
||||
messages=self.messages,
|
||||
logging_obj=self.logging_obj,
|
||||
)
|
||||
except Exception as e:
|
||||
# see sync __next__: a raise from stream_chunk_builder inside this
|
||||
# except handler escapes __anext__ and drops the request from SpendLogs.
|
||||
# Recover best-effort usage from the raw chunks so cost is still tracked
|
||||
verbose_logger.warning(
|
||||
"stream_chunk_builder raised at end-of-stream (%s); logging best-effort usage from chunks.",
|
||||
str(e),
|
||||
)
|
||||
try:
|
||||
complete_streaming_response = self.model_response_creator(
|
||||
chunk={"usage": calculate_total_usage(chunks=self.chunks)}
|
||||
)
|
||||
except Exception:
|
||||
complete_streaming_response = None
|
||||
|
||||
response = self.model_response_creator()
|
||||
if complete_streaming_response is not None:
|
||||
self._propagate_usage_cost_to_hidden_params(complete_streaming_response)
|
||||
|
||||
setattr(
|
||||
response,
|
||||
"usage",
|
||||
getattr(complete_streaming_response, "usage"),
|
||||
)
|
||||
try:
|
||||
_copy = complete_streaming_response.model_copy(deep=True)
|
||||
except RuntimeError:
|
||||
_copy = complete_streaming_response.model_copy()
|
||||
asyncio.create_task(
|
||||
self.async_cache_streaming_response(
|
||||
processed_chunk=_copy,
|
||||
cache_hit=cache_hit,
|
||||
)
|
||||
)
|
||||
# Update hidden_params with final usage from
|
||||
# stream_chunk_builder (see sync __next__ for full comment).
|
||||
if (
|
||||
self.stream_options is None
|
||||
and complete_streaming_response is not None
|
||||
and self._last_returned_hidden_params is not None
|
||||
):
|
||||
final_usage = getattr(complete_streaming_response, "usage", None)
|
||||
if final_usage is not None:
|
||||
self._last_returned_hidden_params["usage"] = final_usage
|
||||
|
||||
if self.sent_stream_usage is False and self.send_stream_usage is True:
|
||||
self.sent_stream_usage = True
|
||||
return response
|
||||
|
||||
_deferred_cb = getattr(
|
||||
self.logging_obj,
|
||||
"_on_deferred_stream_complete",
|
||||
None,
|
||||
)
|
||||
if _deferred_cb is not None:
|
||||
# Proxy has post-call guardrails. Store the assembled
|
||||
# response so the outer streaming consumer
|
||||
# (ProxyLogging.async_post_call_streaming_iterator_hook)
|
||||
# can fire the deferred callback AFTER all guardrail
|
||||
# end-of-stream blocks complete. Scheduling here via
|
||||
# create_task would race with unified_guardrail's
|
||||
# end-of-stream block for short-stream providers.
|
||||
self.logging_obj._deferred_stream_complete_args = ( # type: ignore[attr-defined]
|
||||
complete_streaming_response,
|
||||
cache_hit,
|
||||
)
|
||||
else:
|
||||
# prefer_async_handlers routes CustomLogger to async_success_handler
|
||||
# when consumers use ``async for`` on sync-SDK streams. Legacy string
|
||||
# callbacks still run via executor.submit inside dispatch_success_handlers.
|
||||
asyncio.create_task(
|
||||
self.logging_obj.dispatch_success_handlers(
|
||||
complete_streaming_response,
|
||||
cache_hit=cache_hit,
|
||||
start_time=None,
|
||||
end_time=None,
|
||||
prefer_async_handlers=True,
|
||||
)
|
||||
)
|
||||
|
||||
raise StopAsyncIteration # Re-raise StopIteration
|
||||
else:
|
||||
self.sent_last_chunk = True
|
||||
processed_chunk = self.finish_reason_handler()
|
||||
return processed_chunk
|
||||
return await self._finalize_completed_stream(cache_hit=cache_hit)
|
||||
except httpx.TimeoutException as e: # if httpx read timeout error occues
|
||||
traceback_exception = traceback.format_exc()
|
||||
## ADD DEBUG INFORMATION - E.G. LITELLM REQUEST TIMEOUT
|
||||
|
|
@ -2107,20 +2015,122 @@ class CustomStreamWrapper:
|
|||
# Handle any exceptions that might occur during streaming
|
||||
asyncio.create_task(self.logging_obj.async_failure_handler(e, traceback_exception))
|
||||
self._handle_stream_fallback_error(e)
|
||||
except (httpx.ReadError, httpx.RemoteProtocolError) as e:
|
||||
if self.received_finish_reason is None:
|
||||
self._log_stream_failure_and_raise(e)
|
||||
return await self._finalize_completed_stream(cache_hit=cache_hit)
|
||||
except Exception as e:
|
||||
traceback_exception = traceback.format_exc()
|
||||
if self.logging_obj is not None:
|
||||
self._record_partial_usage_for_failure()
|
||||
## LOGGING
|
||||
threading.Thread(
|
||||
target=self.logging_obj.failure_handler,
|
||||
args=(e, traceback_exception),
|
||||
).start() # log response
|
||||
# Handle any exceptions that might occur during streaming
|
||||
asyncio.create_task(
|
||||
self.logging_obj.async_failure_handler(e, traceback_exception) # type: ignore
|
||||
self._log_stream_failure_and_raise(e)
|
||||
|
||||
async def _finalize_completed_stream(self, cache_hit: bool) -> "ModelResponseStream":
|
||||
if self.sent_last_chunk is True:
|
||||
# log the final chunk with accurate streaming values
|
||||
try:
|
||||
complete_streaming_response = litellm.stream_chunk_builder(
|
||||
chunks=self.chunks,
|
||||
messages=self.messages,
|
||||
logging_obj=self.logging_obj,
|
||||
)
|
||||
self._handle_stream_fallback_error(e)
|
||||
except Exception as e:
|
||||
# see sync __next__: a raise from stream_chunk_builder inside this
|
||||
# except handler escapes __anext__ and drops the request from SpendLogs.
|
||||
# Recover best-effort usage from the raw chunks so cost is still tracked
|
||||
verbose_logger.warning(
|
||||
"stream_chunk_builder raised at end-of-stream (%s); logging best-effort usage from chunks.",
|
||||
str(e),
|
||||
)
|
||||
try:
|
||||
complete_streaming_response = self.model_response_creator(
|
||||
chunk={"usage": calculate_total_usage(chunks=self.chunks)}
|
||||
)
|
||||
except Exception:
|
||||
complete_streaming_response = None
|
||||
|
||||
response = self.model_response_creator()
|
||||
if complete_streaming_response is not None:
|
||||
self._propagate_usage_cost_to_hidden_params(complete_streaming_response)
|
||||
|
||||
setattr(
|
||||
response,
|
||||
"usage",
|
||||
getattr(complete_streaming_response, "usage"),
|
||||
)
|
||||
try:
|
||||
_copy = complete_streaming_response.model_copy(deep=True)
|
||||
except RuntimeError:
|
||||
_copy = complete_streaming_response.model_copy()
|
||||
asyncio.create_task(
|
||||
self.async_cache_streaming_response(
|
||||
processed_chunk=_copy,
|
||||
cache_hit=cache_hit,
|
||||
)
|
||||
)
|
||||
# Update hidden_params with final usage from
|
||||
# stream_chunk_builder (see sync __next__ for full comment).
|
||||
if (
|
||||
self.stream_options is None
|
||||
and complete_streaming_response is not None
|
||||
and self._last_returned_hidden_params is not None
|
||||
):
|
||||
final_usage = getattr(complete_streaming_response, "usage", None)
|
||||
if final_usage is not None:
|
||||
self._last_returned_hidden_params["usage"] = final_usage
|
||||
|
||||
if self.sent_stream_usage is False and self.send_stream_usage is True:
|
||||
self.sent_stream_usage = True
|
||||
return response
|
||||
|
||||
_deferred_cb = getattr(
|
||||
self.logging_obj,
|
||||
"_on_deferred_stream_complete",
|
||||
None,
|
||||
)
|
||||
if _deferred_cb is not None:
|
||||
# Proxy has post-call guardrails. Store the assembled
|
||||
# response so the outer streaming consumer
|
||||
# (ProxyLogging.async_post_call_streaming_iterator_hook)
|
||||
# can fire the deferred callback AFTER all guardrail
|
||||
# end-of-stream blocks complete. Scheduling here via
|
||||
# create_task would race with unified_guardrail's
|
||||
# end-of-stream block for short-stream providers.
|
||||
self.logging_obj._deferred_stream_complete_args = ( # type: ignore[attr-defined]
|
||||
complete_streaming_response,
|
||||
cache_hit,
|
||||
)
|
||||
else:
|
||||
# prefer_async_handlers routes CustomLogger to async_success_handler
|
||||
# when consumers use ``async for`` on sync-SDK streams. Legacy string
|
||||
# callbacks still run via executor.submit inside dispatch_success_handlers.
|
||||
asyncio.create_task(
|
||||
self.logging_obj.dispatch_success_handlers(
|
||||
complete_streaming_response,
|
||||
cache_hit=cache_hit,
|
||||
start_time=None,
|
||||
end_time=None,
|
||||
prefer_async_handlers=True,
|
||||
)
|
||||
)
|
||||
|
||||
raise StopAsyncIteration # Re-raise StopIteration
|
||||
else:
|
||||
self.sent_last_chunk = True
|
||||
processed_chunk = self.finish_reason_handler()
|
||||
return processed_chunk
|
||||
|
||||
def _log_stream_failure_and_raise(self, e: Exception) -> NoReturn:
|
||||
traceback_exception = traceback.format_exc()
|
||||
if self.logging_obj is not None:
|
||||
self._record_partial_usage_for_failure()
|
||||
## LOGGING
|
||||
threading.Thread(
|
||||
target=self.logging_obj.failure_handler,
|
||||
args=(e, traceback_exception),
|
||||
).start() # log response
|
||||
# Handle any exceptions that might occur during streaming
|
||||
asyncio.create_task(
|
||||
self.logging_obj.async_failure_handler(e, traceback_exception) # type: ignore
|
||||
)
|
||||
self._handle_stream_fallback_error(e)
|
||||
|
||||
def _record_partial_usage_for_failure(self) -> None:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -85,29 +85,12 @@ class AiohttpResponseStream(httpx.AsyncByteStream):
|
|||
try:
|
||||
async for chunk in self._aiohttp_response.content.iter_chunked(self.CHUNK_SIZE):
|
||||
yield chunk
|
||||
except (
|
||||
aiohttp.ClientPayloadError,
|
||||
aiohttp.client_exceptions.ClientPayloadError,
|
||||
) as e:
|
||||
# Handle incomplete transfers more gracefully
|
||||
# Log the error but don't re-raise if we've already yielded some data
|
||||
verbose_logger.debug(f"Transfer incomplete, but continuing: {e}")
|
||||
# If the error is due to incomplete transfer encoding, we can still
|
||||
# return what we've received so far, similar to how httpx handles it
|
||||
return
|
||||
except RuntimeError as e:
|
||||
# Some providers (e.g., SSE streams) may close the connection
|
||||
# causing aiohttp StreamReader to raise a generic RuntimeError
|
||||
# with message "Connection closed.". Treat this as a graceful
|
||||
# end-of-stream so downstream consumers don't error.
|
||||
if "Connection closed" in str(e):
|
||||
verbose_logger.debug("Upstream closed streaming connection; ending iterator gracefully")
|
||||
return
|
||||
raise
|
||||
if "Connection closed" not in str(e):
|
||||
raise
|
||||
raise httpx.ReadError(str(e)) from e
|
||||
except aiohttp.http_exceptions.TransferEncodingError as e:
|
||||
# Handle transfer encoding errors gracefully
|
||||
verbose_logger.debug(f"Transfer encoding error, but continuing: {e}")
|
||||
return
|
||||
raise httpx.ReadError(str(e)) from e
|
||||
except Exception:
|
||||
# For other exceptions, use the normal mapping
|
||||
with map_aiohttp_exceptions():
|
||||
|
|
|
|||
|
|
@ -718,6 +718,12 @@ class ResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
except StopAsyncIteration:
|
||||
# Normal end of stream - don't log as failure
|
||||
raise
|
||||
except (httpx.ReadError, httpx.RemoteProtocolError) as e:
|
||||
self.finished = True
|
||||
if self.completed_response is None:
|
||||
self._handle_failure(e)
|
||||
raise
|
||||
raise StopAsyncIteration from e
|
||||
except httpx.HTTPError as e:
|
||||
# Handle HTTP errors
|
||||
self.finished = True
|
||||
|
|
@ -794,6 +800,12 @@ class SyncResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
except StopIteration:
|
||||
# Normal end of stream - don't log as failure
|
||||
raise
|
||||
except (httpx.ReadError, httpx.RemoteProtocolError) as e:
|
||||
self.finished = True
|
||||
if self.completed_response is None:
|
||||
self._handle_failure(e)
|
||||
raise
|
||||
raise StopIteration from e
|
||||
except httpx.HTTPError as e:
|
||||
# Handle HTTP errors
|
||||
self.finished = True
|
||||
|
|
|
|||
|
|
@ -3243,3 +3243,115 @@ async def test_stream_chunk_builder_raise_and_usage_recovery_failure_does_not_cr
|
|||
chunks = [c async for c in response]
|
||||
|
||||
assert len(chunks) > 0
|
||||
|
||||
|
||||
class TransportErrorAfterChunksIterator:
|
||||
"""Yields the given chunks, then raises the given exception once, then StopAsyncIteration."""
|
||||
|
||||
def __init__(self, model_responses, exception):
|
||||
self.model_responses = model_responses
|
||||
self.exception = exception
|
||||
self.index = 0
|
||||
self.raised = False
|
||||
|
||||
def __aiter__(self):
|
||||
return self
|
||||
|
||||
async def __anext__(self):
|
||||
if self.index < len(self.model_responses):
|
||||
chunk = self.model_responses[self.index]
|
||||
self.index += 1
|
||||
return chunk
|
||||
if not self.raised:
|
||||
self.raised = True
|
||||
raise self.exception
|
||||
raise StopAsyncIteration
|
||||
|
||||
|
||||
def _reset_test_chunk(content: Optional[str] = None, finish_reason: Optional[str] = None) -> ModelResponseStream:
|
||||
return ModelResponseStream(
|
||||
id="chatcmpl-reset-test",
|
||||
created=1783458104,
|
||||
model="stub-model",
|
||||
object="chat.completion.chunk",
|
||||
choices=[
|
||||
StreamingChoices(
|
||||
index=0,
|
||||
delta=Delta(content=content),
|
||||
finish_reason=finish_reason,
|
||||
)
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_transport_read_error_after_finish_reason_ends_stream_gracefully(
|
||||
logging_obj: Logging,
|
||||
):
|
||||
"""A trailing connection reset after the provider's finish chunk must not fail the stream."""
|
||||
import httpx
|
||||
|
||||
completion_stream = TransportErrorAfterChunksIterator(
|
||||
model_responses=[
|
||||
_reset_test_chunk(content="Hello"),
|
||||
_reset_test_chunk(finish_reason="stop"),
|
||||
],
|
||||
exception=httpx.ReadError("Response payload is not completed"),
|
||||
)
|
||||
response = CustomStreamWrapper(
|
||||
completion_stream=completion_stream,
|
||||
model="hosted_vllm/stub-model",
|
||||
custom_llm_provider="hosted_vllm",
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
chunks = [chunk async for chunk in response]
|
||||
|
||||
finish_reasons = [
|
||||
chunk.choices[0].finish_reason
|
||||
for chunk in chunks
|
||||
if chunk.choices and chunk.choices[0].finish_reason
|
||||
]
|
||||
contents = [
|
||||
chunk.choices[0].delta.content
|
||||
for chunk in chunks
|
||||
if chunk.choices and chunk.choices[0].delta and chunk.choices[0].delta.content
|
||||
]
|
||||
assert finish_reasons == ["stop"]
|
||||
assert contents == ["Hello"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_transport_read_error_before_finish_reason_raises(logging_obj: Logging):
|
||||
"""A connection reset before any finish chunk must surface, never end as a clean stop.
|
||||
|
||||
Regression test for silent empty/truncated HTTP 200 streams: the aiohttp
|
||||
transport used to swallow mid-stream connection resets, so the wrapper saw a
|
||||
clean end-of-stream and fabricated finish_reason "stop".
|
||||
"""
|
||||
import httpx
|
||||
|
||||
from litellm.exceptions import MidStreamFallbackError
|
||||
|
||||
completion_stream = TransportErrorAfterChunksIterator(
|
||||
model_responses=[_reset_test_chunk(content="Hel")],
|
||||
exception=httpx.ReadError("Response payload is not completed"),
|
||||
)
|
||||
response = CustomStreamWrapper(
|
||||
completion_stream=completion_stream,
|
||||
model="hosted_vllm/stub-model",
|
||||
custom_llm_provider="hosted_vllm",
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
received = []
|
||||
with pytest.raises(MidStreamFallbackError):
|
||||
async for chunk in response:
|
||||
received.append(chunk)
|
||||
|
||||
fabricated_finish_reasons = [
|
||||
chunk.choices[0].finish_reason
|
||||
for chunk in received
|
||||
if chunk.choices and chunk.choices[0].finish_reason
|
||||
]
|
||||
assert fabricated_finish_reasons == []
|
||||
|
|
|
|||
|
|
@ -81,7 +81,7 @@ class MockContent:
|
|||
def __init__(self, chunks=None, exception_to_raise=None, exception_at_chunk=None):
|
||||
self.chunks = chunks or [b"chunk1", b"chunk2", b"chunk3"]
|
||||
self.exception_to_raise = exception_to_raise
|
||||
self.exception_at_chunk = exception_at_chunk or (len(self.chunks) - 1)
|
||||
self.exception_at_chunk = exception_at_chunk if exception_at_chunk is not None else (len(self.chunks) - 1)
|
||||
self.chunk_index = 0
|
||||
|
||||
async def iter_chunked(self, chunk_size):
|
||||
|
|
@ -107,15 +107,11 @@ async def test_aiohttp_response_stream_normal_flow():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_transfer_encoding_error_no_httpx_read_error():
|
||||
"""Test that TransferEncodingError doesn't get converted to httpx.ReadError"""
|
||||
|
||||
# Create a TransferEncodingError wrapped in ClientPayloadError (like in real scenarios)
|
||||
async def test_client_payload_error_mid_stream_raises_read_error():
|
||||
"""A connection reset mid-body must surface as httpx.ReadError, not truncate silently"""
|
||||
transfer_error = aiohttp.http_exceptions.TransferEncodingError(
|
||||
message="400, message: Not enough data for satisfy transfer length header."
|
||||
)
|
||||
|
||||
# Wrap it in ClientPayloadError as aiohttp does
|
||||
client_payload_error = aiohttp.ClientPayloadError(
|
||||
"Response payload is not completed"
|
||||
)
|
||||
|
|
@ -124,47 +120,100 @@ async def test_transfer_encoding_error_no_httpx_read_error():
|
|||
mock_response = MockAiohttpResponse(
|
||||
content_chunks=[b"chunk1", b"chunk2", b"chunk3"],
|
||||
exception_to_raise=client_payload_error,
|
||||
exception_at_chunk=1, # Error occurs at chunk 1
|
||||
exception_at_chunk=1,
|
||||
)
|
||||
|
||||
stream = AiohttpResponseStream(mock_response) # type: ignore
|
||||
received_chunks = []
|
||||
|
||||
# This should NOT raise httpx.ReadError or any other exception
|
||||
# It should handle the error gracefully and just return what was received
|
||||
async for chunk in stream:
|
||||
received_chunks.append(chunk)
|
||||
print(f"received_chunks: {received_chunks}")
|
||||
with pytest.raises(httpx.ReadError):
|
||||
async for chunk in stream:
|
||||
received_chunks.append(chunk)
|
||||
|
||||
# Should have received the first chunk before the error
|
||||
assert received_chunks == [b"chunk1"]
|
||||
assert len(received_chunks) == 1
|
||||
assert mock_response.closed is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_client_payload_error_graceful_handling():
|
||||
"""Test that ClientPayloadError is handled gracefully without stacktrace"""
|
||||
# Create a ClientPayloadError directly
|
||||
async def test_client_payload_error_before_first_chunk_raises_read_error():
|
||||
"""A connection reset before any body byte must surface, not yield an empty 200 body"""
|
||||
client_error = aiohttp.client_exceptions.ClientPayloadError(
|
||||
"Response payload is not completed"
|
||||
)
|
||||
|
||||
mock_response = MockAiohttpResponse(
|
||||
content_chunks=[b"data1", b"data2", b"data3"],
|
||||
content_chunks=[b"data1", b"data2"],
|
||||
exception_to_raise=client_error,
|
||||
exception_at_chunk=2, # Error occurs at chunk 2
|
||||
exception_at_chunk=0,
|
||||
)
|
||||
|
||||
stream = AiohttpResponseStream(mock_response) # type: ignore
|
||||
received_chunks = []
|
||||
|
||||
# This should handle the error gracefully without raising
|
||||
async for chunk in stream:
|
||||
received_chunks.append(chunk)
|
||||
with pytest.raises(httpx.ReadError):
|
||||
async for chunk in stream:
|
||||
received_chunks.append(chunk)
|
||||
|
||||
# Should have received chunks before the error
|
||||
assert received_chunks == [b"data1", b"data2"]
|
||||
assert len(received_chunks) == 2
|
||||
assert received_chunks == []
|
||||
assert mock_response.closed is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_connection_closed_runtime_error_raises_read_error():
|
||||
"""aiohttp's bare RuntimeError('Connection closed.') must surface as httpx.ReadError"""
|
||||
mock_response = MockAiohttpResponse(
|
||||
content_chunks=[b"data1", b"data2"],
|
||||
exception_to_raise=RuntimeError("Connection closed."),
|
||||
exception_at_chunk=1,
|
||||
)
|
||||
|
||||
stream = AiohttpResponseStream(mock_response) # type: ignore
|
||||
received_chunks = []
|
||||
|
||||
with pytest.raises(httpx.ReadError):
|
||||
async for chunk in stream:
|
||||
received_chunks.append(chunk)
|
||||
|
||||
assert received_chunks == [b"data1"]
|
||||
assert mock_response.closed is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unrelated_runtime_error_propagates_unmapped():
|
||||
"""RuntimeErrors other than 'Connection closed' must propagate untouched"""
|
||||
mock_response = MockAiohttpResponse(
|
||||
content_chunks=[b"data1"],
|
||||
exception_to_raise=RuntimeError("something else broke"),
|
||||
exception_at_chunk=0,
|
||||
)
|
||||
|
||||
stream = AiohttpResponseStream(mock_response) # type: ignore
|
||||
|
||||
with pytest.raises(RuntimeError, match="something else broke"):
|
||||
async for _ in stream:
|
||||
pass
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_transfer_encoding_error_raises_read_error():
|
||||
"""A raw TransferEncodingError mid-body must surface as httpx.ReadError"""
|
||||
mock_response = MockAiohttpResponse(
|
||||
content_chunks=[b"data1", b"data2"],
|
||||
exception_to_raise=aiohttp.http_exceptions.TransferEncodingError(
|
||||
message="Not enough data to satisfy transfer length header."
|
||||
),
|
||||
exception_at_chunk=1,
|
||||
)
|
||||
|
||||
stream = AiohttpResponseStream(mock_response) # type: ignore
|
||||
received_chunks = []
|
||||
|
||||
with pytest.raises(httpx.ReadError):
|
||||
async for chunk in stream:
|
||||
received_chunks.append(chunk)
|
||||
|
||||
assert received_chunks == [b"data1"]
|
||||
assert mock_response.closed is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
|
|
@ -5,13 +5,18 @@ completion_start_time = end_time."""
|
|||
|
||||
import json
|
||||
from datetime import datetime
|
||||
from typing import Optional
|
||||
from unittest.mock import Mock
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig
|
||||
from litellm.responses.streaming_iterator import ResponsesAPIStreamingIterator
|
||||
from litellm.responses.streaming_iterator import (
|
||||
ResponsesAPIStreamingIterator,
|
||||
SyncResponsesAPIStreamingIterator,
|
||||
)
|
||||
from litellm.types.llms.openai import (
|
||||
ResponseCompletedEvent,
|
||||
ResponsesAPIResponse,
|
||||
|
|
@ -23,19 +28,7 @@ def _sse_event(payload: dict) -> bytes:
|
|||
return f"data: {json.dumps(payload)}\n\n".encode("utf-8")
|
||||
|
||||
|
||||
def _make_iterator(
|
||||
*,
|
||||
sse_events: list[bytes],
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> ResponsesAPIStreamingIterator:
|
||||
async def aiter_bytes():
|
||||
for evt in sse_events:
|
||||
yield evt
|
||||
|
||||
mock_response = Mock()
|
||||
mock_response.headers = {}
|
||||
mock_response.aiter_bytes = aiter_bytes
|
||||
|
||||
def _mock_config() -> Mock:
|
||||
mock_config = Mock(spec=BaseResponsesAPIConfig)
|
||||
mock_responses_api_response = Mock(spec=ResponsesAPIResponse)
|
||||
mock_responses_api_response.id = "resp_ttft"
|
||||
|
|
@ -52,17 +45,68 @@ def _make_iterator(
|
|||
return stub
|
||||
|
||||
mock_config.transform_streaming_response.side_effect = _transform
|
||||
return mock_config
|
||||
|
||||
|
||||
def _make_iterator(
|
||||
*,
|
||||
sse_events: list[bytes],
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
trailing_error: Optional[Exception] = None,
|
||||
) -> ResponsesAPIStreamingIterator:
|
||||
async def aiter_bytes():
|
||||
for evt in sse_events:
|
||||
yield evt
|
||||
if trailing_error is not None:
|
||||
raise trailing_error
|
||||
|
||||
mock_response = Mock()
|
||||
mock_response.headers = {}
|
||||
mock_response.aiter_bytes = aiter_bytes
|
||||
|
||||
return ResponsesAPIStreamingIterator(
|
||||
response=mock_response,
|
||||
model="gpt-4o-mini",
|
||||
responses_api_provider_config=mock_config,
|
||||
responses_api_provider_config=_mock_config(),
|
||||
logging_obj=logging_obj,
|
||||
litellm_metadata={},
|
||||
custom_llm_provider="openai",
|
||||
)
|
||||
|
||||
|
||||
def _make_sync_iterator(
|
||||
*,
|
||||
sse_events: list[bytes],
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
trailing_error: Optional[Exception] = None,
|
||||
) -> SyncResponsesAPIStreamingIterator:
|
||||
def iter_bytes():
|
||||
for evt in sse_events:
|
||||
yield evt
|
||||
if trailing_error is not None:
|
||||
raise trailing_error
|
||||
|
||||
mock_response = Mock()
|
||||
mock_response.headers = {}
|
||||
mock_response.iter_bytes = iter_bytes
|
||||
|
||||
return SyncResponsesAPIStreamingIterator(
|
||||
response=mock_response,
|
||||
model="gpt-4o-mini",
|
||||
responses_api_provider_config=_mock_config(),
|
||||
logging_obj=logging_obj,
|
||||
litellm_metadata={},
|
||||
custom_llm_provider="openai",
|
||||
)
|
||||
|
||||
|
||||
def _logging_obj_stub() -> Mock:
|
||||
logging_obj = Mock(spec=LiteLLMLoggingObj)
|
||||
logging_obj.completion_start_time = None
|
||||
logging_obj.model_call_details = {"litellm_params": {}}
|
||||
return logging_obj
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_responses_streaming_stamps_completion_start_time_on_first_chunk():
|
||||
"""Without the fix, `logging_obj.completion_start_time` stays None across the
|
||||
|
|
@ -122,3 +166,72 @@ async def test_responses_streaming_does_not_reset_prior_completion_start_time():
|
|||
|
||||
logging_obj._update_completion_start_time.assert_not_called()
|
||||
assert logging_obj.completion_start_time == prior
|
||||
|
||||
|
||||
_COMPLETE_STREAM_EVENTS = [
|
||||
_sse_event({"type": "response.created"}),
|
||||
_sse_event({"type": "response.output_text.delta", "delta": "hi"}),
|
||||
_sse_event({"type": "response.completed"}),
|
||||
]
|
||||
|
||||
_TRAILING_ERRORS = [
|
||||
httpx.ReadError("Response payload is not completed"),
|
||||
httpx.RemoteProtocolError("peer closed connection without sending complete message body"),
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("trailing_error", _TRAILING_ERRORS, ids=type)
|
||||
async def test_transport_error_after_completed_event_ends_stream_cleanly(trailing_error):
|
||||
"""A sloppy connection close after `response.completed` must not turn a
|
||||
complete stream into an error (regression guard for the transport no longer
|
||||
swallowing ClientPayloadError/TransferEncodingError)."""
|
||||
iterator = _make_iterator(
|
||||
sse_events=_COMPLETE_STREAM_EVENTS,
|
||||
logging_obj=_logging_obj_stub(),
|
||||
trailing_error=trailing_error,
|
||||
)
|
||||
|
||||
seen = [event.type async for event in iterator]
|
||||
|
||||
assert ResponsesAPIStreamEvents.RESPONSE_COMPLETED in seen
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_transport_error_before_completed_event_raises():
|
||||
"""A connection lost before any terminal event is a real failure and must
|
||||
surface, not end the stream as if it completed."""
|
||||
iterator = _make_iterator(
|
||||
sse_events=_COMPLETE_STREAM_EVENTS[:-1],
|
||||
logging_obj=_logging_obj_stub(),
|
||||
trailing_error=httpx.ReadError("Response payload is not completed"),
|
||||
)
|
||||
|
||||
with pytest.raises(httpx.ReadError):
|
||||
async for _ in iterator:
|
||||
pass
|
||||
|
||||
|
||||
@pytest.mark.parametrize("trailing_error", _TRAILING_ERRORS, ids=type)
|
||||
def test_sync_transport_error_after_completed_event_ends_stream_cleanly(trailing_error):
|
||||
iterator = _make_sync_iterator(
|
||||
sse_events=_COMPLETE_STREAM_EVENTS,
|
||||
logging_obj=_logging_obj_stub(),
|
||||
trailing_error=trailing_error,
|
||||
)
|
||||
|
||||
seen = [event.type for event in iterator]
|
||||
|
||||
assert ResponsesAPIStreamEvents.RESPONSE_COMPLETED in seen
|
||||
|
||||
|
||||
def test_sync_transport_error_before_completed_event_raises():
|
||||
iterator = _make_sync_iterator(
|
||||
sse_events=_COMPLETE_STREAM_EVENTS[:-1],
|
||||
logging_obj=_logging_obj_stub(),
|
||||
trailing_error=httpx.ReadError("Response payload is not completed"),
|
||||
)
|
||||
|
||||
with pytest.raises(httpx.ReadError):
|
||||
for _ in iterator:
|
||||
pass
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue