diff --git a/litellm/responses/streaming_iterator.py b/litellm/responses/streaming_iterator.py index 3004337f9d1..d6d3e2576f6 100644 --- a/litellm/responses/streaming_iterator.py +++ b/litellm/responses/streaming_iterator.py @@ -427,39 +427,45 @@ class BaseResponsesAPIStreamingIterator: openai_types.ResponsesAPIStreamEvents.RESPONSE_INCOMPLETE, openai_types.ResponsesAPIStreamEvents.RESPONSE_FAILED, ): - self.completed_response = openai_responses_api_chunk _response_obj: Final[object] = getattr(openai_responses_api_chunk, "response", None) - _typed_response: Final[ResponsesAPIResponse | None] = ( - ResponsesAPIResponse.model_construct(**_response_obj) # pyright: ignore[reportUnknownArgumentType] # the model_constructed terminal event leaves response as an untyped dict - if isinstance(_response_obj, dict) - else _response_obj - if isinstance(_response_obj, ResponsesAPIResponse) - else None + _estimate_wanted: Final[bool] = _chunk_type in ( + openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED, + openai_types.ResponsesAPIStreamEvents.RESPONSE_INCOMPLETE, ) - if ( - _typed_response is not None - and _chunk_type - in ( - openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED, - openai_types.ResponsesAPIStreamEvents.RESPONSE_INCOMPLETE, + _billed_response: Final[ResponsesAPIResponse | None] = _billed_terminal_response( + _response_obj, + ( + lambda: ( + _estimate_usage_safely( + self.model or "", + self.request_data.get("input"), + self.request_data, + self._generated_content + self._generated_tool_arguments, + ) + if _estimate_wanted + else None + ) + ), + ) + _terminal_chunk: Final = ( + openai_responses_api_chunk + if _billed_response is None or _billed_response is _response_obj + else ( + openai_responses_api_chunk.model_copy(update={"response": _billed_response}) + if issubclass(type(openai_responses_api_chunk), BaseModel) # pyright: ignore[reportUnnecessaryIsInstance] # test stubs use spec'd Mocks whose __class__ reports BaseModel but whose model_copy returns a Mock + else _replace_response(openai_responses_api_chunk, _billed_response) ) - and _typed_response.usage is None - ): - _typed_response.usage = _estimate_usage_safely( - self.model or "", - self.request_data.get("input"), - self.request_data, - self._generated_content + self._generated_tool_arguments, - ) - if _typed_response is not None and _typed_response is not _response_obj: - openai_responses_api_chunk.response = _typed_response # pyright: ignore[reportAttributeAccessIssue] # reached only on the dict path, which only response-carrying terminal events produce - _stamp_responses_usage_cost(_typed_response, self.logging_obj) + ) + self.completed_response = _terminal_chunk + _stamp_responses_usage_cost(_billed_response, self.logging_obj) if _chunk_type == openai_types.ResponsesAPIStreamEvents.RESPONSE_FAILED: self._handle_logging_failed_response() else: self._handle_logging_completed_response() + return _terminal_chunk + return openai_responses_api_chunk return None @@ -688,7 +694,9 @@ class BaseResponsesAPIStreamingIterator: if cache is None: return - cached_response: Final = response_obj.model_dump_json() + cached_response: Final = _dump_json_safely(response_obj) + if cached_response is None: + return if is_async: from litellm.caching.caching_handler import create_cache_write_task @@ -1334,6 +1342,38 @@ def _add_text_like_part_events( ) +def _billed_terminal_response( + response_obj: object, estimate: Callable[[], ResponseAPIUsage | None] | None +) -> ResponsesAPIResponse | None: + if isinstance(response_obj, ResponsesAPIResponse): + return ( + response_obj + if response_obj.usage is not None or estimate is None + else response_obj.model_copy(update={"usage": estimate()}) + ) + if not isinstance(response_obj, dict): + return None + usage: Final[object] = response_obj.get("usage") # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType] # a model_constructed terminal event leaves response as an untyped dict + return ResponsesAPIResponse.model_construct( + **{**response_obj, "usage": usage if usage is not None or estimate is None else estimate()} # pyright: ignore[reportUnknownArgumentType, reportArgumentType] # same untyped dict spread + ) + + +def _replace_response( + event: ResponsesAPIStreamingResponse, response: ResponsesAPIResponse +) -> ResponsesAPIStreamingResponse: + setattr(event, "response", response) + return event + + +def _dump_json_safely(response: BaseModel) -> str | None: + try: + return response.model_dump_json() + except Exception as exc: + verbose_logger.debug("could not serialize completed response for cache: %s", exc) + return None + + def _logging_copy(event: object) -> object: """Hand logging callbacks a copy, so their usage rewrite (Responses shape to chat shape) never reaches the event the caller is iterating. The round trip through ``model_dump`` sidesteps the diff --git a/tests/test_litellm/responses/test_streaming_iterator.py b/tests/test_litellm/responses/test_streaming_iterator.py index 30a7c4faaed..b33cc4e93c2 100644 --- a/tests/test_litellm/responses/test_streaming_iterator.py +++ b/tests/test_litellm/responses/test_streaming_iterator.py @@ -10,6 +10,7 @@ from unittest.mock import Mock, patch import httpx import pytest +from pydantic_core import PydanticSerializationError import litellm from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj @@ -887,10 +888,11 @@ async def test_completed_event_with_a_dict_response_is_typed_and_billed(): request_data={"input": "count these input tokens please"}, ) - async for _ in iterator: - pass + yielded: Final = [chunk async for chunk in iterator] - completed_response: Final = iterator.completed_response.response + terminal_event: Final = iterator.completed_response + assert yielded[-1] is terminal_event + completed_response: Final = terminal_event.response assert isinstance(completed_response, ResponsesAPIResponse) usage: Final = completed_response.usage assert usage is not None @@ -898,3 +900,45 @@ async def test_completed_event_with_a_dict_response_is_typed_and_billed(): assert usage.output_tokens > 0 assert usage.cost == pytest.approx(0.000704) logging_obj._response_cost_calculator.assert_any_call(result=completed_response) + + +def test_billed_terminal_response_keeps_a_response_that_already_has_usage(): + from litellm.responses.streaming_iterator import _billed_terminal_response + + response: Final = _responses_api_response_with_usage() + + assert _billed_terminal_response(response, None) is response + + +def test_billed_terminal_response_copies_when_estimating_and_leaves_the_original_untouched(): + from litellm.responses.streaming_iterator import _billed_terminal_response + + response: Final = _responses_api_response_without_usage() + estimated: Final = ResponseAPIUsage(input_tokens=3, output_tokens=4, total_tokens=7) + + billed: Final = _billed_terminal_response(response, lambda: estimated) + + assert billed is not response + assert billed.usage is estimated + assert response.usage is None + + +def test_persist_completed_response_to_cache_survives_an_unserializable_response(monkeypatch): + bad_response: Final = ResponsesAPIResponse.model_construct(id="r", output=[object()], usage=None) + with pytest.raises(PydanticSerializationError): + bad_response.model_dump_json() + + logging_obj: Final = _logging_obj_stub() + caching_handler: Final = Mock() + caching_handler.request_kwargs = {"stream": True} + logging_obj._llm_caching_handler = caching_handler + iterator: Final = _make_iterator(sse_events=[], logging_obj=logging_obj) + iterator.completed_response = ResponseCompletedEvent.model_construct( + type="response.completed", response=bad_response + ) + cache: Final = Mock() + monkeypatch.setattr(litellm, "cache", cache) + + iterator._persist_completed_response_to_cache(is_async=False) + + cache.add_cache.assert_not_called()