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:
stevejaker 2026-03-10 10:42:54 -06:00
parent 336fb0cfeb
commit d6ce4e2aad
2 changed files with 274 additions and 1 deletions

View file

@ -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,

View file

@ -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"