feat(models): expand support for additional providers and enhance auto-configuration logic
This commit is contained in:
+53
-39
@@ -1,8 +1,9 @@
|
||||
"""LLM model configuration based on LangChain init_chat_model.
|
||||
|
||||
This module provides a unified interface for creating chat model instances
|
||||
with support for multiple providers (Anthropic, OpenAI) and convenient
|
||||
short names for common models.
|
||||
with support for multiple providers (Anthropic, OpenAI, Google GenAI, NVIDIA,
|
||||
SiliconFlow, OpenRouter, Ollama, and custom OpenAI-compatible endpoints) and
|
||||
convenient short names for common models.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -15,6 +16,14 @@ from langchain.chat_models import init_chat_model
|
||||
_SILICONFLOW_BASE_URL = "https://api.siliconflow.cn/v1"
|
||||
_OPENROUTER_BASE_URL = "https://openrouter.ai/api/v1"
|
||||
|
||||
# Third-party providers routed through the OpenAI provider with a custom base_url.
|
||||
# Maps provider name → (base_url or None, env var for API key).
|
||||
_THIRD_PARTY_PROVIDERS: dict[str, tuple[str | None, str]] = {
|
||||
"siliconflow": (_SILICONFLOW_BASE_URL, "SILICONFLOW_API_KEY"),
|
||||
"openrouter": (_OPENROUTER_BASE_URL, "OPENROUTER_API_KEY"),
|
||||
"custom": (None, "CUSTOM_API_KEY"), # base_url from CUSTOM_BASE_URL env
|
||||
}
|
||||
|
||||
# Model registry: list of (short_name, model_id, provider)
|
||||
# Allows same short_name across different providers.
|
||||
_MODEL_ENTRIES: list[tuple[str, str, str]] = [
|
||||
@@ -89,6 +98,34 @@ def get_models_for_provider(provider: str) -> list[tuple[str, str]]:
|
||||
]
|
||||
|
||||
|
||||
def _apply_auto_config(
|
||||
provider: str,
|
||||
model_id: str,
|
||||
is_third_party: bool,
|
||||
kwargs: dict[str, Any],
|
||||
) -> None:
|
||||
"""Auto-enable provider-specific features (thinking, reasoning, etc.).
|
||||
|
||||
Mutates *kwargs* in place. Only sets keys that the caller hasn't already
|
||||
provided, so explicit user settings are never overridden.
|
||||
"""
|
||||
# Anthropic: extended thinking
|
||||
if provider == "anthropic" and "thinking" not in kwargs:
|
||||
if model_id.endswith("4-6"):
|
||||
kwargs["thinking"] = {"type": "adaptive"}
|
||||
kwargs.setdefault("effort", "max")
|
||||
else:
|
||||
kwargs["thinking"] = {"type": "enabled", "budget_tokens": 10000}
|
||||
|
||||
# OpenAI (native, not third-party routed): reasoning
|
||||
if provider == "openai" and not is_third_party and "reasoning" not in kwargs:
|
||||
kwargs["reasoning"] = {"effort": "high", "summary": "auto"}
|
||||
|
||||
# Google GenAI: surface thinking traces
|
||||
if provider == "google-genai":
|
||||
kwargs.setdefault("include_thoughts", True)
|
||||
|
||||
|
||||
def get_chat_model(
|
||||
model: str | None = None,
|
||||
provider: str | None = None,
|
||||
@@ -140,56 +177,33 @@ def get_chat_model(
|
||||
elif model_id.startswith("ollama:"):
|
||||
provider = "ollama"
|
||||
model_id = model_id.removeprefix("ollama:")
|
||||
elif "/" in model_id:
|
||||
provider = "nvidia"
|
||||
else:
|
||||
provider = "anthropic" # Default fallback
|
||||
|
||||
# SiliconFlow / OpenRouter / Custom → route through OpenAI provider with base_url
|
||||
_is_third_party = provider in ("siliconflow", "openrouter", "custom")
|
||||
if provider == "custom":
|
||||
base_url = os.environ.get("CUSTOM_BASE_URL", "")
|
||||
# Third-party providers → route through OpenAI provider with base_url
|
||||
_is_third_party = provider in _THIRD_PARTY_PROVIDERS
|
||||
if provider in _THIRD_PARTY_PROVIDERS:
|
||||
base_url_default, api_key_env = _THIRD_PARTY_PROVIDERS[provider]
|
||||
if provider == "custom":
|
||||
base_url = os.environ.get("CUSTOM_BASE_URL", "")
|
||||
else:
|
||||
base_url = base_url_default
|
||||
if base_url:
|
||||
kwargs["base_url"] = base_url
|
||||
api_key = os.environ.get("CUSTOM_API_KEY", "")
|
||||
if api_key:
|
||||
kwargs["api_key"] = api_key
|
||||
provider = "openai"
|
||||
elif provider == "siliconflow":
|
||||
kwargs["base_url"] = _SILICONFLOW_BASE_URL
|
||||
api_key = os.environ.get("SILICONFLOW_API_KEY", "")
|
||||
if api_key:
|
||||
kwargs["api_key"] = api_key
|
||||
# Disable thinking — LangChain drops reasoning_content from history,
|
||||
# causing SiliconFlow to reject multi-turn requests (error 20015).
|
||||
kwargs.setdefault("extra_body", {})["enable_thinking"] = False
|
||||
provider = "openai"
|
||||
elif provider == "openrouter":
|
||||
kwargs["base_url"] = _OPENROUTER_BASE_URL
|
||||
api_key = os.environ.get("OPENROUTER_API_KEY", "")
|
||||
api_key = os.environ.get(api_key_env, "")
|
||||
if api_key:
|
||||
kwargs["api_key"] = api_key
|
||||
# SiliconFlow: disable thinking — LangChain drops reasoning_content
|
||||
# from history, causing error 20015 on multi-turn requests.
|
||||
if provider == "siliconflow":
|
||||
kwargs.setdefault("extra_body", {})["enable_thinking"] = False
|
||||
provider = "openai"
|
||||
elif provider == "ollama":
|
||||
base_url = os.environ.get("OLLAMA_BASE_URL", "")
|
||||
if base_url:
|
||||
kwargs["base_url"] = base_url
|
||||
|
||||
# Auto-enable thinking for Anthropic models
|
||||
if provider == "anthropic" and "thinking" not in kwargs:
|
||||
if model_id.endswith("4-6"):
|
||||
kwargs["thinking"] = {"type": "adaptive"}
|
||||
kwargs.setdefault("effort", "max")
|
||||
else:
|
||||
kwargs["thinking"] = {"type": "enabled", "budget_tokens": 10000}
|
||||
|
||||
# Auto-enable reasoning for OpenAI models (not for third-party routed)
|
||||
if provider == "openai" and not _is_third_party and "reasoning" not in kwargs:
|
||||
kwargs["reasoning"] = {"effort": "high", "summary": "auto"}
|
||||
|
||||
# Auto-enable thinking visibility for Google GenAI models
|
||||
if provider == "google-genai":
|
||||
kwargs.setdefault("include_thoughts", True)
|
||||
_apply_auto_config(provider, model_id, _is_third_party, kwargs)
|
||||
|
||||
return init_chat_model(model=model_id, model_provider=provider, **kwargs)
|
||||
|
||||
|
||||
+144
-1
@@ -24,12 +24,14 @@ class TestModelsRegistry:
|
||||
assert isinstance(MODELS, dict)
|
||||
|
||||
def test_entries_has_all_providers(self):
|
||||
"""Test that _MODEL_ENTRIES covers native providers."""
|
||||
"""Test that _MODEL_ENTRIES covers all registered providers."""
|
||||
providers = {p for _, _, p in _MODEL_ENTRIES}
|
||||
assert "anthropic" in providers
|
||||
assert "openai" in providers
|
||||
assert "google-genai" in providers
|
||||
assert "nvidia" in providers
|
||||
assert "siliconflow" in providers
|
||||
assert "openrouter" in providers
|
||||
|
||||
def test_entries_are_valid_tuples(self):
|
||||
"""Test that _MODEL_ENTRIES contains valid (name, model_id, provider) tuples."""
|
||||
@@ -306,3 +308,144 @@ class TestOllamaProvider:
|
||||
assert call_kwargs["model"] == "phi3:mini"
|
||||
assert call_kwargs["model_provider"] == "ollama"
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Test slash model ID no longer routes to nvidia
|
||||
# =============================================================================
|
||||
|
||||
|
||||
class TestSlashModelIdFallback:
|
||||
@patch("EvoScientist.llm.models.init_chat_model")
|
||||
def test_slash_model_id_defaults_to_anthropic(self, mock_init):
|
||||
"""Unregistered model IDs containing '/' should NOT route to nvidia.
|
||||
|
||||
They fall through to the default 'anthropic' provider, consistent
|
||||
with how all other unknown model IDs are handled.
|
||||
"""
|
||||
mock_init.return_value = "mock_model"
|
||||
|
||||
get_chat_model("some-org/some-model")
|
||||
|
||||
call_kwargs = mock_init.call_args[1]
|
||||
assert call_kwargs["model"] == "some-org/some-model"
|
||||
assert call_kwargs["model_provider"] == "anthropic"
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Test third-party provider routing
|
||||
# =============================================================================
|
||||
|
||||
|
||||
class TestThirdPartyRouting:
|
||||
@patch("EvoScientist.llm.models.init_chat_model")
|
||||
def test_siliconflow_routes_through_openai(self, mock_init, monkeypatch):
|
||||
"""SiliconFlow provider should route through OpenAI with correct base_url."""
|
||||
mock_init.return_value = "mock_model"
|
||||
monkeypatch.setenv("SILICONFLOW_API_KEY", "sf-key-123")
|
||||
|
||||
get_chat_model("Pro/zai-org/GLM-5", provider="siliconflow")
|
||||
|
||||
call_kwargs = mock_init.call_args[1]
|
||||
assert call_kwargs["model_provider"] == "openai"
|
||||
assert call_kwargs["base_url"] == "https://api.siliconflow.cn/v1"
|
||||
assert call_kwargs["api_key"] == "sf-key-123"
|
||||
# SiliconFlow should disable thinking
|
||||
assert call_kwargs["extra_body"]["enable_thinking"] is False
|
||||
|
||||
@patch("EvoScientist.llm.models.init_chat_model")
|
||||
def test_openrouter_routes_through_openai(self, mock_init, monkeypatch):
|
||||
"""OpenRouter provider should route through OpenAI with correct base_url."""
|
||||
mock_init.return_value = "mock_model"
|
||||
monkeypatch.setenv("OPENROUTER_API_KEY", "or-key-456")
|
||||
|
||||
get_chat_model("x-ai/grok-4.1-fast", provider="openrouter")
|
||||
|
||||
call_kwargs = mock_init.call_args[1]
|
||||
assert call_kwargs["model_provider"] == "openai"
|
||||
assert call_kwargs["base_url"] == "https://openrouter.ai/api/v1"
|
||||
assert call_kwargs["api_key"] == "or-key-456"
|
||||
|
||||
@patch("EvoScientist.llm.models.init_chat_model")
|
||||
def test_custom_routes_through_openai(self, mock_init, monkeypatch):
|
||||
"""Custom provider should route through OpenAI with env-configured base_url."""
|
||||
mock_init.return_value = "mock_model"
|
||||
monkeypatch.setenv("CUSTOM_BASE_URL", "https://my-llm.example.com/v1")
|
||||
monkeypatch.setenv("CUSTOM_API_KEY", "custom-key-789")
|
||||
|
||||
get_chat_model("my-custom-model", provider="custom")
|
||||
|
||||
call_kwargs = mock_init.call_args[1]
|
||||
assert call_kwargs["model_provider"] == "openai"
|
||||
assert call_kwargs["base_url"] == "https://my-llm.example.com/v1"
|
||||
assert call_kwargs["api_key"] == "custom-key-789"
|
||||
|
||||
@patch("EvoScientist.llm.models.init_chat_model")
|
||||
def test_third_party_no_reasoning(self, mock_init, monkeypatch):
|
||||
"""Third-party providers routed through OpenAI should NOT get auto-reasoning."""
|
||||
mock_init.return_value = "mock_model"
|
||||
monkeypatch.setenv("OPENROUTER_API_KEY", "or-key")
|
||||
|
||||
get_chat_model("x-ai/grok-4.1-fast", provider="openrouter")
|
||||
|
||||
call_kwargs = mock_init.call_args[1]
|
||||
assert "reasoning" not in call_kwargs
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Test _apply_auto_config
|
||||
# =============================================================================
|
||||
|
||||
|
||||
class TestAutoConfig:
|
||||
@patch("EvoScientist.llm.models.init_chat_model")
|
||||
def test_anthropic_4_5_thinking(self, mock_init):
|
||||
"""Anthropic 4-5 models get enabled thinking with budget."""
|
||||
mock_init.return_value = "mock_model"
|
||||
|
||||
get_chat_model("claude-sonnet-4-5")
|
||||
|
||||
call_kwargs = mock_init.call_args[1]
|
||||
assert call_kwargs["thinking"] == {"type": "enabled", "budget_tokens": 10000}
|
||||
|
||||
@patch("EvoScientist.llm.models.init_chat_model")
|
||||
def test_anthropic_4_6_adaptive_thinking(self, mock_init):
|
||||
"""Anthropic 4-6 models get adaptive thinking with max effort."""
|
||||
mock_init.return_value = "mock_model"
|
||||
|
||||
get_chat_model("claude-sonnet-4-6")
|
||||
|
||||
call_kwargs = mock_init.call_args[1]
|
||||
assert call_kwargs["thinking"] == {"type": "adaptive"}
|
||||
assert call_kwargs["effort"] == "max"
|
||||
|
||||
@patch("EvoScientist.llm.models.init_chat_model")
|
||||
def test_anthropic_thinking_not_overridden(self, mock_init):
|
||||
"""User-supplied thinking config should not be overridden."""
|
||||
mock_init.return_value = "mock_model"
|
||||
custom_thinking = {"type": "enabled", "budget_tokens": 500}
|
||||
|
||||
get_chat_model("claude-sonnet-4-6", thinking=custom_thinking)
|
||||
|
||||
call_kwargs = mock_init.call_args[1]
|
||||
assert call_kwargs["thinking"] == custom_thinking
|
||||
|
||||
@patch("EvoScientist.llm.models.init_chat_model")
|
||||
def test_openai_reasoning(self, mock_init):
|
||||
"""Native OpenAI models get auto-reasoning."""
|
||||
mock_init.return_value = "mock_model"
|
||||
|
||||
get_chat_model("gpt-5-nano")
|
||||
|
||||
call_kwargs = mock_init.call_args[1]
|
||||
assert call_kwargs["reasoning"] == {"effort": "high", "summary": "auto"}
|
||||
|
||||
@patch("EvoScientist.llm.models.init_chat_model")
|
||||
def test_google_thoughts(self, mock_init):
|
||||
"""Google GenAI models get include_thoughts=True by default."""
|
||||
mock_init.return_value = "mock_model"
|
||||
|
||||
get_chat_model("gemini-2.5-flash")
|
||||
|
||||
call_kwargs = mock_init.call_args[1]
|
||||
assert call_kwargs["include_thoughts"] is True
|
||||
|
||||
|
||||
Reference in New Issue
Block a user