mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
fix(snowflake): migrate to native Cortex REST API endpoints
Replaces the legacy /api/v2/cortex/inference:complete endpoint with the native OpenAI-compatible /api/v2/cortex/v1/chat/completions endpoint, fixing error 390142 (Incoming request does not contain a valid payload) when using model: snowflake/<model> in LiteLLM proxy. Changes: - litellm/llms/snowflake/chat/transformation.py: route to native /cortex/v1/chat/completions, remove Snowflake-specific tool_spec payload transformation, remove content_list response handling, add stream to supported params - litellm/llms/snowflake/anthropic/transformation.py (new): SnowflakeCortexAnthropicConfig routes Claude models to /cortex/v1/messages with anthropic-version header and Anthropic->OpenAI response transform - tests: 29 unit tests covering URL routing, auth headers, payload format, and response parsing
This commit is contained in:
parent
5ee526d78e
commit
34f4c65909
4 changed files with 858 additions and 224 deletions
0
litellm/llms/snowflake/anthropic/__init__.py
Normal file
0
litellm/llms/snowflake/anthropic/__init__.py
Normal file
284
litellm/llms/snowflake/anthropic/transformation.py
Normal file
284
litellm/llms/snowflake/anthropic/transformation.py
Normal file
|
|
@ -0,0 +1,284 @@
|
|||
"""
|
||||
Snowflake Cortex REST API — Anthropic-Compatible Endpoint
|
||||
|
||||
Routes to the native Anthropic-compatible endpoint:
|
||||
POST /api/v2/cortex/v1/messages
|
||||
|
||||
Use this config for Claude models when you need Claude-specific features:
|
||||
- Extended thinking / reasoning (thinking parameter)
|
||||
- Prompt caching (cache_control on messages/system)
|
||||
- Anthropic tool format
|
||||
|
||||
For standard usage (all models, no Claude-specific features), use the
|
||||
default SnowflakeConfig which routes to /api/v2/cortex/v1/chat/completions.
|
||||
|
||||
Ref: https://docs.snowflake.com/en/user-guide/snowflake-cortex/cortex-rest-api
|
||||
"""
|
||||
|
||||
import json
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.utils import (
|
||||
Choices,
|
||||
Message,
|
||||
ModelResponse,
|
||||
Usage,
|
||||
)
|
||||
|
||||
from ..utils import SnowflakeBaseConfig
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
|
||||
|
||||
LiteLLMLoggingObj = _LiteLLMLoggingObj
|
||||
else:
|
||||
LiteLLMLoggingObj = Any
|
||||
|
||||
ANTHROPIC_VERSION = "2023-06-01"
|
||||
|
||||
_CLAUDE_MODEL_PREFIXES = (
|
||||
"claude-",
|
||||
"claude_",
|
||||
)
|
||||
|
||||
|
||||
def _is_claude_model(model: str) -> bool:
|
||||
"""Return True if model name (after stripping snowflake/ prefix) is a Claude model."""
|
||||
name = model.lower().removeprefix("snowflake/")
|
||||
return any(name.startswith(p) for p in _CLAUDE_MODEL_PREFIXES)
|
||||
|
||||
|
||||
class SnowflakeCortexAnthropicConfig(SnowflakeBaseConfig):
|
||||
"""
|
||||
Snowflake Cortex REST API — Anthropic Messages endpoint.
|
||||
|
||||
Endpoint: POST /api/v2/cortex/v1/messages
|
||||
|
||||
Designed for Claude models. Accepts and returns Anthropic Messages API
|
||||
format. Supports thinking, prompt caching, and extended tool formats.
|
||||
|
||||
Usage in litellm:
|
||||
import litellm
|
||||
response = litellm.completion(
|
||||
model="snowflake/claude-sonnet-4-5",
|
||||
messages=[{"role": "user", "content": "Hello"}],
|
||||
api_key="pat/<your-pat>",
|
||||
api_base="https://<account>.snowflakecomputing.com",
|
||||
custom_llm_provider="snowflake-anthropic", # routes to this config
|
||||
)
|
||||
"""
|
||||
|
||||
def get_supported_openai_params(self, model: str) -> List[str]:
|
||||
return [
|
||||
"temperature",
|
||||
"max_tokens",
|
||||
"top_p",
|
||||
"stream",
|
||||
"tools",
|
||||
"tool_choice",
|
||||
"thinking",
|
||||
]
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: Optional[str],
|
||||
api_key: Optional[str],
|
||||
model: str,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
stream: Optional[bool] = None,
|
||||
) -> str:
|
||||
"""
|
||||
Returns the Anthropic-compatible Cortex REST API endpoint.
|
||||
|
||||
https://{account}.snowflakecomputing.com/api/v2/cortex/v1/messages
|
||||
"""
|
||||
api_base = self._get_api_base(api_base, optional_params)
|
||||
return f"{api_base}/cortex/v1/messages"
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: dict,
|
||||
model: str,
|
||||
messages: List[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
api_key: Optional[str] = None,
|
||||
api_base: Optional[str] = None,
|
||||
) -> dict:
|
||||
"""
|
||||
Set Snowflake auth headers + Anthropic version header.
|
||||
"""
|
||||
headers = super().validate_environment(
|
||||
headers=headers,
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
)
|
||||
headers["anthropic-version"] = ANTHROPIC_VERSION
|
||||
return headers
|
||||
|
||||
def _extract_system_and_messages(
|
||||
self, messages: List[AllMessageValues]
|
||||
) -> tuple[Optional[Union[str, List[Dict]]], List[Dict]]:
|
||||
"""
|
||||
Split messages into system prompt and conversation turns.
|
||||
|
||||
Anthropic's /messages endpoint takes system as a top-level param,
|
||||
not inside the messages array.
|
||||
"""
|
||||
system: Optional[Union[str, List[Dict]]] = None
|
||||
conversation: List[Dict] = []
|
||||
|
||||
for msg in messages:
|
||||
if isinstance(msg, dict):
|
||||
role = msg.get("role", "")
|
||||
content = msg.get("content", "")
|
||||
else:
|
||||
role = getattr(msg, "role", "")
|
||||
content = getattr(msg, "content", "")
|
||||
|
||||
if role == "system":
|
||||
system = content
|
||||
else:
|
||||
conversation.append({"role": role, "content": content})
|
||||
|
||||
return system, conversation
|
||||
|
||||
def transform_request(
|
||||
self,
|
||||
model: str,
|
||||
messages: List[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
headers: dict,
|
||||
) -> dict:
|
||||
"""
|
||||
Transform to Anthropic Messages API format.
|
||||
|
||||
Key differences from OpenAI format:
|
||||
- system message → top-level "system" field
|
||||
- max_tokens is required
|
||||
- tools use Anthropic format (input_schema, not parameters)
|
||||
"""
|
||||
stream: bool = optional_params.pop("stream", False) or False
|
||||
extra_body = optional_params.pop("extra_body", {})
|
||||
|
||||
system, conversation = self._extract_system_and_messages(messages)
|
||||
|
||||
model_name = model.removeprefix("snowflake/")
|
||||
|
||||
body: Dict[str, Any] = {
|
||||
"model": model_name,
|
||||
"messages": conversation,
|
||||
"stream": stream,
|
||||
**optional_params,
|
||||
**extra_body,
|
||||
}
|
||||
|
||||
if system is not None:
|
||||
body["system"] = system
|
||||
|
||||
if "max_tokens" not in body:
|
||||
body["max_tokens"] = 1024
|
||||
|
||||
return body
|
||||
|
||||
def transform_response(
|
||||
self,
|
||||
model: str,
|
||||
raw_response: httpx.Response,
|
||||
model_response: ModelResponse,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
request_data: dict,
|
||||
messages: List[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
api_key: Optional[str] = None,
|
||||
json_mode: Optional[bool] = None,
|
||||
) -> ModelResponse:
|
||||
"""
|
||||
Transform Anthropic Messages response to OpenAI ChatCompletion format.
|
||||
|
||||
Anthropic response:
|
||||
{
|
||||
"id": "msg_...",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "text", "text": "Hello!"}],
|
||||
"stop_reason": "end_turn",
|
||||
"usage": {"input_tokens": 10, "output_tokens": 5}
|
||||
}
|
||||
|
||||
Output: standard OpenAI ModelResponse
|
||||
"""
|
||||
response_json = raw_response.json()
|
||||
|
||||
logging_obj.post_call(
|
||||
input=messages,
|
||||
api_key="",
|
||||
original_response=response_json,
|
||||
additional_args={"complete_input_dict": request_data},
|
||||
)
|
||||
|
||||
text_content = ""
|
||||
tool_calls = []
|
||||
|
||||
for block in response_json.get("content", []):
|
||||
if block.get("type") == "text":
|
||||
text_content += block.get("text", "")
|
||||
elif block.get("type") == "tool_use":
|
||||
tool_calls.append(
|
||||
{
|
||||
"id": block.get("id", ""),
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": block.get("name", ""),
|
||||
"arguments": json.dumps(block.get("input", {})),
|
||||
},
|
||||
}
|
||||
)
|
||||
|
||||
_stop_reason_map = {
|
||||
"end_turn": "stop",
|
||||
"max_tokens": "length",
|
||||
"tool_use": "tool_calls",
|
||||
"stop_sequence": "stop",
|
||||
}
|
||||
finish_reason = _stop_reason_map.get(
|
||||
response_json.get("stop_reason", "end_turn"), "stop"
|
||||
)
|
||||
|
||||
message = Message(content=text_content or None, role="assistant")
|
||||
if tool_calls:
|
||||
message.tool_calls = tool_calls # type: ignore
|
||||
|
||||
choice = Choices(
|
||||
finish_reason=finish_reason,
|
||||
index=0,
|
||||
message=message,
|
||||
)
|
||||
|
||||
usage_data = response_json.get("usage", {})
|
||||
usage = Usage(
|
||||
prompt_tokens=usage_data.get("input_tokens", 0),
|
||||
completion_tokens=usage_data.get("output_tokens", 0),
|
||||
total_tokens=usage_data.get("input_tokens", 0) + usage_data.get("output_tokens", 0),
|
||||
)
|
||||
|
||||
model_response.choices = [choice]
|
||||
model_response.usage = usage
|
||||
model_response.model = "snowflake/" + response_json.get("model", model)
|
||||
model_response.id = response_json.get("id", "")
|
||||
|
||||
if model is not None:
|
||||
model_response._hidden_params["model"] = model
|
||||
|
||||
return model_response
|
||||
|
|
@ -1,17 +1,25 @@
|
|||
"""
|
||||
Support for Snowflake REST API
|
||||
Snowflake Cortex REST API — Chat Transformation
|
||||
|
||||
Routes to the native OpenAI-compatible endpoint:
|
||||
POST /api/v2/cortex/v1/chat/completions
|
||||
|
||||
Previously used the legacy endpoint /api/v2/cortex/inference:complete which
|
||||
required Snowflake-specific payload transformations and did not support the
|
||||
full OpenAI parameter surface. This version uses the native endpoint which
|
||||
accepts standard OpenAI chat completions format directly.
|
||||
|
||||
Ref: https://docs.snowflake.com/en/user-guide/snowflake-cortex/cortex-rest-api
|
||||
"""
|
||||
|
||||
import json
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union
|
||||
from typing import TYPE_CHECKING, Any, List, Optional
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.utils import ChatCompletionMessageToolCall, Function, ModelResponse
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
from ...openai_like.chat.transformation import OpenAIGPTConfig
|
||||
|
||||
from ..utils import SnowflakeBaseConfig
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -24,67 +32,85 @@ else:
|
|||
|
||||
class SnowflakeConfig(SnowflakeBaseConfig, OpenAIGPTConfig):
|
||||
"""
|
||||
Reference: https://docs.snowflake.com/en/user-guide/snowflake-cortex/cortex-llm-rest-api
|
||||
Snowflake Cortex REST API — OpenAI-compatible endpoint.
|
||||
|
||||
Snowflake Cortex LLM REST API supports function calling with specific models (e.g., Claude 3.5 Sonnet).
|
||||
This config handles transformation between OpenAI format and Snowflake's tool_spec format.
|
||||
Endpoint: POST /api/v2/cortex/v1/chat/completions
|
||||
|
||||
Supports all Snowflake Cortex models (Llama, Mistral, DeepSeek, Snowflake
|
||||
Arctic, and Claude series) via the OpenAI-compatible interface.
|
||||
|
||||
For Claude-specific features (thinking, cache_control, extended context),
|
||||
use SnowflakeCortexAnthropicConfig which routes to /api/v2/cortex/v1/messages.
|
||||
|
||||
Auth:
|
||||
PAT: api_key="pat/<token>" → X-Snowflake-Authorization-Token-Type: PROGRAMMATIC_ACCESS_TOKEN
|
||||
JWT: api_key="<jwt>" → X-Snowflake-Authorization-Token-Type: KEYPAIR_JWT
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
def get_config(cls):
|
||||
return super().get_config()
|
||||
|
||||
def _transform_tool_calls_from_snowflake_to_openai(
|
||||
self, content_list: List[Dict[str, Any]]
|
||||
) -> Tuple[str, Optional[List[ChatCompletionMessageToolCall]]]:
|
||||
def get_supported_openai_params(self, model: str) -> List[str]:
|
||||
return [
|
||||
"temperature",
|
||||
"max_tokens",
|
||||
"top_p",
|
||||
"stream",
|
||||
"response_format",
|
||||
"tools",
|
||||
"tool_choice",
|
||||
]
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: Optional[str],
|
||||
api_key: Optional[str],
|
||||
model: str,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
stream: Optional[bool] = None,
|
||||
) -> str:
|
||||
"""
|
||||
Transform Snowflake tool calls to OpenAI format.
|
||||
Returns the native OpenAI-compatible Cortex REST API endpoint.
|
||||
|
||||
Args:
|
||||
content_list: Snowflake's content_list array containing text and tool_use items
|
||||
_get_api_base normalizes api_base to:
|
||||
https://{account}.snowflakecomputing.com/api/v2
|
||||
|
||||
Returns:
|
||||
Tuple of (text_content, tool_calls)
|
||||
We append:
|
||||
/cortex/v1/chat/completions
|
||||
|
||||
Snowflake format in content_list:
|
||||
{
|
||||
"type": "tool_use",
|
||||
"tool_use": {
|
||||
"tool_use_id": "tooluse_...",
|
||||
"name": "get_weather",
|
||||
"input": {"location": "Paris"}
|
||||
}
|
||||
Resulting in:
|
||||
https://{account}.snowflakecomputing.com/api/v2/cortex/v1/chat/completions
|
||||
"""
|
||||
api_base = self._get_api_base(api_base, optional_params)
|
||||
return f"{api_base}/cortex/v1/chat/completions"
|
||||
|
||||
def transform_request(
|
||||
self,
|
||||
model: str,
|
||||
messages: List[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
headers: dict,
|
||||
) -> dict:
|
||||
"""
|
||||
Transform to OpenAI chat completions format.
|
||||
|
||||
The native /chat/completions endpoint accepts standard OpenAI format
|
||||
directly — no Snowflake-specific tool_spec transformation required.
|
||||
"""
|
||||
stream: bool = optional_params.pop("stream", False) or False
|
||||
extra_body = optional_params.pop("extra_body", {})
|
||||
|
||||
return {
|
||||
"model": model,
|
||||
"messages": messages,
|
||||
"stream": stream,
|
||||
**optional_params,
|
||||
**extra_body,
|
||||
}
|
||||
|
||||
OpenAI format (returned tool_calls):
|
||||
ChatCompletionMessageToolCall(
|
||||
id="tooluse_...",
|
||||
type="function",
|
||||
function=Function(name="get_weather", arguments='{"location": "Paris"}')
|
||||
)
|
||||
"""
|
||||
text_content = ""
|
||||
tool_calls: List[ChatCompletionMessageToolCall] = []
|
||||
|
||||
for idx, content_item in enumerate(content_list):
|
||||
if content_item.get("type") == "text":
|
||||
text_content += content_item.get("text", "")
|
||||
|
||||
## TOOL CALLING
|
||||
elif content_item.get("type") == "tool_use":
|
||||
tool_use_data = content_item.get("tool_use", {})
|
||||
tool_call = ChatCompletionMessageToolCall(
|
||||
id=tool_use_data.get("tool_use_id", ""),
|
||||
type="function",
|
||||
function=Function(
|
||||
name=tool_use_data.get("name", ""),
|
||||
arguments=json.dumps(tool_use_data.get("input", {})),
|
||||
),
|
||||
)
|
||||
tool_calls.append(tool_call)
|
||||
|
||||
return text_content, tool_calls if tool_calls else None
|
||||
|
||||
def transform_response(
|
||||
self,
|
||||
model: str,
|
||||
|
|
@ -99,6 +125,12 @@ class SnowflakeConfig(SnowflakeBaseConfig, OpenAIGPTConfig):
|
|||
api_key: Optional[str] = None,
|
||||
json_mode: Optional[bool] = None,
|
||||
) -> ModelResponse:
|
||||
"""
|
||||
Transform from standard OpenAI chat completions response.
|
||||
|
||||
The native endpoint returns standard OpenAI format — no content_list
|
||||
transformation required.
|
||||
"""
|
||||
response_json = raw_response.json()
|
||||
|
||||
logging_obj.post_call(
|
||||
|
|
@ -108,180 +140,10 @@ class SnowflakeConfig(SnowflakeBaseConfig, OpenAIGPTConfig):
|
|||
additional_args={"complete_input_dict": request_data},
|
||||
)
|
||||
|
||||
## RESPONSE TRANSFORMATION
|
||||
# Snowflake returns content_list (not content) with tool_use objects
|
||||
# We need to transform this to OpenAI's format with content + tool_calls
|
||||
if "choices" in response_json and len(response_json["choices"]) > 0:
|
||||
choice = response_json["choices"][0]
|
||||
if "message" in choice and "content_list" in choice["message"]:
|
||||
content_list = choice["message"]["content_list"]
|
||||
(
|
||||
text_content,
|
||||
tool_calls,
|
||||
) = self._transform_tool_calls_from_snowflake_to_openai(content_list)
|
||||
|
||||
# Update the choice message with OpenAI format
|
||||
choice["message"]["content"] = text_content
|
||||
if tool_calls:
|
||||
choice["message"]["tool_calls"] = tool_calls
|
||||
|
||||
# Remove Snowflake-specific content_list
|
||||
del choice["message"]["content_list"]
|
||||
|
||||
returned_response = ModelResponse(**response_json)
|
||||
|
||||
returned_response.model = "snowflake/" + (returned_response.model or "")
|
||||
|
||||
if model is not None:
|
||||
returned_response._hidden_params["model"] = model
|
||||
|
||||
return returned_response
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: Optional[str],
|
||||
api_key: Optional[str],
|
||||
model: str,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
stream: Optional[bool] = None,
|
||||
) -> str:
|
||||
"""
|
||||
If api_base is not provided, use the default DeepSeek /chat/completions endpoint.
|
||||
"""
|
||||
|
||||
api_base = self._get_api_base(api_base, optional_params)
|
||||
|
||||
return f"{api_base}/cortex/inference:complete"
|
||||
|
||||
def _transform_tools(self, tools: List[Dict[str, Any]]) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Transform OpenAI tool format to Snowflake tool format.
|
||||
|
||||
Args:
|
||||
tools: List of tools in OpenAI format
|
||||
|
||||
Returns:
|
||||
List of tools in Snowflake format
|
||||
|
||||
OpenAI format:
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "get_weather",
|
||||
"description": "...",
|
||||
"parameters": {...}
|
||||
}
|
||||
}
|
||||
|
||||
Snowflake format:
|
||||
{
|
||||
"tool_spec": {
|
||||
"type": "generic",
|
||||
"name": "get_weather",
|
||||
"description": "...",
|
||||
"input_schema": {...}
|
||||
}
|
||||
}
|
||||
"""
|
||||
snowflake_tools: List[Dict[str, Any]] = []
|
||||
for tool in tools:
|
||||
if tool.get("type") == "function":
|
||||
function = tool.get("function", {})
|
||||
snowflake_tool: Dict[str, Any] = {
|
||||
"tool_spec": {
|
||||
"type": "generic",
|
||||
"name": function.get("name"),
|
||||
"input_schema": function.get(
|
||||
"parameters",
|
||||
{"type": "object", "properties": {}},
|
||||
),
|
||||
}
|
||||
}
|
||||
# Add description if present
|
||||
if "description" in function:
|
||||
snowflake_tool["tool_spec"]["description"] = function["description"]
|
||||
|
||||
snowflake_tools.append(snowflake_tool)
|
||||
|
||||
return snowflake_tools
|
||||
|
||||
def _transform_tool_choice(
|
||||
self, tool_choice: Union[str, Dict[str, Any]]
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Transform OpenAI tool_choice format to Snowflake format.
|
||||
|
||||
Snowflake requires tool_choice to be an object, not a string.
|
||||
Ref: https://docs.snowflake.com/en/developer-guide/snowflake-rest-api/reference/cortex-inference#post--api-v2-cortex-inference-complete-req-body-schema
|
||||
|
||||
Args:
|
||||
tool_choice: Tool choice in OpenAI format (str or dict)
|
||||
|
||||
Returns:
|
||||
Tool choice in Snowflake format (always an object, never a string)
|
||||
|
||||
OpenAI format (string):
|
||||
"auto", "required", "none"
|
||||
|
||||
OpenAI format (dict):
|
||||
{"type": "function", "function": {"name": "get_weather"}}
|
||||
|
||||
Snowflake format:
|
||||
{"type": "auto"} / {"type": "any"} / {"type": "none"}
|
||||
{"type": "tool", "name": ["get_weather"]}
|
||||
|
||||
Snowflake's API (like Anthropic) requires tool_choice as an object
|
||||
with a "type" field, not as a bare string.
|
||||
"""
|
||||
if isinstance(tool_choice, str):
|
||||
# Snowflake requires object format, not string.
|
||||
# Map OpenAI string values to Snowflake object format.
|
||||
# "required" maps to "any" (Snowflake/Anthropic convention).
|
||||
_type_map = {
|
||||
"auto": "auto",
|
||||
"required": "any",
|
||||
"none": "none",
|
||||
}
|
||||
mapped_type = _type_map.get(tool_choice, tool_choice)
|
||||
return {"type": mapped_type}
|
||||
|
||||
if isinstance(tool_choice, dict):
|
||||
if tool_choice.get("type") == "function":
|
||||
function_name = tool_choice.get("function", {}).get("name")
|
||||
if function_name:
|
||||
return {
|
||||
"type": "tool",
|
||||
"name": [function_name], # Snowflake expects array
|
||||
}
|
||||
|
||||
return tool_choice
|
||||
|
||||
def transform_request(
|
||||
self,
|
||||
model: str,
|
||||
messages: List[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
headers: dict,
|
||||
) -> dict:
|
||||
stream: bool = optional_params.pop("stream", None) or False
|
||||
extra_body = optional_params.pop("extra_body", {})
|
||||
|
||||
## TOOL CALLING
|
||||
# Transform tools from OpenAI format to Snowflake's tool_spec format
|
||||
tools = optional_params.pop("tools", None)
|
||||
if tools:
|
||||
optional_params["tools"] = self._transform_tools(tools)
|
||||
|
||||
# Transform tool_choice from OpenAI format to Snowflake's tool name array format
|
||||
tool_choice = optional_params.pop("tool_choice", None)
|
||||
if tool_choice:
|
||||
optional_params["tool_choice"] = self._transform_tool_choice(tool_choice)
|
||||
|
||||
return {
|
||||
"model": model,
|
||||
"messages": messages,
|
||||
"stream": stream,
|
||||
**optional_params,
|
||||
**extra_body,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -0,0 +1,488 @@
|
|||
"""
|
||||
Tests for Snowflake Cortex native endpoint migration.
|
||||
|
||||
Covers:
|
||||
- SnowflakeConfig (OpenAI-compatible /chat/completions)
|
||||
- SnowflakeCortexAnthropicConfig (Anthropic-compatible /messages)
|
||||
|
||||
Run:
|
||||
pytest tests/test_litellm/llms/snowflake/test_snowflake_native_endpoints.py -v
|
||||
"""
|
||||
|
||||
import json
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
from litellm.llms.snowflake.chat.transformation import SnowflakeConfig
|
||||
from litellm.llms.snowflake.anthropic.transformation import (
|
||||
SnowflakeCortexAnthropicConfig,
|
||||
_is_claude_model,
|
||||
)
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
|
||||
# ─── Fixtures ──────────────────────────────────────────────────────────────
|
||||
|
||||
ACCOUNT_ID = "myaccount"
|
||||
API_BASE = f"https://{ACCOUNT_ID}.snowflakecomputing.com"
|
||||
PAT_TOKEN = "pat/my-secret-pat-token"
|
||||
JWT_TOKEN = "eyJhbGciOiJSUzI1NiJ9.test"
|
||||
|
||||
|
||||
def _mock_logging():
|
||||
m = MagicMock()
|
||||
m.post_call = MagicMock()
|
||||
return m
|
||||
|
||||
|
||||
def _make_openai_response(content: str = "Hello!") -> httpx.Response:
|
||||
body = {
|
||||
"id": "chatcmpl-abc123",
|
||||
"object": "chat.completion",
|
||||
"model": "llama3.1-70b",
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {"role": "assistant", "content": content},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
],
|
||||
"usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15},
|
||||
}
|
||||
return httpx.Response(200, json=body)
|
||||
|
||||
|
||||
def _make_anthropic_response(content: str = "Hello!") -> httpx.Response:
|
||||
body = {
|
||||
"id": "msg_abc123",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"model": "claude-sonnet-4-5",
|
||||
"content": [{"type": "text", "text": content}],
|
||||
"stop_reason": "end_turn",
|
||||
"usage": {"input_tokens": 10, "output_tokens": 5},
|
||||
}
|
||||
return httpx.Response(200, json=body)
|
||||
|
||||
|
||||
# ─── SnowflakeConfig (OpenAI-compatible) ───────────────────────────────────
|
||||
|
||||
class TestSnowflakeConfigURL:
|
||||
def setup_method(self):
|
||||
self.cfg = SnowflakeConfig()
|
||||
|
||||
def test_url_with_account_id_in_optional_params(self):
|
||||
optional_params = {"account_id": ACCOUNT_ID}
|
||||
url = self.cfg.get_complete_url(
|
||||
api_base=None,
|
||||
api_key=JWT_TOKEN,
|
||||
model="snowflake/llama3.1-70b",
|
||||
optional_params=optional_params,
|
||||
litellm_params={},
|
||||
)
|
||||
assert url == f"https://{ACCOUNT_ID}.snowflakecomputing.com/api/v2/cortex/v1/chat/completions"
|
||||
|
||||
def test_url_with_explicit_api_base(self):
|
||||
url = self.cfg.get_complete_url(
|
||||
api_base=API_BASE,
|
||||
api_key=JWT_TOKEN,
|
||||
model="snowflake/llama3.1-70b",
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
)
|
||||
assert url.endswith("/api/v2/cortex/v1/chat/completions")
|
||||
assert "cortex/inference:complete" not in url
|
||||
|
||||
def test_url_never_uses_legacy_endpoint(self):
|
||||
url = self.cfg.get_complete_url(
|
||||
api_base=API_BASE,
|
||||
api_key=JWT_TOKEN,
|
||||
model="snowflake/claude-sonnet-4-5",
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
)
|
||||
assert "inference:complete" not in url
|
||||
assert "/v1/chat/completions" in url
|
||||
|
||||
def test_url_works_for_claude_models(self):
|
||||
url = self.cfg.get_complete_url(
|
||||
api_base=API_BASE,
|
||||
api_key=JWT_TOKEN,
|
||||
model="snowflake/claude-sonnet-4-5",
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
)
|
||||
assert "/cortex/v1/chat/completions" in url
|
||||
|
||||
def test_url_works_for_llama_models(self):
|
||||
url = self.cfg.get_complete_url(
|
||||
api_base=API_BASE,
|
||||
api_key=JWT_TOKEN,
|
||||
model="snowflake/llama3.1-70b",
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
)
|
||||
assert "/cortex/v1/chat/completions" in url
|
||||
|
||||
|
||||
class TestSnowflakeConfigAuth:
|
||||
def setup_method(self):
|
||||
self.cfg = SnowflakeConfig()
|
||||
|
||||
def test_pat_auth_strips_prefix_and_sets_header(self):
|
||||
headers = self.cfg.validate_environment(
|
||||
headers={},
|
||||
model="snowflake/llama3.1-70b",
|
||||
messages=[],
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
api_key=PAT_TOKEN,
|
||||
)
|
||||
assert headers["X-Snowflake-Authorization-Token-Type"] == "PROGRAMMATIC_ACCESS_TOKEN"
|
||||
assert headers["Authorization"] == "Bearer my-secret-pat-token"
|
||||
|
||||
def test_jwt_auth_sets_keypair_header(self):
|
||||
headers = self.cfg.validate_environment(
|
||||
headers={},
|
||||
model="snowflake/llama3.1-70b",
|
||||
messages=[],
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
api_key=JWT_TOKEN,
|
||||
)
|
||||
assert headers["X-Snowflake-Authorization-Token-Type"] == "KEYPAIR_JWT"
|
||||
assert headers["Authorization"] == f"Bearer {JWT_TOKEN}"
|
||||
|
||||
def test_missing_api_key_raises(self):
|
||||
with pytest.raises(ValueError, match="Missing Snowflake JWT key"):
|
||||
self.cfg.validate_environment(
|
||||
headers={},
|
||||
model="snowflake/llama3.1-70b",
|
||||
messages=[],
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
api_key=None,
|
||||
)
|
||||
|
||||
|
||||
class TestSnowflakeConfigRequest:
|
||||
def setup_method(self):
|
||||
self.cfg = SnowflakeConfig()
|
||||
self.messages = [{"role": "user", "content": "hello"}]
|
||||
|
||||
def test_request_uses_openai_tool_format(self):
|
||||
tools = [
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "get_weather",
|
||||
"description": "Get weather",
|
||||
"parameters": {"type": "object", "properties": {"city": {"type": "string"}}},
|
||||
},
|
||||
}
|
||||
]
|
||||
body = self.cfg.transform_request(
|
||||
model="snowflake/llama3.1-70b",
|
||||
messages=self.messages,
|
||||
optional_params={"tools": tools},
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
assert body["tools"] == tools
|
||||
assert "tool_spec" not in json.dumps(body)
|
||||
|
||||
def test_stream_defaults_to_false(self):
|
||||
body = self.cfg.transform_request(
|
||||
model="snowflake/llama3.1-70b",
|
||||
messages=self.messages,
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
assert body["stream"] is False
|
||||
|
||||
def test_stream_true_passes_through(self):
|
||||
body = self.cfg.transform_request(
|
||||
model="snowflake/llama3.1-70b",
|
||||
messages=self.messages,
|
||||
optional_params={"stream": True},
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
assert body["stream"] is True
|
||||
|
||||
def test_supported_params_includes_stream(self):
|
||||
params = self.cfg.get_supported_openai_params("snowflake/llama3.1-70b")
|
||||
assert "stream" in params
|
||||
|
||||
def test_no_content_list_in_request(self):
|
||||
body = self.cfg.transform_request(
|
||||
model="snowflake/llama3.1-70b",
|
||||
messages=self.messages,
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
assert "content_list" not in body
|
||||
|
||||
|
||||
class TestSnowflakeConfigResponse:
|
||||
def setup_method(self):
|
||||
self.cfg = SnowflakeConfig()
|
||||
|
||||
def test_standard_response_parsed(self):
|
||||
raw = _make_openai_response("Hello from Snowflake!")
|
||||
result = self.cfg.transform_response(
|
||||
model="snowflake/llama3.1-70b",
|
||||
raw_response=raw,
|
||||
model_response=ModelResponse(),
|
||||
logging_obj=_mock_logging(),
|
||||
request_data={},
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
encoding=None,
|
||||
)
|
||||
assert result.choices[0].message.content == "Hello from Snowflake!"
|
||||
assert result.model.startswith("snowflake/")
|
||||
|
||||
def test_model_prefixed_with_snowflake(self):
|
||||
raw = _make_openai_response()
|
||||
result = self.cfg.transform_response(
|
||||
model="snowflake/llama3.1-70b",
|
||||
raw_response=raw,
|
||||
model_response=ModelResponse(),
|
||||
logging_obj=_mock_logging(),
|
||||
request_data={},
|
||||
messages=[],
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
encoding=None,
|
||||
)
|
||||
assert result.model.startswith("snowflake/")
|
||||
|
||||
|
||||
# ─── SnowflakeCortexAnthropicConfig ────────────────────────────────────────
|
||||
|
||||
class TestAnthropicConfigURL:
|
||||
def setup_method(self):
|
||||
self.cfg = SnowflakeCortexAnthropicConfig()
|
||||
|
||||
def test_url_routes_to_messages_endpoint(self):
|
||||
url = self.cfg.get_complete_url(
|
||||
api_base=API_BASE,
|
||||
api_key=PAT_TOKEN,
|
||||
model="snowflake/claude-sonnet-4-5",
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
)
|
||||
assert url.endswith("/api/v2/cortex/v1/messages")
|
||||
assert "chat/completions" not in url
|
||||
assert "inference:complete" not in url
|
||||
|
||||
def test_url_with_account_id(self):
|
||||
url = self.cfg.get_complete_url(
|
||||
api_base=None,
|
||||
api_key=PAT_TOKEN,
|
||||
model="snowflake/claude-sonnet-4-5",
|
||||
optional_params={"account_id": ACCOUNT_ID},
|
||||
litellm_params={},
|
||||
)
|
||||
assert f"https://{ACCOUNT_ID}.snowflakecomputing.com/api/v2/cortex/v1/messages" == url
|
||||
|
||||
|
||||
class TestAnthropicConfigAuth:
|
||||
def setup_method(self):
|
||||
self.cfg = SnowflakeCortexAnthropicConfig()
|
||||
|
||||
def test_anthropic_version_header_set(self):
|
||||
headers = self.cfg.validate_environment(
|
||||
headers={},
|
||||
model="snowflake/claude-sonnet-4-5",
|
||||
messages=[],
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
api_key=PAT_TOKEN,
|
||||
)
|
||||
assert headers["anthropic-version"] == "2023-06-01"
|
||||
|
||||
def test_pat_auth_and_anthropic_version_combined(self):
|
||||
headers = self.cfg.validate_environment(
|
||||
headers={},
|
||||
model="snowflake/claude-sonnet-4-5",
|
||||
messages=[],
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
api_key=PAT_TOKEN,
|
||||
)
|
||||
assert headers["X-Snowflake-Authorization-Token-Type"] == "PROGRAMMATIC_ACCESS_TOKEN"
|
||||
assert headers["anthropic-version"] == "2023-06-01"
|
||||
assert "Bearer" in headers["Authorization"]
|
||||
|
||||
|
||||
class TestAnthropicConfigRequest:
|
||||
def setup_method(self):
|
||||
self.cfg = SnowflakeCortexAnthropicConfig()
|
||||
|
||||
def test_system_message_extracted_to_top_level(self):
|
||||
messages = [
|
||||
{"role": "system", "content": "You are helpful."},
|
||||
{"role": "user", "content": "Hello"},
|
||||
]
|
||||
body = self.cfg.transform_request(
|
||||
model="snowflake/claude-sonnet-4-5",
|
||||
messages=messages,
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
assert body["system"] == "You are helpful."
|
||||
assert all(m["role"] != "system" for m in body["messages"])
|
||||
assert body["messages"][0] == {"role": "user", "content": "Hello"}
|
||||
|
||||
def test_model_prefix_stripped(self):
|
||||
body = self.cfg.transform_request(
|
||||
model="snowflake/claude-sonnet-4-5",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
assert body["model"] == "claude-sonnet-4-5"
|
||||
assert "snowflake/" not in body["model"]
|
||||
|
||||
def test_max_tokens_defaulted_when_missing(self):
|
||||
body = self.cfg.transform_request(
|
||||
model="snowflake/claude-sonnet-4-5",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
assert "max_tokens" in body
|
||||
assert body["max_tokens"] == 1024
|
||||
|
||||
def test_max_tokens_not_overridden_when_provided(self):
|
||||
body = self.cfg.transform_request(
|
||||
model="snowflake/claude-sonnet-4-5",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
optional_params={"max_tokens": 500},
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
assert body["max_tokens"] == 500
|
||||
|
||||
def test_no_system_key_when_no_system_message(self):
|
||||
body = self.cfg.transform_request(
|
||||
model="snowflake/claude-sonnet-4-5",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
assert "system" not in body
|
||||
|
||||
|
||||
class TestAnthropicConfigResponse:
|
||||
def setup_method(self):
|
||||
self.cfg = SnowflakeCortexAnthropicConfig()
|
||||
|
||||
def test_anthropic_response_to_openai_format(self):
|
||||
raw = _make_anthropic_response("Hi there!")
|
||||
result = self.cfg.transform_response(
|
||||
model="snowflake/claude-sonnet-4-5",
|
||||
raw_response=raw,
|
||||
model_response=ModelResponse(),
|
||||
logging_obj=_mock_logging(),
|
||||
request_data={},
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
encoding=None,
|
||||
)
|
||||
assert result.choices[0].message.content == "Hi there!"
|
||||
assert result.choices[0].finish_reason == "stop"
|
||||
|
||||
def test_usage_tokens_mapped(self):
|
||||
raw = _make_anthropic_response()
|
||||
result = self.cfg.transform_response(
|
||||
model="snowflake/claude-sonnet-4-5",
|
||||
raw_response=raw,
|
||||
model_response=ModelResponse(),
|
||||
logging_obj=_mock_logging(),
|
||||
request_data={},
|
||||
messages=[],
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
encoding=None,
|
||||
)
|
||||
assert result.usage.prompt_tokens == 10
|
||||
assert result.usage.completion_tokens == 5
|
||||
assert result.usage.total_tokens == 15
|
||||
|
||||
def test_stop_reason_end_turn_maps_to_stop(self):
|
||||
raw = _make_anthropic_response()
|
||||
result = self.cfg.transform_response(
|
||||
model="snowflake/claude-sonnet-4-5",
|
||||
raw_response=raw,
|
||||
model_response=ModelResponse(),
|
||||
logging_obj=_mock_logging(),
|
||||
request_data={},
|
||||
messages=[],
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
encoding=None,
|
||||
)
|
||||
assert result.choices[0].finish_reason == "stop"
|
||||
|
||||
def test_tool_use_block_mapped_to_tool_calls(self):
|
||||
body = {
|
||||
"id": "msg_tool",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"model": "claude-sonnet-4-5",
|
||||
"content": [
|
||||
{
|
||||
"type": "tool_use",
|
||||
"id": "toolu_01",
|
||||
"name": "get_weather",
|
||||
"input": {"city": "Paris"},
|
||||
}
|
||||
],
|
||||
"stop_reason": "tool_use",
|
||||
"usage": {"input_tokens": 20, "output_tokens": 10},
|
||||
}
|
||||
raw = httpx.Response(200, json=body)
|
||||
result = self.cfg.transform_response(
|
||||
model="snowflake/claude-sonnet-4-5",
|
||||
raw_response=raw,
|
||||
model_response=ModelResponse(),
|
||||
logging_obj=_mock_logging(),
|
||||
request_data={},
|
||||
messages=[],
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
encoding=None,
|
||||
)
|
||||
assert result.choices[0].finish_reason == "tool_calls"
|
||||
tool_calls = result.choices[0].message.tool_calls
|
||||
assert len(tool_calls) == 1
|
||||
assert tool_calls[0]["function"]["name"] == "get_weather"
|
||||
assert json.loads(tool_calls[0]["function"]["arguments"]) == {"city": "Paris"}
|
||||
|
||||
|
||||
# ─── Model detection helper ────────────────────────────────────────────────
|
||||
|
||||
class TestIsClaudeModel:
|
||||
def test_claude_model_detected(self):
|
||||
assert _is_claude_model("snowflake/claude-sonnet-4-5") is True
|
||||
assert _is_claude_model("claude-3-haiku") is True
|
||||
assert _is_claude_model("snowflake/claude-opus-4") is True
|
||||
|
||||
def test_non_claude_not_detected(self):
|
||||
assert _is_claude_model("snowflake/llama3.1-70b") is False
|
||||
assert _is_claude_model("snowflake/mistral-large") is False
|
||||
assert _is_claude_model("snowflake/deepseek-r1") is False
|
||||
assert _is_claude_model("snowflake/snowflake-arctic") is False
|
||||
Loading…
Add table
Reference in a new issue