mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Merge pull request #19390 from BerriAI/litellm_consistent_id_streaming_responses
Fix: ID mismatch between text-start and text-delta
This commit is contained in:
commit
11dbae85d1
2 changed files with 243 additions and 8 deletions
|
|
@ -81,6 +81,8 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
|
|||
Union[ModelResponse, TextCompletionResponse]
|
||||
] = None
|
||||
self.final_text: str = ""
|
||||
self._cached_item_id: Optional[str] = None
|
||||
self._cached_response_id: Optional[str] = None
|
||||
self._pending_tool_events: List[BaseLiteLLMOpenAIResponseObject] = []
|
||||
self._tool_output_index_by_call_id: dict[str, int] = {}
|
||||
self._tool_args_by_call_id: dict[str, str] = {}
|
||||
|
|
@ -307,12 +309,15 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
|
|||
)
|
||||
|
||||
def create_output_item_added_event(self) -> OutputItemAddedEvent:
|
||||
if self._cached_item_id is None:
|
||||
self._cached_item_id = f"msg_{str(uuid.uuid4())}"
|
||||
|
||||
return OutputItemAddedEvent(
|
||||
type=ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED,
|
||||
output_index=0,
|
||||
item=BaseLiteLLMOpenAIResponseObject(
|
||||
**{
|
||||
"id": f"msg_{str(uuid.uuid4())}",
|
||||
"id": self._cached_item_id,
|
||||
"type": "message",
|
||||
"status": "in_progress",
|
||||
"role": "assistant",
|
||||
|
|
@ -322,9 +327,12 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
|
|||
)
|
||||
|
||||
def create_content_part_added_event(self) -> ContentPartAddedEvent:
|
||||
if self._cached_item_id is None:
|
||||
self._cached_item_id = f"msg_{str(uuid.uuid4())}"
|
||||
|
||||
return ContentPartAddedEvent(
|
||||
type=ResponsesAPIStreamEvents.CONTENT_PART_ADDED,
|
||||
item_id=f"msg_{str(uuid.uuid4())}",
|
||||
item_id=self._cached_item_id,
|
||||
output_index=0,
|
||||
content_index=0,
|
||||
part=BaseLiteLLMOpenAIResponseObject(
|
||||
|
|
@ -346,9 +354,12 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
|
|||
def create_output_text_done_event(
|
||||
self, litellm_complete_object: ModelResponse
|
||||
) -> OutputTextDoneEvent:
|
||||
if self._cached_item_id is None:
|
||||
self._cached_item_id = f"msg_{str(uuid.uuid4())}"
|
||||
|
||||
return OutputTextDoneEvent(
|
||||
type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DONE,
|
||||
item_id=f"msg_{str(uuid.uuid4())}",
|
||||
item_id=self._cached_item_id,
|
||||
output_index=0,
|
||||
content_index=0,
|
||||
text=getattr(litellm_complete_object.choices[0].message, "content", "") # type: ignore
|
||||
|
|
@ -358,6 +369,8 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
|
|||
def create_output_content_part_done_event(
|
||||
self, litellm_complete_object: ModelResponse
|
||||
) -> ContentPartDoneEvent:
|
||||
if self._cached_item_id is None:
|
||||
self._cached_item_id = f"msg_{str(uuid.uuid4())}"
|
||||
|
||||
text = getattr(litellm_complete_object.choices[0].message, "content", "") or "" # type: ignore
|
||||
reasoning_content = getattr(litellm_complete_object.choices[0].message, "reasoning_content", "") or "" # type: ignore
|
||||
|
|
@ -383,7 +396,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
|
|||
|
||||
return ContentPartDoneEvent(
|
||||
type=ResponsesAPIStreamEvents.CONTENT_PART_DONE,
|
||||
item_id=f"msg_{str(uuid.uuid4())}",
|
||||
item_id=self._cached_item_id,
|
||||
output_index=0,
|
||||
content_index=0,
|
||||
part=part,
|
||||
|
|
@ -392,6 +405,9 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
|
|||
def create_output_item_done_event(
|
||||
self, litellm_complete_object: ModelResponse
|
||||
) -> OutputItemDoneEvent:
|
||||
if self._cached_item_id is None:
|
||||
self._cached_item_id = f"msg_{str(uuid.uuid4())}"
|
||||
|
||||
text = self.litellm_model_response.choices[0].message.content or "" # type: ignore
|
||||
annotations = getattr(self.litellm_model_response.choices[0].message, "annotations", None) # type: ignore
|
||||
|
||||
|
|
@ -404,7 +420,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
|
|||
sequence_number=1,
|
||||
item=BaseLiteLLMOpenAIResponseObject(
|
||||
**{
|
||||
"id": f"msg_{str(uuid.uuid4())}",
|
||||
"id": self._cached_item_id,
|
||||
"status": "completed",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
|
|
@ -576,6 +592,10 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
|
|||
and the ReasoningSummaryTextDeltaEvent, which is used by the responses API to emit reasoning content.
|
||||
It also handles emitting annotation.added events when annotations are detected in the chunk.
|
||||
"""
|
||||
if self._cached_item_id is None and chunk.id:
|
||||
self._cached_item_id = chunk.id
|
||||
item_id = self._cached_item_id or chunk.id
|
||||
|
||||
# Check if this chunk has annotations first (before processing text/reasoning)
|
||||
# This ensures we detect and queue annotation events from the annotation chunk
|
||||
if chunk.choices and hasattr(chunk.choices[0].delta, "annotations"):
|
||||
|
|
@ -593,7 +613,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
|
|||
annotation_dict = annotation.model_dump() if hasattr(annotation, 'model_dump') else dict(annotation)
|
||||
event = OutputTextAnnotationAddedEvent(
|
||||
type=ResponsesAPIStreamEvents.OUTPUT_TEXT_ANNOTATION_ADDED,
|
||||
item_id=chunk.id,
|
||||
item_id=item_id,
|
||||
output_index=0,
|
||||
content_index=0,
|
||||
annotation_index=idx,
|
||||
|
|
@ -620,7 +640,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
|
|||
if delta_content:
|
||||
return OutputTextDeltaEvent(
|
||||
type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA,
|
||||
item_id=chunk.id,
|
||||
item_id=item_id,
|
||||
output_index=0,
|
||||
content_index=0,
|
||||
delta=delta_content,
|
||||
|
|
|
|||
|
|
@ -1351,4 +1351,219 @@ class TestUsageTransformation:
|
|||
assert response_usage.output_tokens == 27
|
||||
assert response_usage.total_tokens == 36
|
||||
assert response_usage.input_tokens_details is None
|
||||
assert response_usage.output_tokens_details is None
|
||||
assert response_usage.output_tokens_details is None
|
||||
|
||||
|
||||
class TestStreamingIDConsistency:
|
||||
"""Test cases for consistent IDs across streaming events (issue #14962)"""
|
||||
|
||||
def test_streaming_iterator_uses_consistent_item_ids(self):
|
||||
"""
|
||||
Test that all streaming events use the same item_id throughout the stream.
|
||||
This fixes the issue where text-start, text-delta, and text-end events
|
||||
had different IDs, breaking SDK text accumulation.
|
||||
|
||||
Reproduces: https://github.com/BerriAI/litellm/issues/14962
|
||||
"""
|
||||
from unittest.mock import Mock
|
||||
|
||||
import litellm
|
||||
from litellm.responses.litellm_completion_transformation.streaming_iterator import (
|
||||
LiteLLMCompletionStreamingIterator,
|
||||
)
|
||||
from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices
|
||||
|
||||
# Create a mock stream wrapper
|
||||
mock_stream_wrapper = Mock(spec=litellm.CustomStreamWrapper)
|
||||
mock_logging_obj = Mock()
|
||||
mock_stream_wrapper.logging_obj = mock_logging_obj
|
||||
|
||||
# Create the streaming iterator
|
||||
iterator = LiteLLMCompletionStreamingIterator(
|
||||
model="gemini/gemini-2.5-flash-lite",
|
||||
litellm_custom_stream_wrapper=mock_stream_wrapper,
|
||||
request_input="Say Hello World",
|
||||
responses_api_request={},
|
||||
custom_llm_provider="gemini",
|
||||
)
|
||||
|
||||
# Simulate streaming chunks with different IDs (as Gemini does)
|
||||
chunk1 = ModelResponseStream(
|
||||
id="chatcmpl-first-id",
|
||||
choices=[
|
||||
StreamingChoices(
|
||||
index=0,
|
||||
delta=Delta(content="Hello", role="assistant"),
|
||||
finish_reason=None,
|
||||
)
|
||||
],
|
||||
created=1234567890,
|
||||
model="gemini-2.5-flash-lite",
|
||||
object="chat.completion.chunk",
|
||||
)
|
||||
|
||||
chunk2 = ModelResponseStream(
|
||||
id="chatcmpl-second-id", # Different ID from chunk1
|
||||
choices=[
|
||||
StreamingChoices(
|
||||
index=0,
|
||||
delta=Delta(content=" World", role=None),
|
||||
finish_reason=None,
|
||||
)
|
||||
],
|
||||
created=1234567890,
|
||||
model="gemini-2.5-flash-lite",
|
||||
object="chat.completion.chunk",
|
||||
)
|
||||
|
||||
chunk3 = ModelResponseStream(
|
||||
id="chatcmpl-third-id", # Different ID from chunk1 and chunk2
|
||||
choices=[
|
||||
StreamingChoices(
|
||||
index=0,
|
||||
delta=Delta(content="", role=None),
|
||||
finish_reason="stop",
|
||||
)
|
||||
],
|
||||
created=1234567890,
|
||||
model="gemini-2.5-flash-lite",
|
||||
object="chat.completion.chunk",
|
||||
)
|
||||
|
||||
# Transform chunks to response API events
|
||||
event1 = iterator._transform_chat_completion_chunk_to_response_api_chunk(chunk1)
|
||||
event2 = iterator._transform_chat_completion_chunk_to_response_api_chunk(chunk2)
|
||||
event3 = iterator._transform_chat_completion_chunk_to_response_api_chunk(chunk3)
|
||||
|
||||
# Assert: All events should use the same item_id (from the first chunk)
|
||||
assert event1 is not None, "First event should not be None"
|
||||
assert event2 is not None, "Second event should not be None"
|
||||
|
||||
# Extract item_ids from events
|
||||
item_id_1 = getattr(event1, "item_id", None)
|
||||
item_id_2 = getattr(event2, "item_id", None)
|
||||
|
||||
assert item_id_1 is not None, "First event should have an item_id"
|
||||
assert item_id_2 is not None, "Second event should have an item_id"
|
||||
|
||||
# The critical assertion: IDs should match across all events
|
||||
assert item_id_1 == item_id_2, (
|
||||
f"Item IDs should be consistent across streaming events. "
|
||||
f"Got {item_id_1} and {item_id_2}. "
|
||||
f"This breaks SDK text accumulation (issue #14962)."
|
||||
)
|
||||
|
||||
# Verify the cached ID is set and matches
|
||||
assert iterator._cached_item_id is not None, "Iterator should cache the item_id"
|
||||
assert iterator._cached_item_id == item_id_1, "Cached ID should match event IDs"
|
||||
assert iterator._cached_item_id == "chatcmpl-first-id", "Should use the first chunk's ID"
|
||||
|
||||
def test_streaming_iterator_initial_events_use_cached_id(self):
|
||||
"""
|
||||
Test that initial events (output_item_added, content_part_added) also use the cached ID.
|
||||
"""
|
||||
from unittest.mock import Mock
|
||||
|
||||
import litellm
|
||||
from litellm.responses.litellm_completion_transformation.streaming_iterator import (
|
||||
LiteLLMCompletionStreamingIterator,
|
||||
)
|
||||
|
||||
# Create a mock stream wrapper
|
||||
mock_stream_wrapper = Mock(spec=litellm.CustomStreamWrapper)
|
||||
mock_logging_obj = Mock()
|
||||
mock_stream_wrapper.logging_obj = mock_logging_obj
|
||||
|
||||
# Create the streaming iterator
|
||||
iterator = LiteLLMCompletionStreamingIterator(
|
||||
model="gemini/gemini-2.5-flash-lite",
|
||||
litellm_custom_stream_wrapper=mock_stream_wrapper,
|
||||
request_input="Test",
|
||||
responses_api_request={},
|
||||
)
|
||||
|
||||
# Create initial events
|
||||
output_item_event = iterator.create_output_item_added_event()
|
||||
content_part_event = iterator.create_content_part_added_event()
|
||||
|
||||
# Extract IDs
|
||||
output_item_id = getattr(output_item_event.item, "id", None)
|
||||
content_part_id = getattr(content_part_event, "item_id", None)
|
||||
|
||||
# Assert: Both should use the same cached ID
|
||||
assert output_item_id is not None, "Output item should have an ID"
|
||||
assert content_part_id is not None, "Content part should have an item_id"
|
||||
assert output_item_id == content_part_id, (
|
||||
f"Initial events should use consistent IDs. "
|
||||
f"Got output_item_id={output_item_id}, content_part_id={content_part_id}"
|
||||
)
|
||||
|
||||
# Verify it matches the cached ID
|
||||
assert iterator._cached_item_id is not None
|
||||
assert iterator._cached_item_id == output_item_id
|
||||
|
||||
def test_streaming_iterator_done_events_use_cached_id(self):
|
||||
"""
|
||||
Test that done events (output_text_done, content_part_done, output_item_done) use the cached ID.
|
||||
"""
|
||||
from unittest.mock import Mock
|
||||
|
||||
import litellm
|
||||
from litellm.responses.litellm_completion_transformation.streaming_iterator import (
|
||||
LiteLLMCompletionStreamingIterator,
|
||||
)
|
||||
from litellm.types.utils import Choices, Message, ModelResponse
|
||||
|
||||
# Create a mock stream wrapper
|
||||
mock_stream_wrapper = Mock(spec=litellm.CustomStreamWrapper)
|
||||
mock_logging_obj = Mock()
|
||||
mock_stream_wrapper.logging_obj = mock_logging_obj
|
||||
mock_logging_obj._response_cost_calculator = Mock(return_value=0.001)
|
||||
|
||||
# Create the streaming iterator
|
||||
iterator = LiteLLMCompletionStreamingIterator(
|
||||
model="gemini/gemini-2.5-flash-lite",
|
||||
litellm_custom_stream_wrapper=mock_stream_wrapper,
|
||||
request_input="Test",
|
||||
responses_api_request={},
|
||||
)
|
||||
|
||||
# Set up a complete model response
|
||||
complete_response = ModelResponse(
|
||||
id="test-response-id",
|
||||
created=1234567890,
|
||||
model="gemini-2.5-flash-lite",
|
||||
object="chat.completion",
|
||||
choices=[
|
||||
Choices(
|
||||
finish_reason="stop",
|
||||
index=0,
|
||||
message=Message(content="Hello World", role="assistant"),
|
||||
)
|
||||
],
|
||||
)
|
||||
iterator.litellm_model_response = complete_response
|
||||
|
||||
# Create done events
|
||||
text_done_event = iterator.create_output_text_done_event(complete_response)
|
||||
content_done_event = iterator.create_output_content_part_done_event(complete_response)
|
||||
item_done_event = iterator.create_output_item_done_event(complete_response)
|
||||
|
||||
# Extract IDs
|
||||
text_done_id = getattr(text_done_event, "item_id", None)
|
||||
content_done_id = getattr(content_done_event, "item_id", None)
|
||||
item_done_id = getattr(item_done_event.item, "id", None)
|
||||
|
||||
# Assert: All done events should use the same cached ID
|
||||
assert text_done_id is not None, "Text done event should have an item_id"
|
||||
assert content_done_id is not None, "Content done event should have an item_id"
|
||||
assert item_done_id is not None, "Item done event should have an id"
|
||||
|
||||
assert text_done_id == content_done_id == item_done_id, (
|
||||
f"All done events should use consistent IDs. "
|
||||
f"Got text_done={text_done_id}, content_done={content_done_id}, item_done={item_done_id}"
|
||||
)
|
||||
|
||||
# Verify it matches the cached ID
|
||||
assert iterator._cached_item_id is not None
|
||||
assert iterator._cached_item_id == text_done_id
|
||||
Loading…
Add table
Reference in a new issue