forked from HAL9000/cleveragents-core
a074b4846f
Remove FakeListLLM as a silent fallback in agent graph constructors (plan_generation.py, context_analysis.py, auto_debug.py). All three now raise ValueError when llm=None, making missing-provider errors explicit. Add Settings.mock_providers flag and validate_provider_availability() method. Update container.get_ai_provider() to check Settings.mock_providers first, with env-var fallback for backward compatibility. Add resolve_provider_by_name() helper to the provider registry and export it from cleveragents.providers. Add structlog trace logging to ProviderRegistry.get_default_provider_type() to record selection reasoning. Update all existing behave step files, robot tests, and benchmarks that relied on the implicit FakeListLLM default to pass an explicit LLM instance instead. Add new BDD tests (features/provider_fixes.feature with 17 scenarios), Robot Framework integration tests (robot/provider_detection_smoke.robot), and ASV benchmarks (benchmarks/provider_selection_bench.py). ISSUES CLOSED: #323
728 lines
26 KiB
Python
728 lines
26 KiB
Python
"""Provider Registry for AI Model Provider Discovery and Selection.
|
|
|
|
This module implements the provider registry that discovers configured providers
|
|
from Settings, reports capability metadata, and selects defaults based on
|
|
CLEVERAGENTS_DEFAULT_PROVIDER / CLEVERAGENTS_DEFAULT_MODEL environment variables.
|
|
|
|
Following ADR-008 (Provider Plugin Architecture), this registry provides:
|
|
- Dynamic provider discovery based on configured API keys
|
|
- Capability metadata for feature support (streaming, tool calls, etc.)
|
|
- Default selection based on configuration
|
|
- Fallback chains with graceful degradation
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import os
|
|
from dataclasses import dataclass
|
|
from enum import StrEnum
|
|
from typing import TYPE_CHECKING, Any, ClassVar
|
|
|
|
import structlog
|
|
from pydantic import BaseModel, ConfigDict, Field
|
|
|
|
from cleveragents.config.settings import Settings, get_settings
|
|
from cleveragents.domain.providers.ai_provider import AIProviderInterface
|
|
|
|
if TYPE_CHECKING:
|
|
from langchain_core.language_models import BaseLanguageModel
|
|
|
|
|
|
logger: structlog.stdlib.BoundLogger = structlog.get_logger(__name__)
|
|
|
|
|
|
def _coerce_optional_str(value: object | None) -> str | None:
|
|
if value is None:
|
|
return None
|
|
if isinstance(value, str):
|
|
return value
|
|
return str(value)
|
|
|
|
|
|
class ProviderType(StrEnum):
|
|
"""Supported AI provider types."""
|
|
|
|
OPENAI = "openai"
|
|
ANTHROPIC = "anthropic"
|
|
GOOGLE = "google"
|
|
AZURE = "azure"
|
|
OPENROUTER = "openrouter"
|
|
GEMINI = "gemini"
|
|
COHERE = "cohere"
|
|
GROQ = "groq"
|
|
TOGETHER = "together"
|
|
MOCK = "mock"
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class ProviderCapabilities:
|
|
"""Capability metadata for a provider.
|
|
|
|
Attributes:
|
|
supports_streaming: Whether the provider supports streaming responses
|
|
supports_tool_calls: Whether the provider supports tool/function calling
|
|
supports_vision: Whether the provider supports image inputs
|
|
max_context_length: Maximum context length in tokens
|
|
supports_json_mode: Whether the provider supports structured JSON output
|
|
"""
|
|
|
|
supports_streaming: bool = True
|
|
supports_tool_calls: bool = False
|
|
supports_vision: bool = False
|
|
max_context_length: int = 4096
|
|
supports_json_mode: bool = False
|
|
|
|
|
|
class ProviderInfo(BaseModel):
|
|
"""Information about a registered provider."""
|
|
|
|
provider_type: ProviderType
|
|
name: str
|
|
api_key_env_var: str
|
|
default_model: str
|
|
capabilities: ProviderCapabilities = Field(default_factory=ProviderCapabilities)
|
|
is_configured: bool = False
|
|
|
|
model_config = ConfigDict(str_strip_whitespace=True)
|
|
|
|
|
|
class ProviderRegistry:
|
|
"""Registry for discovering and managing AI providers.
|
|
|
|
The registry discovers configured providers based on environment variables
|
|
and provides methods for selecting providers and models.
|
|
|
|
Example:
|
|
registry = ProviderRegistry()
|
|
provider = registry.get_default_provider()
|
|
if provider:
|
|
response = provider.generate_changes(project, plan, contexts)
|
|
"""
|
|
|
|
# Default provider capabilities for each provider type
|
|
DEFAULT_CAPABILITIES: ClassVar[dict[ProviderType, ProviderCapabilities]] = {
|
|
ProviderType.OPENAI: ProviderCapabilities(
|
|
supports_streaming=True,
|
|
supports_tool_calls=True,
|
|
supports_vision=True,
|
|
max_context_length=128000,
|
|
supports_json_mode=True,
|
|
),
|
|
ProviderType.ANTHROPIC: ProviderCapabilities(
|
|
supports_streaming=True,
|
|
supports_tool_calls=True,
|
|
supports_vision=True,
|
|
max_context_length=200000,
|
|
supports_json_mode=False,
|
|
),
|
|
ProviderType.GOOGLE: ProviderCapabilities(
|
|
supports_streaming=True,
|
|
supports_tool_calls=True,
|
|
supports_vision=True,
|
|
max_context_length=1000000,
|
|
supports_json_mode=True,
|
|
),
|
|
ProviderType.GEMINI: ProviderCapabilities(
|
|
supports_streaming=True,
|
|
supports_tool_calls=True,
|
|
supports_vision=True,
|
|
max_context_length=1000000,
|
|
supports_json_mode=True,
|
|
),
|
|
ProviderType.AZURE: ProviderCapabilities(
|
|
supports_streaming=True,
|
|
supports_tool_calls=True,
|
|
supports_vision=True,
|
|
max_context_length=128000,
|
|
supports_json_mode=True,
|
|
),
|
|
ProviderType.OPENROUTER: ProviderCapabilities(
|
|
supports_streaming=True,
|
|
supports_tool_calls=True,
|
|
supports_vision=True,
|
|
max_context_length=128000,
|
|
supports_json_mode=True,
|
|
),
|
|
ProviderType.COHERE: ProviderCapabilities(
|
|
supports_streaming=True,
|
|
supports_tool_calls=True,
|
|
supports_vision=False,
|
|
max_context_length=128000,
|
|
supports_json_mode=False,
|
|
),
|
|
ProviderType.GROQ: ProviderCapabilities(
|
|
supports_streaming=True,
|
|
supports_tool_calls=True,
|
|
supports_vision=False,
|
|
max_context_length=32000,
|
|
supports_json_mode=True,
|
|
),
|
|
ProviderType.TOGETHER: ProviderCapabilities(
|
|
supports_streaming=True,
|
|
supports_tool_calls=True,
|
|
supports_vision=False,
|
|
max_context_length=32000,
|
|
supports_json_mode=False,
|
|
),
|
|
ProviderType.MOCK: ProviderCapabilities(
|
|
supports_streaming=True,
|
|
supports_tool_calls=False,
|
|
supports_vision=False,
|
|
max_context_length=4096,
|
|
supports_json_mode=False,
|
|
),
|
|
}
|
|
|
|
# Default models for each provider type
|
|
DEFAULT_MODELS: ClassVar[dict[ProviderType, str]] = {
|
|
ProviderType.OPENAI: "gpt-4o",
|
|
ProviderType.ANTHROPIC: "claude-sonnet-4-20250514",
|
|
ProviderType.GOOGLE: "gemini-2.0-flash",
|
|
ProviderType.GEMINI: "gemini-2.0-flash",
|
|
ProviderType.AZURE: "gpt-4o",
|
|
ProviderType.OPENROUTER: "anthropic/claude-sonnet-4-20250514",
|
|
ProviderType.COHERE: "command-r-plus",
|
|
ProviderType.GROQ: "llama-3.1-70b-versatile",
|
|
ProviderType.TOGETHER: "meta-llama/Llama-3.1-70B-Instruct-Turbo",
|
|
ProviderType.MOCK: "mock-gpt",
|
|
}
|
|
|
|
# Mapping of provider type to Settings attribute name for API key
|
|
PROVIDER_KEY_ATTRS: ClassVar[dict[ProviderType, str]] = {
|
|
ProviderType.OPENAI: "openai_api_key",
|
|
ProviderType.ANTHROPIC: "anthropic_api_key",
|
|
ProviderType.GOOGLE: "google_api_key",
|
|
ProviderType.GEMINI: "gemini_api_key",
|
|
ProviderType.AZURE: "azure_api_key",
|
|
ProviderType.OPENROUTER: "openrouter_api_key",
|
|
ProviderType.COHERE: "cohere_api_key",
|
|
ProviderType.GROQ: "groq_api_key",
|
|
ProviderType.TOGETHER: "together_api_key",
|
|
}
|
|
|
|
# Provider priority order for fallback selection
|
|
FALLBACK_ORDER: ClassVar[list[ProviderType]] = [
|
|
ProviderType.OPENAI,
|
|
ProviderType.ANTHROPIC,
|
|
ProviderType.GOOGLE,
|
|
ProviderType.AZURE,
|
|
ProviderType.OPENROUTER,
|
|
ProviderType.GROQ,
|
|
ProviderType.TOGETHER,
|
|
ProviderType.COHERE,
|
|
]
|
|
|
|
def __init__(self, settings: Settings | None = None) -> None:
|
|
"""Initialize the provider registry.
|
|
|
|
Args:
|
|
settings: Optional settings instance. Uses global settings if not provided.
|
|
"""
|
|
self._settings = settings or get_settings()
|
|
self._providers: dict[ProviderType, ProviderInfo] = {}
|
|
self._discover_providers()
|
|
|
|
def _discover_providers(self) -> None:
|
|
"""Discover configured providers from settings."""
|
|
for provider_type, key_attr in self.PROVIDER_KEY_ATTRS.items():
|
|
api_key = getattr(self._settings, key_attr, None)
|
|
is_configured = bool(api_key)
|
|
|
|
provider_info = ProviderInfo(
|
|
provider_type=provider_type,
|
|
name=provider_type.value.title(),
|
|
api_key_env_var=key_attr.upper(),
|
|
default_model=self.DEFAULT_MODELS.get(provider_type, "unknown"),
|
|
capabilities=self.DEFAULT_CAPABILITIES.get(
|
|
provider_type, ProviderCapabilities()
|
|
),
|
|
is_configured=is_configured,
|
|
)
|
|
self._providers[provider_type] = provider_info
|
|
|
|
def get_configured_providers(self) -> list[ProviderInfo]:
|
|
"""Get list of all configured providers.
|
|
|
|
Returns:
|
|
List of ProviderInfo for providers with valid credentials.
|
|
"""
|
|
return [p for p in self._providers.values() if p.is_configured]
|
|
|
|
def get_all_providers(self) -> list[ProviderInfo]:
|
|
"""Get list of all known providers (configured or not).
|
|
|
|
Returns:
|
|
List of all ProviderInfo instances.
|
|
"""
|
|
return list(self._providers.values())
|
|
|
|
def get_provider_info(
|
|
self, provider_type: ProviderType | str
|
|
) -> ProviderInfo | None:
|
|
"""Get information about a specific provider.
|
|
|
|
Args:
|
|
provider_type: The provider type to look up.
|
|
|
|
Returns:
|
|
ProviderInfo if found, None otherwise.
|
|
"""
|
|
if not isinstance(provider_type, ProviderType):
|
|
try:
|
|
provider_type = ProviderType(provider_type.lower())
|
|
except ValueError:
|
|
return None
|
|
return self._providers.get(provider_type)
|
|
|
|
def is_provider_configured(self, provider_type: ProviderType | str) -> bool:
|
|
"""Check if a provider is configured with valid credentials.
|
|
|
|
Args:
|
|
provider_type: The provider type to check.
|
|
|
|
Returns:
|
|
True if the provider has valid credentials.
|
|
"""
|
|
info = self.get_provider_info(provider_type)
|
|
return info.is_configured if info else False
|
|
|
|
def get_default_provider_type(self) -> ProviderType | None:
|
|
"""Get the default provider type based on configuration.
|
|
|
|
Order of precedence:
|
|
1. CLEVERAGENTS_DEFAULT_PROVIDER environment variable
|
|
2. First configured provider in fallback order
|
|
|
|
Returns:
|
|
The default provider type, or None if no provider is configured.
|
|
"""
|
|
# Check environment variable first
|
|
env_provider = os.environ.get("CLEVERAGENTS_DEFAULT_PROVIDER", "").lower()
|
|
if env_provider:
|
|
try:
|
|
provider_type = ProviderType(env_provider)
|
|
if self.is_provider_configured(provider_type):
|
|
logger.debug(
|
|
"provider_selected",
|
|
provider=provider_type.value,
|
|
reason="environment_variable",
|
|
)
|
|
return provider_type
|
|
logger.debug(
|
|
"provider_skipped",
|
|
provider=env_provider,
|
|
reason="not_configured",
|
|
)
|
|
except ValueError:
|
|
logger.debug(
|
|
"provider_skipped",
|
|
provider=env_provider,
|
|
reason="invalid_provider_type",
|
|
)
|
|
|
|
# Use settings default when configured
|
|
settings_default = (self._settings.default_provider or "").lower()
|
|
if settings_default:
|
|
try:
|
|
provider_type = ProviderType(settings_default)
|
|
if self.is_provider_configured(provider_type):
|
|
logger.debug(
|
|
"provider_selected",
|
|
provider=provider_type.value,
|
|
reason="settings_default",
|
|
)
|
|
return provider_type
|
|
logger.debug(
|
|
"provider_skipped",
|
|
provider=settings_default,
|
|
reason="not_configured",
|
|
)
|
|
except ValueError:
|
|
logger.debug(
|
|
"provider_skipped",
|
|
provider=settings_default,
|
|
reason="invalid_provider_type",
|
|
)
|
|
|
|
# Fall back to first configured provider in priority order
|
|
skipped: list[str] = []
|
|
for provider_type in self.FALLBACK_ORDER:
|
|
if self.is_provider_configured(provider_type):
|
|
logger.debug(
|
|
"provider_selected",
|
|
provider=provider_type.value,
|
|
reason="fallback_order",
|
|
skipped_providers=skipped,
|
|
)
|
|
return provider_type
|
|
skipped.append(provider_type.value)
|
|
|
|
logger.debug(
|
|
"no_provider_selected",
|
|
reason="none_configured",
|
|
skipped_providers=skipped,
|
|
)
|
|
return None
|
|
|
|
def get_default_model(
|
|
self, provider_type: ProviderType | str | None = None
|
|
) -> str | None:
|
|
"""Get the default model for a provider.
|
|
|
|
Order of precedence:
|
|
1. CLEVERAGENTS_DEFAULT_MODEL environment variable
|
|
2. Provider's default model
|
|
|
|
Args:
|
|
provider_type: Optional provider type. Uses default if not specified.
|
|
|
|
Returns:
|
|
The default model ID, or None if no provider is configured.
|
|
"""
|
|
# Check environment variable first
|
|
env_model = os.environ.get("CLEVERAGENTS_DEFAULT_MODEL", "")
|
|
if env_model:
|
|
return env_model
|
|
|
|
settings_model = self._settings.default_model
|
|
if settings_model:
|
|
return settings_model
|
|
|
|
# Get provider type
|
|
if provider_type is None:
|
|
provider_type = self.get_default_provider_type()
|
|
if provider_type is None:
|
|
return None
|
|
elif not isinstance(provider_type, ProviderType):
|
|
try:
|
|
provider_type = ProviderType(provider_type.lower())
|
|
except ValueError:
|
|
return None
|
|
|
|
return self.DEFAULT_MODELS.get(provider_type)
|
|
|
|
def create_llm(
|
|
self,
|
|
provider_type: ProviderType | str | None = None,
|
|
model_id: str | None = None,
|
|
**kwargs: object,
|
|
) -> BaseLanguageModel:
|
|
"""Create a LangChain LLM instance for a provider.
|
|
|
|
Args:
|
|
provider_type: The provider type. Uses default if not specified.
|
|
model_id: The model ID. Uses default for provider if not specified.
|
|
**kwargs: Additional arguments passed to the LLM constructor.
|
|
|
|
Returns:
|
|
A LangChain BaseLanguageModel instance.
|
|
|
|
Raises:
|
|
ValueError: If no provider is configured or provider is invalid.
|
|
"""
|
|
# Resolve provider type
|
|
if provider_type is None:
|
|
provider_type = self.get_default_provider_type()
|
|
if provider_type is None:
|
|
raise ValueError(
|
|
"No AI provider configured. Please set one of the following "
|
|
"environment variables: OPENAI_API_KEY, ANTHROPIC_API_KEY, "
|
|
"GOOGLE_API_KEY, or another supported provider key."
|
|
)
|
|
elif not isinstance(provider_type, ProviderType):
|
|
try:
|
|
provider_type = ProviderType(provider_type.lower())
|
|
except ValueError as e:
|
|
raise ValueError(f"Unknown provider type: {provider_type}") from e
|
|
|
|
# Resolve model ID
|
|
if model_id is None:
|
|
model_id = self.get_default_model(provider_type)
|
|
|
|
# Get API key
|
|
key_attr = self.PROVIDER_KEY_ATTRS.get(provider_type)
|
|
if key_attr:
|
|
api_key = getattr(self._settings, key_attr, None)
|
|
if not api_key:
|
|
raise ValueError(
|
|
f"Provider {provider_type.value} is not configured. "
|
|
f"Please set the {key_attr.upper()} environment variable."
|
|
)
|
|
|
|
# Create the appropriate LLM
|
|
return self._create_provider_llm(provider_type, model_id, **kwargs)
|
|
|
|
def _create_provider_llm(
|
|
self,
|
|
provider_type: ProviderType,
|
|
model_id: str | None,
|
|
**kwargs: Any,
|
|
) -> BaseLanguageModel:
|
|
"""Create a LangChain LLM for the specified provider.
|
|
|
|
Args:
|
|
provider_type: The provider type.
|
|
model_id: The model ID.
|
|
**kwargs: Additional arguments for the LLM.
|
|
|
|
Returns:
|
|
A LangChain BaseLanguageModel instance.
|
|
"""
|
|
# Import the appropriate LangChain module based on provider
|
|
if provider_type == ProviderType.OPENAI:
|
|
from langchain_openai import ChatOpenAI
|
|
|
|
return ChatOpenAI(model=model_id or "gpt-4o", **kwargs) # type: ignore[arg-type]
|
|
|
|
if provider_type == ProviderType.ANTHROPIC:
|
|
from langchain_anthropic import ChatAnthropic
|
|
|
|
return ChatAnthropic(
|
|
model=model_id or "claude-sonnet-4-20250514",
|
|
**kwargs, # type: ignore[arg-type]
|
|
)
|
|
|
|
if provider_type in (ProviderType.GOOGLE, ProviderType.GEMINI):
|
|
from langchain_google_genai import ChatGoogleGenerativeAI
|
|
|
|
return ChatGoogleGenerativeAI(
|
|
model=model_id or "gemini-2.0-flash",
|
|
**kwargs, # type: ignore[arg-type]
|
|
)
|
|
|
|
if provider_type == ProviderType.AZURE:
|
|
from langchain_openai import AzureChatOpenAI
|
|
|
|
deployment_override = _coerce_optional_str(
|
|
kwargs.pop("deployment_name", None)
|
|
)
|
|
deployment_name = (
|
|
deployment_override
|
|
or model_id
|
|
or self._settings.azure_openai_deployment
|
|
or "gpt-4o"
|
|
)
|
|
|
|
endpoint_override = _coerce_optional_str(
|
|
kwargs.pop("azure_endpoint", self._settings.azure_openai_endpoint)
|
|
)
|
|
azure_endpoint = endpoint_override or self._settings.azure_openai_endpoint
|
|
if not azure_endpoint:
|
|
raise ValueError(
|
|
"Azure OpenAI endpoint not configured. "
|
|
"Set AZURE_OPENAI_ENDPOINT or CLEVERAGENTS_AZURE_OPENAI_ENDPOINT."
|
|
)
|
|
|
|
api_version_override = _coerce_optional_str(
|
|
kwargs.pop("api_version", self._settings.azure_openai_api_version)
|
|
)
|
|
api_version = api_version_override or "2024-05-01-preview"
|
|
|
|
return AzureChatOpenAI(
|
|
deployment_name=deployment_name,
|
|
azure_endpoint=azure_endpoint,
|
|
api_version=api_version,
|
|
**kwargs, # type: ignore[arg-type]
|
|
)
|
|
|
|
if provider_type == ProviderType.GROQ:
|
|
from langchain_groq import ChatGroq # type: ignore[import-untyped]
|
|
|
|
return ChatGroq(
|
|
model=model_id or "llama-3.1-70b-versatile",
|
|
**kwargs, # type: ignore[arg-type]
|
|
)
|
|
|
|
if provider_type == ProviderType.TOGETHER:
|
|
from langchain_together import ChatTogether # type: ignore[import-untyped]
|
|
|
|
return ChatTogether(
|
|
model=model_id or "meta-llama/Llama-3.1-70B-Instruct-Turbo",
|
|
**kwargs, # type: ignore[arg-type]
|
|
)
|
|
|
|
if provider_type == ProviderType.COHERE:
|
|
from langchain_cohere import ChatCohere # type: ignore[import-untyped]
|
|
|
|
return ChatCohere(
|
|
model=model_id or "command-r-plus",
|
|
**kwargs, # type: ignore[arg-type]
|
|
)
|
|
|
|
raise ValueError(f"Unsupported provider type: {provider_type}")
|
|
|
|
def create_ai_provider(
|
|
self,
|
|
provider_type: ProviderType | str | None = None,
|
|
model_id: str | None = None,
|
|
max_retries: int = 3,
|
|
) -> AIProviderInterface:
|
|
"""Create an AIProviderInterface instance for a provider.
|
|
|
|
This creates a LangChainChatProvider that wraps the LangChain LLM
|
|
with the AIProviderInterface protocol.
|
|
|
|
Args:
|
|
provider_type: The provider type. Uses default if not specified.
|
|
model_id: The model ID. Uses default for provider if not specified.
|
|
max_retries: Maximum retry attempts for failed operations.
|
|
|
|
Returns:
|
|
An AIProviderInterface implementation.
|
|
|
|
Raises:
|
|
ValueError: If no provider is configured or provider is invalid.
|
|
"""
|
|
from cleveragents.providers.llm.langchain_chat_provider import (
|
|
LangChainChatProvider,
|
|
)
|
|
|
|
# Resolve types
|
|
if provider_type is None:
|
|
provider_type = self.get_default_provider_type()
|
|
if provider_type is None:
|
|
raise ValueError(
|
|
"No AI provider configured. Please set one of the following "
|
|
"environment variables: OPENAI_API_KEY, ANTHROPIC_API_KEY, "
|
|
"GOOGLE_API_KEY, or another supported provider key."
|
|
)
|
|
elif not isinstance(provider_type, ProviderType):
|
|
try:
|
|
provider_type = ProviderType(provider_type.lower())
|
|
except ValueError as e:
|
|
raise ValueError(f"Unknown provider type: {provider_type}") from e
|
|
|
|
if model_id is None:
|
|
model_id = self.get_default_model(provider_type)
|
|
|
|
provider_info = self.get_provider_info(provider_type)
|
|
capabilities = (
|
|
provider_info.capabilities if provider_info else ProviderCapabilities()
|
|
)
|
|
|
|
if provider_type == ProviderType.GOOGLE:
|
|
from cleveragents.providers.llm.google_provider import GoogleChatProvider
|
|
|
|
key_attr = self.PROVIDER_KEY_ATTRS.get(provider_type)
|
|
api_key = getattr(self._settings, key_attr, None) if key_attr else None
|
|
if not api_key:
|
|
missing_env = (
|
|
key_attr.upper() if key_attr else provider_type.value.upper()
|
|
)
|
|
raise ValueError(
|
|
f"Provider {provider_type.value} is not configured. "
|
|
f"Please set the {missing_env} environment variable."
|
|
)
|
|
|
|
return GoogleChatProvider(
|
|
api_key=api_key,
|
|
model=model_id
|
|
or self.DEFAULT_MODELS.get(provider_type, "gemini-2.0-flash"),
|
|
max_retries=max_retries,
|
|
)
|
|
|
|
if provider_type == ProviderType.OPENROUTER:
|
|
from cleveragents.providers.llm.openrouter_provider import (
|
|
OpenRouterChatProvider,
|
|
)
|
|
|
|
key_attr = self.PROVIDER_KEY_ATTRS.get(provider_type)
|
|
api_key = getattr(self._settings, key_attr, None) if key_attr else None
|
|
if not api_key:
|
|
missing_env = (
|
|
key_attr.upper() if key_attr else provider_type.value.upper()
|
|
)
|
|
raise ValueError(
|
|
f"Provider {provider_type.value} is not configured. "
|
|
f"Please set the {missing_env} environment variable."
|
|
)
|
|
|
|
return OpenRouterChatProvider(
|
|
api_key=api_key,
|
|
model=model_id
|
|
or self.DEFAULT_MODELS.get(
|
|
provider_type, "anthropic/claude-sonnet-4-20250514"
|
|
),
|
|
organization=self._settings.openrouter_organization,
|
|
max_retries=max_retries,
|
|
)
|
|
|
|
# Create factory function for the LLM
|
|
def llm_factory(mid: str) -> BaseLanguageModel:
|
|
factory_kwargs: dict[str, Any] = {"max_retries": max_retries}
|
|
return self._create_provider_llm(
|
|
provider_type,
|
|
mid,
|
|
**factory_kwargs,
|
|
) # type: ignore[arg-type]
|
|
|
|
return LangChainChatProvider(
|
|
name=provider_type.value,
|
|
model_id=model_id or self.DEFAULT_MODELS.get(provider_type, "unknown"),
|
|
llm_factory=llm_factory,
|
|
max_retries=max_retries,
|
|
supports_streaming=capabilities.supports_streaming,
|
|
)
|
|
|
|
|
|
def resolve_provider_by_name(
|
|
name: str,
|
|
settings: Settings | None = None,
|
|
) -> AIProviderInterface:
|
|
"""Resolve a provider by name with an explicit error when not configured.
|
|
|
|
Args:
|
|
name: Provider name (e.g. ``"openai"``, ``"anthropic"``).
|
|
settings: Optional settings override.
|
|
|
|
Returns:
|
|
An ``AIProviderInterface`` implementation for the requested provider.
|
|
|
|
Raises:
|
|
ValueError: If the provider is not configured or the name is unknown.
|
|
"""
|
|
registry = get_provider_registry(settings)
|
|
try:
|
|
provider_type = ProviderType(name.lower())
|
|
except ValueError as exc:
|
|
available = ", ".join(pt.value for pt in ProviderType)
|
|
raise ValueError(
|
|
f"Unknown provider '{name}'. Available providers: {available}"
|
|
) from exc
|
|
|
|
if not registry.is_provider_configured(provider_type):
|
|
raise ValueError(
|
|
f"Provider '{name}' is not configured. "
|
|
f"Set the required API key environment variable."
|
|
)
|
|
logger.debug("provider_resolved_by_name", provider=name)
|
|
return registry.create_ai_provider(provider_type=provider_type)
|
|
|
|
|
|
# Global registry instance
|
|
_registry: ProviderRegistry | None = None
|
|
|
|
|
|
def get_provider_registry(settings: Settings | None = None) -> ProviderRegistry:
|
|
"""Get the global provider registry instance.
|
|
|
|
Args:
|
|
settings: Optional settings instance. Creates new registry if provided.
|
|
|
|
Returns:
|
|
The global ProviderRegistry instance.
|
|
"""
|
|
global _registry
|
|
if _registry is None or settings is not None:
|
|
_registry = ProviderRegistry(settings)
|
|
return _registry
|
|
|
|
|
|
def reset_provider_registry() -> None:
|
|
"""Reset the global provider registry.
|
|
|
|
Useful for testing to ensure clean state between tests.
|
|
"""
|
|
global _registry
|
|
_registry = None
|