mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(ollama): only reconstruct streaming tool calls when JSON output was requested
The /api/generate streaming buffer reconstructed a prompted function call from any
response that started with {"name", with no check that the caller had asked for one.
A plain stream with no tools and no response_format whose output happened to parse as
{"name": ..., "arguments": ...} lost its content entirely and came back as a
synthesized tool call with finish_reason=tool_calls.
The non-streaming path already gates that reconstruction on format=json, which litellm
sets for ollama whenever tools are passed. Carry the same signal from the request
transform into the streaming iterator so both paths agree.
The flag rides on the config rather than the iterator's json_mode argument because the
async streaming handler builds the iterator without forwarding json_mode, so a request
through the proxy would never see it.
This commit is contained in:
parent
5377128bb2
commit
d32e9748c0
2 changed files with 90 additions and 24 deletions
|
|
@ -103,6 +103,7 @@ class OllamaConfig(BaseConfig):
|
|||
top_p: float | None = None
|
||||
system: str | None = None
|
||||
template: str | None = None
|
||||
json_output_requested: bool = False
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
|
|
@ -380,6 +381,7 @@ class OllamaConfig(BaseConfig):
|
|||
ollama_prompt = modified_prompt
|
||||
stream: Final = optional_params.pop("stream", False)
|
||||
format: Final = optional_params.pop("format", None)
|
||||
self.json_output_requested = format == "json"
|
||||
images = optional_params.pop("images", None)
|
||||
think: Final = optional_params.pop("think", None)
|
||||
data: Final = {
|
||||
|
|
@ -444,7 +446,7 @@ class OllamaConfig(BaseConfig):
|
|||
return OllamaTextCompletionResponseIterator(
|
||||
streaming_response=streaming_response,
|
||||
sync_stream=sync_stream,
|
||||
json_mode=json_mode,
|
||||
json_mode=json_mode or self.json_output_requested,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -454,7 +456,7 @@ class OllamaTextCompletionResponseIterator(BaseModelResponseIterator):
|
|||
self.started_reasoning_content: bool = False
|
||||
self.finished_reasoning_content: bool = False
|
||||
self.buffered_json_content: str | None = None
|
||||
self.function_call_buffering_disabled: bool = False
|
||||
self.function_call_buffering_enabled: bool = bool(json_mode)
|
||||
|
||||
def _handle_string_chunk(self, str_line: str) -> GenericStreamingChunk | ModelResponseStream:
|
||||
return self.chunk_parser(json.loads(str_line))
|
||||
|
|
@ -464,6 +466,23 @@ class OllamaTextCompletionResponseIterator(BaseModelResponseIterator):
|
|||
prefix: Final = '{"name"'
|
||||
return normalized.startswith(prefix) or prefix.startswith(normalized)
|
||||
|
||||
def _released_text(self, response_text: str) -> str | None:
|
||||
"""None while a fragment is held back because it may still complete a prompted function call."""
|
||||
if self.buffered_json_content is None and not (
|
||||
self.function_call_buffering_enabled
|
||||
and not self.started_reasoning_content
|
||||
and response_text.lstrip().startswith("{")
|
||||
):
|
||||
self.function_call_buffering_enabled = False
|
||||
return response_text
|
||||
candidate: Final = (self.buffered_json_content or "") + response_text
|
||||
if self._could_be_function_call(candidate):
|
||||
self.buffered_json_content = candidate
|
||||
return None
|
||||
self.buffered_json_content = None
|
||||
self.function_call_buffering_enabled = False
|
||||
return candidate
|
||||
|
||||
def _parse_buffered_function_call(self) -> ChatCompletionDeltaToolCall | None:
|
||||
if self.buffered_json_content is None:
|
||||
return None
|
||||
|
|
@ -537,26 +556,14 @@ class OllamaTextCompletionResponseIterator(BaseModelResponseIterator):
|
|||
usage=usage,
|
||||
)
|
||||
elif chunk["response"]:
|
||||
text = chunk["response"]
|
||||
if self.buffered_json_content is not None or (
|
||||
not self.function_call_buffering_disabled
|
||||
and not self.started_reasoning_content
|
||||
and text.lstrip().startswith("{")
|
||||
):
|
||||
candidate: Final = (self.buffered_json_content or "") + text
|
||||
if self._could_be_function_call(candidate):
|
||||
self.buffered_json_content = candidate
|
||||
return ModelResponseStream(
|
||||
choices=[ # mutable-ok: ModelResponseStream only accepts a list of choices
|
||||
StreamingChoices(index=0, delta=Delta())
|
||||
],
|
||||
usage=None,
|
||||
)
|
||||
self.buffered_json_content = None
|
||||
self.function_call_buffering_disabled = True
|
||||
text = candidate
|
||||
else:
|
||||
self.function_call_buffering_disabled = True
|
||||
text = self._released_text(chunk["response"])
|
||||
if text is None:
|
||||
return ModelResponseStream(
|
||||
choices=[ # mutable-ok: ModelResponseStream only accepts a list of choices
|
||||
StreamingChoices(index=0, delta=Delta())
|
||||
],
|
||||
usage=None,
|
||||
)
|
||||
reasoning_content: str | None = None
|
||||
content: str | None = None
|
||||
if text is not None:
|
||||
|
|
|
|||
|
|
@ -507,8 +507,10 @@ class TestOllamaTextCompletionResponseIterator:
|
|||
class TestOllamaTextCompletionStreamingToolCalls:
|
||||
"""Regression tests for https://github.com/BerriAI/litellm/issues/35711"""
|
||||
|
||||
def _stream(self, responses):
|
||||
iterator = OllamaTextCompletionResponseIterator(streaming_response=iter([]), sync_stream=True)
|
||||
def _stream(self, responses, json_mode=True):
|
||||
iterator = OllamaTextCompletionResponseIterator(
|
||||
streaming_response=iter([]), sync_stream=True, json_mode=json_mode
|
||||
)
|
||||
chunks = [
|
||||
iterator.chunk_parser({"model": "qwen3", "created_at": "t", "done": False, "response": r})
|
||||
for r in responses
|
||||
|
|
@ -571,3 +573,60 @@ class TestOllamaTextCompletionStreamingToolCalls:
|
|||
assert chunks[0].choices[0].delta.content == "Hello"
|
||||
assert chunks[1].choices[0].delta.content == " world"
|
||||
assert done["finish_reason"] == "stop"
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"arguments_fragment",
|
||||
[' "arguments": {"location": "Paris"}}', ' "arguments": "{\\"location\\": \\"Paris\\"}"}'],
|
||||
)
|
||||
def test_no_tool_call_reconstruction_when_json_was_not_requested(self, arguments_fragment):
|
||||
"""A caller that sent no tools and no response_format must never get a synthesized tool call,
|
||||
and must never lose the content it did ask for."""
|
||||
chunks, done = self._stream(['{"name": "get_weather",', arguments_fragment], json_mode=False)
|
||||
|
||||
streamed = "".join(c.choices[0].delta.content or "" for c in chunks)
|
||||
assert streamed == '{"name": "get_weather",' + arguments_fragment
|
||||
for chunk in chunks:
|
||||
assert chunk.choices[0].delta.tool_calls is None
|
||||
assert done["finish_reason"] == "stop"
|
||||
|
||||
|
||||
class TestOllamaStreamGating:
|
||||
"""`utils.py` sets format=json for ollama whenever tools are passed, and the prompted function call
|
||||
only makes sense for those requests. The iterator learns about it from the request transform."""
|
||||
|
||||
def _iterator_for(self, optional_params):
|
||||
config = OllamaConfig()
|
||||
config.transform_request(
|
||||
model="qwen3",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
optional_params=optional_params,
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
return config.get_model_response_iterator(streaming_response=iter([]), sync_stream=True, json_mode=False)
|
||||
|
||||
def test_json_format_request_buffers_a_possible_function_call(self):
|
||||
iterator = self._iterator_for({"format": "json"})
|
||||
|
||||
assert iterator.function_call_buffering_enabled is True
|
||||
|
||||
def test_plain_request_never_buffers(self):
|
||||
iterator = self._iterator_for({"temperature": 0.5})
|
||||
|
||||
assert iterator.function_call_buffering_enabled is False
|
||||
|
||||
@pytest.mark.parametrize("sync_stream", [True, False])
|
||||
def test_gate_survives_both_sync_and_async_streaming(self, sync_stream):
|
||||
"""The async handler builds the iterator without forwarding json_mode, so the flag has to ride
|
||||
on the config rather than on that argument."""
|
||||
config = OllamaConfig()
|
||||
config.transform_request(
|
||||
model="qwen3",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
optional_params={"format": "json"},
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
iterator = config.get_model_response_iterator(streaming_response=iter([]), sync_stream=sync_stream)
|
||||
|
||||
assert iterator.function_call_buffering_enabled is True
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue