feat: add OpenRouter API key support and enhance model handling
This commit is contained in:
@@ -64,6 +64,7 @@ class EvoScientistConfig:
|
||||
nvidia_api_key: str = ""
|
||||
google_api_key: str = ""
|
||||
siliconflow_api_key: str = ""
|
||||
openrouter_api_key: str = ""
|
||||
tavily_api_key: str = ""
|
||||
|
||||
# LLM Settings
|
||||
@@ -215,6 +216,7 @@ _ENV_MAPPINGS = {
|
||||
"nvidia_api_key": "NVIDIA_API_KEY",
|
||||
"google_api_key": "GOOGLE_API_KEY",
|
||||
"siliconflow_api_key": "SILICONFLOW_API_KEY",
|
||||
"openrouter_api_key": "OPENROUTER_API_KEY",
|
||||
"tavily_api_key": "TAVILY_API_KEY",
|
||||
"default_mode": "EVOSCIENTIST_DEFAULT_MODE",
|
||||
"default_workdir": "EVOSCIENTIST_WORKSPACE_DIR",
|
||||
@@ -285,5 +287,7 @@ def apply_config_to_env(config: EvoScientistConfig) -> None:
|
||||
os.environ["GOOGLE_API_KEY"] = config.google_api_key
|
||||
if config.siliconflow_api_key and not os.environ.get("SILICONFLOW_API_KEY"):
|
||||
os.environ["SILICONFLOW_API_KEY"] = config.siliconflow_api_key
|
||||
if config.openrouter_api_key and not os.environ.get("OPENROUTER_API_KEY"):
|
||||
os.environ["OPENROUTER_API_KEY"] = config.openrouter_api_key
|
||||
if config.tavily_api_key and not os.environ.get("TAVILY_API_KEY"):
|
||||
os.environ["TAVILY_API_KEY"] = config.tavily_api_key
|
||||
|
||||
@@ -8,6 +8,7 @@ from .models import (
|
||||
MODELS,
|
||||
DEFAULT_MODEL,
|
||||
get_chat_model,
|
||||
get_models_for_provider,
|
||||
list_models,
|
||||
get_model_info,
|
||||
)
|
||||
@@ -16,6 +17,7 @@ __all__ = [
|
||||
"MODELS",
|
||||
"DEFAULT_MODEL",
|
||||
"get_chat_model",
|
||||
"get_models_for_provider",
|
||||
"list_models",
|
||||
"get_model_info",
|
||||
]
|
||||
|
||||
+79
-37
@@ -13,43 +13,65 @@ from typing import Any
|
||||
from langchain.chat_models import init_chat_model
|
||||
|
||||
_SILICONFLOW_BASE_URL = "https://api.siliconflow.cn/v1"
|
||||
_OPENROUTER_BASE_URL = "https://openrouter.ai/api/v1"
|
||||
|
||||
# Model registry: short_name -> (model_id, provider)
|
||||
MODELS: dict[str, tuple[str, str]] = {
|
||||
# Model registry: list of (short_name, model_id, provider)
|
||||
# Allows same short_name across different providers.
|
||||
_MODEL_ENTRIES: list[tuple[str, str, str]] = [
|
||||
# Anthropic (ordered by capability)
|
||||
"claude-opus-4-6": ("claude-opus-4-6", "anthropic"),
|
||||
"claude-opus-4-5": ("claude-opus-4-5-20251101", "anthropic"),
|
||||
"claude-sonnet-4-5": ("claude-sonnet-4-5-20250929", "anthropic"),
|
||||
"claude-haiku-4-5": ("claude-haiku-4-5-20251001", "anthropic"),
|
||||
("claude-opus-4-6", "claude-opus-4-6", "anthropic"),
|
||||
("claude-opus-4-5", "claude-opus-4-5-20251101", "anthropic"),
|
||||
("claude-sonnet-4-5", "claude-sonnet-4-5-20250929", "anthropic"),
|
||||
("claude-haiku-4-5", "claude-haiku-4-5-20251001", "anthropic"),
|
||||
# OpenAI
|
||||
"gpt-5.2-codex": ("gpt-5.2-codex", "openai"),
|
||||
"gpt-5.2": ("gpt-5.2-2025-12-11", "openai"),
|
||||
"gpt-5.1": ("gpt-5.1-2025-11-13", "openai"),
|
||||
"gpt-5": ("gpt-5-2025-08-07", "openai"),
|
||||
"gpt-5-mini": ("gpt-5-mini-2025-08-07", "openai"),
|
||||
"gpt-5-nano": ("gpt-5-nano-2025-08-07", "openai"),
|
||||
("gpt-5.2-codex", "gpt-5.2-codex", "openai"),
|
||||
("gpt-5.2", "gpt-5.2-2025-12-11", "openai"),
|
||||
("gpt-5.1", "gpt-5.1-2025-11-13", "openai"),
|
||||
("gpt-5", "gpt-5-2025-08-07", "openai"),
|
||||
("gpt-5-mini", "gpt-5-mini-2025-08-07", "openai"),
|
||||
("gpt-5-nano", "gpt-5-nano-2025-08-07", "openai"),
|
||||
# Google GenAI
|
||||
"gemini-3-pro": ("gemini-3-pro-preview", "google-genai"),
|
||||
"gemini-3-flash": ("gemini-3-flash-preview", "google-genai"),
|
||||
"gemini-2.5-flash": ("gemini-2.5-flash", "google-genai"),
|
||||
"gemini-2.5-flash-lite": ("gemini-2.5-flash-lite", "google-genai"),
|
||||
"gemini-2.5-pro": ("gemini-2.5-pro", "google-genai"),
|
||||
("gemini-3-pro", "gemini-3-pro-preview", "google-genai"),
|
||||
("gemini-3-flash", "gemini-3-flash-preview", "google-genai"),
|
||||
("gemini-2.5-flash", "gemini-2.5-flash", "google-genai"),
|
||||
("gemini-2.5-flash-lite", "gemini-2.5-flash-lite", "google-genai"),
|
||||
("gemini-2.5-pro", "gemini-2.5-pro", "google-genai"),
|
||||
# NVIDIA
|
||||
"glm4.7": ("z-ai/glm4.7", "nvidia"),
|
||||
"deepseek-v3.2": ("deepseek-ai/deepseek-v3.2", "nvidia"),
|
||||
"deepseek-v3.1": ("deepseek-ai/deepseek-v3.1-terminus", "nvidia"),
|
||||
"kimi-k2.5": ("moonshotai/kimi-k2.5", "nvidia"),
|
||||
"kimi-k2-thinking": ("moonshotai/kimi-k2-thinking", "nvidia"),
|
||||
"minimax-m2.1": ("minimaxai/minimax-m2.1", "nvidia"),
|
||||
"step-3.5-flash": ("stepfun-ai/step-3.5-flash", "nvidia"),
|
||||
"nemotron-nano": ("nvidia/nemotron-3-nano-30b-a3b", "nvidia"),
|
||||
# SiliconFlow (OpenAI-compatible API)
|
||||
"glm-5": ("Pro/zai-org/GLM-5", "siliconflow"),
|
||||
("glm4.7", "z-ai/glm4.7", "nvidia"),
|
||||
("deepseek-v3.2", "deepseek-ai/deepseek-v3.2", "nvidia"),
|
||||
("deepseek-v3.1", "deepseek-ai/deepseek-v3.1-terminus", "nvidia"),
|
||||
("kimi-k2.5", "moonshotai/kimi-k2.5", "nvidia"),
|
||||
("kimi-k2-thinking", "moonshotai/kimi-k2-thinking", "nvidia"),
|
||||
("minimax-m2.1", "minimaxai/minimax-m2.1", "nvidia"),
|
||||
("step-3.5-flash", "stepfun-ai/step-3.5-flash", "nvidia"),
|
||||
("nemotron-nano", "nvidia/nemotron-3-nano-30b-a3b", "nvidia"),
|
||||
]
|
||||
|
||||
# Public dict for simple lookups (last entry wins for duplicate names).
|
||||
# Use get_models_for_provider() for provider-aware lookups.
|
||||
MODELS: dict[str, tuple[str, str]] = {
|
||||
name: (model_id, provider) for name, model_id, provider in _MODEL_ENTRIES
|
||||
}
|
||||
|
||||
DEFAULT_MODEL = "claude-sonnet-4-5"
|
||||
|
||||
|
||||
def get_models_for_provider(provider: str) -> list[tuple[str, str]]:
|
||||
"""Get all models for a specific provider.
|
||||
|
||||
Args:
|
||||
provider: Provider name (e.g., 'anthropic', 'openrouter').
|
||||
|
||||
Returns:
|
||||
List of (short_name, model_id) tuples for the provider.
|
||||
"""
|
||||
return [
|
||||
(name, model_id)
|
||||
for name, model_id, p in _MODEL_ENTRIES
|
||||
if p == provider
|
||||
]
|
||||
|
||||
|
||||
def get_chat_model(
|
||||
model: str | None = None,
|
||||
provider: str | None = None,
|
||||
@@ -75,11 +97,19 @@ def get_chat_model(
|
||||
"""
|
||||
model = model or DEFAULT_MODEL
|
||||
|
||||
# Look up short name in registry
|
||||
if model in MODELS:
|
||||
# Look up short name in registry (provider-aware)
|
||||
model_id = None
|
||||
if provider:
|
||||
# Try exact match with provider first
|
||||
for name, mid, p in _MODEL_ENTRIES:
|
||||
if name == model and p == provider:
|
||||
model_id = mid
|
||||
break
|
||||
if model_id is None and model in MODELS:
|
||||
model_id, default_provider = MODELS[model]
|
||||
provider = provider or default_provider
|
||||
else:
|
||||
|
||||
if model_id is None:
|
||||
# Assume it's a full model ID
|
||||
model_id = model
|
||||
# Try to infer provider from model ID prefix
|
||||
@@ -95,14 +125,20 @@ def get_chat_model(
|
||||
else:
|
||||
provider = "anthropic" # Default fallback
|
||||
|
||||
# SiliconFlow → route through OpenAI provider with base_url
|
||||
_is_siliconflow = provider == "siliconflow"
|
||||
if _is_siliconflow:
|
||||
# SiliconFlow / OpenRouter → route through OpenAI provider with base_url
|
||||
_is_third_party = provider in ("siliconflow", "openrouter")
|
||||
if provider == "siliconflow":
|
||||
kwargs["base_url"] = _SILICONFLOW_BASE_URL
|
||||
api_key = os.environ.get("SILICONFLOW_API_KEY", "")
|
||||
if api_key:
|
||||
kwargs["api_key"] = api_key
|
||||
provider = "openai"
|
||||
elif provider == "openrouter":
|
||||
kwargs["base_url"] = _OPENROUTER_BASE_URL
|
||||
api_key = os.environ.get("OPENROUTER_API_KEY", "")
|
||||
if api_key:
|
||||
kwargs["api_key"] = api_key
|
||||
provider = "openai"
|
||||
|
||||
# Auto-enable thinking for Anthropic models
|
||||
if provider == "anthropic" and "thinking" not in kwargs:
|
||||
@@ -111,8 +147,8 @@ def get_chat_model(
|
||||
else:
|
||||
kwargs["thinking"] = {"type": "enabled", "budget_tokens": 2000}
|
||||
|
||||
# Auto-enable reasoning for OpenAI models (not for SiliconFlow)
|
||||
if provider == "openai" and not _is_siliconflow and "reasoning" not in kwargs:
|
||||
# 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": "medium", "summary": "auto"}
|
||||
|
||||
# Auto-enable thinking visibility for Google GenAI models
|
||||
@@ -126,9 +162,15 @@ def list_models() -> list[str]:
|
||||
"""List all available model short names.
|
||||
|
||||
Returns:
|
||||
List of model short names that can be passed to get_chat_model().
|
||||
List of unique model short names that can be passed to get_chat_model().
|
||||
"""
|
||||
return list(MODELS.keys())
|
||||
seen = set()
|
||||
result = []
|
||||
for name, _, _ in _MODEL_ENTRIES:
|
||||
if name not in seen:
|
||||
seen.add(name)
|
||||
result.append(name)
|
||||
return result
|
||||
|
||||
|
||||
def get_model_info(model: str) -> tuple[str, str] | None:
|
||||
|
||||
+87
-12
@@ -26,7 +26,7 @@ from .config import (
|
||||
save_config,
|
||||
get_config_path,
|
||||
)
|
||||
from .llm import MODELS
|
||||
from .llm import get_models_for_provider
|
||||
|
||||
console = Console()
|
||||
|
||||
@@ -261,6 +261,27 @@ def validate_siliconflow_key(api_key: str) -> tuple[bool, str]:
|
||||
return False, f"Error: {e}"
|
||||
|
||||
|
||||
def validate_openrouter_key(api_key: str) -> tuple[bool, str]:
|
||||
"""Validate an OpenRouter API key by making a test request.
|
||||
|
||||
Returns:
|
||||
Tuple of (is_valid, message).
|
||||
"""
|
||||
if not api_key:
|
||||
return True, "Skipped (no key provided)"
|
||||
|
||||
try:
|
||||
import openai
|
||||
client = openai.OpenAI(api_key=api_key, base_url="https://openrouter.ai/api/v1")
|
||||
client.models.list()
|
||||
return True, "Valid"
|
||||
except Exception as e:
|
||||
error_str = str(e).lower()
|
||||
if "401" in error_str or "unauthorized" in error_str or "invalid" in error_str or "authentication" in error_str:
|
||||
return False, "Invalid API key"
|
||||
return False, f"Error: {e}"
|
||||
|
||||
|
||||
def validate_tavily_key(api_key: str) -> tuple[bool, str]:
|
||||
"""Validate a Tavily API key by making a test request.
|
||||
|
||||
@@ -344,11 +365,12 @@ def _step_provider(config: EvoScientistConfig) -> str:
|
||||
Choice(title="OpenAI (GPT models)", value="openai"),
|
||||
Choice(title="Google GenAI (Gemini models)", value="google-genai"),
|
||||
Choice(title="NVIDIA (DeepSeek, Kimi, GLM, MiniMax, Step, etc.)", value="nvidia"),
|
||||
Choice(title="SiliconFlow (GLM, DeepSeek, Qwen, etc.)", value="siliconflow"),
|
||||
Choice(title="SiliconFlow (third party)", value="siliconflow"),
|
||||
Choice(title="OpenRouter (third party)", value="openrouter"),
|
||||
]
|
||||
|
||||
# Set default based on current config
|
||||
default = config.provider if config.provider in ["anthropic", "openai", "google-genai", "nvidia", "siliconflow"] else "anthropic"
|
||||
default = config.provider if config.provider in ["anthropic", "openai", "google-genai", "nvidia", "siliconflow", "openrouter"] else "anthropic"
|
||||
|
||||
provider = questionary.select(
|
||||
"Select your LLM provider:",
|
||||
@@ -372,6 +394,7 @@ def _provider_key_info(config: EvoScientistConfig, provider: str):
|
||||
"nvidia": ("NVIDIA", config.nvidia_api_key or os.environ.get("NVIDIA_API_KEY", ""), validate_nvidia_key),
|
||||
"google-genai": ("Google", config.google_api_key or os.environ.get("GOOGLE_API_KEY", ""), validate_google_key),
|
||||
"siliconflow": ("SiliconFlow", config.siliconflow_api_key or os.environ.get("SILICONFLOW_API_KEY", ""), validate_siliconflow_key),
|
||||
"openrouter": ("OpenRouter", config.openrouter_api_key or os.environ.get("OPENROUTER_API_KEY", ""), validate_openrouter_key),
|
||||
}
|
||||
return mapping.get(provider, ("OpenAI", config.openai_api_key or os.environ.get("OPENAI_API_KEY", ""), validate_openai_key))
|
||||
|
||||
@@ -457,6 +480,18 @@ def _step_provider_api_key(
|
||||
)
|
||||
|
||||
|
||||
_THIRD_PARTY_EXAMPLES: dict[str, list[tuple[str, str]]] = {
|
||||
"openrouter": [
|
||||
("minimax/minimax-m2.5", "MiniMax M2.5"),
|
||||
("x-ai/grok-4.1-fast", "Grok 4.1 Fast"),
|
||||
],
|
||||
"siliconflow": [
|
||||
("Pro/zai-org/GLM-5", "GLM 5"),
|
||||
("Pro/moonshotai/Kimi-K2.5", "Kimi K2.5"),
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
def _step_model(config: EvoScientistConfig, provider: str) -> str:
|
||||
"""Step 3: Select model for the provider.
|
||||
|
||||
@@ -467,13 +502,48 @@ def _step_model(config: EvoScientistConfig, provider: str) -> str:
|
||||
Returns:
|
||||
Selected model name.
|
||||
"""
|
||||
# Get models for the selected provider
|
||||
provider_models = [
|
||||
name for name, (model_id, p) in MODELS.items()
|
||||
if p == provider
|
||||
]
|
||||
# Third-party providers: select from examples or type custom model name
|
||||
if provider in _THIRD_PARTY_EXAMPLES:
|
||||
examples = _THIRD_PARTY_EXAMPLES[provider]
|
||||
_CUSTOM_SENTINEL = "__custom__"
|
||||
choices = [
|
||||
Choice(title=f"{label} ({mid})", value=mid)
|
||||
for mid, label in examples
|
||||
]
|
||||
choices.append(Choice(title="Customize your model...", value=_CUSTOM_SENTINEL))
|
||||
|
||||
if not provider_models:
|
||||
selected = questionary.select(
|
||||
"Select model:",
|
||||
choices=choices,
|
||||
default=choices[0].value,
|
||||
style=WIZARD_STYLE,
|
||||
qmark=QMARK,
|
||||
use_indicator=True,
|
||||
).ask()
|
||||
if selected is None:
|
||||
raise KeyboardInterrupt()
|
||||
|
||||
if selected != _CUSTOM_SENTINEL:
|
||||
return selected
|
||||
|
||||
model = questionary.text(
|
||||
"Model name:",
|
||||
style=WIZARD_STYLE,
|
||||
qmark=QMARK,
|
||||
placeholder=FormattedText([("fg:#858585", " e.g. owner/model-name")]),
|
||||
).ask()
|
||||
if model is None:
|
||||
raise KeyboardInterrupt()
|
||||
model = model.strip()
|
||||
if not model:
|
||||
model = examples[0][0]
|
||||
console.print(f" [dim]Using default: {model}[/dim]")
|
||||
return model
|
||||
|
||||
# Get models for the selected provider
|
||||
entries = get_models_for_provider(provider)
|
||||
|
||||
if not entries:
|
||||
# Fallback if no models for provider
|
||||
console.print(f" [yellow]No registered models for {provider}[/yellow]")
|
||||
model = questionary.text(
|
||||
@@ -486,10 +556,11 @@ def _step_model(config: EvoScientistConfig, provider: str) -> str:
|
||||
raise KeyboardInterrupt()
|
||||
return model
|
||||
|
||||
provider_models = [name for name, _ in entries]
|
||||
|
||||
# Create choices with model IDs as hints
|
||||
choices = []
|
||||
for name in provider_models:
|
||||
model_id, _ = MODELS[name]
|
||||
for name, model_id in entries:
|
||||
choices.append(Choice(title=f"{name} ({model_id})", value=name))
|
||||
|
||||
# Determine default
|
||||
@@ -533,7 +604,7 @@ def _step_tavily_key(
|
||||
|
||||
return _prompt_and_validate_api_key(
|
||||
prompt_text, current, validate_tavily_key, skip_validation,
|
||||
placeholder=FormattedText([("fg:#858585", "(recommended for web search)")]),
|
||||
placeholder=FormattedText([("fg:#858585", " (recommended for web search)")]),
|
||||
)
|
||||
|
||||
|
||||
@@ -1362,6 +1433,8 @@ def run_onboard(skip_validation: bool = False) -> bool:
|
||||
config.google_api_key = new_key
|
||||
elif provider == "siliconflow":
|
||||
config.siliconflow_api_key = new_key
|
||||
elif provider == "openrouter":
|
||||
config.openrouter_api_key = new_key
|
||||
else:
|
||||
config.openai_api_key = new_key
|
||||
else:
|
||||
@@ -1373,6 +1446,8 @@ def run_onboard(skip_validation: bool = False) -> bool:
|
||||
current = config.google_api_key
|
||||
elif provider == "siliconflow":
|
||||
current = config.siliconflow_api_key
|
||||
elif provider == "openrouter":
|
||||
current = config.openrouter_api_key
|
||||
else:
|
||||
current = config.openai_api_key
|
||||
if not current:
|
||||
|
||||
+24
-28
@@ -6,9 +6,11 @@ from EvoScientist.llm import (
|
||||
MODELS,
|
||||
DEFAULT_MODEL,
|
||||
get_chat_model,
|
||||
get_models_for_provider,
|
||||
list_models,
|
||||
get_model_info,
|
||||
)
|
||||
from EvoScientist.llm.models import _MODEL_ENTRIES
|
||||
|
||||
|
||||
# =============================================================================
|
||||
@@ -21,41 +23,35 @@ class TestModelsRegistry:
|
||||
"""Test that MODELS is a dictionary."""
|
||||
assert isinstance(MODELS, dict)
|
||||
|
||||
def test_models_has_all_providers(self):
|
||||
"""Test that MODELS covers all supported providers."""
|
||||
providers = {p for _, (_, p) in MODELS.items()}
|
||||
def test_entries_has_all_providers(self):
|
||||
"""Test that _MODEL_ENTRIES covers native 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
|
||||
|
||||
def test_models_values_are_tuples(self):
|
||||
"""Test that MODELS values are (model_id, provider) tuples."""
|
||||
for name, value in MODELS.items():
|
||||
assert isinstance(value, tuple), f"MODELS['{name}'] is not a tuple"
|
||||
assert len(value) == 2, f"MODELS['{name}'] doesn't have 2 elements"
|
||||
model_id, provider = value
|
||||
assert isinstance(model_id, str), f"model_id for '{name}' is not a string"
|
||||
assert isinstance(provider, str), f"provider for '{name}' is not a string"
|
||||
assert provider in ("anthropic", "openai", "google-genai", "nvidia", "siliconflow"), f"Unknown provider for '{name}': {provider}"
|
||||
def test_entries_are_valid_tuples(self):
|
||||
"""Test that _MODEL_ENTRIES contains valid (name, model_id, provider) tuples."""
|
||||
valid_providers = {"anthropic", "openai", "google-genai", "nvidia"}
|
||||
for entry in _MODEL_ENTRIES:
|
||||
assert len(entry) == 3, f"Entry {entry} doesn't have 3 elements"
|
||||
name, model_id, provider = entry
|
||||
assert isinstance(name, str)
|
||||
assert isinstance(model_id, str)
|
||||
assert provider in valid_providers, f"Unknown provider '{provider}' for '{name}'"
|
||||
|
||||
def test_anthropic_models_have_anthropic_provider(self):
|
||||
"""Test that claude models use anthropic provider."""
|
||||
for name, (model_id, provider) in MODELS.items():
|
||||
if name.startswith("claude"):
|
||||
assert provider == "anthropic", f"Claude model '{name}' doesn't use anthropic provider"
|
||||
def test_get_models_for_provider(self):
|
||||
"""Test that get_models_for_provider returns correct models."""
|
||||
anthropic_models = get_models_for_provider("anthropic")
|
||||
assert len(anthropic_models) > 0
|
||||
for name, model_id in anthropic_models:
|
||||
assert isinstance(name, str)
|
||||
assert isinstance(model_id, str)
|
||||
|
||||
def test_openai_models_have_openai_provider(self):
|
||||
"""Test that gpt models use openai provider."""
|
||||
for name, (model_id, provider) in MODELS.items():
|
||||
if name.startswith("gpt"):
|
||||
assert provider == "openai", f"OpenAI model '{name}' doesn't use openai provider"
|
||||
|
||||
def test_google_models_have_google_provider(self):
|
||||
"""Test that gemini models use google-genai provider."""
|
||||
for name, (model_id, provider) in MODELS.items():
|
||||
if name.startswith("gemini"):
|
||||
assert provider == "google-genai", f"Google model '{name}' doesn't use google-genai provider"
|
||||
# Third-party providers have no registered models (user types model name)
|
||||
openrouter_models = get_models_for_provider("openrouter")
|
||||
assert len(openrouter_models) == 0
|
||||
|
||||
|
||||
# =============================================================================
|
||||
|
||||
Reference in New Issue
Block a user