mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(ollama): fix streaming tool call finish_reason, arguments format, and parallel index
Ollama sends tool calls in a chunk with done: false, then a separate final chunk with done_reason: stop. The streaming iterator had three defects that broke spec-strict OpenAI clients: 1. finish_reason stayed stop because the override only checked the current chunk. Track saw_tool_calls across chunks and upgrade finish_reason to tool_calls on the done chunk, gated on stop so length and other reasons pass through (matching streaming_handler.py). 2. function.arguments arrived as a dict instead of a JSON string. The non-streaming path already serialized via json.dumps; the streaming path now does the same. 3. Every parallel tool call got index: 0 because Delta.__init__ resets its counter per chunk. Track a persistent _tool_call_index on the iterator so each call gets a sequential index across chunks. Streaming chunks also now share a stable response_id and tool call ids use the call_ prefix, matching the non-streaming format.
This commit is contained in:
parent
ff02d5cfc0
commit
acc96b2b90
2 changed files with 181 additions and 5 deletions
|
|
@ -421,6 +421,12 @@ class OllamaChatConfig(BaseConfig):
|
|||
class OllamaChatCompletionResponseIterator(BaseModelResponseIterator):
|
||||
started_reasoning_content: bool = False
|
||||
finished_reasoning_content: bool = False
|
||||
saw_tool_calls: bool = False
|
||||
|
||||
def __init__(self, streaming_response, sync_stream: bool, json_mode: bool | None = False) -> None:
|
||||
super().__init__(streaming_response, sync_stream, json_mode)
|
||||
self.response_id: str = str(uuid.uuid4())
|
||||
self._tool_call_index: int = 0
|
||||
|
||||
def _is_function_call_complete(self, function_args: str | dict) -> bool:
|
||||
if isinstance(function_args, dict):
|
||||
|
|
@ -466,12 +472,17 @@ class OllamaChatCompletionResponseIterator(BaseModelResponseIterator):
|
|||
# process tool calls - if complete function arg - add id to tool call
|
||||
tool_calls: Final = chunk["message"].get("tool_calls")
|
||||
if tool_calls is not None:
|
||||
self.saw_tool_calls = True
|
||||
for tool_call in tool_calls:
|
||||
tool_call["index"] = self._tool_call_index
|
||||
self._tool_call_index += 1
|
||||
function_args = tool_call.get("function").get("arguments")
|
||||
if function_args is not None and len(function_args) > 0:
|
||||
if isinstance(function_args, dict):
|
||||
tool_call["function"]["arguments"] = json.dumps(function_args)
|
||||
is_function_call_complete = self._is_function_call_complete(function_args)
|
||||
if is_function_call_complete:
|
||||
tool_call["id"] = str(uuid.uuid4())
|
||||
tool_call["id"] = f"call_{uuid.uuid4()}"
|
||||
|
||||
# PROCESS REASONING CONTENT
|
||||
reasoning_content: str | None = None
|
||||
|
|
@ -506,9 +517,7 @@ class OllamaChatCompletionResponseIterator(BaseModelResponseIterator):
|
|||
|
||||
if chunk["done"] is True:
|
||||
finish_reason = chunk.get("done_reason") or "stop"
|
||||
# Override finish_reason when tool_calls are present
|
||||
# Fixes: https://github.com/BerriAI/litellm/issues/18922
|
||||
if tool_calls is not None:
|
||||
if finish_reason == "stop" and (tool_calls is not None or self.saw_tool_calls):
|
||||
finish_reason = "tool_calls"
|
||||
choices = [
|
||||
StreamingChoices(
|
||||
|
|
@ -530,7 +539,7 @@ class OllamaChatCompletionResponseIterator(BaseModelResponseIterator):
|
|||
)
|
||||
|
||||
return ModelResponseStream(
|
||||
id=str(uuid.uuid4()),
|
||||
id=self.response_id,
|
||||
object="chat.completion.chunk",
|
||||
created=int(time.time()), # ollama created_at is in UTC
|
||||
usage=usage,
|
||||
|
|
|
|||
|
|
@ -1,4 +1,5 @@
|
|||
import inspect
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
from typing import cast
|
||||
|
|
@ -906,3 +907,169 @@ class TestOllamaToolCallTransformation:
|
|||
assert tool_msg["content"] == "Sunny, 72°F"
|
||||
assert "tool_call_id" in tool_msg, "tool_call_id must be forwarded to Ollama"
|
||||
assert tool_msg["tool_call_id"] == "call_abc123"
|
||||
|
||||
|
||||
class TestOllamaStreamingToolCalls:
|
||||
"""Regression tests for ollama_chat streaming tool call defects."""
|
||||
|
||||
@staticmethod
|
||||
def _make_tool_call_chunk(tool_name: str, arguments: dict, done: bool = False, done_reason: str = "stop") -> dict:
|
||||
return {
|
||||
"model": "qwen3:14b",
|
||||
"created_at": "2025-01-11T00:00:00.000000Z",
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"content": "",
|
||||
"tool_calls": [
|
||||
{"function": {"name": tool_name, "arguments": arguments}}
|
||||
],
|
||||
},
|
||||
"done": done,
|
||||
"done_reason": done_reason if done else None,
|
||||
"prompt_eval_count": 10,
|
||||
"eval_count": 5,
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def _make_done_chunk(done_reason: str = "stop") -> dict:
|
||||
return {
|
||||
"model": "qwen3:14b",
|
||||
"created_at": "2025-01-11T00:00:00.000000Z",
|
||||
"message": {"role": "assistant", "content": ""},
|
||||
"done": True,
|
||||
"done_reason": done_reason,
|
||||
"prompt_eval_count": 10,
|
||||
"eval_count": 5,
|
||||
}
|
||||
|
||||
def test_streaming_chunks_have_consistent_id(self):
|
||||
iterator = OllamaChatCompletionResponseIterator(
|
||||
streaming_response=iter([]),
|
||||
sync_stream=True,
|
||||
)
|
||||
expected_id = iterator.response_id
|
||||
|
||||
chunk1 = {
|
||||
"model": "qwen3:14b",
|
||||
"created_at": "2025-01-11T00:00:00.000000Z",
|
||||
"message": {"role": "assistant", "content": "Hello"},
|
||||
"done": False,
|
||||
}
|
||||
chunk2 = {
|
||||
"model": "qwen3:14b",
|
||||
"created_at": "2025-01-11T00:00:00.000000Z",
|
||||
"message": {"role": "assistant", "content": " world"},
|
||||
"done": True,
|
||||
"done_reason": "stop",
|
||||
"prompt_eval_count": 5,
|
||||
"eval_count": 2,
|
||||
}
|
||||
|
||||
result1 = iterator.chunk_parser(chunk1)
|
||||
result2 = iterator.chunk_parser(chunk2)
|
||||
|
||||
assert result1.id == expected_id
|
||||
assert result2.id == expected_id
|
||||
|
||||
def test_streaming_tool_call_id_has_call_prefix(self):
|
||||
iterator = OllamaChatCompletionResponseIterator(
|
||||
streaming_response=iter([]),
|
||||
sync_stream=True,
|
||||
)
|
||||
|
||||
result = iterator.chunk_parser(self._make_tool_call_chunk("get_weather", {"location": "Tokyo"}, done=True))
|
||||
tool_call = result.choices[0].delta.tool_calls[0]
|
||||
assert tool_call["id"].startswith("call_")
|
||||
|
||||
def test_streaming_arguments_converted_to_json_string(self):
|
||||
iterator = OllamaChatCompletionResponseIterator(
|
||||
streaming_response=iter([]),
|
||||
sync_stream=True,
|
||||
)
|
||||
|
||||
result = iterator.chunk_parser(self._make_tool_call_chunk("get_weather", {"location": "Tokyo"}, done=True))
|
||||
arguments = result.choices[0].delta.tool_calls[0]["function"]["arguments"]
|
||||
assert isinstance(arguments, str)
|
||||
assert json.loads(arguments) == {"location": "Tokyo"}
|
||||
|
||||
def test_streaming_finish_reason_tool_calls_in_done_chunk(self):
|
||||
iterator = OllamaChatCompletionResponseIterator(
|
||||
streaming_response=iter([]),
|
||||
sync_stream=True,
|
||||
)
|
||||
|
||||
result = iterator.chunk_parser(self._make_tool_call_chunk("get_weather", {"location": "Tokyo"}, done=True))
|
||||
assert result.choices[0].finish_reason == "tool_calls"
|
||||
|
||||
def test_streaming_saw_tool_calls_propagates_to_done_chunk(self):
|
||||
iterator = OllamaChatCompletionResponseIterator(
|
||||
streaming_response=iter([]),
|
||||
sync_stream=True,
|
||||
)
|
||||
|
||||
iterator.chunk_parser(self._make_tool_call_chunk("get_weather", {"location": "Tokyo"}))
|
||||
result = iterator.chunk_parser(self._make_done_chunk())
|
||||
assert result.choices[0].finish_reason == "tool_calls"
|
||||
|
||||
def test_streaming_length_finish_reason_preserved_with_tool_calls(self):
|
||||
iterator = OllamaChatCompletionResponseIterator(
|
||||
streaming_response=iter([]),
|
||||
sync_stream=True,
|
||||
)
|
||||
|
||||
iterator.chunk_parser(self._make_tool_call_chunk("get_weather", {"location": "Tokyo"}))
|
||||
result = iterator.chunk_parser(self._make_done_chunk(done_reason="length"))
|
||||
assert result.choices[0].finish_reason == "length"
|
||||
|
||||
def test_streaming_parallel_tool_calls_get_unique_indices(self):
|
||||
iterator = OllamaChatCompletionResponseIterator(
|
||||
streaming_response=iter([]),
|
||||
sync_stream=True,
|
||||
)
|
||||
|
||||
chunk1 = {
|
||||
"model": "qwen3:14b",
|
||||
"created_at": "2025-01-11T00:00:00.000000Z",
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"content": "",
|
||||
"tool_calls": [
|
||||
{
|
||||
"function": {
|
||||
"name": "get_weather",
|
||||
"arguments": {"city": "Tokyo"},
|
||||
}
|
||||
}
|
||||
],
|
||||
},
|
||||
"done": False,
|
||||
}
|
||||
chunk2 = {
|
||||
"model": "qwen3:14b",
|
||||
"created_at": "2025-01-11T00:00:00.000000Z",
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"content": "",
|
||||
"tool_calls": [
|
||||
{
|
||||
"function": {
|
||||
"name": "get_time",
|
||||
"arguments": {"timezone": "America/New_York"},
|
||||
}
|
||||
}
|
||||
],
|
||||
},
|
||||
"done": False,
|
||||
}
|
||||
|
||||
result1 = iterator.chunk_parser(chunk1)
|
||||
result2 = iterator.chunk_parser(chunk2)
|
||||
|
||||
tc1 = result1.choices[0].delta.tool_calls[0]
|
||||
tc2 = result2.choices[0].delta.tool_calls[0]
|
||||
assert tc1["index"] == 0
|
||||
assert tc2["index"] == 1
|
||||
assert tc1["function"]["name"] == "get_weather"
|
||||
assert tc2["function"]["name"] == "get_time"
|
||||
assert json.loads(tc1["function"]["arguments"]) == {"city": "Tokyo"}
|
||||
assert json.loads(tc2["function"]["arguments"]) == {"timezone": "America/New_York"}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue