diff --git a/litellm/llms/ollama/chat/transformation.py b/litellm/llms/ollama/chat/transformation.py index de626b468f0..48f74137954 100644 --- a/litellm/llms/ollama/chat/transformation.py +++ b/litellm/llms/ollama/chat/transformation.py @@ -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( diff --git a/litellm/llms/ollama/completion/transformation.py b/litellm/llms/ollama/completion/transformation.py index dccc83efed4..d885592df94 100644 --- a/litellm/llms/ollama/completion/transformation.py +++ b/litellm/llms/ollama/completion/transformation.py @@ -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: diff --git a/tests/test_litellm/llms/ollama/test_ollama_chat_transformation.py b/tests/test_litellm/llms/ollama/test_ollama_chat_transformation.py index 8f3dbf7b0d9..51e9590e65f 100644 --- a/tests/test_litellm/llms/ollama/test_ollama_chat_transformation.py +++ b/tests/test_litellm/llms/ollama/test_ollama_chat_transformation.py @@ -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) diff --git a/tests/test_litellm/llms/ollama/test_ollama_completion_transformation.py b/tests/test_litellm/llms/ollama/test_ollama_completion_transformation.py index eadc2bc9541..aa5fbd0faf5 100644 --- a/tests/test_litellm/llms/ollama/test_ollama_completion_transformation.py +++ b/tests/test_litellm/llms/ollama/test_ollama_completion_transformation.py @@ -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"