mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix(responses): recover streamed output before logging
This commit is contained in:
parent
9196098e9e
commit
be73c28576
2 changed files with 122 additions and 3 deletions
|
|
@ -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 (
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue