fix(ollama): turn streamed prompt-based JSON tool calls into real tool calls (#45053)

* fix(ollama): turn streamed prompt-based JSON tool calls into real tool calls

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(ollama): separate replayed tool calls from text and tighten parser typing

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

---------

Co-authored-by: Mubashir Osmani <mubashir@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-10-07 11:05:50 -04:00 • committed by GitHub
parent 736ff14f11
commit 2b8e53b5ca
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 191 additions and 20 deletions

View file

@ -250,26 +250,18 @@ def ollama_pt(
assistant_content_str += convert_content_list_to_str(messages[msg_i])
tool_calls = messages[msg_i].get("tool_calls")
ollama_tool_calls = []
if tool_calls:
for call in tool_calls:
call_id: str = call["id"]
function_name: str = call["function"]["name"]
arguments = json.loads(call["function"]["arguments"])
ollama_tool_calls.append(
if assistant_content_str:
assistant_content_str += "\n"
assistant_content_str += "\n".join(
json.dumps(
{
"id": call_id,
"type": "function",
"function": {
"name": function_name,
"arguments": arguments,
},
"name": call["function"]["name"],
"arguments": json.loads(call["function"]["arguments"]),
}
)
if ollama_tool_calls:
assistant_content_str += f"Tool Calls: {json.dumps(ollama_tool_calls, indent=2)}"
for call in tool_calls
)
msg_i += 1

View file

@ -1,4 +1,5 @@
import json
import re
import time
from collections.abc import AsyncIterator, Iterator
from typing import TYPE_CHECKING, Any, Final
@ -25,7 +26,12 @@ from litellm.litellm_core_utils.prompt_templates.image_handling import (
from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator
from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException
from litellm.types.llms.base import LiteLLMBaseModel
from litellm.types.llms.openai import AllMessageValues, ChatCompletionUsageBlock
from litellm.types.llms.openai import (
AllMessageValues,
ChatCompletionToolCallChunk,
ChatCompletionToolCallFunctionChunk,
ChatCompletionUsageBlock,
)
from litellm.types.utils import (
Delta,
GenericStreamingChunk,
@ -90,6 +96,14 @@ class _OllamaGenerateResponse(TypedDict):
_OLLAMA_GENERATE_RESPONSE: Final = TypeAdapter(_OllamaGenerateResponse)
_JSON_CODE_FENCE: Final = re.compile(r"^```(?:json)?\s*(.*?)\s*(?:```)?$", re.DOTALL)
_JSON_OBJECT_START: Final = re.compile(r"(?:```(?:json)?\s*)?\{")
_JSON_OBJECT_PARTIAL_START: Final = re.compile(r"`{0,3}|```(?:j|js|jso|json)?\s*")
def _strip_json_code_fence(text: str) -> str:
fenced: Final = _JSON_CODE_FENCE.match(text.strip())
return fenced.group(1) if fenced else text
class OllamaConfig(BaseConfig):
@ -319,7 +333,7 @@ class OllamaConfig(BaseConfig):
model_response.choices[0].finish_reason = "stop"
else:
try:
response_content: Final[object] = json.loads(response_text)
response_content: Final[object] = json.loads(_strip_json_code_fence(response_text))
# Check if this is a function call format with name/arguments structure
if (
@ -512,6 +526,50 @@ class OllamaTextCompletionResponseIterator(BaseModelResponseIterator):
super().__init__(streaming_response, sync_stream, json_mode)
self.started_reasoning_content: bool = False
self.finished_reasoning_content: bool = False
self.streamed_content: bool = False
self.held_content: str = ""
self.holding_json_object: bool = False
def _hold_json_object_start(self, content: str) -> str | None:
if self.streamed_content:
return content
self.held_content += content
if self.holding_json_object:
return None
candidate: Final = self.held_content.lstrip()
if _JSON_OBJECT_START.match(candidate):
self.holding_json_object = True
return None
if _JSON_OBJECT_PARTIAL_START.fullmatch(candidate):
return None
self.streamed_content = True
released: Final = self.held_content
self.held_content = ""
return released
def _flush_held_content(self, usage: ChatCompletionUsageBlock | None) -> GenericStreamingChunk:
held: Final = self.held_content
self.held_content = ""
try:
parsed: Final[object] = json.loads(_strip_json_code_fence(held))
except json.JSONDecodeError:
return GenericStreamingChunk(text=held, is_finished=True, finish_reason="stop", usage=usage)
if isinstance(parsed, dict) and "name" in parsed and "arguments" in parsed:
return GenericStreamingChunk(
text="",
tool_use=ChatCompletionToolCallChunk(
id=f"call_{uuid.uuid4()}",
type="function",
function=ChatCompletionToolCallFunctionChunk(
name=parsed["name"], arguments=json.dumps(parsed["arguments"])
),
index=0,
),
is_finished=True,
finish_reason="tool_calls",
usage=usage,
)
return GenericStreamingChunk(text=held, is_finished=True, finish_reason="stop", usage=usage)
def _handle_string_chunk(self, str_line: str) -> GenericStreamingChunk | ModelResponseStream:
return self.chunk_parser(json.loads(str_line))
@ -538,6 +596,8 @@ class OllamaTextCompletionResponseIterator(BaseModelResponseIterator):
completion_tokens=eval_count,
total_tokens=prompt_eval_count + eval_count,
)
if self.held_content:
return self._flush_held_content(usage)
return GenericStreamingChunk(
text=text,
is_finished=is_finished,
@ -559,7 +619,7 @@ class OllamaTextCompletionResponseIterator(BaseModelResponseIterator):
if self.started_reasoning_content and not self.finished_reasoning_content:
reasoning_content = text
else:
content = text
content = self._hold_json_object_start(text)
return ModelResponseStream(
choices=[

View file

@ -134,6 +134,36 @@ def test_ollama_pt_simple_messages():
assert result["images"] == []
@pytest.mark.parametrize(
("assistant_content", "rendered_prefix"),
[("", ""), ("Checking calc.py", "Checking calc.py\n")],
)
def test_ollama_pt_renders_tool_calls_in_function_prompt_format(assistant_content: str, rendered_prefix: str):
messages: Final = [
{"role": "user", "content": "Fix calc.py"},
{
"role": "assistant",
"content": assistant_content,
"tool_calls": [
{
"id": "call_1",
"type": "function",
"function": {"name": "read", "arguments": '{"filePath": "calc.py"}'},
}
],
},
{"role": "tool", "tool_call_id": "call_1", "content": "def add(a, b): return a - b"},
]
result: Final = ollama_pt(model="gemma4:31b", messages=messages)
assert result["prompt"] == (
"### User:\nFix calc.py\n\n"
f'### Assistant:\n{rendered_prefix}{{"name": "read", "arguments": {{"filePath": "calc.py"}}}}\n\n'
"### User:\ndef add(a, b): return a - b\n\n"
)
def test_ollama_pt_consecutive_user_messages():
"""Test handling consecutive user messages"""
messages = [

View file

@ -2,13 +2,14 @@ import base64
import io
import json
import sys
from litellm._uuid import uuid
from typing import Final
from unittest.mock import MagicMock, patch
import httpx
import pytest
import litellm
from litellm._uuid import uuid
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.llms.ollama.completion.transformation import (
OllamaConfig,
@ -116,6 +117,28 @@ class TestOllamaConfig:
) == {"location": "San Francisco"}
# No usage assertions here as we don't need to test them in every case
def test_transform_response_json_function_call_in_code_fence(self):
raw_response: Final = MagicMock()
raw_response.json.return_value = {
"response": '```json\n{"name": "read", "arguments": {"filePath": "calc.py"}}\n```'
}
result: Final = OllamaConfig().transform_response(
model="gemma4:31b",
raw_response=raw_response,
model_response=ModelResponse(choices=[{"message": Message(content="")}]),
logging_obj=MagicMock(),
request_data={"format": "json"},
messages=[],
optional_params={},
litellm_params={},
encoding=MagicMock(),
)
tool_call: Final = result.choices[0].message.tool_calls[0]
assert (result.choices[0].finish_reason, tool_call.function.name) == ("tool_calls", "read")
assert json.loads(tool_call.function.arguments) == {"filePath": "calc.py"}
def test_transform_response_regular_json(self):
# Initialize config
config = OllamaConfig()
@ -665,6 +688,72 @@ class TestOllamaTextCompletionResponseIterator:
assert result["usage"]["completion_tokens"] == 5
assert result["usage"]["total_tokens"] == 15
@pytest.mark.parametrize(
("response_chunks", "expected_text", "expected_tool_name", "expected_finish_reason"),
[
(
['{"name": "write_file", ', '"arguments": {"path": "hello.py"}}'],
"",
"write_file",
"tool_calls",
),
(
["```", "json", "\n", '{"name": "write_file", "arguments": {"path": "hello.py"}}', "\n```"],
"",
"write_file",
"tool_calls",
),
(['{"answer": ', "42}"], '{"answer": 42}', None, "stop"),
(["{not json"], "{not json", None, "stop"),
],
)
def test_chunk_parser_turns_streamed_json_function_call_into_tool_call(
self,
response_chunks: list[str],
expected_text: str,
expected_tool_name: str | None,
expected_finish_reason: str,
):
iterator: Final = OllamaTextCompletionResponseIterator(
streaming_response=iter([]), sync_stream=True, json_mode=False
)
streamed: Final = [iterator.chunk_parser({"response": text, "done": False}) for text in response_chunks]
done: Final = iterator.chunk_parser({"response": "", "done": True, "prompt_eval_count": 3, "eval_count": 2})
assert [chunk.choices[0].delta.content for chunk in streamed] == [None] * len(response_chunks)
assert done["text"] == expected_text
assert done["finish_reason"] == expected_finish_reason
assert done["usage"]["total_tokens"] == 5
tool_use: Final = done.get("tool_use")
assert (tool_use["function"]["name"] if tool_use else None) == expected_tool_name
if tool_use:
assert json.loads(tool_use["function"]["arguments"]) == {"path": "hello.py"}
def test_chunk_parser_streams_text_that_only_later_contains_braces(self):
iterator: Final = OllamaTextCompletionResponseIterator(
streaming_response=iter([]), sync_stream=True, json_mode=False
)
first: Final = iterator.chunk_parser({"response": "Here is code: ", "done": False})
second: Final = iterator.chunk_parser({"response": '{"a": 1}', "done": False})
done: Final = iterator.chunk_parser({"response": "", "done": True})
assert (first.choices[0].delta.content, second.choices[0].delta.content) == ("Here is code: ", '{"a": 1}')
assert (done["text"], done["finish_reason"], done.get("tool_use")) == ("", "stop", None)
def test_chunk_parser_releases_held_code_fence_that_is_not_json(self):
iterator: Final = OllamaTextCompletionResponseIterator(
streaming_response=iter([]), sync_stream=True, json_mode=False
)
contents: Final = [
iterator.chunk_parser({"response": text, "done": False}).choices[0].delta.content
for text in ("```", "python\n", "print(1)")
]
assert contents == [None, "```python\n", "print(1)"]
async def test_ollama_async_completion_inlines_remote_images_off_the_event_loop(async_only_image_fetch):
image_url = f"https://img.example/{uuid.uuid4()}.png"