mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-14 23:21:35 +00:00
fix(snowflake): transform messages for tool calling round trips
Transform OpenAI message format to Snowflake format for tool calling: - role: "tool" messages -> role: "user" with content_list containing tool_results - assistant messages with tool_calls -> content_list with tool_use blocks - Multiple consecutive tool messages are combined into a single user message Snowflake uses a Bedrock-style format where tool results must be in user messages with content_list containing tool_results blocks (not role: "tool"). Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
parent
336fb0cfeb
commit
d6ce4e2aad
2 changed files with 274 additions and 1 deletions
|
|
@ -154,6 +154,166 @@ class SnowflakeConfig(SnowflakeBaseConfig, OpenAIGPTConfig):
|
|||
|
||||
return f"{api_base}/cortex/inference:complete"
|
||||
|
||||
def _transform_messages(
|
||||
self, messages: List[AllMessageValues]
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Transform OpenAI messages to Snowflake format.
|
||||
|
||||
Key transformations:
|
||||
1. Assistant messages with tool_calls -> content_list with tool_use blocks
|
||||
2. Tool messages (role: "tool") -> User messages with content_list containing tool_results
|
||||
|
||||
Snowflake uses a format similar to Anthropic/Bedrock where:
|
||||
- tool_use blocks are in assistant message content_list
|
||||
- tool_results are in user message content_list (not role: "tool")
|
||||
"""
|
||||
# Build a map of tool_call_id -> tool_call for looking up function names
|
||||
tool_calls_map: Dict[str, Dict[str, Any]] = {}
|
||||
for message in messages:
|
||||
if isinstance(message, dict) and message.get("role") == "assistant":
|
||||
for tc in message.get("tool_calls") or []:
|
||||
if isinstance(tc, dict):
|
||||
tool_calls_map[tc.get("id", "")] = tc
|
||||
|
||||
transformed: List[Dict[str, Any]] = []
|
||||
pending_tool_messages: List[Dict[str, Any]] = []
|
||||
|
||||
for message in messages:
|
||||
if not isinstance(message, dict):
|
||||
continue
|
||||
|
||||
role = message.get("role", "")
|
||||
|
||||
# Flush pending tool messages before any non-tool message
|
||||
if role != "tool" and pending_tool_messages:
|
||||
transformed.append(
|
||||
self._convert_tool_messages_to_user_message(
|
||||
pending_tool_messages, tool_calls_map
|
||||
)
|
||||
)
|
||||
pending_tool_messages = []
|
||||
|
||||
if role == "tool":
|
||||
# Collect tool messages to combine into a single user message
|
||||
pending_tool_messages.append(message)
|
||||
|
||||
elif role == "assistant" and message.get("tool_calls"):
|
||||
# Transform assistant message with tool_calls to content_list format
|
||||
transformed.append(self._convert_assistant_tool_message(message))
|
||||
|
||||
else:
|
||||
# Pass through other messages as-is
|
||||
transformed.append(message)
|
||||
|
||||
# Flush any remaining tool messages
|
||||
if pending_tool_messages:
|
||||
transformed.append(
|
||||
self._convert_tool_messages_to_user_message(
|
||||
pending_tool_messages, tool_calls_map
|
||||
)
|
||||
)
|
||||
|
||||
return transformed
|
||||
|
||||
def _convert_assistant_tool_message(
|
||||
self, message: Dict[str, Any]
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Convert assistant message with tool_calls to Snowflake's content_list format.
|
||||
|
||||
OpenAI format:
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "I'll check that for you.",
|
||||
"tool_calls": [{"id": "...", "function": {"name": "...", "arguments": "..."}}]
|
||||
}
|
||||
|
||||
Snowflake format:
|
||||
{
|
||||
"role": "assistant",
|
||||
"content_list": [
|
||||
{"type": "text", "text": "I'll check that for you."},
|
||||
{"type": "tool_use", "tool_use": {"tool_use_id": "...", "name": "...", "input": {...}}}
|
||||
]
|
||||
}
|
||||
"""
|
||||
content_list: List[Dict[str, Any]] = []
|
||||
|
||||
# Add text content if present
|
||||
text_content = message.get("content")
|
||||
if text_content:
|
||||
content_list.append({"type": "text", "text": text_content})
|
||||
|
||||
# Add tool_use blocks
|
||||
for tool_call in message.get("tool_calls") or []:
|
||||
if isinstance(tool_call, dict):
|
||||
function = tool_call.get("function", {})
|
||||
# Parse arguments from JSON string to dict
|
||||
arguments_str = function.get("arguments", "{}")
|
||||
try:
|
||||
arguments = json.loads(arguments_str) if arguments_str else {}
|
||||
except json.JSONDecodeError:
|
||||
arguments = {}
|
||||
|
||||
content_list.append({
|
||||
"type": "tool_use",
|
||||
"tool_use": {
|
||||
"tool_use_id": tool_call.get("id", ""),
|
||||
"name": function.get("name", ""),
|
||||
"input": arguments,
|
||||
},
|
||||
})
|
||||
|
||||
return {"role": "assistant", "content_list": content_list}
|
||||
|
||||
def _convert_tool_messages_to_user_message(
|
||||
self,
|
||||
tool_messages: List[Dict[str, Any]],
|
||||
tool_calls_map: Dict[str, Dict[str, Any]],
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Convert tool result messages to a single Snowflake user message with tool_results.
|
||||
|
||||
OpenAI format (multiple messages):
|
||||
[
|
||||
{"role": "tool", "tool_call_id": "...", "content": "result1"},
|
||||
{"role": "tool", "tool_call_id": "...", "content": "result2"}
|
||||
]
|
||||
|
||||
Snowflake format (single user message):
|
||||
{
|
||||
"role": "user",
|
||||
"content_list": [
|
||||
{"type": "tool_results", "tool_results": {"tool_use_id": "...", "name": "...", "content": [{"type": "text", "text": "result1"}]}},
|
||||
{"type": "tool_results", "tool_results": {"tool_use_id": "...", "name": "...", "content": [{"type": "text", "text": "result2"}]}}
|
||||
]
|
||||
}
|
||||
"""
|
||||
content_list: List[Dict[str, Any]] = []
|
||||
|
||||
for tool_msg in tool_messages:
|
||||
tool_call_id = tool_msg.get("tool_call_id", "")
|
||||
tool_call = tool_calls_map.get(tool_call_id, {})
|
||||
function = tool_call.get("function", {})
|
||||
function_name = function.get("name", "")
|
||||
|
||||
# Get content - could be string or None
|
||||
content = tool_msg.get("content")
|
||||
if content is None:
|
||||
content = "null"
|
||||
|
||||
content_list.append({
|
||||
"type": "tool_results",
|
||||
"tool_results": {
|
||||
"tool_use_id": tool_call_id,
|
||||
"name": function_name,
|
||||
"content": [{"type": "text", "text": content}],
|
||||
},
|
||||
})
|
||||
|
||||
return {"role": "user", "content_list": content_list}
|
||||
|
||||
def _transform_tools(self, tools: List[Dict[str, Any]]) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Transform OpenAI tool format to Snowflake tool format.
|
||||
|
|
@ -263,9 +423,14 @@ class SnowflakeConfig(SnowflakeBaseConfig, OpenAIGPTConfig):
|
|||
if tool_choice:
|
||||
optional_params["tool_choice"] = self._transform_tool_choice(tool_choice)
|
||||
|
||||
# Transform messages from OpenAI format to Snowflake format
|
||||
# This handles role: "tool" -> role: "user" with tool_results content_list
|
||||
# and assistant messages with tool_calls -> content_list with tool_use blocks
|
||||
transformed_messages = self._transform_messages(messages)
|
||||
|
||||
return {
|
||||
"model": model,
|
||||
"messages": messages,
|
||||
"messages": transformed_messages,
|
||||
"stream": stream,
|
||||
**optional_params,
|
||||
**extra_body,
|
||||
|
|
|
|||
|
|
@ -331,6 +331,114 @@ class TestSnowflakeToolTransformation:
|
|||
assert "temperature" in supported_params
|
||||
assert "max_tokens" in supported_params
|
||||
|
||||
def test_transform_messages_with_tool_results(self):
|
||||
"""
|
||||
Test that OpenAI role: "tool" messages are transformed to Snowflake format.
|
||||
|
||||
OpenAI sends tool results as:
|
||||
{"role": "tool", "tool_call_id": "...", "content": "result"}
|
||||
|
||||
Snowflake expects:
|
||||
{"role": "user", "content_list": [{"type": "tool_results", "tool_results": {...}}]}
|
||||
"""
|
||||
config = SnowflakeConfig()
|
||||
|
||||
messages = [
|
||||
{"role": "user", "content": "What's the weather in Paris?"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "I'll check that for you.",
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "call_123",
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "get_weather",
|
||||
"arguments": '{"location": "Paris"}',
|
||||
},
|
||||
}
|
||||
],
|
||||
},
|
||||
{
|
||||
"role": "tool",
|
||||
"tool_call_id": "call_123",
|
||||
"content": "72°F and sunny",
|
||||
},
|
||||
]
|
||||
|
||||
transformed = config._transform_messages(messages)
|
||||
|
||||
# Should have 3 messages
|
||||
assert len(transformed) == 3
|
||||
|
||||
# First message unchanged
|
||||
assert transformed[0]["role"] == "user"
|
||||
assert transformed[0]["content"] == "What's the weather in Paris?"
|
||||
|
||||
# Second message (assistant) should have content_list with tool_use
|
||||
assert transformed[1]["role"] == "assistant"
|
||||
assert "content_list" in transformed[1]
|
||||
content_list = transformed[1]["content_list"]
|
||||
assert len(content_list) == 2
|
||||
assert content_list[0]["type"] == "text"
|
||||
assert content_list[0]["text"] == "I'll check that for you."
|
||||
assert content_list[1]["type"] == "tool_use"
|
||||
assert content_list[1]["tool_use"]["tool_use_id"] == "call_123"
|
||||
assert content_list[1]["tool_use"]["name"] == "get_weather"
|
||||
assert content_list[1]["tool_use"]["input"] == {"location": "Paris"}
|
||||
|
||||
# Third message (tool) should become user with tool_results
|
||||
assert transformed[2]["role"] == "user"
|
||||
assert "content_list" in transformed[2]
|
||||
tool_results = transformed[2]["content_list"]
|
||||
assert len(tool_results) == 1
|
||||
assert tool_results[0]["type"] == "tool_results"
|
||||
assert tool_results[0]["tool_results"]["tool_use_id"] == "call_123"
|
||||
assert tool_results[0]["tool_results"]["name"] == "get_weather"
|
||||
assert tool_results[0]["tool_results"]["content"] == [
|
||||
{"type": "text", "text": "72°F and sunny"}
|
||||
]
|
||||
|
||||
def test_transform_messages_multiple_tool_results(self):
|
||||
"""
|
||||
Test that multiple consecutive tool messages are combined into one user message.
|
||||
"""
|
||||
config = SnowflakeConfig()
|
||||
|
||||
messages = [
|
||||
{"role": "user", "content": "Get weather for Paris and London"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "",
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "call_1",
|
||||
"type": "function",
|
||||
"function": {"name": "get_weather", "arguments": '{"location": "Paris"}'},
|
||||
},
|
||||
{
|
||||
"id": "call_2",
|
||||
"type": "function",
|
||||
"function": {"name": "get_weather", "arguments": '{"location": "London"}'},
|
||||
},
|
||||
],
|
||||
},
|
||||
{"role": "tool", "tool_call_id": "call_1", "content": "72°F"},
|
||||
{"role": "tool", "tool_call_id": "call_2", "content": "55°F"},
|
||||
]
|
||||
|
||||
transformed = config._transform_messages(messages)
|
||||
|
||||
# Should have 3 messages (user, assistant, combined tool results)
|
||||
assert len(transformed) == 3
|
||||
|
||||
# Third message should have both tool results
|
||||
assert transformed[2]["role"] == "user"
|
||||
tool_results = transformed[2]["content_list"]
|
||||
assert len(tool_results) == 2
|
||||
assert tool_results[0]["tool_results"]["tool_use_id"] == "call_1"
|
||||
assert tool_results[1]["tool_results"]["tool_use_id"] == "call_2"
|
||||
|
||||
|
||||
class TestSnowFlakeCompletion:
|
||||
model_name = "mistral"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue