mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
Fix: tool call streaming in chat completino brigde
This commit is contained in:
parent
a5ea08a0bf
commit
3a0d166eb0
2 changed files with 242 additions and 74 deletions
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue