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:
Sameer Kankute 2026-01-20 20:46:34 +05:30 • committed by GitHub
commit 11dbae85d1
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 243 additions and 8 deletions

View file

@ -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,

View file

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