litellm/tests/test_litellm/responses/test_streaming_iterator.py

328 lines
11 KiB
Python

"""Regression tests for LIT-4185 — /v1/responses streaming must stamp
completion_start_time on the first chunk so downstream TTFT consumers
(Prometheus, OTEL, SpendLogs completionStartTime) do not fall back to
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,
SyncResponsesAPIStreamingIterator,
)
from litellm.types.llms.openai import (
ResponseCompletedEvent,
ResponsesAPIResponse,
ResponsesAPIStreamEvents,
)
def _sse_event(payload: dict) -> bytes:
return f"data: {json.dumps(payload)}\n\n".encode("utf-8")
def _mock_config() -> Mock:
mock_config = Mock(spec=BaseResponsesAPIConfig)
mock_responses_api_response = Mock(spec=ResponsesAPIResponse)
mock_responses_api_response.id = "resp_ttft"
def _transform(model, parsed_chunk, logging_obj):
evt_type = parsed_chunk.get("type")
if evt_type == "response.completed":
completed = Mock(spec=ResponseCompletedEvent)
completed.type = ResponsesAPIStreamEvents.RESPONSE_COMPLETED
completed.response = mock_responses_api_response
return completed
stub = Mock()
stub.type = evt_type
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(),
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
entire stream and _success_handler_helper_fn falls back to end_time — collapsing
the reported TTFT to full generation time."""
logging_obj = Mock(spec=LiteLLMLoggingObj)
logging_obj.completion_start_time = None
logging_obj.model_call_details = {"litellm_params": {}}
stamped: list[datetime] = []
def _update(*, completion_start_time):
stamped.append(completion_start_time)
logging_obj.completion_start_time = completion_start_time
logging_obj.model_call_details["completion_start_time"] = completion_start_time
logging_obj._update_completion_start_time.side_effect = _update
iterator = _make_iterator(
sse_events=[
_sse_event({"type": "response.created"}),
_sse_event({"type": "response.output_text.delta", "delta": "hi"}),
_sse_event({"type": "response.completed"}),
],
logging_obj=logging_obj,
)
async for _ in iterator:
pass
assert len(stamped) == 1, (
f"Expected exactly one first-chunk stamp; got {len(stamped)}. "
"Later chunks must not re-stamp completion_start_time."
)
assert isinstance(stamped[0], datetime)
@pytest.mark.asyncio
async def test_responses_streaming_does_not_reset_prior_completion_start_time():
"""If `completion_start_time` is already set (e.g. by an outer wrapper), the
iterator must not overwrite it — otherwise TTFT would collapse to
time-to-last-chunk under contention."""
prior = datetime(2020, 1, 1, 0, 0, 0)
logging_obj = Mock(spec=LiteLLMLoggingObj)
logging_obj.completion_start_time = prior
logging_obj.model_call_details = {"litellm_params": {}}
iterator = _make_iterator(
sse_events=[
_sse_event({"type": "response.created"}),
_sse_event({"type": "response.completed"}),
],
logging_obj=logging_obj,
)
async for _ in iterator:
pass
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
def test_stream_cache_write_completes_when_asyncio_run_closes_the_loop(monkeypatch):
"""
Regression test for LIT-6184 on the /v1/responses streaming surface: the
completed-stream cache write was dispatched as a bare fire-and-forget task,
so asyncio.run cancelled it at loop close before the write landed. The
write must survive loop shutdown just like the chat-completions one.
"""
import asyncio
from types import SimpleNamespace
import litellm
from litellm.types.utils import CallTypes
writes = []
class _SlowWriteCache:
async def async_add_cache(self, result, dynamic_cache_object=None, **kwargs):
await asyncio.sleep(0.2)
writes.append(result)
def add_cache(self, *args, **kwargs):
raise AssertionError("sync write must not run on the async path")
caching_handler = SimpleNamespace(
request_kwargs={
"model": "test-model",
"input": "hello",
"stream": True,
"caching": True,
"metadata": None,
"custom_llm_provider": "openai",
},
preset_cache_key="responses-stream-cache-key",
original_function=litellm.aresponses,
dual_cache=None,
_should_store_result_in_cache=lambda original_function, kwargs: True,
)
logging_obj = SimpleNamespace(
model_call_details={"litellm_params": {}},
_llm_caching_handler=caching_handler,
)
iterator = ResponsesAPIStreamingIterator(
response=httpx.Response(200),
model="test-model",
responses_api_provider_config=Mock(spec=BaseResponsesAPIConfig),
logging_obj=logging_obj,
request_data=caching_handler.request_kwargs,
call_type=CallTypes.aresponses.value,
)
iterator.completed_response = ResponseCompletedEvent(
type=ResponsesAPIStreamEvents.RESPONSE_COMPLETED,
response=ResponsesAPIResponse(
id="resp_lit6184",
created_at=int(datetime.now().timestamp()),
status="completed",
model="test-model",
object="response",
output=[],
),
)
monkeypatch.setattr(litellm, "cache", _SlowWriteCache())
async def _short_lived_script():
iterator._persist_completed_response_to_cache(is_async=True)
asyncio.run(_short_lived_script())
assert len(writes) == 1
def test_run_post_success_hooks_does_not_report_generation_time_as_overhead():
"""LIT-5466: the provider call is timed to first byte, so at stream completion the total minus
that duration is token generation, not LiteLLM overhead."""
logging_obj = _logging_obj_stub()
logging_obj.model_call_details = {"litellm_params": {}, "llm_api_duration_ms": 200.0}
logging_obj.caching_details = None
class _CompletedEvent:
def __init__(self) -> None:
self._hidden_params: dict = {}
iterator = _make_iterator(sse_events=[], logging_obj=logging_obj)
iterator.completed_response = _CompletedEvent()
iterator.start_time = datetime(2025, 1, 1, 0, 0, 0)
iterator._run_post_success_hooks(datetime(2025, 1, 1, 0, 0, 10))
assert iterator.completed_response._hidden_params["_response_ms"] == 10000.0
assert "litellm_overhead_time_ms" not in iterator.completed_response._hidden_params