Files
temp/src/cleveragents/providers/registry.py
freemo a074b4846f fix(provider): remove FakeListLLM defaults
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
2026-02-27 09:47:10 -05:00

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