mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
parent
1249f84b10
commit
fa7b3f74d9
4 changed files with 226 additions and 6 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue