Add native Snowflake Cortex LLM provider (#6005)

This commit is contained in:
alex-clawd
2026-06-02 04:10:13 -07:00
committed by GitHub
parent fee5b3e395
commit 4a0769d97c
8 changed files with 635 additions and 4 deletions

View File

@@ -288,6 +288,7 @@ SUPPORTED_NATIVE_PROVIDERS: Final[list[str]] = [
"hosted_vllm",
"cerebras",
"dashscope",
"snowflake",
]
@@ -376,6 +377,7 @@ class LLM(BaseLLM):
"hosted_vllm": "hosted_vllm",
"cerebras": "cerebras",
"dashscope": "dashscope",
"snowflake": "snowflake",
}
canonical_provider = provider_mapping.get(prefix.lower())
@@ -494,6 +496,9 @@ class LLM(BaseLLM):
# OpenRouter uses org/model format but accepts anything
return True
if provider == "snowflake":
return True
return False
@classmethod
@@ -592,6 +597,11 @@ class LLM(BaseLLM):
return BedrockCompletion
if provider == "snowflake":
from crewai.llms.providers.snowflake.completion import SnowflakeCompletion
return SnowflakeCompletion
openai_compatible_providers = {
"openrouter",
"deepseek",

View File

@@ -0,0 +1,183 @@
from __future__ import annotations
import os
from typing import Any, Literal
from pydantic import model_validator
from crewai.llms.providers.openai.completion import OpenAICompletion
from crewai.utilities.types import LLMMessage
SNOWFLAKE_CORTEX_PATH = "/api/v2/cortex/v1"
SNOWFLAKE_TOKEN_ENV_VARS = (
"SNOWFLAKE_PAT",
"SNOWFLAKE_TOKEN",
"SNOWFLAKE_JWT",
)
def _normalize_snowflake_base_url(value: str) -> str:
"""Return a Snowflake Cortex REST OpenAI-compatible base URL."""
base_url = value.strip().rstrip("/")
if not base_url:
raise ValueError("Snowflake account URL cannot be empty")
if "://" not in base_url:
base_url = f"https://{base_url}"
if base_url.endswith(SNOWFLAKE_CORTEX_PATH):
return base_url
if "/api/v2/cortex" in base_url:
raise ValueError(
"Snowflake base URL must be the account URL or Cortex API root "
f"ending in {SNOWFLAKE_CORTEX_PATH}; do not include endpoint paths."
)
return f"{base_url}{SNOWFLAKE_CORTEX_PATH}"
def _base_url_from_account_identifier(account_identifier: str) -> str:
account = account_identifier.strip()
if not account:
raise ValueError("Snowflake account identifier cannot be empty")
return _normalize_snowflake_base_url(f"{account}.snowflakecomputing.com")
class SnowflakeCompletion(OpenAICompletion):
"""Snowflake Cortex REST API native completion implementation.
Snowflake exposes an OpenAI-compatible Chat Completions endpoint at
``/api/v2/cortex/v1/chat/completions``. This provider reuses CrewAI's
native OpenAI transport while applying Snowflake-specific authentication,
endpoint normalization, and Claude-family message constraints.
"""
provider: str = "snowflake"
api: Literal["completions"] = "completions"
account_url: str | None = None
account_identifier: str | None = None
database: str | None = None
schema_name: str | None = None
warehouse: str | None = None
role: str | None = None
@model_validator(mode="before")
@classmethod
def _normalize_snowflake_fields(cls, data: Any) -> Any:
if not isinstance(data, dict):
return data
data["provider"] = "snowflake"
api = data.get("api")
if api and api != "completions":
raise ValueError(
"Snowflake Cortex native provider supports only the Chat Completions API"
)
data["api"] = "completions"
data["api_key"] = cls._resolve_token(data.get("api_key"))
resolved_base_url = cls._resolve_base_url(data)
data["base_url"] = resolved_base_url
data["account_url"] = resolved_base_url
return data
@staticmethod
def _resolve_token(api_key: str | None) -> str:
token = api_key
if not token:
for env_var in SNOWFLAKE_TOKEN_ENV_VARS:
token = os.getenv(env_var)
if token:
break
if not token:
raise ValueError(
"Snowflake token is required. Set SNOWFLAKE_PAT, SNOWFLAKE_TOKEN, "
"or SNOWFLAKE_JWT, or pass api_key."
)
if token.startswith("pat/"):
token = token.removeprefix("pat/")
return token
@classmethod
def _resolve_base_url(cls, data: dict[str, Any]) -> str:
explicit_base_url = data.get("base_url") or data.get("api_base")
if explicit_base_url:
return _normalize_snowflake_base_url(explicit_base_url)
account_url = data.get("account_url") or os.getenv("SNOWFLAKE_ACCOUNT_URL")
if account_url:
return _normalize_snowflake_base_url(account_url)
account_identifier = (
data.get("account_identifier")
or data.get("account")
or data.get("snowflake_account")
or os.getenv("SNOWFLAKE_ACCOUNT")
or os.getenv("SNOWFLAKE_ACCOUNT_ID")
or os.getenv("SNOWFLAKE_ACCOUNT_IDENTIFIER")
)
if account_identifier:
return _base_url_from_account_identifier(account_identifier)
raise ValueError(
"Snowflake account URL is required. Set SNOWFLAKE_ACCOUNT_URL or "
"SNOWFLAKE_ACCOUNT, or pass account_url/base_url/account_identifier."
)
def _format_messages(self, messages: str | list[LLMMessage]) -> list[LLMMessage]:
formatted_messages = super()._format_messages(messages)
if self._is_claude_model():
return self._ensure_claude_conversation_ends_with_user(formatted_messages)
return formatted_messages
def _is_claude_model(self) -> bool:
model = self.model.lower()
return model.startswith(("claude-", "anthropic."))
@staticmethod
def _ensure_claude_conversation_ends_with_user(
messages: list[LLMMessage],
) -> list[LLMMessage]:
if not messages:
return [{"role": "user", "content": "Hello"}]
if messages[-1].get("role") == "assistant" and not messages[-1].get(
"tool_calls"
):
messages = messages[:-1]
if not messages:
return [{"role": "user", "content": "Hello"}]
if messages[-1].get("role") == "user":
return messages
return [
*messages,
{
"role": "user",
"content": "Please continue and provide your final answer.",
},
]
def _prepare_completion_params(
self, messages: list[LLMMessage], tools: list[dict[str, Any]] | None = None
) -> dict[str, Any]:
params = super()._prepare_completion_params(messages=messages, tools=tools)
if self._is_claude_model() and "max_tokens" in params:
params["max_completion_tokens"] = params.pop("max_tokens")
return params
def supports_function_calling(self) -> bool:
model = self.model.lower()
return model.startswith(("openai-", "claude-", "anthropic."))
def supports_multimodal(self) -> bool:
model = self.model.lower()
return model.startswith(("openai-", "claude-", "anthropic."))

View File

@@ -0,0 +1,242 @@
from __future__ import annotations
from types import SimpleNamespace
from unittest.mock import Mock, patch
import pytest
from crewai.llm import LLM
from crewai.llms.providers.snowflake.completion import (
SNOWFLAKE_CORTEX_PATH,
SnowflakeCompletion,
_normalize_snowflake_base_url,
)
def _snowflake_env(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("SNOWFLAKE_PAT", "test-pat")
monkeypatch.setenv("SNOWFLAKE_ACCOUNT_URL", "https://org-account.snowflakecomputing.com")
monkeypatch.delenv("SNOWFLAKE_TOKEN", raising=False)
monkeypatch.delenv("SNOWFLAKE_JWT", raising=False)
monkeypatch.delenv("SNOWFLAKE_ACCOUNT", raising=False)
monkeypatch.delenv("SNOWFLAKE_ACCOUNT_ID", raising=False)
monkeypatch.delenv("SNOWFLAKE_ACCOUNT_IDENTIFIER", raising=False)
class TestSnowflakeConfig:
def test_normalizes_account_url_to_cortex_base_url(self):
assert (
_normalize_snowflake_base_url("https://org-account.snowflakecomputing.com")
== f"https://org-account.snowflakecomputing.com{SNOWFLAKE_CORTEX_PATH}"
)
def test_preserves_existing_cortex_base_url(self):
base_url = f"https://org-account.snowflakecomputing.com{SNOWFLAKE_CORTEX_PATH}"
assert _normalize_snowflake_base_url(base_url) == base_url
def test_rejects_endpoint_path_in_base_url(self):
with pytest.raises(ValueError, match="do not include endpoint paths"):
_normalize_snowflake_base_url(
"https://org-account.snowflakecomputing.com"
f"{SNOWFLAKE_CORTEX_PATH}/chat/completions"
)
def test_empty_api_key_falls_back_to_env_token(
self, monkeypatch: pytest.MonkeyPatch
):
_snowflake_env(monkeypatch)
llm = SnowflakeCompletion(model="openai-gpt-4.1", api_key="")
assert llm.api_key == "test-pat"
def test_uses_env_token_and_account_url(self, monkeypatch: pytest.MonkeyPatch):
_snowflake_env(monkeypatch)
llm = SnowflakeCompletion(model="openai-gpt-4.1")
assert llm.api_key == "test-pat"
assert llm.base_url == (
f"https://org-account.snowflakecomputing.com{SNOWFLAKE_CORTEX_PATH}"
)
assert llm.account_url == llm.base_url
def test_strips_litellm_pat_prefix_for_compatibility(
self, monkeypatch: pytest.MonkeyPatch
):
monkeypatch.setenv("SNOWFLAKE_PAT", "pat/test-pat")
monkeypatch.setenv("SNOWFLAKE_ACCOUNT", "org-account")
llm = SnowflakeCompletion(model="openai-gpt-4.1")
assert llm.api_key == "test-pat"
def test_missing_token_raises_clear_error(self, monkeypatch: pytest.MonkeyPatch):
monkeypatch.delenv("SNOWFLAKE_PAT", raising=False)
monkeypatch.delenv("SNOWFLAKE_TOKEN", raising=False)
monkeypatch.delenv("SNOWFLAKE_JWT", raising=False)
monkeypatch.setenv("SNOWFLAKE_ACCOUNT_URL", "https://org-account.snowflakecomputing.com")
with pytest.raises(ValueError, match="Snowflake token is required"):
SnowflakeCompletion(model="openai-gpt-4.1")
def test_missing_account_raises_clear_error(self, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setenv("SNOWFLAKE_PAT", "test-pat")
monkeypatch.delenv("SNOWFLAKE_ACCOUNT_URL", raising=False)
monkeypatch.delenv("SNOWFLAKE_ACCOUNT", raising=False)
monkeypatch.delenv("SNOWFLAKE_ACCOUNT_ID", raising=False)
monkeypatch.delenv("SNOWFLAKE_ACCOUNT_IDENTIFIER", raising=False)
with pytest.raises(ValueError, match="Snowflake account URL is required"):
SnowflakeCompletion(model="openai-gpt-4.1")
def test_responses_api_is_rejected(self, monkeypatch: pytest.MonkeyPatch):
_snowflake_env(monkeypatch)
with pytest.raises(ValueError, match="supports only the Chat Completions API"):
SnowflakeCompletion(model="openai-gpt-4.1", api="responses")
class TestSnowflakeFactory:
def test_llm_creates_native_snowflake_provider(
self, monkeypatch: pytest.MonkeyPatch
):
_snowflake_env(monkeypatch)
llm = LLM(model="snowflake/openai-gpt-4.1")
assert isinstance(llm, SnowflakeCompletion)
assert llm.provider == "snowflake"
assert llm.model == "openai-gpt-4.1"
assert llm.is_litellm is False
def test_explicit_provider_creates_native_snowflake_provider(
self, monkeypatch: pytest.MonkeyPatch
):
_snowflake_env(monkeypatch)
llm = LLM(model="claude-sonnet-4-5", provider="snowflake")
assert isinstance(llm, SnowflakeCompletion)
assert llm.model == "claude-sonnet-4-5"
class TestSnowflakeRequests:
def test_prepare_completion_params_uses_snowflake_model_name(
self, monkeypatch: pytest.MonkeyPatch
):
_snowflake_env(monkeypatch)
llm = SnowflakeCompletion(
model="openai-gpt-4.1",
temperature=0.2,
max_completion_tokens=128,
)
params = llm._prepare_completion_params(
[{"role": "user", "content": "Hello"}]
)
assert params["model"] == "openai-gpt-4.1"
assert params["temperature"] == 0.2
assert params["max_completion_tokens"] == 128
assert params["messages"] == [{"role": "user", "content": "Hello"}]
def test_claude_model_removes_trailing_assistant_prefill(
self, monkeypatch: pytest.MonkeyPatch
):
_snowflake_env(monkeypatch)
llm = SnowflakeCompletion(model="claude-sonnet-4-5")
messages = llm._format_messages(
[
{"role": "user", "content": "Write a summary."},
{"role": "assistant", "content": "Here is"},
]
)
assert messages == [{"role": "user", "content": "Write a summary."}]
def test_claude_model_adds_user_turn_after_tool_call_assistant_message(
self, monkeypatch: pytest.MonkeyPatch
):
_snowflake_env(monkeypatch)
llm = SnowflakeCompletion(model="claude-sonnet-4-5")
messages = llm._format_messages(
[
{"role": "user", "content": "Use the tool."},
{
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "call_1",
"type": "function",
"function": {"name": "lookup", "arguments": "{}"},
}
],
},
]
)
assert messages[-2]["role"] == "assistant"
assert messages[-2]["tool_calls"][0]["id"] == "call_1"
assert messages[-1]["role"] == "user"
def test_claude_model_maps_max_tokens_to_max_completion_tokens(
self, monkeypatch: pytest.MonkeyPatch
):
_snowflake_env(monkeypatch)
llm = SnowflakeCompletion(model="claude-sonnet-4-5", max_tokens=256)
params = llm._prepare_completion_params(
[{"role": "user", "content": "Hello"}]
)
assert "max_tokens" not in params
assert params["max_completion_tokens"] == 256
def test_streaming_params_include_usage(self, monkeypatch: pytest.MonkeyPatch):
_snowflake_env(monkeypatch)
llm = SnowflakeCompletion(model="openai-gpt-4.1", stream=True)
params = llm._prepare_completion_params(
[{"role": "user", "content": "Hello"}]
)
assert params["stream"] is True
assert params["stream_options"] == {"include_usage": True}
def test_non_streaming_call_uses_native_openai_client(
self, monkeypatch: pytest.MonkeyPatch
):
_snowflake_env(monkeypatch)
llm = SnowflakeCompletion(model="openai-gpt-4.1")
fake_response = SimpleNamespace(
usage=SimpleNamespace(
prompt_tokens=3,
completion_tokens=2,
total_tokens=5,
prompt_tokens_details=None,
completion_tokens_details=None,
),
choices=[
SimpleNamespace(
message=SimpleNamespace(content="Snowflake response", tool_calls=None)
)
],
)
create = Mock(return_value=fake_response)
fake_client = SimpleNamespace(
chat=SimpleNamespace(completions=SimpleNamespace(create=create))
)
with patch.object(llm, "_get_sync_client", return_value=fake_client):
response = llm.call([{"role": "user", "content": "Hello"}])
assert response == "Snowflake response"
create.assert_called_once()
assert create.call_args.kwargs["model"] == "openai-gpt-4.1"
assert create.call_args.kwargs["messages"] == [
{"role": "user", "content": "Hello"}
]