This commit is contained in:
Arthi Arumugam 2026-08-27 18:50:40 -05:00 committed by GitHub
commit a6c804fe88
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 419 additions and 4 deletions

View file

@ -4,6 +4,7 @@ from collections.abc import AsyncIterator, Iterator
from typing import TYPE_CHECKING, Any, Final
from httpx._models import Headers, Response
from pydantic import BaseModel, ConfigDict, ValidationError
import litellm
from litellm._logging import verbose_proxy_logger
@ -18,7 +19,12 @@ from litellm.litellm_core_utils.prompt_templates.factory 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.openai import AllMessageValues, ChatCompletionUsageBlock
from litellm.types.llms.openai import (
AllMessageValues,
ChatCompletionToolCallChunk,
ChatCompletionToolCallFunctionChunk,
ChatCompletionUsageBlock,
)
from litellm.types.utils import (
Delta,
GenericStreamingChunk,
@ -100,6 +106,8 @@ class OllamaConfig(BaseConfig):
system: str | None = None
template: str | None = None
_prompted_tool_names: frozenset[str] = frozenset()
def __init__(
self,
mirostat: int | None = None,
@ -377,6 +385,7 @@ class OllamaConfig(BaseConfig):
format: Final = optional_params.pop("format", None)
images = optional_params.pop("images", None)
think: Final = optional_params.pop("think", None)
self._prompted_tool_names = _prompted_tool_names(optional_params.pop("prompted_tool_calls", None))
data: Final = {
"model": model,
"prompt": ollama_prompt,
@ -440,14 +449,89 @@ class OllamaConfig(BaseConfig):
streaming_response=streaming_response,
sync_stream=sync_stream,
json_mode=json_mode,
tool_names=self._prompted_tool_names,
)
def _prompted_tool_names(prompted_tools: object) -> frozenset[str]:
"""Names of the functions `get_optional_params` rewrote into this request's prompt."""
if not isinstance(prompted_tools, list):
return frozenset()
schemas: Final = ((tool.get("function") or tool) for tool in prompted_tools if isinstance(tool, dict))
return frozenset(
schema["name"] for schema in schemas if isinstance(schema, dict) and isinstance(schema.get("name"), str)
)
class _OllamaJsonModeToolCall(BaseModel):
"""
The exact object litellm's own tool prompt asks ollama for, in `factory.function_call_prompt`.
Extra keys are rejected so a JSON body that merely happens to carry a `name` stays content.
"""
model_config = ConfigDict(extra="forbid")
name: str
arguments: dict[str, object] # mutable-ok: mirrors Ollama's arbitrary JSON argument object
def _build_tool_call(response_text: str, tool_names: frozenset[str]) -> ChatCompletionToolCallChunk | None:
"""
Build a tool call out of a completed response body, but only for a function this request
actually offered. This is a stricter test than the one `OllamaConfig.transform_response`
applies to the same body, so it never converts anything that path would leave as content.
"""
try:
function_call: Final = _OllamaJsonModeToolCall.model_validate_json(response_text)
except ValidationError:
return None
if function_call.name not in tool_names:
return None
return ChatCompletionToolCallChunk(
id=f"call_{uuid.uuid4()}",
type="function",
function=ChatCompletionToolCallFunctionChunk(
name=function_call.name,
arguments=json.dumps(function_call.arguments),
),
index=0,
)
class OllamaTextCompletionResponseIterator(BaseModelResponseIterator):
def __init__(self, streaming_response, sync_stream: bool, json_mode: bool | None = False):
def __init__(
self,
streaming_response,
sync_stream: bool,
json_mode: bool | None = False,
tool_names: frozenset[str] = frozenset(),
):
super().__init__(streaming_response, sync_stream, json_mode)
self.started_reasoning_content: bool = False
self.finished_reasoning_content: bool = False
self.tool_names: frozenset[str] = tool_names
self.tool_call_buffer: str | None = "" if tool_names else None
def _hold_back_or_release(self, content: str) -> str | None:
"""
`/api/generate` returns a tool call as a JSON object inside the response text, so it can
only be recognised once that object is complete. Hold text back while it can still turn
out to be one, and release everything held as soon as it cannot.
"""
if self.tool_call_buffer is None:
return content
self.tool_call_buffer += content
stripped: Final = self.tool_call_buffer.lstrip()
if not stripped or stripped.startswith("{"):
return None
released: Final = self.tool_call_buffer
self.tool_call_buffer = None
return released
def _handle_string_chunk(self, str_line: str) -> GenericStreamingChunk | ModelResponseStream:
return self.chunk_parser(json.loads(str_line))
@ -474,8 +558,21 @@ class OllamaTextCompletionResponseIterator(BaseModelResponseIterator):
completion_tokens=eval_count,
total_tokens=prompt_eval_count + eval_count,
)
held_back: Final = self.tool_call_buffer
self.tool_call_buffer = None
tool_call: Final = _build_tool_call(held_back, self.tool_names) if held_back else None
if tool_call is not None:
return GenericStreamingChunk(
text="",
tool_use=tool_call,
is_finished=True,
finish_reason="tool_calls",
usage=usage,
)
return GenericStreamingChunk(
text=text,
text=held_back or text,
is_finished=is_finished,
finish_reason=finish_reason,
usage=usage,
@ -495,7 +592,7 @@ class OllamaTextCompletionResponseIterator(BaseModelResponseIterator):
if self.started_reasoning_content and not self.finished_reasoning_content:
reasoning_content = text
else:
content = text
content = self._hold_back_or_release(text)
return ModelResponseStream(
choices=[

View file

@ -4051,6 +4051,7 @@ def pre_process_optional_params(passed_params: dict, non_default_params: dict, c
non_default_params.pop("tool_choice", None) # causes ollama requests to hang
elif "functions" in non_default_params:
optional_params["functions_unsupported_model"] = non_default_params.pop("functions")
optional_params["prompted_tool_calls"] = optional_params.get("functions_unsupported_model") or [] # mutable-ok: fresh per-request list; empty means no tools offered
elif litellm.add_function_to_prompt: # if user opts to add it to prompt instead
optional_params["functions_unsupported_model"] = non_default_params.pop(
"tools", non_default_params.pop("functions", None)

View file

@ -502,3 +502,320 @@ class TestOllamaTextCompletionResponseIterator:
assert result["usage"]["prompt_tokens"] == 10
assert result["usage"]["completion_tokens"] == 5
assert result["usage"]["total_tokens"] == 15
def test_chunk_parser_streams_tool_call_instead_of_raw_json(self):
"""A format=json tool call must arrive as tool_calls, not as JSON in the content."""
iterator = OllamaTextCompletionResponseIterator(
streaming_response=iter([]),
sync_stream=True,
json_mode=False,
tool_names=frozenset({"lookup_account", "lookup"}),
)
fragments = [
'{"',
"name",
'": "',
"lookup_account",
'", "',
"arguments",
'": {"',
"email",
'": "',
"maya.iyer@example.com",
'"}}',
]
streamed_content = ""
for fragment in fragments:
result = iterator.chunk_parser({"model": "qwen2.5:3b", "response": fragment, "done": False})
assert isinstance(result, ModelResponseStream)
assert result.choices and result.choices[0].delta is not None
streamed_content += result.choices[0].delta.content or ""
assert streamed_content == ""
result = iterator.chunk_parser(
{
"model": "qwen2.5:3b",
"response": "",
"done": True,
"prompt_eval_count": 120,
"eval_count": 24,
}
)
assert result["text"] == ""
assert result["finish_reason"] == "tool_calls"
assert result["is_finished"] is True
assert result["tool_use"] is not None
assert result["tool_use"]["type"] == "function"
assert result["tool_use"]["index"] == 0
assert result["tool_use"]["id"]
assert result["tool_use"]["function"]["name"] == "lookup_account"
assert json.loads(result["tool_use"]["function"]["arguments"]) == {"email": "maya.iyer@example.com"}
def test_chunk_parser_streams_prose_incrementally(self):
"""Prose must keep streaming fragment by fragment and carry no tool call."""
iterator = OllamaTextCompletionResponseIterator(
streaming_response=iter([]),
sync_stream=True,
json_mode=False,
tool_names=frozenset({"lookup_account", "lookup"}),
)
fragments = ["I ", "will ", "look ", "that ", "up", "."]
deltas = []
for fragment in fragments:
result = iterator.chunk_parser({"model": "qwen2.5:3b", "response": fragment, "done": False})
assert isinstance(result, ModelResponseStream)
assert result.choices and result.choices[0].delta is not None
deltas.append(result.choices[0].delta.content)
assert deltas == fragments
result = iterator.chunk_parser(
{
"model": "qwen2.5:3b",
"response": "",
"done": True,
"prompt_eval_count": 10,
"eval_count": 6,
}
)
assert result["text"] == ""
assert result["finish_reason"] == "stop"
assert result.get("tool_use") is None
def test_chunk_parser_releases_non_tool_call_json_as_content(self):
"""A JSON body that is not a tool call must still be delivered as content."""
iterator = OllamaTextCompletionResponseIterator(
streaming_response=iter([]),
sync_stream=True,
json_mode=False,
tool_names=frozenset({"lookup_account", "lookup"}),
)
fragments = ['{"', "city", '": "', "Chennai", '"}']
streamed_content = ""
for fragment in fragments:
result = iterator.chunk_parser({"model": "qwen2.5:3b", "response": fragment, "done": False})
assert isinstance(result, ModelResponseStream)
assert result.choices and result.choices[0].delta is not None
streamed_content += result.choices[0].delta.content or ""
result = iterator.chunk_parser(
{
"model": "qwen2.5:3b",
"response": "",
"done": True,
"prompt_eval_count": 10,
"eval_count": 6,
}
)
assert streamed_content + result["text"] == '{"city": "Chennai"}'
assert result["finish_reason"] == "stop"
assert result.get("tool_use") is None
def test_chunk_parser_leaves_json_carrying_extra_keys_as_content(self):
"""A JSON body that merely happens to carry a name must not be mistaken for a tool call."""
iterator = OllamaTextCompletionResponseIterator(
streaming_response=iter([]),
sync_stream=True,
json_mode=False,
tool_names=frozenset({"lookup_account", "lookup"}),
)
body = '{"name": "Chennai", "arguments": {"x": 1}, "population": 7000000}'
streamed_content = ""
for fragment in [body[i : i + 6] for i in range(0, len(body), 6)]:
result = iterator.chunk_parser({"model": "qwen2.5:3b", "response": fragment, "done": False})
assert isinstance(result, ModelResponseStream)
assert result.choices and result.choices[0].delta is not None
streamed_content += result.choices[0].delta.content or ""
result = iterator.chunk_parser(
{"model": "qwen2.5:3b", "response": "", "done": True, "prompt_eval_count": 5, "eval_count": 9}
)
assert streamed_content + result["text"] == body
assert result["finish_reason"] == "stop"
assert result.get("tool_use") is None
def test_chunk_parser_leaves_non_object_arguments_as_content(self):
"""`arguments` must be the object the tool prompt asks for, not any value that happens to be there."""
iterator = OllamaTextCompletionResponseIterator(
streaming_response=iter([]),
sync_stream=True,
json_mode=False,
tool_names=frozenset({"lookup_account", "lookup"}),
)
body = '{"name": "lookup", "arguments": "not an object"}'
streamed_content = ""
for fragment in [body[i : i + 6] for i in range(0, len(body), 6)]:
result = iterator.chunk_parser({"model": "qwen2.5:3b", "response": fragment, "done": False})
assert isinstance(result, ModelResponseStream)
assert result.choices and result.choices[0].delta is not None
streamed_content += result.choices[0].delta.content or ""
result = iterator.chunk_parser(
{"model": "qwen2.5:3b", "response": "", "done": True, "prompt_eval_count": 5, "eval_count": 9}
)
assert streamed_content + result["text"] == body
assert result["finish_reason"] == "stop"
assert result.get("tool_use") is None
TOOL_CALL_BODY = '{"name": "lookup_account", "arguments": {"email": "maya.iyer@example.com"}}'
LOOKUP_ACCOUNT_TOOL = {
"type": "function",
"function": {
"name": "lookup_account",
"description": "Look up a customer account by email address",
"parameters": {"type": "object", "properties": {"email": {"type": "string"}}, "required": ["email"]},
},
}
def _generate_transport():
"""Stands in for ollama's /api/generate, fragmenting the body the way it really does."""
import httpx
done = {"model": "qwen2.5:3b", "response": "", "done": True, "prompt_eval_count": 146, "eval_count": 21}
def handler(request: "httpx.Request") -> "httpx.Response":
fragments = [TOOL_CALL_BODY[i : i + 4] for i in range(0, len(TOOL_CALL_BODY), 4)]
lines = [json.dumps({"model": "qwen2.5:3b", "response": f, "done": False}) for f in fragments]
lines.append(json.dumps(done))
return httpx.Response(200, content=("\n".join(lines) + "\n").encode())
return httpx.MockTransport(handler)
class TestOllamaStreamingToolCallsEndToEnd:
def test_streamed_tool_call_is_not_delivered_as_content(self):
"""The reported bug: a streamed tool call reaching the caller as raw JSON in the content."""
import httpx
from litellm import completion
from litellm.llms.custom_httpx.http_handler import HTTPHandler
client = HTTPHandler(client=httpx.Client(transport=_generate_transport()))
response = completion(
model="ollama/qwen2.5:3b",
messages=[{"role": "user", "content": "Cancel the subscription for maya.iyer@example.com"}],
tools=[LOOKUP_ACCOUNT_TOOL],
api_base="http://localhost:11434",
stream=True,
client=client,
)
streamed_content = ""
tool_calls = []
finish_reason = None
for chunk in response:
delta = chunk.choices[0].delta
streamed_content += getattr(delta, "content", None) or ""
tool_calls += getattr(delta, "tool_calls", None) or []
finish_reason = chunk.choices[0].finish_reason or finish_reason
assert streamed_content == ""
assert finish_reason == "tool_calls"
assert len(tool_calls) == 1
assert tool_calls[0].function.name == "lookup_account"
assert json.loads(tool_calls[0].function.arguments) == {"email": "maya.iyer@example.com"}
def test_detection_is_off_for_a_request_that_offered_no_tools(self):
"""No tools were rewritten into a prompt, so a tool-shaped body is just content."""
iterator = OllamaTextCompletionResponseIterator(streaming_response=iter([]), sync_stream=True, json_mode=False)
streamed_content = ""
for fragment in [TOOL_CALL_BODY[i : i + 6] for i in range(0, len(TOOL_CALL_BODY), 6)]:
result = iterator.chunk_parser({"model": "qwen2.5:3b", "response": fragment, "done": False})
assert isinstance(result, ModelResponseStream)
assert result.choices and result.choices[0].delta is not None
streamed_content += result.choices[0].delta.content or ""
result = iterator.chunk_parser(
{"model": "qwen2.5:3b", "response": "", "done": True, "prompt_eval_count": 5, "eval_count": 9}
)
assert streamed_content == TOOL_CALL_BODY
assert result["finish_reason"] == "stop"
assert result.get("tool_use") is None
def test_detection_is_armed_only_with_the_functions_the_request_offered(self):
"""`get_optional_params` passes on the functions it rewrote; those are the only ones accepted."""
from litellm.utils import get_optional_params
config = OllamaConfig()
base = {
"model": "qwen2.5:3b",
"messages": [{"role": "user", "content": "hi"}],
"litellm_params": {},
"headers": {},
}
json_mode_only = get_optional_params(
model="qwen2.5:3b",
custom_llm_provider="ollama",
stream=True,
response_format={"type": "json_object"},
)
assert "prompted_tool_calls" not in json_mode_only
config.transform_request(optional_params=json_mode_only, **base)
assert config.get_model_response_iterator(iter([]), sync_stream=True).tool_names == frozenset()
with_tools = get_optional_params(
model="qwen2.5:3b",
custom_llm_provider="ollama",
stream=True,
tools=[LOOKUP_ACCOUNT_TOOL],
)
assert with_tools["prompted_tool_calls"] == [LOOKUP_ACCOUNT_TOOL]
config.transform_request(optional_params=with_tools, **base)
assert config.get_model_response_iterator(iter([]), sync_stream=True).tool_names == frozenset(
{"lookup_account"}
)
def test_a_function_the_request_never_offered_stays_content(self):
"""A body naming some other function must not be synthesised into a tool call."""
iterator = OllamaTextCompletionResponseIterator(
streaming_response=iter([]), sync_stream=True, json_mode=False, tool_names=frozenset({"lookup_account"})
)
body = '{"name": "delete_account", "arguments": {"email": "maya.iyer@example.com"}}'
streamed_content = ""
for fragment in [body[i : i + 6] for i in range(0, len(body), 6)]:
result = iterator.chunk_parser({"model": "qwen2.5:3b", "response": fragment, "done": False})
assert isinstance(result, ModelResponseStream)
assert result.choices and result.choices[0].delta is not None
streamed_content += result.choices[0].delta.content or ""
result = iterator.chunk_parser(
{"model": "qwen2.5:3b", "response": "", "done": True, "prompt_eval_count": 5, "eval_count": 9}
)
assert streamed_content + result["text"] == body
assert result["finish_reason"] == "stop"
assert result.get("tool_use") is None
def test_prompted_tool_calls_marker_is_not_sent_to_ollama(self):
"""The marker is litellm-internal; it must not leak into the request body."""
config = OllamaConfig()
data = config.transform_request(
model="qwen2.5:3b",
messages=[{"role": "user", "content": "hi"}],
optional_params={"stream": True, "format": "json", "prompted_tool_calls": True},
litellm_params={},
headers={},
)
assert "prompted_tool_calls" not in data
assert "prompted_tool_calls" not in data["options"]