mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix(tests): update chat transformation tests for native OpenAI-compatible endpoint
This commit is contained in:
parent
70d3fbb568
commit
f5c7a1f174
1 changed files with 58 additions and 89 deletions
|
|
@ -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):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue