From 3a0d166eb0c521961482fed10cf753991a9563b0 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Wed, 21 Jan 2026 12:14:02 +0530 Subject: [PATCH] Fix: tool call streaming in chat completino brigde --- .../streaming_iterator.py | 181 +++++++++++------- ...test_tool_call_streaming_transformation.py | 135 ++++++++++++- 2 files changed, 242 insertions(+), 74 deletions(-) diff --git a/litellm/responses/litellm_completion_transformation/streaming_iterator.py b/litellm/responses/litellm_completion_transformation/streaming_iterator.py index b128452d9a4..eeb728721fe 100644 --- a/litellm/responses/litellm_completion_transformation/streaming_iterator.py +++ b/litellm/responses/litellm_completion_transformation/streaming_iterator.py @@ -16,13 +16,13 @@ from litellm.types.llms.openai import ( ContentPartDoneEvent, ContentPartDonePartOutputText, ContentPartDonePartReasoningText, + FunctionCallArgumentsDeltaEvent, + FunctionCallArgumentsDoneEvent, OutputItemAddedEvent, OutputItemDoneEvent, OutputTextAnnotationAddedEvent, OutputTextDeltaEvent, OutputTextDoneEvent, - FunctionCallArgumentsDeltaEvent, - FunctionCallArgumentsDoneEvent, ReasoningSummaryTextDeltaEvent, ResponseCompletedEvent, ResponseCreatedEvent, @@ -88,6 +88,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): self._tool_args_by_call_id: dict[str, str] = {} self._next_tool_output_index: int = 1 # output_index=0 reserved for the message item self._final_tool_events_queued: bool = False + self._sequence_number: int = 0 def _get_or_assign_tool_output_index(self, call_id: str) -> int: existing = self._tool_output_index_by_call_id.get(call_id) @@ -104,7 +105,10 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): We emit: - response.output_item.added (function_call) - - response.function_call_arguments.delta + - response.function_call_arguments.delta (split into smaller chunks to match OpenAI behavior) + + Note: Some providers (like Bedrock) send tool call arguments in one large chunk. + We split these into smaller deltas to match OpenAI's token-by-token streaming behavior. """ if not isinstance(tool_calls, list): return @@ -129,33 +133,42 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): if call_id not in self._tool_args_by_call_id: self._tool_args_by_call_id[call_id] = "" - self._pending_tool_events.append( - OutputItemAddedEvent( - type=ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED, - output_index=output_index, - item=BaseLiteLLMOpenAIResponseObject( - **{ - "type": "function_call", - "id": call_id, - "call_id": call_id, - "name": fn_name, - "arguments": "", - "status": "in_progress", - } - ), - ) + self._sequence_number += 1 + event = OutputItemAddedEvent( + type=ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED, + output_index=output_index, + item=BaseLiteLLMOpenAIResponseObject( + **{ + "type": "function_call", + "id": call_id, + "call_id": call_id, + "name": fn_name, + "arguments": "", + "status": "in_progress", + } + ), ) + event.__dict__['sequence_number'] = self._sequence_number + self._pending_tool_events.append(event) if fn_args_delta: self._tool_args_by_call_id[call_id] += fn_args_delta - self._pending_tool_events.append( - FunctionCallArgumentsDeltaEvent( + + # Split large argument deltas into smaller chunks to match OpenAI's streaming behavior + # This is especially important for providers like Bedrock that send complete arguments at once + chunk_size = 10 # Match typical OpenAI delta size + for i in range(0, len(fn_args_delta), chunk_size): + delta_chunk = fn_args_delta[i:i + chunk_size] + self._sequence_number += 1 + event = FunctionCallArgumentsDeltaEvent( type=ResponsesAPIStreamEvents.FUNCTION_CALL_ARGUMENTS_DELTA, item_id=call_id, output_index=output_index, - delta=fn_args_delta, + delta=delta_chunk, ) - ) + # Add sequence_number as extra field (BaseLiteLLMOpenAIResponseObject allows extra fields) + event.__dict__['sequence_number'] = self._sequence_number + self._pending_tool_events.append(event) def _queue_final_tool_call_done_events(self, litellm_complete_object: ModelResponse) -> None: """ @@ -191,53 +204,79 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): fn_name = str(getattr(fn, "name", "") or "") fn_args = str(getattr(fn, "arguments", "") or "") + # Track if this is a new tool call that wasn't streamed + is_new_tool_call = call_id not in self._tool_args_by_call_id + # If we never sent output_item.added for this call_id, emit it now. - if call_id not in self._tool_args_by_call_id: + if is_new_tool_call: self._tool_args_by_call_id[call_id] = "" - self._pending_tool_events.append( - OutputItemAddedEvent( - type=ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED, - output_index=output_index, - item=BaseLiteLLMOpenAIResponseObject( - **{ - "type": "function_call", - "id": call_id, - "call_id": call_id, - "name": fn_name, - "arguments": "", - "status": "in_progress", - } - ), - ) - ) - - final_args = fn_args or self._tool_args_by_call_id.get(call_id, "") - self._pending_tool_events.append( - FunctionCallArgumentsDoneEvent( - type=ResponsesAPIStreamEvents.FUNCTION_CALL_ARGUMENTS_DONE, - item_id=call_id, + self._sequence_number += 1 + event = OutputItemAddedEvent( + type=ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED, output_index=output_index, - arguments=final_args, - ) - ) - - self._pending_tool_events.append( - OutputItemDoneEvent( - type=ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE, - output_index=output_index, - sequence_number=1, item=BaseLiteLLMOpenAIResponseObject( **{ "type": "function_call", "id": call_id, "call_id": call_id, "name": fn_name, - "arguments": final_args, - "status": "completed", + "arguments": "", + "status": "in_progress", } ), ) + event.__dict__['sequence_number'] = self._sequence_number + self._pending_tool_events.append(event) + + final_args = fn_args or self._tool_args_by_call_id.get(call_id, "") + + # Emit delta events for arguments that weren't streamed yet + # This handles cases where Bedrock sends the complete tool call at the end + already_streamed = self._tool_args_by_call_id.get(call_id, "") + remaining_args = final_args[len(already_streamed):] if final_args else "" + + if remaining_args: + # Split into smaller chunks to match OpenAI's streaming behavior + chunk_size = 10 # Match typical OpenAI delta size + for i in range(0, len(remaining_args), chunk_size): + delta_chunk = remaining_args[i:i + chunk_size] + self._sequence_number += 1 + delta_event = FunctionCallArgumentsDeltaEvent( + type=ResponsesAPIStreamEvents.FUNCTION_CALL_ARGUMENTS_DELTA, + item_id=call_id, + output_index=output_index, + delta=delta_chunk, + ) + delta_event.__dict__['sequence_number'] = self._sequence_number + self._pending_tool_events.append(delta_event) + + self._sequence_number += 1 + done_event = FunctionCallArgumentsDoneEvent( + type=ResponsesAPIStreamEvents.FUNCTION_CALL_ARGUMENTS_DONE, + item_id=call_id, + output_index=output_index, + arguments=final_args, ) + done_event.__dict__['sequence_number'] = self._sequence_number + self._pending_tool_events.append(done_event) + + self._sequence_number += 1 + item_done_event = OutputItemDoneEvent( + type=ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE, + output_index=output_index, + sequence_number=self._sequence_number, + item=BaseLiteLLMOpenAIResponseObject( + **{ + "type": "function_call", + "id": call_id, + "call_id": call_id, + "name": fn_name, + "arguments": final_args, + "status": "completed", + } + ), + ) + self._pending_tool_events.append(item_done_event) def _default_response_created_event_data(self) -> dict: response_created_event_data = { @@ -295,24 +334,31 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): """ response_created_event_data = self._default_response_created_event_data() - return ResponseCreatedEvent( + self._sequence_number += 1 + event = ResponseCreatedEvent( type=ResponsesAPIStreamEvents.RESPONSE_CREATED, response=ResponsesAPIResponse(**response_created_event_data), ) + event.__dict__['sequence_number'] = self._sequence_number + return event def create_response_in_progress_event(self) -> ResponseInProgressEvent: response_in_progress_event_data = self._default_response_created_event_data() response_in_progress_event_data["status"] = "in_progress" - return ResponseInProgressEvent( + self._sequence_number += 1 + event = ResponseInProgressEvent( type=ResponsesAPIStreamEvents.RESPONSE_IN_PROGRESS, response=ResponsesAPIResponse(**response_in_progress_event_data), ) + event.__dict__['sequence_number'] = self._sequence_number + return event 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( + self._sequence_number += 1 + event = OutputItemAddedEvent( type=ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED, output_index=0, item=BaseLiteLLMOpenAIResponseObject( @@ -325,12 +371,15 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): } ), ) + event.__dict__['sequence_number'] = self._sequence_number + return event 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( + self._sequence_number += 1 + event = ContentPartAddedEvent( type=ResponsesAPIStreamEvents.CONTENT_PART_ADDED, item_id=self._cached_item_id, output_index=0, @@ -339,6 +388,8 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): **{"type": "output_text", "text": "", "annotations": []} ), ) + event.__dict__['sequence_number'] = self._sequence_number + return event def create_litellm_model_response( self, @@ -458,9 +509,6 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): elif self.sent_response_in_progress_event is False: self.sent_response_in_progress_event = True return self.create_response_in_progress_event() - elif self.sent_output_item_added_event is False: - self.sent_output_item_added_event = True - return self.create_output_item_added_event() elif self.sent_content_part_added_event is False: self.sent_content_part_added_event = True return self.create_content_part_added_event() @@ -638,21 +686,26 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): # Priority 2: Handle text deltas delta_content = self._get_delta_string_from_streaming_choices(chunk.choices) if delta_content: - return OutputTextDeltaEvent( + self._sequence_number += 1 + event = OutputTextDeltaEvent( type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA, item_id=item_id, output_index=0, content_index=0, delta=delta_content, ) + event.__dict__['sequence_number'] = self._sequence_number + return event # Priority 3: Handle tool call deltas (if any) -> queue events and emit them + # For each tool call delta, we emit events one at a time to match OpenAI's streaming behavior if ( chunk.choices and hasattr(chunk.choices[0].delta, "tool_calls") and chunk.choices[0].delta.tool_calls ): self._queue_tool_call_delta_events(chunk.choices[0].delta.tool_calls) + # Return one pending tool event at a time if self._pending_tool_events: return self._pending_tool_events.pop(0) diff --git a/tests/test_litellm/responses/litellm_completion_transformation/test_tool_call_streaming_transformation.py b/tests/test_litellm/responses/litellm_completion_transformation/test_tool_call_streaming_transformation.py index 51150383b01..8d324bea611 100644 --- a/tests/test_litellm/responses/litellm_completion_transformation/test_tool_call_streaming_transformation.py +++ b/tests/test_litellm/responses/litellm_completion_transformation/test_tool_call_streaming_transformation.py @@ -14,7 +14,12 @@ from litellm.responses.litellm_completion_transformation.streaming_iterator impo LiteLLMCompletionStreamingIterator, ) from litellm.types.llms.openai import ResponsesAPIStreamEvents -from litellm.types.utils import Delta, ModelResponse, ModelResponseStream, StreamingChoices +from litellm.types.utils import ( + Delta, + ModelResponse, + ModelResponseStream, + StreamingChoices, +) def test_tool_call_delta_is_emitted_as_responses_events(): @@ -55,12 +60,14 @@ def test_tool_call_delta_is_emitted_as_responses_events(): assert evt1.type == ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED assert evt1.output_index == 1 + # The arguments are now chunked, so we get the first delta chunk evt2 = iterator._transform_chat_completion_chunk_to_response_api_chunk(chunk) assert evt2 is not None assert evt2.type == ResponsesAPIStreamEvents.FUNCTION_CALL_ARGUMENTS_DELTA assert evt2.item_id == "call_1" assert evt2.output_index == 1 - assert evt2.delta == '{"x":1}' + # The delta will be a chunk of the arguments, not the full arguments + assert len(evt2.delta) <= 10 # Chunks are max 10 characters def test_tool_calls_present_only_in_final_response_are_emitted_before_completed(): @@ -104,13 +111,121 @@ def test_tool_calls_present_only_in_final_response_are_emitted_before_completed( assert evt1.type == ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED assert evt1.output_index == 1 - evt2 = iterator.common_done_event_logic(sync_mode=True) - assert evt2.type == ResponsesAPIStreamEvents.FUNCTION_CALL_ARGUMENTS_DONE - assert evt2.item_id == "call_2" - assert evt2.output_index == 1 - assert evt2.arguments == '{"y":2}' + # Now delta events are emitted (arguments split into chunks) + # Collect all delta events + delta_events = [] + while True: + evt = iterator.common_done_event_logic(sync_mode=True) + if evt.type == ResponsesAPIStreamEvents.FUNCTION_CALL_ARGUMENTS_DELTA: + delta_events.append(evt) + else: + break + + # Verify we got delta events + assert len(delta_events) > 0 + # Verify they reconstruct the original arguments + concatenated_args = ''.join(evt.delta for evt in delta_events) + assert concatenated_args == '{"y":2}' - evt3 = iterator.common_done_event_logic(sync_mode=True) - assert evt3.type == ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE - assert evt3.output_index == 1 + # The last event should be FUNCTION_CALL_ARGUMENTS_DONE + assert evt.type == ResponsesAPIStreamEvents.FUNCTION_CALL_ARGUMENTS_DONE + assert evt.item_id == "call_2" + assert evt.output_index == 1 + assert evt.arguments == '{"y":2}' + + evt_final = iterator.common_done_event_logic(sync_mode=True) + assert evt_final.type == ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE + assert evt_final.output_index == 1 + + +def test_tool_call_arguments_are_chunked_to_match_openai_behavior(): + """ + Test that large tool call arguments are split into smaller chunks (size 10) + to replicate OpenAI's native streaming behavior. + + This is especially important for providers like Bedrock that send complete + arguments at once, which need to be split to match OpenAI's token-by-token streaming. + """ + iterator = LiteLLMCompletionStreamingIterator( + model="test-model", + litellm_custom_stream_wrapper=AsyncMock(), + request_input="Test input", + responses_api_request={}, + ) + + # Create a chunk with a large arguments string that should be split + large_arguments = '{"param1": "value1", "param2": "value2", "param3": "value3"}' # 67 chars + chunk = ModelResponseStream( + id="chunk-1", + created=123, + model="test-model", + object="chat.completion.chunk", + choices=[ + StreamingChoices( + finish_reason=None, + index=0, + delta=Delta( + role="assistant", + content="", + tool_calls=[ + { + "id": "call_test", + "type": "function", + "function": {"name": "test_function", "arguments": large_arguments}, + } + ], + ), + ) + ], + ) + + # Process the chunk once - it queues all events internally + evt = iterator._transform_chat_completion_chunk_to_response_api_chunk(chunk) + + # First event should be OUTPUT_ITEM_ADDED + assert evt is not None + assert evt.type == ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED + assert evt.output_index == 1 + assert hasattr(evt, '__dict__') and 'sequence_number' in evt.__dict__ + + # Collect all remaining delta events from the pending queue by creating empty chunks + delta_events = [] + empty_chunk = ModelResponseStream( + id="chunk-1", + created=123, + model="test-model", + object="chat.completion.chunk", + choices=[ + StreamingChoices( + finish_reason=None, + index=0, + delta=Delta(role="assistant", content=""), + ) + ], + ) + + # Keep draining pending events (expected: ceil(67 / 10) = 7 delta events) + while iterator._pending_tool_events: + evt = iterator._transform_chat_completion_chunk_to_response_api_chunk(empty_chunk) + if evt and evt.type == ResponsesAPIStreamEvents.FUNCTION_CALL_ARGUMENTS_DELTA: + delta_events.append(evt) + + # Verify multiple delta events were created (at least 6 chunks for 67 chars) + assert len(delta_events) >= 6 # 67 chars split into chunks of max 10 chars each + + # Verify each delta is at most 10 characters + for evt in delta_events: + assert len(evt.delta) <= 10 + assert evt.item_id == "call_test" + assert evt.output_index == 1 + assert hasattr(evt, '__dict__') and 'sequence_number' in evt.__dict__ + + # Verify all deltas concatenated equal the original arguments + concatenated = ''.join(evt.delta for evt in delta_events) + assert concatenated == large_arguments + + # Verify sequence numbers are increasing + sequence_numbers = [evt.__dict__['sequence_number'] for evt in delta_events] + assert sequence_numbers == sorted(sequence_numbers) + assert len(set(sequence_numbers)) == len(sequence_numbers) # All unique