mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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
This commit is contained in:
parent
2654d3b0b1
commit
8dd8615a54
7 changed files with 161 additions and 29 deletions
|
|
@ -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}"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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": (
|
||||
"'<think>\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</think>\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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue