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