fix(responses): recover streamed output before logging

This commit is contained in:
DanielMaly 2026-06-04 16:11:10 +02:00
parent 9196098e9e
commit be73c28576
2 changed files with 122 additions and 3 deletions

View file

@ -25,6 +25,10 @@ from litellm.litellm_core_utils.llm_response_utils.response_metadata import (
)
from litellm.litellm_core_utils.thread_pool_executor import executor
from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig
from litellm.responses.sse_output_recovery import (
record_output_item_chunk,
record_output_text_chunk,
)
from litellm.responses.utils import ResponsesAPIRequestUtils
from litellm.types.llms.openai import ResponsesAPIStreamEvents
from litellm.types.utils import CallTypes
@ -77,6 +81,8 @@ class BaseResponsesAPIStreamingIterator:
self._completed_response_cache_hit: Optional[bool] = None
self._persist_completed_response_before_logging = True
self._stream_created_time: float = time.time()
self._recovered_output_items: Dict[int, Dict[str, Any]] = {}
self._recovered_text_only_items: Dict[int, Dict[str, Any]] = {}
# track request context for hooks
self.litellm_metadata = litellm_metadata
@ -116,6 +122,33 @@ class BaseResponsesAPIStreamingIterator:
llm_provider=self.custom_llm_provider or "",
)
def _record_recoverable_output_chunk(self, parsed_chunk: Dict[str, Any]) -> None:
event_type = parsed_chunk.get("type")
if event_type == ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE:
record_output_item_chunk(
parsed_chunk=parsed_chunk,
output_items=self._recovered_output_items,
)
elif event_type == ResponsesAPIStreamEvents.OUTPUT_TEXT_DONE:
record_output_text_chunk(
parsed_chunk=parsed_chunk,
output_items=self._recovered_output_items,
text_only_items=self._recovered_text_only_items,
)
def _recover_output_on_completed_response(self, response_obj: Any) -> None:
existing_output = getattr(response_obj, "output", None)
if existing_output:
return
merged_items: Dict[int, Dict[str, Any]] = {**self._recovered_text_only_items}
merged_items.update(self._recovered_output_items)
if not merged_items:
return
recovered_output = [item for _, item in sorted(merged_items.items())]
setattr(response_obj, "output", recovered_output)
def _process_chunk(self, chunk) -> Optional[Any]:
"""Process a single chunk of data from the stream"""
if not chunk:
@ -137,6 +170,7 @@ class BaseResponsesAPIStreamingIterator:
# Format as ResponsesAPIStreamingResponse
if isinstance(parsed_chunk, dict):
self._record_recoverable_output_chunk(parsed_chunk)
if self.responses_api_provider_config is None:
raise ValueError(
"responses_api_provider_config is required to process live streaming chunks"
@ -245,6 +279,13 @@ class BaseResponsesAPIStreamingIterator:
openai_types.ResponsesAPIStreamEvents.RESPONSE_INCOMPLETE,
openai_types.ResponsesAPIStreamEvents.RESPONSE_FAILED,
):
response_obj_for_recovery = getattr(
openai_responses_api_chunk, "response", None
)
if response_obj_for_recovery is not None:
self._recover_output_on_completed_response(
response_obj_for_recovery
)
self.completed_response = openai_responses_api_chunk
# Add cost to usage object if include_cost_in_streaming_usage is True
if (

View file

@ -342,6 +342,86 @@ def test_process_chunk_wraps_encrypted_content_with_model_id():
assert event.item.encrypted_content.endswith(";ciphertext")
def test_process_chunk_recovers_output_from_stream_events_before_logging(monkeypatch):
openai_types = streaming_module._get_openai_response_types()
class _RecoveringConfig:
def transform_streaming_response(self, **kwargs):
parsed_chunk = kwargs["parsed_chunk"]
event_type = parsed_chunk.get("type")
if event_type == openai_types.ResponsesAPIStreamEvents.OUTPUT_TEXT_DONE:
return openai_types.OutputTextDoneEvent(
type=openai_types.ResponsesAPIStreamEvents.OUTPUT_TEXT_DONE,
item_id="msg_1",
output_index=0,
content_index=0,
text="bridge-ok",
)
return openai_types.ResponseCompletedEvent(
type=openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED,
response=ResponsesAPIResponse(
id="resp_live",
created_at=int(datetime.now().timestamp()),
status="completed",
model="test-model",
object="response",
output=[],
usage=openai_types.ResponseAPIUsage(
input_tokens=10,
output_tokens=3,
total_tokens=13,
),
),
)
iterator = ResponsesAPIStreamingIterator(
response=httpx.Response(200),
model="test-model",
responses_api_provider_config=_RecoveringConfig(),
logging_obj=_FakeLoggingObj(),
request_data={"foo": "bar"},
call_type=CallTypes.responses.value,
)
completion_handler = MagicMock()
monkeypatch.setattr(
iterator, "_handle_logging_completed_response", completion_handler
)
iterator._process_chunk(
json.dumps(
{
"type": "response.output_text.done",
"item_id": "msg_1",
"output_index": 0,
"content_index": 0,
"text": "bridge-ok",
}
)
)
event = iterator._process_chunk(
json.dumps(
{
"type": "response.completed",
"response": {
"id": "resp_live",
"output": [],
"usage": {
"input_tokens": 10,
"output_tokens": 3,
"total_tokens": 13,
},
},
}
)
)
assert event.response.output != []
output_item = event.response.output[0]
assert output_item["content"][0]["text"] == "bridge-ok"
assert iterator.completed_response is event
completion_handler.assert_called_once()
def test_process_chunk_completed_response_updates_id_and_usage_cost(monkeypatch):
original_include_cost = litellm.include_cost_in_streaming_usage
litellm.include_cost_in_streaming_usage = True
@ -387,9 +467,7 @@ def test_process_chunk_completed_response_updates_id_and_usage_cost(monkeypatch)
# Chunk must include a top-level "response" key so BaseResponsesAPIStreamingIterator
# runs _update_responses_api_response_id_with_model_id (see streaming_iterator.py).
event = iterator._process_chunk(
json.dumps(
{"type": "response.completed", "response": {"id": "resp_live"}}
)
json.dumps({"type": "response.completed", "response": {"id": "resp_live"}})
)
finally:
litellm.include_cost_in_streaming_usage = original_include_cost