fix(tests): update chat transformation tests for native OpenAI-compatible endpoint

This commit is contained in:
Navnit Shukla 2026-06-08 14:13:55 -07:00
parent 70d3fbb568
commit f5c7a1f174
No known key found for this signature in database
GPG key ID: 46D87897A91C215F

View file

@ -26,11 +26,13 @@ class TestSnowflakeToolTransformation:
def test_transform_request_with_tools(self):
"""
Test that OpenAI tool format is correctly transformed to Snowflake's tool_spec format.
Test that OpenAI tool format is passed through as-is to the native endpoint.
The native /chat/completions endpoint accepts standard OpenAI tool format
directly — no Snowflake-specific tool_spec transformation needed.
"""
config = SnowflakeConfig()
# OpenAI format tools
tools = [
{
"type": "function",
@ -65,29 +67,17 @@ class TestSnowflakeToolTransformation:
headers={},
)
# Verify tools were transformed to Snowflake format
assert "tools" in transformed_request
assert len(transformed_request["tools"]) == 1
snowflake_tool = transformed_request["tools"][0]
assert "tool_spec" in snowflake_tool
assert snowflake_tool["tool_spec"]["type"] == "generic"
assert snowflake_tool["tool_spec"]["name"] == "get_weather"
assert (
snowflake_tool["tool_spec"]["description"]
== "Get the current weather in a given location"
)
assert "input_schema" in snowflake_tool["tool_spec"]
assert snowflake_tool["tool_spec"]["input_schema"]["type"] == "object"
assert "location" in snowflake_tool["tool_spec"]["input_schema"]["properties"]
assert transformed_request["tools"] == tools
assert "tool_spec" not in json.dumps(transformed_request)
def test_transform_request_with_tool_choice(self):
"""
Test that OpenAI tool_choice format is correctly transformed to Snowflake format.
Test that OpenAI tool_choice format is passed through as-is to the native endpoint.
"""
config = SnowflakeConfig()
# OpenAI format tool_choice
tool_choice = {"type": "function", "function": {"name": "get_weather"}}
optional_params = {"tool_choice": tool_choice}
@ -100,30 +90,19 @@ class TestSnowflakeToolTransformation:
headers={},
)
# Verify tool_choice was transformed to Snowflake format
assert "tool_choice" in transformed_request
assert transformed_request["tool_choice"]["type"] == "tool"
assert transformed_request["tool_choice"]["name"] == [
"get_weather"
] # Array format
assert transformed_request["tool_choice"] == tool_choice
def test_transform_request_with_string_tool_choice(self):
"""
Test that string tool_choice values are transformed to Snowflake object format.
Test that string tool_choice values are passed through as-is to the native endpoint.
Snowflake's API (like Anthropic) requires tool_choice as an object
with a "type" field, not as a bare string. OpenAI's "required" maps
to Snowflake's "any".
The native /chat/completions endpoint accepts OpenAI-style string
tool_choice values directly ("auto", "required", "none").
"""
config = SnowflakeConfig()
expected_mappings = {
"auto": {"type": "auto"},
"required": {"type": "any"},
"none": {"type": "none"},
}
for value, expected in expected_mappings.items():
for value in ["auto", "required", "none"]:
optional_params = {"tool_choice": value}
transformed_request = config.transform_request(
@ -134,37 +113,41 @@ class TestSnowflakeToolTransformation:
headers={},
)
assert transformed_request["tool_choice"] == expected, (
f"tool_choice='{value}' should be transformed to {expected}, "
assert transformed_request["tool_choice"] == value, (
f"tool_choice='{value}' should pass through unchanged, "
f"got {transformed_request['tool_choice']}"
)
def test_transform_response_with_tool_calls(self):
"""
Test that Snowflake's content_list with tool_use is transformed to OpenAI format.
Test that standard OpenAI tool_calls response format is parsed correctly.
The native /chat/completions endpoint returns standard OpenAI format.
"""
config = SnowflakeConfig()
# Mock Snowflake response with tool call
mock_snowflake_response = {
mock_response = {
"id": "chatcmpl-123",
"object": "chat.completion",
"model": "claude-3-5-sonnet",
"choices": [
{
"index": 0,
"message": {
"content_list": [
{"type": "text", "text": ""},
"role": "assistant",
"content": None,
"tool_calls": [
{
"type": "tool_use",
"tool_use": {
"tool_use_id": "tooluse_abc123",
"id": "call_abc123",
"type": "function",
"function": {
"name": "get_weather",
"input": {
"location": "Paris, France",
"unit": "celsius",
},
"arguments": json.dumps({"location": "Paris, France", "unit": "celsius"}),
},
},
]
}
}
],
},
"finish_reason": "tool_calls",
}
],
"usage": {"prompt_tokens": 10, "completion_tokens": 20, "total_tokens": 30},
@ -172,7 +155,7 @@ class TestSnowflakeToolTransformation:
response = httpx.Response(
status_code=200,
json=mock_snowflake_response,
json=mock_response,
headers={"Content-Type": "application/json"},
)
@ -194,61 +177,50 @@ class TestSnowflakeToolTransformation:
encoding={},
)
# General assertions
assert isinstance(result, ModelResponse)
assert len(result.choices) == 1
choice = result.choices[0]
assert isinstance(choice, litellm.Choices)
# Message and tool_calls assertions
message = choice.message
assert isinstance(message, litellm.Message)
assert hasattr(message, "tool_calls")
assert isinstance(message.tool_calls, list)
message = result.choices[0].message
assert message.tool_calls is not None
assert len(message.tool_calls) == 1
# Specific tool_call assertions
tool_call = message.tool_calls[0]
assert isinstance(tool_call, litellm.utils.ChatCompletionMessageToolCall)
assert tool_call.id == "tooluse_abc123"
assert tool_call.id == "call_abc123"
assert tool_call.type == "function"
assert tool_call.function.name == "get_weather"
# Verify arguments are properly JSON serialized
arguments = json.loads(tool_call.function.arguments)
assert arguments["location"] == "Paris, France"
assert arguments["unit"] == "celsius"
# Verify content_list was removed and content was set
assert message.content == ""
def test_transform_response_with_mixed_content(self):
"""
Test that responses with both text and tool calls are handled correctly.
Test that responses with both text content and tool calls are parsed correctly.
"""
config = SnowflakeConfig()
# Mock Snowflake response with text and tool call
mock_snowflake_response = {
mock_response = {
"id": "chatcmpl-456",
"object": "chat.completion",
"model": "claude-3-5-sonnet",
"choices": [
{
"index": 0,
"message": {
"content_list": [
"role": "assistant",
"content": "Let me check the weather for you.",
"tool_calls": [
{
"type": "text",
"text": "Let me check the weather for you. ",
},
{
"type": "tool_use",
"tool_use": {
"tool_use_id": "tooluse_xyz789",
"id": "call_xyz789",
"type": "function",
"function": {
"name": "get_weather",
"input": {"location": "Tokyo, Japan"},
"arguments": json.dumps({"location": "Tokyo, Japan"}),
},
},
]
}
}
],
},
"finish_reason": "tool_calls",
}
],
"usage": {"prompt_tokens": 15, "completion_tokens": 25, "total_tokens": 40},
@ -256,7 +228,7 @@ class TestSnowflakeToolTransformation:
response = httpx.Response(
status_code=200,
json=mock_snowflake_response,
json=mock_response,
headers={"Content-Type": "application/json"},
)
@ -278,11 +250,8 @@ class TestSnowflakeToolTransformation:
encoding={},
)
# Verify text content was extracted
message = result.choices[0].message
assert message.content == "Let me check the weather for you. "
# Verify tool call was also extracted
assert message.content == "Let me check the weather for you."
assert len(message.tool_calls) == 1
assert message.tool_calls[0].function.name == "get_weather"
@ -392,8 +361,8 @@ class TestSnowFlakeCompletion:
assert "00000" in post_kwargs["headers"]["Authorization"]
# account id was used
assert "AAAA-BBBB" in post_kwargs["url"]
# is completion
assert post_kwargs["url"].endswith("cortex/inference:complete")
# uses native endpoint
assert post_kwargs["url"].endswith("cortex/v1/chat/completions")
@patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post")
def test_snowflake_pat_key_account_id(self, mock_post):