mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
736ff14f11
commit
2b8e53b5ca
4 changed files with 191 additions and 20 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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=[
|
||||
|
|
|
|||
|
|
@ -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 = [
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue