From 8dd8615a548bb1542ad0ead5ed26e6d3a9959a2a Mon Sep 17 00:00:00 2001 From: Krish Dholakia Date: Sat, 7 Jun 2025 20:50:07 -0700 Subject: [PATCH] Ensure consistent 'created' across all chunks + set tool call id for ollama streaming calls (#11528) * fix(streaming_handler.py): maintain same 'created' across all chunks Fixes https://github.com/BerriAI/litellm/issues/11437 * test: add unit test to ensure created is always the same across all chunks * fix(types/utils.py): set a tool call id, if missing in delta tool call Ensures stream chunk builder can reconstruct tool calls correctly Fixes https://github.com/BerriAI/litellm/issues/11262 * fix(responses/transformation.py): support passing mcp server tool call to anthropic allows switching between openai and anthropic for mcp tool calling * fix(ollama/chat/transformation.py): set tool call id's when missing --- .../litellm_core_utils/streaming_handler.py | 28 ++++++++----- litellm/llms/ollama/chat/transformation.py | 23 ++++++++++- .../transformation.py | 41 +++++++++++-------- tests/llm_translation/base_llm_unit_tests.py | 9 +++- .../test_anthropic_completion.py | 22 ++++++++++ tests/local_testing/test_ollama.py | 39 ++++++++++++++++++ .../test_streaming_handler.py | 28 +++++++++++++ 7 files changed, 161 insertions(+), 29 deletions(-) diff --git a/litellm/litellm_core_utils/streaming_handler.py b/litellm/litellm_core_utils/streaming_handler.py index 25799dc2dc5..4a78ff6accc 100644 --- a/litellm/litellm_core_utils/streaming_handler.py +++ b/litellm/litellm_core_utils/streaming_handler.py @@ -85,9 +85,9 @@ class CustomStreamWrapper: self.system_fingerprint: Optional[str] = None self.received_finish_reason: Optional[str] = None - self.intermittent_finish_reason: Optional[str] = ( - None # finish reasons that show up mid-stream - ) + self.intermittent_finish_reason: Optional[ + str + ] = None # finish reasons that show up mid-stream self.special_tokens = [ "<|assistant|>", "<|system|>", @@ -135,6 +135,7 @@ class CustomStreamWrapper: [] ) # keep track of the returned chunks - used for calculating the input/output tokens for stream options self.is_function_call = self.check_is_function_call(logging_obj=logging_obj) + self.created: Optional[int] = None def __iter__(self): return self @@ -621,6 +622,13 @@ class CustomStreamWrapper: model_response.id = self.response_id if self.system_fingerprint is not None: model_response.system_fingerprint = self.system_fingerprint + + if ( + self.created is not None + ): # maintain same 'created' across all chunks - https://github.com/BerriAI/litellm/issues/11437 + model_response.created = self.created + else: + self.created = model_response.created if hidden_params is not None: model_response._hidden_params = hidden_params model_response._hidden_params["custom_llm_provider"] = _logging_obj_llm_provider @@ -914,7 +922,6 @@ class CustomStreamWrapper: def chunk_creator(self, chunk: Any): # type: ignore # noqa: PLR0915 model_response = self.model_response_creator() response_obj: Dict[str, Any] = {} - try: # return this for all models completion_obj: Dict[str, Any] = {"content": ""} @@ -1309,9 +1316,9 @@ class CustomStreamWrapper: _json_delta = delta.model_dump() print_verbose(f"_json_delta: {_json_delta}") if "role" not in _json_delta or _json_delta["role"] is None: - _json_delta["role"] = ( - "assistant" # mistral's api returns role as None - ) + _json_delta[ + "role" + ] = "assistant" # mistral's api returns role as None if "tool_calls" in _json_delta and isinstance( _json_delta["tool_calls"], list ): @@ -1480,6 +1487,7 @@ class CustomStreamWrapper: try: if self.completion_stream is None: self.fetch_sync_stream() + while True: if ( isinstance(self.completion_stream, str) @@ -1701,9 +1709,9 @@ class CustomStreamWrapper: chunk = next(self.completion_stream) if chunk is not None and chunk != b"": print_verbose(f"PROCESSED CHUNK PRE CHUNK CREATOR: {chunk}") - processed_chunk: Optional[ModelResponseStream] = ( - self.chunk_creator(chunk=chunk) - ) + processed_chunk: Optional[ + ModelResponseStream + ] = self.chunk_creator(chunk=chunk) print_verbose( f"PROCESSED CHUNK POST CHUNK CREATOR: {processed_chunk}" ) diff --git a/litellm/llms/ollama/chat/transformation.py b/litellm/llms/ollama/chat/transformation.py index e415704a762..dd0b42dd6c8 100644 --- a/litellm/llms/ollama/chat/transformation.py +++ b/litellm/llms/ollama/chat/transformation.py @@ -406,6 +406,15 @@ class OllamaChatConfig(BaseConfig): class OllamaChatCompletionResponseIterator(BaseModelResponseIterator): + def _is_function_call_complete(self, function_args: Union[str, dict]) -> bool: + if isinstance(function_args, dict): + return True + try: + json.loads(function_args) + return True + except Exception: + return False + def chunk_parser(self, chunk: dict) -> ModelResponseStream: try: """ @@ -438,9 +447,21 @@ class OllamaChatCompletionResponseIterator(BaseModelResponseIterator): """ from litellm.types.utils import Delta, StreamingChoices + # process tool calls - if complete function arg - add id to tool call + tool_calls = 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") + 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: + tool_call["id"] = str(uuid.uuid4()) + delta = Delta( content=chunk["message"].get("content", ""), - tool_calls=chunk["message"].get("tool_calls"), + tool_calls=tool_calls, ) if chunk["done"] is True: diff --git a/litellm/responses/litellm_completion_transformation/transformation.py b/litellm/responses/litellm_completion_transformation/transformation.py index d5bf7c606b8..be0ee6806cd 100644 --- a/litellm/responses/litellm_completion_transformation/transformation.py +++ b/litellm/responses/litellm_completion_transformation/transformation.py @@ -2,7 +2,7 @@ Handles transforming from Responses API -> LiteLLM completion (Chat Completion API) """ -from typing import Any, Dict, List, Optional, Union +from typing import Any, Dict, List, Optional, Union, cast from openai.types.responses.tool_param import FunctionToolParam from typing_extensions import TypedDict @@ -31,6 +31,7 @@ from litellm.types.llms.openai import ( ChatCompletionToolParamFunctionChunk, ChatCompletionUserMessage, GenericChatCompletionMessage, + OpenAIMcpServerTool, Reasoning, ResponseAPIUsage, ResponseInputParam, @@ -208,9 +209,9 @@ class LiteLLMCompletionResponsesConfig: _messages = litellm_completion_request.get("messages") or [] session_messages = chat_completion_session.get("messages") or [] litellm_completion_request["messages"] = session_messages + _messages - litellm_completion_request["litellm_trace_id"] = ( - chat_completion_session.get("litellm_session_id") - ) + litellm_completion_request[ + "litellm_trace_id" + ] = chat_completion_session.get("litellm_session_id") return litellm_completion_request @staticmethod @@ -466,26 +467,32 @@ class LiteLLMCompletionResponsesConfig: @staticmethod def transform_responses_api_tools_to_chat_completion_tools( - tools: Optional[List[FunctionToolParam]], - ) -> List[ChatCompletionToolParam]: + tools: Optional[List[Union[FunctionToolParam, OpenAIMcpServerTool]]], + ) -> List[Union[ChatCompletionToolParam, OpenAIMcpServerTool]]: """ Transform a Responses API tools into a Chat Completion tools """ if tools is None: return [] - chat_completion_tools: List[ChatCompletionToolParam] = [] + chat_completion_tools: List[ + Union[ChatCompletionToolParam, OpenAIMcpServerTool] + ] = [] for tool in tools: - chat_completion_tools.append( - ChatCompletionToolParam( - type="function", - function=ChatCompletionToolParamFunctionChunk( - name=tool["name"], - description=tool.get("description") or "", - parameters=dict(tool.get("parameters", {}) or {}), - strict=tool.get("strict", False) or False, - ), + if tool.get("type") == "mcp": + chat_completion_tools.append(cast(OpenAIMcpServerTool, tool)) + else: + typed_tool = cast(FunctionToolParam, tool) + chat_completion_tools.append( + ChatCompletionToolParam( + type="function", + function=ChatCompletionToolParamFunctionChunk( + name=typed_tool["name"], + description=typed_tool.get("description") or "", + parameters=dict(typed_tool.get("parameters", {}) or {}), + strict=typed_tool.get("strict", False) or False, + ), + ) ) - ) return chat_completion_tools @staticmethod diff --git a/tests/llm_translation/base_llm_unit_tests.py b/tests/llm_translation/base_llm_unit_tests.py index 4d3dbee28ff..04953e32577 100644 --- a/tests/llm_translation/base_llm_unit_tests.py +++ b/tests/llm_translation/base_llm_unit_tests.py @@ -143,8 +143,10 @@ class BaseLLMChatTest(ABC): def test_streaming(self): """Check if litellm handles streaming correctly""" + from litellm.types.utils import ModelResponseStream + from typing import Optional base_completion_call_args = self.get_base_completion_call_args() - litellm.set_verbose = True + # litellm.set_verbose = True messages = [ { "role": "user", @@ -164,9 +166,14 @@ class BaseLLMChatTest(ABC): # for OpenAI the content contains the JSON schema, so we need to assert that the content is not None chunks = [] + created_at: Optional[int] = None for chunk in response: print(chunk) chunks.append(chunk) + if isinstance(chunk, ModelResponseStream): + if created_at is None: + created_at = chunk.created + assert chunk.created == created_at resp = litellm.stream_chunk_builder(chunks=chunks) print(resp) diff --git a/tests/llm_translation/test_anthropic_completion.py b/tests/llm_translation/test_anthropic_completion.py index 74677cee26c..7bc89b60fe1 100644 --- a/tests/llm_translation/test_anthropic_completion.py +++ b/tests/llm_translation/test_anthropic_completion.py @@ -1307,3 +1307,25 @@ def test_anthropic_mcp_server_tool_use(spec: str): print(e) assert response is not None + +@pytest.mark.parametrize("model", ["openai/gpt-4.1", "anthropic/claude-sonnet-4-20250514"]) +def test_anthropic_mcp_server_responses_api(model: str): + from litellm import responses + + tools=[ + { + "type": "mcp", + "server_label": "deepwiki", + "server_url": "https://mcp.deepwiki.com/mcp", + "require_approval": "never", + }, + ] + + response = litellm.responses( + model=model, + input="Who won the World Cup in 2022?", + max_output_tokens=100, + tools=tools + ) + + assert response is not None diff --git a/tests/local_testing/test_ollama.py b/tests/local_testing/test_ollama.py index 169db483f28..e7660ddc24c 100644 --- a/tests/local_testing/test_ollama.py +++ b/tests/local_testing/test_ollama.py @@ -291,4 +291,43 @@ async def test_async_ollama_ssl_verify(stream): assert litellm_created_session.connector._ssl is False assert litellm_created_session.connector._ssl == aiohttp_session.connector._ssl +@pytest.mark.skip(reason="local only test") +def test_ollama_streaming_with_chunk_builder(): + from litellm.main import stream_chunk_builder + tools = [ + { + "type": "function", + "function": { + "name": "get_weather", + "description": "Get the weather for a location", + "parameters": { + "type": "object", + "properties": {"location": {"type": "string"}}, + "required": ["location"], + }, + }, + } + ] + completion_kwargs = { + "model": "ollama_chat/qwen2.5:0.5b", # Important: use `ollama_chat` instead of `ollama` + "messages": [ + {"role": "user", "content": "What's the weather like in New York?"}, + { + "role": "assistant", + "content": ( + "'\nOkay, the user is asking about the weather in New York. " + "Let me check the tools available. " + "There's a function called get_weather that takes a location parameter. " + "So I need to call that function with 'New York' as the location. " + "I should make sure the arguments are correctly formatted in JSON. " + "Let me structure the tool call accordingly.\n\n\n" + ), + }, + ], + "tools": tools, + "stream": True, + } + response = litellm.completion(**completion_kwargs) + response = stream_chunk_builder(list(response)) + assert response.choices[0].message.tool_calls, "No tool call detected" diff --git a/tests/test_litellm/litellm_core_utils/test_streaming_handler.py b/tests/test_litellm/litellm_core_utils/test_streaming_handler.py index 8027db94428..6c08467f0eb 100644 --- a/tests/test_litellm/litellm_core_utils/test_streaming_handler.py +++ b/tests/test_litellm/litellm_core_utils/test_streaming_handler.py @@ -686,3 +686,31 @@ async def test_streaming_completion_start_time(logging_obj: Logging): logging_obj.model_call_details["completion_start_time"] < logging_obj.model_call_details["end_time"] ) + + +def test_streaming_handler_with_created_time_propagation( + initialized_custom_stream_wrapper: CustomStreamWrapper, logging_obj: Logging +): + """Test that the created time is consistent across chunks""" + import time + + bad_chunk = ModelResponseStream( + choices=[], created=int(time.time()) + ) # chunk with different created time + + completion_stream = ModelResponseListIterator( + model_responses=bedrock_chunks + [bad_chunk] + ) + + response = CustomStreamWrapper( + completion_stream=completion_stream, + model="bedrock/claude-3-5-sonnet-20240620-v1:0", + logging_obj=logging_obj, + ) + + created: Optional[int] = None + for chunk in response: + if created is None: + created = chunk.created + else: + assert created == chunk.created