Fix: tool call streaming in chat completino brigde

This commit is contained in:
Sameer Kankute 2026-01-21 12:14:02 +05:30
parent a5ea08a0bf
commit 3a0d166eb0
2 changed files with 242 additions and 74 deletions

View file

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

View file

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