From 2b8e53b5cafc1119345b406ee4928d6069ef975c Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Wed, 7 Oct 2026 11:05:50 -0400 Subject: [PATCH] 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 Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../prompt_templates/factory.py | 24 ++--- .../llms/ollama/completion/transformation.py | 66 +++++++++++++- ...llm_core_utils_prompt_templates_factory.py | 30 ++++++ .../test_ollama_completion_transformation.py | 91 ++++++++++++++++++- 4 files changed, 191 insertions(+), 20 deletions(-) diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py index f1eb9a0aef4..017a75dca2e 100644 --- a/litellm/litellm_core_utils/prompt_templates/factory.py +++ b/litellm/litellm_core_utils/prompt_templates/factory.py @@ -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 diff --git a/litellm/llms/ollama/completion/transformation.py b/litellm/llms/ollama/completion/transformation.py index 309589b367a..e307b8ca07a 100644 --- a/litellm/llms/ollama/completion/transformation.py +++ b/litellm/llms/ollama/completion/transformation.py @@ -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=[ diff --git a/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py b/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py index 29f66c09f16..a2987730851 100644 --- a/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py +++ b/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py @@ -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 = [ diff --git a/tests/unit/llms/ollama/test_ollama_completion_transformation.py b/tests/unit/llms/ollama/test_ollama_completion_transformation.py index d6215a742f0..8558bf50bb9 100644 --- a/tests/unit/llms/ollama/test_ollama_completion_transformation.py +++ b/tests/unit/llms/ollama/test_ollama_completion_transformation.py @@ -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"