mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
fix(responses): correlate completed tools by call id
This commit is contained in:
parent
6ca3966377
commit
3fb1128a5f
3 changed files with 116 additions and 3 deletions
|
|
@ -119,6 +119,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
|
|||
self._tool_call_id_by_index: dict[int, str] = {}
|
||||
self._streamed_tool_call_ids_in_order: list[str] = [] # mutable-ok: accumulates call ids across stream chunks
|
||||
self._resolved_tool_call_id_by_position: dict[int, str] = {} # mutable-ok: terminal correlation state
|
||||
self._streamed_call_id_by_terminal_id: dict[str, str] = {} # mutable-ok: terminal identity correlation
|
||||
self._ambiguous_tool_call_indexes: set[int] = set()
|
||||
self._next_tool_output_index: int = 1 # output_index=0 reserved for the message item
|
||||
self._final_tool_events_queued: bool = False
|
||||
|
|
@ -334,8 +335,10 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
|
|||
call_id_raw = tc.get("id") if isinstance(tc, dict) else getattr(tc, "id", None)
|
||||
if not call_id_raw:
|
||||
continue
|
||||
call_id = self._streamed_tool_call_id_for_terminal_call(tc, position) or str(call_id_raw)
|
||||
terminal_call_id = str(call_id_raw)
|
||||
call_id = self._streamed_tool_call_id_for_terminal_call(tc, position) or terminal_call_id
|
||||
self._resolved_tool_call_id_by_position[position] = call_id
|
||||
self._streamed_call_id_by_terminal_id[terminal_call_id] = call_id
|
||||
output_index = self._get_or_assign_tool_output_index(call_id)
|
||||
|
||||
fn = tc.get("function") if isinstance(tc, dict) else getattr(tc, "function", None)
|
||||
|
|
@ -1211,9 +1214,19 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
|
|||
return chat_completion_delta.content or ""
|
||||
|
||||
def _output_item_with_streamed_tool_identity(self, item: object, tool_position: int) -> object:
|
||||
resolved_call_id: Final = self._resolved_tool_call_id_by_position.get(tool_position)
|
||||
terminal_call_id: Final = getattr(item, "call_id", None)
|
||||
resolved_by_call_id: Final = (
|
||||
self._streamed_call_id_by_terminal_id.get(terminal_call_id) if isinstance(terminal_call_id, str) else None
|
||||
)
|
||||
resolved_by_position: Final = self._resolved_tool_call_id_by_position.get(tool_position)
|
||||
streamed_call_id: Final = (
|
||||
resolved_call_id if resolved_call_id is not None else self._streamed_tool_call_id_at_position(tool_position)
|
||||
resolved_by_call_id
|
||||
if resolved_by_call_id is not None
|
||||
else (
|
||||
resolved_by_position
|
||||
if resolved_by_position is not None
|
||||
else self._streamed_tool_call_id_at_position(tool_position)
|
||||
)
|
||||
)
|
||||
if streamed_call_id is None:
|
||||
return item
|
||||
|
|
|
|||
|
|
@ -3412,6 +3412,7 @@ class TestEnsureOutputItemContentPartAdded:
|
|||
iterator._tool_call_id_by_index = {}
|
||||
iterator._streamed_tool_call_ids_in_order = []
|
||||
iterator._resolved_tool_call_id_by_position = {}
|
||||
iterator._streamed_call_id_by_terminal_id = {}
|
||||
iterator._ambiguous_tool_call_indexes = set()
|
||||
iterator._next_tool_output_index = 1
|
||||
iterator._final_tool_events_queued = False
|
||||
|
|
|
|||
|
|
@ -705,6 +705,105 @@ def test_terminal_only_call_is_not_conflated_with_later_streamed_call():
|
|||
]
|
||||
|
||||
|
||||
def test_completed_snapshot_correlates_function_after_server_tool_replacement():
|
||||
iterator = LiteLLMCompletionStreamingIterator(
|
||||
model="test-model",
|
||||
litellm_custom_stream_wrapper=AsyncMock(),
|
||||
request_input="Test input",
|
||||
responses_api_request={},
|
||||
)
|
||||
iterator._queue_tool_call_delta_events(
|
||||
[
|
||||
{
|
||||
"index": 0,
|
||||
"id": "call_exec_stream",
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "bash_code_execution",
|
||||
"arguments": '{"command":"printf server"}',
|
||||
},
|
||||
},
|
||||
{
|
||||
"index": 1,
|
||||
"id": "call_regular_stream",
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "lookup_weather",
|
||||
"arguments": '{"city":"Paris"}',
|
||||
},
|
||||
},
|
||||
]
|
||||
)
|
||||
iterator._pending_tool_events.clear()
|
||||
terminal_response = ModelResponse(
|
||||
id="chatcmpl-terminal",
|
||||
created=123,
|
||||
model="test-model",
|
||||
object="chat.completion",
|
||||
choices=[
|
||||
{
|
||||
"index": 0,
|
||||
"finish_reason": "tool_calls",
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"content": None,
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "srvtoolu_exec_terminal",
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "bash_code_execution",
|
||||
"arguments": '{"command":"printf server"}',
|
||||
},
|
||||
},
|
||||
{
|
||||
"id": "call_regular_terminal",
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "lookup_weather",
|
||||
"arguments": '{"city":"Paris"}',
|
||||
},
|
||||
},
|
||||
],
|
||||
"provider_specific_fields": {
|
||||
"code_interpreter_results": [
|
||||
{
|
||||
"type": "code_interpreter_call",
|
||||
"id": "srvtoolu_exec_terminal",
|
||||
"code": "printf server",
|
||||
"container_id": None,
|
||||
"status": "completed",
|
||||
"outputs": [{"type": "logs", "logs": "server"}],
|
||||
}
|
||||
]
|
||||
},
|
||||
},
|
||||
}
|
||||
],
|
||||
)
|
||||
|
||||
iterator._queue_final_tool_call_done_events(terminal_response)
|
||||
completed = iterator._emit_response_completed_event(terminal_response)
|
||||
|
||||
assert completed is not None
|
||||
code_calls = [item for item in completed.response.output if item.type == "code_interpreter_call"]
|
||||
function_calls = [item for item in completed.response.output if item.type == "function_call"]
|
||||
assert len(code_calls) == 1
|
||||
assert len(function_calls) == 1
|
||||
code_call = code_calls[0]
|
||||
assert code_call.id == "srvtoolu_exec_terminal"
|
||||
assert code_call.code == "printf server"
|
||||
assert code_call.container_id is None
|
||||
assert code_call.outputs[0].logs == "server"
|
||||
function_call = function_calls[0]
|
||||
assert (function_call.id, function_call.call_id) == (
|
||||
"fc_call_regular_stream",
|
||||
"call_regular_stream",
|
||||
)
|
||||
assert function_call.name == "lookup_weather"
|
||||
assert function_call.arguments == '{"city":"Paris"}'
|
||||
|
||||
|
||||
def test_reused_index_with_new_call_id_marks_fallback_ambiguous():
|
||||
iterator = LiteLLMCompletionStreamingIterator(
|
||||
model="test-model",
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue