fix(ollama): reconstruct streaming tool calls, indices, finish_reason and error chunks

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
Devin AI 2026-08-31 20:06:25 +00:00
parent 1249f84b10
commit fa7b3f74d9
4 changed files with 226 additions and 6 deletions

View file

@ -338,6 +338,13 @@ class OllamaChatConfig(BaseConfig):
response_json: Final = raw_response.json()
if "error" in response_json:
raise OllamaError(
message=str(response_json["error"]),
status_code=raw_response.status_code if raw_response.status_code >= 400 else 400,
headers=dict(raw_response.headers),
)
## RESPONSE OBJECT
_done_reason: Final = map_finish_reason(response_json.get("done_reason") or "stop")
model_response.choices[0].finish_reason = _done_reason
@ -380,7 +387,7 @@ class OllamaChatConfig(BaseConfig):
model_response.choices[0].message = message
model_response.choices[0].finish_reason = "tool_calls"
else:
_message: Final = litellm.Message(**response_json_message)
_message: Final = litellm.Message(**(response_json_message or {}))
model_response.choices[0].message = _message
# Set finish_reason to "tool_calls" when tool_calls are present
# Fixes: https://github.com/BerriAI/litellm/issues/18922
@ -389,9 +396,8 @@ class OllamaChatConfig(BaseConfig):
model_response.created = int(time.time())
model_response.model = "ollama_chat/" + model
prompt_tokens = response_json.get("prompt_eval_count", litellm.token_counter(messages=messages))
completion_tokens: Final = response_json.get(
"eval_count",
litellm.token_counter(text=response_json["message"]["content"]),
completion_tokens: Final = response_json.get("eval_count") or litellm.token_counter(
text=(response_json_message or {}).get("content") or ""
)
setattr(
model_response,
@ -423,6 +429,7 @@ class OllamaChatConfig(BaseConfig):
class OllamaChatCompletionResponseIterator(BaseModelResponseIterator):
started_reasoning_content: bool = False
finished_reasoning_content: bool = False
stream_tool_call_count: int = 0
def _is_function_call_complete(self, function_args: str | dict) -> bool:
if isinstance(function_args, dict):
@ -465,11 +472,23 @@ class OllamaChatCompletionResponseIterator(BaseModelResponseIterator):
"""
from litellm.types.utils import Delta, StreamingChoices
if "error" in chunk:
raise OllamaError(
message=str(chunk["error"]),
status_code=400,
headers={"Content-Type": "application/json"},
)
# 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:
for tool_call in tool_calls:
function_args = tool_call.get("function").get("arguments")
function = tool_call.get("function") or {}
function_index = function.pop("index", None)
if tool_call.get("index") is None:
tool_call["index"] = function_index if function_index is not None else self.stream_tool_call_count
self.stream_tool_call_count = max(self.stream_tool_call_count + 1, tool_call["index"] + 1)
function_args = function.get("arguments")
if function_args is not None and len(function_args) > 0:
is_function_call_complete = self._is_function_call_complete(function_args)
if is_function_call_complete:
@ -510,7 +529,7 @@ class OllamaChatCompletionResponseIterator(BaseModelResponseIterator):
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 (tool_calls is not None or self.stream_tool_call_count > 0) and finish_reason != "length":
finish_reason = "tool_calls"
choices = [
StreamingChoices(

View file

@ -451,10 +451,32 @@ class OllamaTextCompletionResponseIterator(BaseModelResponseIterator):
super().__init__(streaming_response, sync_stream, json_mode)
self.started_reasoning_content: bool = False
self.finished_reasoning_content: bool = False
self.buffered_json_content: str | None = None
def _handle_string_chunk(self, str_line: str) -> GenericStreamingChunk | ModelResponseStream:
return self.chunk_parser(json.loads(str_line))
def _parse_buffered_function_call(self) -> list[dict] | None:
if self.buffered_json_content is None:
return None
try:
parsed = json.loads(self.buffered_json_content)
except json.JSONDecodeError:
return None
if isinstance(parsed, dict) and "name" in parsed and "arguments" in parsed:
return [
{
"id": f"call_{uuid.uuid4()}",
"index": 0,
"function": {
"name": parsed["name"],
"arguments": json.dumps(parsed["arguments"]),
},
"type": "function",
}
]
return None
def chunk_parser(self, chunk: dict) -> GenericStreamingChunk | ModelResponseStream:
try:
if "error" in chunk:
@ -477,6 +499,29 @@ class OllamaTextCompletionResponseIterator(BaseModelResponseIterator):
completion_tokens=eval_count,
total_tokens=prompt_eval_count + eval_count,
)
tool_calls: Final = self._parse_buffered_function_call()
if tool_calls is not None:
return ModelResponseStream(
choices=[
StreamingChoices(
index=0,
delta=Delta(content=None, tool_calls=tool_calls),
finish_reason="tool_calls",
)
],
usage=usage,
)
if self.buffered_json_content is not None:
return ModelResponseStream(
choices=[
StreamingChoices(
index=0,
delta=Delta(content=self.buffered_json_content),
finish_reason=finish_reason,
)
],
usage=usage,
)
return GenericStreamingChunk(
text=text,
is_finished=is_finished,
@ -485,6 +530,16 @@ class OllamaTextCompletionResponseIterator(BaseModelResponseIterator):
)
elif chunk["response"]:
text = chunk["response"]
if self.buffered_json_content is not None or (
self.buffered_json_content is None
and not self.started_reasoning_content
and text.lstrip().startswith("{")
):
self.buffered_json_content = (self.buffered_json_content or "") + text
return ModelResponseStream(
choices=[StreamingChoices(index=0, delta=Delta())],
usage=None,
)
reasoning_content: str | None = None
content: str | None = None
if text is not None:

View file

@ -904,3 +904,97 @@ 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 TestOllamaStreamingToolCallCluster:
"""Regression tests for https://github.com/BerriAI/litellm/issues/33678,
https://github.com/BerriAI/litellm/issues/35663 and https://github.com/BerriAI/litellm/issues/33622"""
def _tool_call_chunk(self, name, arguments, function_index=None):
function = {"name": name, "arguments": arguments}
if function_index is not None:
function["index"] = function_index
return {
"model": "qwen3",
"message": {"role": "assistant", "content": "", "tool_calls": [{"function": function}]},
"done": False,
}
def test_parallel_tool_calls_get_distinct_indices_from_function_index(self):
iterator = OllamaChatCompletionResponseIterator(streaming_response=iter([]), sync_stream=True)
first = iterator.chunk_parser(self._tool_call_chunk("read_file", {"path": "a.rs"}, function_index=0))
second = iterator.chunk_parser(self._tool_call_chunk("read_file", {"path": "b.rs"}, function_index=1))
assert first.choices[0].delta.tool_calls[0].index == 0
assert second.choices[0].delta.tool_calls[0].index == 1
def test_parallel_tool_calls_get_distinct_indices_without_function_index(self):
iterator = OllamaChatCompletionResponseIterator(streaming_response=iter([]), sync_stream=True)
first = iterator.chunk_parser(self._tool_call_chunk("read_file", {"path": "a.rs"}))
second = iterator.chunk_parser(self._tool_call_chunk("read_file", {"path": "b.rs"}))
assert first.choices[0].delta.tool_calls[0].index == 0
assert second.choices[0].delta.tool_calls[0].index == 1
def test_finish_reason_tool_calls_when_tool_call_arrives_before_done_chunk(self):
iterator = OllamaChatCompletionResponseIterator(streaming_response=iter([]), sync_stream=True)
iterator.chunk_parser(self._tool_call_chunk("get_weather", {"location": "Paris"}, function_index=0))
done = iterator.chunk_parser(
{
"model": "qwen3",
"message": {"role": "assistant", "content": ""},
"done": True,
"done_reason": "stop",
}
)
assert done.choices[0].finish_reason == "tool_calls"
def test_finish_reason_length_not_overridden_by_tool_calls(self):
iterator = OllamaChatCompletionResponseIterator(streaming_response=iter([]), sync_stream=True)
iterator.chunk_parser(self._tool_call_chunk("get_weather", {"location": "Paris"}, function_index=0))
done = iterator.chunk_parser(
{
"model": "qwen3",
"message": {"role": "assistant", "content": ""},
"done": True,
"done_reason": "length",
}
)
assert done.choices[0].finish_reason == "length"
def test_error_chunk_raises_ollama_error_with_provider_message(self):
from litellm.llms.ollama.common_utils import OllamaError
iterator = OllamaChatCompletionResponseIterator(streaming_response=iter([]), sync_stream=True)
with pytest.raises(OllamaError) as exc_info:
iterator.chunk_parser({"error": "error parsing tool call: invalid character ']'"})
assert "error parsing tool call" in str(exc_info.value)
assert "KeyError" not in str(exc_info.value)
def test_transform_response_error_dict_raises_ollama_error(self):
import httpx
from litellm.llms.ollama.common_utils import OllamaError
raw_response = httpx.Response(
200,
json={"error": "error parsing tool call: bad json"},
request=httpx.Request("POST", "http://localhost:11434/api/chat"),
)
with pytest.raises(OllamaError) as exc_info:
OllamaChatConfig().transform_response(
model="qwen3",
raw_response=raw_response,
model_response=ModelResponse(),
logging_obj=MagicMock(),
request_data={},
messages=[{"role": "user", "content": "hi"}],
optional_params={},
litellm_params={},
encoding=None,
)
assert "error parsing tool call" in str(exc_info.value)

View file

@ -502,3 +502,55 @@ class TestOllamaTextCompletionResponseIterator:
assert result["usage"]["prompt_tokens"] == 10
assert result["usage"]["completion_tokens"] == 5
assert result["usage"]["total_tokens"] == 15
class TestOllamaTextCompletionStreamingToolCalls:
"""Regression tests for https://github.com/BerriAI/litellm/issues/35711"""
def _stream(self, responses):
iterator = OllamaTextCompletionResponseIterator(streaming_response=iter([]), sync_stream=True)
chunks = [
iterator.chunk_parser({"model": "qwen3", "created_at": "t", "done": False, "response": r})
for r in responses
]
done = iterator.chunk_parser(
{
"model": "qwen3",
"created_at": "t",
"done": True,
"done_reason": "stop",
"response": "",
"prompt_eval_count": 10,
"eval_count": 5,
}
)
return chunks, done
def test_streamed_function_call_json_reconstructed_as_tool_call(self):
chunks, done = self._stream(['{"name": "get_weather",', ' "arguments": {"location": "Paris"}}'])
for chunk in chunks:
assert isinstance(chunk, ModelResponseStream)
assert not chunk.choices[0].delta.content
assert isinstance(done, ModelResponseStream)
tool_calls = done.choices[0].delta.tool_calls
assert tool_calls is not None and len(tool_calls) == 1
assert tool_calls[0].function.name == "get_weather"
assert json.loads(tool_calls[0].function.arguments) == {"location": "Paris"}
assert done.choices[0].finish_reason == "tool_calls"
def test_streamed_regular_json_emitted_as_content_on_done(self):
chunks, done = self._stream(['{"answer":', ' 42}'])
assert isinstance(done, ModelResponseStream)
assert done.choices[0].delta.tool_calls is None
assert done.choices[0].delta.content == '{"answer": 42}'
assert done.choices[0].finish_reason == "stop"
def test_plain_text_still_streams_incrementally(self):
chunks, done = self._stream(["Hello", " world"])
assert chunks[0].choices[0].delta.content == "Hello"
assert chunks[1].choices[0].delta.content == " world"
assert done["finish_reason"] == "stop"