mirror of
https://github.com/crewAIInc/crewAI.git
synced 2026-07-29 02:29:31 +00:00
Add native Snowflake Cortex LLM provider (#6005)
This commit is contained in:
@@ -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",
|
||||
|
||||
183
lib/crewai/src/crewai/llms/providers/snowflake/completion.py
Normal file
183
lib/crewai/src/crewai/llms/providers/snowflake/completion.py
Normal 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."))
|
||||
242
lib/crewai/tests/llms/snowflake/test_snowflake.py
Normal file
242
lib/crewai/tests/llms/snowflake/test_snowflake.py
Normal 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"}
|
||||
]
|
||||
Reference in New Issue
Block a user