mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
refactor(ollama): build streaming tool calls without mutating the provider chunk
The index/id normalizer rewrote the caller's tool call dict in place and popped the nested index out of it. Return a typed ChatCompletionDeltaToolCall instead, which clears the two LIT011 parameter-mutation violations the helper introduced and leaves the emitted chunk shape unchanged. Also stop a legitimate eval_count of 0 from falling back to a token estimate, and drop the text-is-not-None guard the /api/generate branch no longer needs now that it returns early when a fragment is held back.
This commit is contained in:
parent
d32e9748c0
commit
f18be3f883
3 changed files with 91 additions and 31 deletions
|
|
@ -1,6 +1,6 @@
|
|||
import json
|
||||
import time
|
||||
from collections.abc import AsyncIterator, Iterator
|
||||
from collections.abc import AsyncIterator, Iterator, Mapping
|
||||
from typing import TYPE_CHECKING, Any, Final, cast
|
||||
|
||||
from httpx._models import Headers, Response
|
||||
|
|
@ -26,7 +26,12 @@ from litellm.types.llms.openai import (
|
|||
ChatCompletionAssistantToolCall,
|
||||
ChatCompletionUsageBlock,
|
||||
)
|
||||
from litellm.types.utils import ModelResponse, ModelResponseStream
|
||||
from litellm.types.utils import (
|
||||
ChatCompletionDeltaToolCall,
|
||||
Function,
|
||||
ModelResponse,
|
||||
ModelResponseStream,
|
||||
)
|
||||
|
||||
from ..common_utils import OllamaError
|
||||
|
||||
|
|
@ -401,7 +406,10 @@ class OllamaChatConfig(BaseConfig):
|
|||
model_response.model = "ollama_chat/" + model
|
||||
prompt_tokens = response_json.get("prompt_eval_count", litellm.token_counter(messages=messages))
|
||||
_message_content: Final = response_json_message.get("content") if response_json_message is not None else None
|
||||
completion_tokens: Final = response_json.get("eval_count") or litellm.token_counter(text=_message_content or "")
|
||||
_eval_count: Final = response_json.get("eval_count")
|
||||
completion_tokens: Final = (
|
||||
_eval_count if _eval_count is not None else litellm.token_counter(text=_message_content or "")
|
||||
)
|
||||
setattr(
|
||||
model_response,
|
||||
"usage",
|
||||
|
|
@ -434,20 +442,31 @@ class OllamaChatCompletionResponseIterator(BaseModelResponseIterator):
|
|||
finished_reasoning_content: bool = False
|
||||
stream_tool_call_count: int = 0
|
||||
|
||||
def _assign_tool_call_index_and_id(
|
||||
self,
|
||||
tool_call: dict, # mutable-ok: normalizes the provider chunk's tool call dict in place
|
||||
) -> None:
|
||||
def _normalized_tool_call(self, tool_call: Mapping[str, Any]) -> ChatCompletionDeltaToolCall | None:
|
||||
"""Ollama nests the parallel-call ordinal under `function.index`, where OpenAI clients expect it
|
||||
on the tool call itself. Returns None for a chunk carrying no function to normalize."""
|
||||
function: Final = tool_call.get("function")
|
||||
if function is None:
|
||||
return
|
||||
function_index: Final = 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)
|
||||
if not isinstance(function, Mapping):
|
||||
return None
|
||||
index: Final = self._resolve_tool_call_index(tool_call.get("index"), function.get("index"))
|
||||
self.stream_tool_call_count = max(self.stream_tool_call_count + 1, index + 1)
|
||||
function_args: Final = function.get("arguments")
|
||||
if function_args is not None and len(function_args) > 0 and self._is_function_call_complete(function_args):
|
||||
tool_call["id"] = str(uuid.uuid4())
|
||||
is_complete: Final = (
|
||||
function_args is not None and len(function_args) > 0 and self._is_function_call_complete(function_args)
|
||||
)
|
||||
return ChatCompletionDeltaToolCall(
|
||||
id=str(uuid.uuid4()) if is_complete else None,
|
||||
index=index,
|
||||
type="function",
|
||||
function=Function(name=function.get("name"), arguments=function_args),
|
||||
)
|
||||
|
||||
def _resolve_tool_call_index(self, index: int | None, nested_index: int | None) -> int:
|
||||
if index is not None:
|
||||
return index
|
||||
if nested_index is not None:
|
||||
return nested_index
|
||||
return self.stream_tool_call_count
|
||||
|
||||
def _is_function_call_complete(self, function_args: str | dict) -> bool:
|
||||
if isinstance(function_args, dict):
|
||||
|
|
@ -497,11 +516,16 @@ class OllamaChatCompletionResponseIterator(BaseModelResponseIterator):
|
|||
headers=Headers(),
|
||||
)
|
||||
|
||||
# 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:
|
||||
self._assign_tool_call_index_and_id(tool_call)
|
||||
raw_tool_calls: Final = chunk["message"].get("tool_calls")
|
||||
tool_calls: Final = (
|
||||
tuple(
|
||||
normalized
|
||||
for normalized in (self._normalized_tool_call(tool_call) for tool_call in raw_tool_calls)
|
||||
if normalized is not None
|
||||
)
|
||||
if raw_tool_calls is not None
|
||||
else None
|
||||
)
|
||||
|
||||
# PROCESS REASONING CONTENT
|
||||
reasoning_content: str | None = None
|
||||
|
|
|
|||
|
|
@ -566,18 +566,17 @@ class OllamaTextCompletionResponseIterator(BaseModelResponseIterator):
|
|||
)
|
||||
reasoning_content: str | None = None
|
||||
content: str | None = None
|
||||
if text is not None:
|
||||
if "<think>" in text:
|
||||
text = text.replace("<think>", "")
|
||||
self.started_reasoning_content = True
|
||||
elif "</think>" in text:
|
||||
text = text.replace("</think>", "")
|
||||
self.finished_reasoning_content = True
|
||||
if "<think>" in text:
|
||||
text = text.replace("<think>", "")
|
||||
self.started_reasoning_content = True
|
||||
elif "</think>" in text:
|
||||
text = text.replace("</think>", "")
|
||||
self.finished_reasoning_content = True
|
||||
|
||||
if self.started_reasoning_content and not self.finished_reasoning_content:
|
||||
reasoning_content = text
|
||||
else:
|
||||
content = text
|
||||
if self.started_reasoning_content and not self.finished_reasoning_content:
|
||||
reasoning_content = text
|
||||
else:
|
||||
content = text
|
||||
|
||||
return ModelResponseStream(
|
||||
choices=[
|
||||
|
|
|
|||
|
|
@ -974,6 +974,43 @@ class TestOllamaStreamingToolCallCluster:
|
|||
assert "error parsing tool call" in str(exc_info.value)
|
||||
assert "KeyError" not in str(exc_info.value)
|
||||
|
||||
def test_normalizing_a_tool_call_does_not_mutate_the_provider_chunk(self):
|
||||
iterator = OllamaChatCompletionResponseIterator(streaming_response=iter([]), sync_stream=True)
|
||||
chunk = self._tool_call_chunk("read_file", {"path": "a.rs"}, function_index=3)
|
||||
|
||||
iterator.chunk_parser(chunk)
|
||||
|
||||
assert chunk["message"]["tool_calls"][0] == {"function": {"name": "read_file", "arguments": {"path": "a.rs"}, "index": 3}}
|
||||
|
||||
def test_zero_eval_count_is_not_replaced_by_a_token_estimate(self):
|
||||
import httpx
|
||||
|
||||
raw_response = httpx.Response(
|
||||
200,
|
||||
json={
|
||||
"model": "qwen3",
|
||||
"message": {"role": "assistant", "content": "some words the estimator would count"},
|
||||
"done": True,
|
||||
"done_reason": "stop",
|
||||
"prompt_eval_count": 7,
|
||||
"eval_count": 0,
|
||||
},
|
||||
request=httpx.Request("POST", "http://localhost:11434/api/chat"),
|
||||
)
|
||||
response = 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 response.usage.completion_tokens == 0
|
||||
|
||||
def test_transform_response_error_dict_raises_ollama_error(self):
|
||||
import httpx
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue