mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-01 02:02:20 +00:00
Add SiliconFlow as an OpenAI-compatible chat provider
SiliconFlow exposes an OpenAI-compatible API and has been requested several times (#12888, #8263). This wires it in as a first-class provider: a SiliconFlowConfig subclass of OpenAIGPTConfig, registration in the provider enum and OpenAI-compatible provider/endpoint lists, and api_base/api_key resolution (default https://api.siliconflow.com/v1, overridable via SILICONFLOW_API_BASE for the mainland endpoint). Adds mocked unit tests covering config precedence and end-to-end provider resolution.
This commit is contained in:
parent
12d29a38a7
commit
ca8e9f1bcf
9 changed files with 193 additions and 0 deletions
|
|
@ -616,6 +616,7 @@ snowflake_models: Set = set()
|
|||
gradient_ai_models: Set = set()
|
||||
llama_models: Set = set()
|
||||
nscale_models: Set = set()
|
||||
siliconflow_models: Set = set()
|
||||
nebius_models: Set = set()
|
||||
nebius_embedding_models: Set = set()
|
||||
aiml_models: Set = set()
|
||||
|
|
@ -805,6 +806,8 @@ def add_known_models(model_cost_map: Optional[Dict] = None):
|
|||
llama_models.add(key)
|
||||
elif value.get("litellm_provider") == "nscale":
|
||||
nscale_models.add(key)
|
||||
elif value.get("litellm_provider") == "siliconflow":
|
||||
siliconflow_models.add(key)
|
||||
elif value.get("litellm_provider") == "azure_ai":
|
||||
azure_ai_models.add(key)
|
||||
elif value.get("litellm_provider") == "voyage":
|
||||
|
|
@ -1009,6 +1012,7 @@ model_list = list(
|
|||
| llama_models
|
||||
| featherless_ai_models
|
||||
| nscale_models
|
||||
| siliconflow_models
|
||||
| deepgram_models
|
||||
| elevenlabs_models
|
||||
| dashscope_models
|
||||
|
|
@ -1107,6 +1111,7 @@ models_by_provider: dict = {
|
|||
"gradient_ai": gradient_ai_models,
|
||||
"meta_llama": llama_models,
|
||||
"nscale": nscale_models,
|
||||
"siliconflow": siliconflow_models,
|
||||
"featherless_ai": featherless_ai_models,
|
||||
"deepgram": deepgram_models,
|
||||
"elevenlabs": elevenlabs_models,
|
||||
|
|
@ -1787,6 +1792,9 @@ if TYPE_CHECKING:
|
|||
PerplexityChatConfig as _PerplexityChatConfig,
|
||||
)
|
||||
from .llms.nscale.chat.transformation import NscaleConfig as _NscaleConfig
|
||||
from .llms.siliconflow.chat.transformation import (
|
||||
SiliconFlowConfig as _SiliconFlowConfig,
|
||||
)
|
||||
from .llms.watsonx.chat.transformation import (
|
||||
IBMWatsonXChatConfig as _IBMWatsonXChatConfig,
|
||||
)
|
||||
|
|
@ -1821,6 +1829,7 @@ if TYPE_CHECKING:
|
|||
AzureOpenAIO1Config: Type[_AzureOpenAIO1Config]
|
||||
PerplexityChatConfig: Type[_PerplexityChatConfig]
|
||||
NscaleConfig: Type[_NscaleConfig]
|
||||
SiliconFlowConfig: Type[_SiliconFlowConfig]
|
||||
IBMWatsonXChatConfig: Type[_IBMWatsonXChatConfig]
|
||||
IBMWatsonXAIConfig: Type[_IBMWatsonXAIConfig]
|
||||
LiteLLMProxyChatConfig: Type[_LiteLLMProxyChatConfig]
|
||||
|
|
|
|||
|
|
@ -285,6 +285,7 @@ LLM_CONFIG_NAMES = (
|
|||
"LMStudioChatConfig",
|
||||
"LmStudioEmbeddingConfig",
|
||||
"NscaleConfig",
|
||||
"SiliconFlowConfig",
|
||||
"PerplexityChatConfig",
|
||||
"AzureOpenAIO1Config",
|
||||
"IBMWatsonXAIConfig",
|
||||
|
|
@ -1088,6 +1089,10 @@ _LLM_CONFIGS_IMPORT_MAP = {
|
|||
"LmStudioEmbeddingConfig",
|
||||
),
|
||||
"NscaleConfig": (".llms.nscale.chat.transformation", "NscaleConfig"),
|
||||
"SiliconFlowConfig": (
|
||||
".llms.siliconflow.chat.transformation",
|
||||
"SiliconFlowConfig",
|
||||
),
|
||||
"PerplexityChatConfig": (
|
||||
".llms.perplexity.chat.transformation",
|
||||
"PerplexityChatConfig",
|
||||
|
|
|
|||
|
|
@ -611,6 +611,7 @@ LITELLM_CHAT_PROVIDERS = [
|
|||
"meta_llama",
|
||||
"featherless_ai",
|
||||
"nscale",
|
||||
"siliconflow",
|
||||
"nebius",
|
||||
"dashscope",
|
||||
"moonshot",
|
||||
|
|
@ -766,6 +767,7 @@ openai_compatible_endpoints: List = [
|
|||
"api.llama.com/compat/v1/",
|
||||
"api.featherless.ai/v1",
|
||||
"inference.api.nscale.com/v1",
|
||||
"api.siliconflow.com/v1",
|
||||
"api.studio.nebius.ai/v1",
|
||||
"https://dashscope-intl.aliyuncs.com/compatible-mode/v1",
|
||||
"https://api.moonshot.ai/v1",
|
||||
|
|
@ -826,6 +828,7 @@ openai_compatible_providers: List = [
|
|||
"chutes", # Chutes - JSON-configured provider
|
||||
"featherless_ai",
|
||||
"nscale",
|
||||
"siliconflow",
|
||||
"nebius",
|
||||
"dashscope",
|
||||
"moonshot",
|
||||
|
|
|
|||
|
|
@ -311,6 +311,9 @@ def get_llm_provider( # noqa: PLR0915
|
|||
elif endpoint == litellm.NscaleConfig.API_BASE_URL:
|
||||
custom_llm_provider = "nscale"
|
||||
dynamic_api_key = litellm.NscaleConfig.get_api_key()
|
||||
elif endpoint == "api.siliconflow.com/v1":
|
||||
custom_llm_provider = "siliconflow"
|
||||
dynamic_api_key = litellm.SiliconFlowConfig.get_api_key()
|
||||
elif endpoint == "dashscope-intl.aliyuncs.com/compatible-mode/v1":
|
||||
custom_llm_provider = "dashscope"
|
||||
dynamic_api_key = get_secret_str("DASHSCOPE_API_KEY")
|
||||
|
|
@ -881,6 +884,13 @@ def _get_openai_compatible_provider_info( # noqa: PLR0915
|
|||
) = litellm.NscaleConfig()._get_openai_compatible_provider_info(
|
||||
api_base=api_base, api_key=api_key
|
||||
)
|
||||
elif custom_llm_provider == "siliconflow":
|
||||
(
|
||||
api_base,
|
||||
dynamic_api_key,
|
||||
) = litellm.SiliconFlowConfig()._get_openai_compatible_provider_info(
|
||||
api_base=api_base, api_key=api_key
|
||||
)
|
||||
elif custom_llm_provider == "heroku":
|
||||
(
|
||||
api_base,
|
||||
|
|
|
|||
|
|
@ -261,6 +261,8 @@ def get_supported_openai_params( # noqa: PLR0915
|
|||
return litellm.PerplexityChatConfig().get_supported_openai_params(model=model)
|
||||
elif custom_llm_provider == "nscale":
|
||||
return litellm.NscaleConfig().get_supported_openai_params(model=model)
|
||||
elif custom_llm_provider == "siliconflow":
|
||||
return litellm.SiliconFlowConfig().get_supported_openai_params(model=model)
|
||||
elif custom_llm_provider == "anyscale":
|
||||
return [
|
||||
"temperature",
|
||||
|
|
|
|||
59
litellm/llms/siliconflow/chat/transformation.py
Normal file
59
litellm/llms/siliconflow/chat/transformation.py
Normal file
|
|
@ -0,0 +1,59 @@
|
|||
from typing import Optional, Tuple
|
||||
|
||||
from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
|
||||
|
||||
class SiliconFlowConfig(OpenAIGPTConfig):
|
||||
"""
|
||||
Reference: SiliconFlow is OpenAI compatible.
|
||||
|
||||
API Key: SILICONFLOW_API_KEY
|
||||
Default API Base: https://api.siliconflow.com/v1
|
||||
|
||||
Users on the China mainland endpoint can set
|
||||
SILICONFLOW_API_BASE=https://api.siliconflow.cn/v1
|
||||
"""
|
||||
|
||||
API_BASE_URL = "https://api.siliconflow.com/v1"
|
||||
|
||||
@property
|
||||
def custom_llm_provider(self) -> Optional[str]:
|
||||
return "siliconflow"
|
||||
|
||||
@staticmethod
|
||||
def get_api_key(api_key: Optional[str] = None) -> Optional[str]:
|
||||
return api_key or get_secret_str("SILICONFLOW_API_KEY")
|
||||
|
||||
@staticmethod
|
||||
def get_api_base(api_base: Optional[str] = None) -> Optional[str]:
|
||||
return (
|
||||
api_base
|
||||
or get_secret_str("SILICONFLOW_API_BASE")
|
||||
or SiliconFlowConfig.API_BASE_URL
|
||||
)
|
||||
|
||||
def _get_openai_compatible_provider_info(
|
||||
self, api_base: Optional[str], api_key: Optional[str]
|
||||
) -> Tuple[Optional[str], Optional[str]]:
|
||||
resolved_api_base = SiliconFlowConfig.get_api_base(api_base)
|
||||
resolved_api_key = SiliconFlowConfig.get_api_key(api_key)
|
||||
return resolved_api_base, resolved_api_key
|
||||
|
||||
def get_supported_openai_params(self, model: str) -> list:
|
||||
return [
|
||||
"max_tokens",
|
||||
"n",
|
||||
"temperature",
|
||||
"top_p",
|
||||
"stream",
|
||||
"logprobs",
|
||||
"top_logprobs",
|
||||
"frequency_penalty",
|
||||
"presence_penalty",
|
||||
"response_format",
|
||||
"stop",
|
||||
"logit_bias",
|
||||
"tools",
|
||||
"tool_choice",
|
||||
]
|
||||
|
|
@ -3335,6 +3335,7 @@ class LlmProviders(str, Enum):
|
|||
GRADIENT_AI = "gradient_ai"
|
||||
LLAMA = "meta_llama"
|
||||
NSCALE = "nscale"
|
||||
SILICONFLOW = "siliconflow"
|
||||
PG_VECTOR = "pg_vector"
|
||||
S3_VECTORS = "s3_vectors"
|
||||
HELICONE = "helicone"
|
||||
|
|
|
|||
|
|
@ -8331,6 +8331,10 @@ class ProviderConfigManager:
|
|||
),
|
||||
LlmProviders.GRADIENT_AI: (lambda: litellm.GradientAIConfig(), False),
|
||||
LlmProviders.NSCALE: (lambda: litellm.NscaleConfig(), False),
|
||||
LlmProviders.SILICONFLOW: (
|
||||
lambda: litellm.SiliconFlowConfig(),
|
||||
False,
|
||||
),
|
||||
LlmProviders.HEROKU: (lambda: litellm.HerokuChatConfig(), False),
|
||||
LlmProviders.OCI: (lambda: litellm.OCIChatConfig(), False),
|
||||
LlmProviders.HYPERBOLIC: (lambda: litellm.HyperbolicChatConfig(), False),
|
||||
|
|
|
|||
|
|
@ -0,0 +1,100 @@
|
|||
import os
|
||||
import sys
|
||||
from unittest.mock import patch
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../../../../.."))
|
||||
|
||||
import litellm
|
||||
from litellm import get_llm_provider, get_supported_openai_params
|
||||
from litellm.llms.siliconflow.chat.transformation import SiliconFlowConfig
|
||||
|
||||
|
||||
class TestSiliconFlowConfig:
|
||||
def setup_method(self):
|
||||
self.config = SiliconFlowConfig()
|
||||
|
||||
def test_custom_llm_provider(self):
|
||||
assert self.config.custom_llm_provider == "siliconflow"
|
||||
|
||||
def test_get_api_key(self):
|
||||
assert self.config.get_api_key("test-key") == "test-key"
|
||||
|
||||
with patch(
|
||||
"litellm.llms.siliconflow.chat.transformation.get_secret_str",
|
||||
return_value="env-key",
|
||||
):
|
||||
assert self.config.get_api_key() == "env-key"
|
||||
|
||||
with patch.dict(os.environ, {"SILICONFLOW_API_KEY": "env-key"}, clear=False):
|
||||
assert self.config.get_api_key() == "env-key"
|
||||
|
||||
def test_get_api_base_precedence(self):
|
||||
# Explicit argument wins over everything.
|
||||
assert (
|
||||
self.config.get_api_base("https://custom-base.com/v1")
|
||||
== "https://custom-base.com/v1"
|
||||
)
|
||||
|
||||
# SILICONFLOW_API_BASE override (e.g. the China mainland endpoint).
|
||||
with patch(
|
||||
"litellm.llms.siliconflow.chat.transformation.get_secret_str",
|
||||
return_value="https://api.siliconflow.cn/v1",
|
||||
):
|
||||
assert self.config.get_api_base() == "https://api.siliconflow.cn/v1"
|
||||
|
||||
# Falls back to the default global endpoint.
|
||||
with patch(
|
||||
"litellm.llms.siliconflow.chat.transformation.get_secret_str",
|
||||
return_value=None,
|
||||
):
|
||||
assert self.config.get_api_base() == SiliconFlowConfig.API_BASE_URL
|
||||
assert SiliconFlowConfig.API_BASE_URL == "https://api.siliconflow.com/v1"
|
||||
|
||||
def test_get_openai_compatible_provider_info(self):
|
||||
with patch.dict(os.environ, {"SILICONFLOW_API_KEY": "sk-secret"}, clear=False):
|
||||
api_base, api_key = self.config._get_openai_compatible_provider_info(
|
||||
api_base=None, api_key=None
|
||||
)
|
||||
assert api_base == "https://api.siliconflow.com/v1"
|
||||
assert api_key == "sk-secret"
|
||||
|
||||
def test_supported_params_include_tools(self):
|
||||
params = self.config.get_supported_openai_params(
|
||||
model="deepseek-ai/DeepSeek-V3"
|
||||
)
|
||||
for expected in ("temperature", "stream", "tools", "tool_choice"):
|
||||
assert expected in params
|
||||
|
||||
|
||||
class TestSiliconFlowProviderResolution:
|
||||
def test_get_llm_provider_resolves_prefixed_model(self):
|
||||
with patch.dict(os.environ, {"SILICONFLOW_API_KEY": "sk-secret"}, clear=False):
|
||||
model, provider, api_key, api_base = get_llm_provider(
|
||||
model="siliconflow/deepseek-ai/DeepSeek-V3"
|
||||
)
|
||||
assert model == "deepseek-ai/DeepSeek-V3"
|
||||
assert provider == "siliconflow"
|
||||
assert api_key == "sk-secret"
|
||||
assert api_base == "https://api.siliconflow.com/v1"
|
||||
|
||||
def test_get_llm_provider_detects_provider_from_api_base(self):
|
||||
_, provider, _, _ = get_llm_provider(
|
||||
model="deepseek-ai/DeepSeek-V3",
|
||||
api_base="https://api.siliconflow.com/v1",
|
||||
api_key="sk-secret",
|
||||
)
|
||||
assert provider == "siliconflow"
|
||||
|
||||
def test_get_supported_openai_params_routes_to_config(self):
|
||||
params = get_supported_openai_params(
|
||||
model="deepseek-ai/DeepSeek-V3", custom_llm_provider="siliconflow"
|
||||
)
|
||||
assert params is not None
|
||||
assert "tools" in params
|
||||
|
||||
def test_provider_registered_in_enum_and_lists(self):
|
||||
from litellm.types.utils import LlmProviders
|
||||
|
||||
assert LlmProviders.SILICONFLOW.value == "siliconflow"
|
||||
assert "siliconflow" in litellm.openai_compatible_providers
|
||||
assert "api.siliconflow.com/v1" in litellm.openai_compatible_endpoints
|
||||
Loading…
Add table
Reference in a new issue