use std::sync::Arc;
use tracing::warn;
use crate::application_context::{ApplicationContext, AttributionPolicy, AttributionProviderKind};
use crate::error::{LlmError, Result};
use crate::http::attribution::supports_attribution;
use crate::model_config::{ProviderConfig, ProviderType as ConfigProviderType};
use crate::providers::anthropic::AnthropicProvider;
use crate::providers::azure_openai::AzureOpenAIProvider;
use crate::providers::gemini::GeminiProvider;
use crate::providers::huggingface::HuggingFaceProvider;
use crate::providers::jina::JinaProvider;
use crate::providers::lmstudio::LMStudioProvider;
use crate::providers::mistral::MistralProvider;
use crate::providers::nvidia::NvidiaProvider;
use crate::providers::openai_compatible::OpenAICompatibleProvider;
use crate::providers::openrouter::OpenRouterProvider;
use crate::providers::xai::XAIProvider;
use crate::traits::{EmbeddingProvider, LLMProvider};
use crate::{MockProvider, OllamaProvider, OpenAIProvider, VsCodeCopilotProvider};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ProviderType {
OpenAI,
Anthropic,
Gemini,
VertexAI,
OpenRouter,
XAI,
HuggingFace,
OpenAICompatible,
Ollama,
LMStudio,
VsCodeCopilot,
Mock,
Mistral,
AzureOpenAI,
#[cfg(feature = "bedrock")]
Bedrock,
Nvidia,
Cohere,
}
impl ProviderType {
#[allow(clippy::should_implement_trait)]
pub fn from_str(s: &str) -> Option<Self> {
match s.to_lowercase().as_str() {
"openai" => Some(Self::OpenAI),
"anthropic" | "claude" => Some(Self::Anthropic),
"gemini" | "google" => Some(Self::Gemini),
"vertex" | "vertexai" => Some(Self::VertexAI),
"openrouter" | "open-router" => Some(Self::OpenRouter),
"xai" | "grok" => Some(Self::XAI),
"huggingface" | "hf" | "hugging-face" | "hugging_face" => Some(Self::HuggingFace),
"openai-compatible" | "openai_compatible" | "openaicompatible" | "compatible" => {
Some(Self::OpenAICompatible)
}
"ollama" => Some(Self::Ollama),
"lmstudio" | "lm-studio" | "lm_studio" => Some(Self::LMStudio),
"vscode" | "vscode-copilot" | "copilot" => Some(Self::VsCodeCopilot),
"mock" => Some(Self::Mock),
"mistral" | "mistral-ai" | "mistralai" => Some(Self::Mistral),
"azure" | "azure-openai" | "azure_openai" | "azureopenai" => Some(Self::AzureOpenAI),
#[cfg(feature = "bedrock")]
"bedrock" | "aws-bedrock" | "aws_bedrock" => Some(Self::Bedrock),
"nvidia" | "nvidia-nim" | "nim" => Some(Self::Nvidia),
"cohere" | "cohere-ai" => Some(Self::Cohere),
_ => crate::provider_catalog::ProviderCatalog::resolve_id(s)
.and_then(|id| Self::all().iter().find(|t| t.canonical_id() == id).copied()),
}
}
pub fn all() -> &'static [Self] {
&[
Self::OpenAI,
Self::Anthropic,
Self::Gemini,
Self::VertexAI,
Self::OpenRouter,
Self::XAI,
Self::HuggingFace,
Self::OpenAICompatible,
Self::Ollama,
Self::LMStudio,
Self::VsCodeCopilot,
Self::Mock,
Self::Mistral,
Self::AzureOpenAI,
Self::Nvidia,
Self::Cohere,
#[cfg(feature = "bedrock")]
Self::Bedrock,
]
}
pub fn canonical_id(self) -> &'static str {
match self {
Self::OpenAI => "openai",
Self::Anthropic => "anthropic",
Self::Gemini => "gemini",
Self::VertexAI => "vertexai",
Self::OpenRouter => "openrouter",
Self::XAI => "xai",
Self::HuggingFace => "huggingface",
Self::OpenAICompatible => "openai-compatible",
Self::Ollama => "ollama",
Self::LMStudio => "lmstudio",
Self::VsCodeCopilot => "vscode-copilot",
Self::Mock => "mock",
Self::Mistral => "mistral",
Self::AzureOpenAI => "azure",
Self::Nvidia => "nvidia",
Self::Cohere => "cohere",
#[cfg(feature = "bedrock")]
Self::Bedrock => "bedrock",
}
}
pub fn descriptor(self) -> Option<&'static crate::provider_catalog::ProviderDescriptor> {
crate::provider_catalog::ProviderCatalog::get(self.canonical_id())
}
}
impl From<ProviderType> for AttributionProviderKind {
fn from(value: ProviderType) -> Self {
match value {
ProviderType::OpenAI => Self::OpenAI,
ProviderType::AzureOpenAI => Self::AzureOpenAI,
ProviderType::Anthropic => Self::Anthropic,
ProviderType::Gemini => Self::Gemini,
ProviderType::VertexAI => Self::VertexAI,
ProviderType::OpenRouter => Self::OpenRouter,
ProviderType::OpenAICompatible => Self::OpenAICompatible,
ProviderType::Mistral => Self::Mistral,
ProviderType::Nvidia => Self::Nvidia,
ProviderType::Cohere => Self::Cohere,
#[cfg(feature = "bedrock")]
ProviderType::Bedrock => Self::Bedrock,
ProviderType::XAI => Self::XAI,
ProviderType::HuggingFace => Self::HuggingFace,
ProviderType::LMStudio => Self::LMStudio,
ProviderType::Ollama => Self::Ollama,
ProviderType::VsCodeCopilot => Self::VsCodeCopilot,
ProviderType::Mock => Self::Mock,
}
}
}
pub struct ProviderFactory;
impl ProviderFactory {
pub fn list_providers() -> Vec<&'static str> {
crate::provider_catalog::ProviderCatalog::list_llm_providers()
}
pub fn list_embedding_providers() -> Vec<&'static str> {
crate::provider_catalog::ProviderCatalog::list_embedding_providers()
}
pub fn list_discovery_providers() -> Vec<&'static str> {
crate::provider_catalog::ProviderCatalog::list_discovery_providers()
}
pub fn from_env() -> Result<(Arc<dyn LLMProvider>, Arc<dyn EmbeddingProvider>)> {
if let Ok(provider_str) = std::env::var("EDGEQUAKE_LLM_PROVIDER") {
if let Some(provider_type) = ProviderType::from_str(&provider_str) {
return Self::create(provider_type);
}
return Err(LlmError::ConfigError(format!(
"Unknown provider type: {}. Valid options: openai, anthropic, gemini, vertexai, openrouter, xai, huggingface, openai-compatible, ollama, lmstudio, vscode-copilot, mistral, azure, bedrock, mock",
provider_str
)));
}
if std::env::var("OLLAMA_HOST").is_ok()
|| std::env::var("OLLAMA_MODEL").is_ok()
|| std::env::var("OLLAMA_API_KEY").is_ok()
{
return Self::create(ProviderType::Ollama);
}
if std::env::var("LMSTUDIO_HOST").is_ok() || std::env::var("LMSTUDIO_MODEL").is_ok() {
return Self::create(ProviderType::LMStudio);
}
if AnthropicProvider::resolve_api_key_from_env().is_ok() {
return Self::create(ProviderType::Anthropic);
}
if let Ok(api_key) =
std::env::var("GEMINI_API_KEY").or_else(|_| std::env::var("GOOGLE_API_KEY"))
{
if !api_key.is_empty() {
return Self::create(ProviderType::Gemini);
}
}
if let Ok(api_key) = std::env::var("MISTRAL_API_KEY") {
if !api_key.is_empty() {
return Self::create(ProviderType::Mistral);
}
}
let azure_key = std::env::var("AZURE_OPENAI_CONTENTGEN_API_KEY")
.or_else(|_| std::env::var("AZURE_OPENAI_API_KEY"));
if let Ok(api_key) = azure_key {
if !api_key.is_empty() {
let azure_endpoint = std::env::var("AZURE_OPENAI_CONTENTGEN_API_ENDPOINT")
.or_else(|_| std::env::var("AZURE_OPENAI_ENDPOINT"));
if azure_endpoint.is_ok() {
return Self::create(ProviderType::AzureOpenAI);
}
}
}
if let Ok(api_key) = std::env::var("XAI_API_KEY") {
if !api_key.is_empty() {
return Self::create(ProviderType::XAI);
}
}
if let Ok(api_key) =
std::env::var("HF_TOKEN").or_else(|_| std::env::var("HUGGINGFACE_TOKEN"))
{
if !api_key.is_empty() {
return Self::create(ProviderType::HuggingFace);
}
}
if let Ok(api_key) = std::env::var("OPENROUTER_API_KEY") {
if !api_key.is_empty() {
return Self::create(ProviderType::OpenRouter);
}
}
if let Ok(api_key) = std::env::var("NVIDIA_API_KEY") {
if !api_key.is_empty() {
return Self::create(ProviderType::Nvidia);
}
}
if let Ok(api_key) = std::env::var("COHERE_API_KEY") {
if !api_key.is_empty() {
return Self::create(ProviderType::Cohere);
}
}
if let Ok(api_key) = std::env::var("OPENAI_API_KEY") {
if !api_key.is_empty() && api_key != "test-key" {
return Self::create(ProviderType::OpenAI);
}
}
Ok(Self::create_mock())
}
pub fn create(
provider_type: ProviderType,
) -> Result<(Arc<dyn LLMProvider>, Arc<dyn EmbeddingProvider>)> {
match provider_type {
ProviderType::OpenAI => Self::create_openai(),
ProviderType::Anthropic => Self::create_anthropic(),
ProviderType::Gemini => Self::create_gemini(),
ProviderType::VertexAI => Self::create_vertex_ai(),
ProviderType::OpenRouter => Self::create_openrouter(),
ProviderType::XAI => Self::create_xai(),
ProviderType::HuggingFace => Self::create_huggingface(),
ProviderType::OpenAICompatible => Self::create_openai_compatible_from_env(),
ProviderType::Ollama => Self::create_ollama(),
ProviderType::LMStudio => Self::create_lmstudio(),
ProviderType::VsCodeCopilot => Self::create_vscode_copilot(),
ProviderType::Mock => Ok(Self::create_mock()),
ProviderType::Mistral => Self::create_mistral(),
ProviderType::AzureOpenAI => Self::create_azure_openai(),
#[cfg(feature = "bedrock")]
ProviderType::Bedrock => Self::create_bedrock(),
ProviderType::Nvidia => Self::create_nvidia(),
ProviderType::Cohere => Self::create_cohere(),
}
}
pub fn create_with_model(
provider_type: ProviderType,
model: Option<&str>,
) -> Result<(Arc<dyn LLMProvider>, Arc<dyn EmbeddingProvider>)> {
match model {
Some(m) => match provider_type {
ProviderType::OpenRouter => Self::create_openrouter_with_model(m),
ProviderType::Anthropic => Self::create_anthropic_with_model(m),
ProviderType::Gemini => Self::create_gemini_with_model(m),
ProviderType::VertexAI => Self::create_vertex_ai_with_model(m),
ProviderType::XAI => Self::create_xai_with_model(m),
ProviderType::OpenAI => Self::create_openai_with_model(m),
ProviderType::OpenAICompatible => {
Self::create_openai_compatible_from_env_with_model(m)
}
ProviderType::Ollama => Self::create_ollama_with_model(m),
ProviderType::LMStudio => Self::create_lmstudio_with_model(m),
ProviderType::HuggingFace => Self::create_huggingface(),
ProviderType::VsCodeCopilot => Self::create_vscode_copilot(),
ProviderType::Mock => Ok(Self::create_mock()),
ProviderType::Mistral => Self::create_mistral_with_model(m),
ProviderType::AzureOpenAI => Self::create_azure_openai_with_deployment(m),
#[cfg(feature = "bedrock")]
ProviderType::Bedrock => Self::create_bedrock_with_model(m),
ProviderType::Nvidia => Self::create_nvidia_with_model(m),
ProviderType::Cohere => Self::create_cohere_with_model(m),
},
None => Self::create(provider_type),
}
}
pub fn from_config(
config: &ProviderConfig,
) -> Result<(Arc<dyn LLMProvider>, Arc<dyn EmbeddingProvider>)> {
Self::from_config_with_model(config, None)
}
pub fn from_config_with_model(
config: &ProviderConfig,
model_name: Option<&str>,
) -> Result<(Arc<dyn LLMProvider>, Arc<dyn EmbeddingProvider>)> {
match config.provider_type {
ConfigProviderType::OpenAI => Self::create_openai(),
ConfigProviderType::Ollama => Self::create_ollama(),
ConfigProviderType::LMStudio => Self::create_lmstudio(),
ConfigProviderType::Mock => Ok(Self::create_mock()),
ConfigProviderType::OpenAICompatible => {
Self::create_openai_compatible_with_model(config, model_name)
}
ConfigProviderType::Azure => {
Self::create_azure_openai()
}
ConfigProviderType::Anthropic => Self::create_anthropic_from_config(config, model_name),
ConfigProviderType::OpenRouter => {
Self::create_openrouter_from_config(config, model_name)
}
ConfigProviderType::Mistral => Self::create_mistral_from_config(config, model_name),
}
}
#[allow(dead_code)]
fn create_openai_compatible(
config: &ProviderConfig,
) -> Result<(Arc<dyn LLMProvider>, Arc<dyn EmbeddingProvider>)> {
Self::create_openai_compatible_with_model(config, None)
}
fn create_openai_compatible_with_model(
config: &ProviderConfig,
model_name: Option<&str>,
) -> Result<(Arc<dyn LLMProvider>, Arc<dyn EmbeddingProvider>)> {
let mut provider_instance = OpenAICompatibleProvider::from_config(config.clone())?;
if let Some(model) = model_name {
provider_instance = provider_instance.with_model(model);
}
let provider = Arc::new(provider_instance);
let has_embedding = config.default_embedding_model.is_some();
if has_embedding {
Ok((provider.clone(), provider))
} else {
match Self::create_openai() {
Ok((_, embedding)) => Ok((provider, embedding)),
Err(_) => {
Ok((provider.clone(), provider))
}
}
}
}
fn create_openai() -> Result<(Arc<dyn LLMProvider>, Arc<dyn EmbeddingProvider>)> {
let api_key = std::env::var("OPENAI_API_KEY").map_err(|_| {
LlmError::ConfigError("OPENAI_API_KEY not set for OpenAI provider".to_string())
})?;
if api_key.is_empty() || api_key == "test-key" {
return Err(LlmError::ConfigError(
"OPENAI_API_KEY is empty or invalid".to_string(),
));
}
let provider = Arc::new(OpenAIProvider::new(api_key));
Ok((provider.clone(), provider))
}
fn create_openai_compatible_from_env(
) -> Result<(Arc<dyn LLMProvider>, Arc<dyn EmbeddingProvider>)> {
Self::create_openai_compatible_from_env_with_model(
&std::env::var("OPENAI_COMPATIBLE_MODEL").unwrap_or_else(|_| "default".to_string()),
)
}
fn create_openai_compatible_from_env_with_model(
model: &str,
) -> Result<(Arc<dyn LLMProvider>, Arc<dyn EmbeddingProvider>)> {
let base_url = std::env::var("OPENAI_COMPATIBLE_BASE_URL").map_err(|_| {
LlmError::ConfigError(
"OPENAI_COMPATIBLE_BASE_URL not set for OpenAI-compatible provider".to_string(),
)
})?;
let mut config = ProviderConfig {
name: "openai-compatible".to_string(),
display_name: "OpenAI Compatible".to_string(),
provider_type: ConfigProviderType::OpenAICompatible,
base_url: Some(base_url),
default_llm_model: Some(model.to_string()),
default_embedding_model: std::env::var("OPENAI_COMPATIBLE_EMBEDDING_MODEL").ok(),
..Default::default()
};
if let Ok(api_key) = std::env::var("OPENAI_COMPATIBLE_API_KEY") {
if !api_key.is_empty() {
config.api_key = Some(api_key);
}
}
Self::create_openai_compatible_with_model(&config, Some(model))
}
fn create_anthropic() -> Result<(Arc<dyn LLMProvider>, Arc<dyn EmbeddingProvider>)> {
let provider = Arc::new(AnthropicProvider::from_env()?);
let embedding: Arc<dyn EmbeddingProvider> = match Self::create_openai() {
Ok((_, embedding)) => embedding,
Err(_) => Arc::new(MockProvider::new()),
};
Ok((provider, embedding))
}
fn create_anthropic_from_config(
config: &ProviderConfig,
model_name: Option<&str>,
) -> Result<(Arc<dyn LLMProvider>, Arc<dyn EmbeddingProvider>)> {
let api_key_var = config.api_key_env.as_deref().unwrap_or("ANTHROPIC_API_KEY");
let api_key = if matches!(api_key_var, "ANTHROPIC_API_KEY" | "ANTHROPIC_AUTH_TOKEN") {
AnthropicProvider::resolve_api_key_from_env()?
} else {
let value = std::env::var(api_key_var).map_err(|_| {
LlmError::ConfigError(format!("{} not set for Anthropic provider", api_key_var))
})?;
let trimmed = value.trim();
if trimmed.is_empty() {
return Err(LlmError::ConfigError(format!("{} is empty", api_key_var)));
}
trimmed.to_string()
};
let model = model_name
.map(|s| s.to_string())
.or_else(|| config.default_llm_model.clone())
.unwrap_or_else(|| "claude-sonnet-4-5-20250929".to_string());
let mut provider = AnthropicProvider::new(api_key).with_model(model);
if let Some(base_url) = &config.base_url {
provider = provider.with_base_url(base_url);
}
let llm_provider = Arc::new(provider);
let embedding: Arc<dyn EmbeddingProvider> = match Self::create_openai() {
Ok((_, embedding)) => embedding,
Err(_) => Arc::new(MockProvider::new()),
};
Ok((llm_provider, embedding))
}
fn create_gemini() -> Result<(Arc<dyn LLMProvider>, Arc<dyn EmbeddingProvider>)> {
use crate::GeminiProvider;
let provider = GeminiProvider::from_env()?;
let llm_provider: Arc<dyn LLMProvider> = Arc::new(provider);
let embedding: Arc<dyn EmbeddingProvider> = Arc::new(GeminiProvider::from_env()?);
Ok((llm_provider, embedding))
}
fn create_vertex_ai() -> Result<(Arc<dyn LLMProvider>, Arc<dyn EmbeddingProvider>)> {
use crate::GeminiProvider;
let provider = GeminiProvider::from_env_vertex_ai()?;
let llm_provider: Arc<dyn LLMProvider> = Arc::new(provider);
let embedding: Arc<dyn EmbeddingProvider> = Arc::new(GeminiProvider::from_env_vertex_ai()?);
Ok((llm_provider, embedding))
}
fn create_vertex_ai_with_model(
model: &str,
) -> Result<(Arc<dyn LLMProvider>, Arc<dyn EmbeddingProvider>)> {
use crate::GeminiProvider;
let actual = model.strip_prefix("vertexai:").unwrap_or(model);
let provider = Arc::new(GeminiProvider::from_env_vertex_ai()?.with_model(actual));
let embedding: Arc<dyn EmbeddingProvider> = Arc::new(GeminiProvider::from_env_vertex_ai()?);
Ok((provider, embedding))
}
fn create_openrouter() -> Result<(Arc<dyn LLMProvider>, Arc<dyn EmbeddingProvider>)> {
let api_key = std::env::var("OPENROUTER_API_KEY").map_err(|_| {
LlmError::ConfigError("OPENROUTER_API_KEY not set for OpenRouter provider".to_string())
})?;
if api_key.is_empty() {
return Err(LlmError::ConfigError(
"OPENROUTER_API_KEY is empty".to_string(),
));
}
let model = std::env::var("OPENROUTER_MODEL")
.unwrap_or_else(|_| "anthropic/claude-3.5-sonnet".to_string());
let mut provider = OpenRouterProvider::new(api_key).with_model(model);
if let Ok(url) = std::env::var("OPENROUTER_SITE_URL") {
provider = provider.with_site_url(url);
}
if let Ok(name) = std::env::var("OPENROUTER_SITE_NAME") {
provider = provider.with_site_name(name);
}
let llm_provider = Arc::new(provider);
let embedding: Arc<dyn EmbeddingProvider> = match Self::create_openai() {
Ok((_, embedding)) => embedding,
Err(_) => Arc::new(MockProvider::new()),
};
Ok((llm_provider, embedding))
}
fn create_openrouter_with_model(
model: &str,
) -> Result<(Arc<dyn LLMProvider>, Arc<dyn EmbeddingProvider>)> {
let api_key = std::env::var("OPENROUTER_API_KEY").map_err(|_| {
LlmError::ConfigError("OPENROUTER_API_KEY not set for OpenRouter provider".to_string())
})?;
if api_key.is_empty() {
return Err(LlmError::ConfigError(
"OPENROUTER_API_KEY is empty".to_string(),
));
}
let mut provider = OpenRouterProvider::new(api_key).with_model(model);
if let Ok(url) = std::env::var("OPENROUTER_SITE_URL") {
provider = provider.with_site_url(url);
}
if let Ok(name) = std::env::var("OPENROUTER_SITE_NAME") {
provider = provider.with_site_name(name);
}
let llm_provider = Arc::new(provider);
let embedding: Arc<dyn EmbeddingProvider> = match Self::create_openai() {
Ok((_, embedding)) => embedding,
Err(_) => Arc::new(MockProvider::new()),
};
Ok((llm_provider, embedding))
}
fn create_openai_with_model(
model: &str,
) -> Result<(Arc<dyn LLMProvider>, Arc<dyn EmbeddingProvider>)> {
let api_key = std::env::var("OPENAI_API_KEY").map_err(|_| {
LlmError::ConfigError("OPENAI_API_KEY not set for OpenAI provider".to_string())
})?;
if api_key.is_empty() || api_key == "test-key" {
return Err(LlmError::ConfigError(
"OPENAI_API_KEY is empty or invalid".to_string(),
));
}
let provider = Arc::new(OpenAIProvider::new(api_key).with_model(model));
Ok((provider.clone(), provider))
}
fn create_anthropic_with_model(
model: &str,
) -> Result<(Arc<dyn LLMProvider>, Arc<dyn EmbeddingProvider>)> {
let api_key = AnthropicProvider::resolve_api_key_from_env()?;
let mut provider = AnthropicProvider::new(&api_key);
if let Ok(base_url) = std::env::var("ANTHROPIC_BASE_URL") {
provider = provider.with_base_url(&base_url);
}
provider = provider.with_model(model);
let llm_provider = Arc::new(provider);
let embedding: Arc<dyn EmbeddingProvider> = match Self::create_openai() {
Ok((_, embedding)) => embedding,
Err(_) => Arc::new(MockProvider::new()),
};
Ok((llm_provider, embedding))
}
fn create_gemini_with_model(
model: &str,
) -> Result<(Arc<dyn LLMProvider>, Arc<dyn EmbeddingProvider>)> {
if model.starts_with("vertexai:") {
let actual_model = model.strip_prefix("vertexai:").unwrap_or(model);
let provider = Arc::new(GeminiProvider::from_env_vertex_ai()?.with_model(actual_model));
let embedding: Arc<dyn EmbeddingProvider> =
Arc::new(GeminiProvider::from_env_vertex_ai()?);
return Ok((provider, embedding));
}
let api_key = std::env::var("GEMINI_API_KEY")
.or_else(|_| std::env::var("GOOGLE_API_KEY"))
.map_err(|_| {
LlmError::ConfigError(
"GEMINI_API_KEY or GOOGLE_API_KEY not set for Gemini provider".to_string(),
)
})?;
let provider = Arc::new(GeminiProvider::new(&api_key).with_model(model));
let embedding: Arc<dyn EmbeddingProvider> = Arc::new(GeminiProvider::new(&api_key));
Ok((provider, embedding))
}
fn create_xai_with_model(
model: &str,
) -> Result<(Arc<dyn LLMProvider>, Arc<dyn EmbeddingProvider>)> {
std::env::var("XAI_API_KEY").map_err(|_| {
LlmError::ConfigError("XAI_API_KEY not set for xAI provider".to_string())
})?;
std::env::set_var("XAI_MODEL", model);
let provider = XAIProvider::from_env()?;
let llm_provider = Arc::new(provider);
let embedding: Arc<dyn EmbeddingProvider> = match Self::create_openai() {
Ok((_, embedding)) => embedding,
Err(_) => Arc::new(MockProvider::new()),
};
Ok((llm_provider, embedding))
}
fn create_ollama_with_model(
model: &str,
) -> Result<(Arc<dyn LLMProvider>, Arc<dyn EmbeddingProvider>)> {
let provider = Arc::new(OllamaProvider::from_env_with_model(model)?);
Ok((provider.clone(), provider))
}
fn create_lmstudio_with_model(
model: &str,
) -> Result<(Arc<dyn LLMProvider>, Arc<dyn EmbeddingProvider>)> {
std::env::set_var("LMSTUDIO_MODEL", model);
let provider = Arc::new(LMStudioProvider::from_env()?);
Ok((provider.clone(), provider))
}
fn create_openrouter_from_config(
config: &ProviderConfig,
model_name: Option<&str>,
) -> Result<(Arc<dyn LLMProvider>, Arc<dyn EmbeddingProvider>)> {
let api_key_var = config
.api_key_env
.as_deref()
.unwrap_or("OPENROUTER_API_KEY");
let api_key = std::env::var(api_key_var).map_err(|_| {
LlmError::ConfigError(format!("{} not set for OpenRouter provider", api_key_var))
})?;
if api_key.is_empty() {
return Err(LlmError::ConfigError(format!("{} is empty", api_key_var)));
}
let model = model_name
.map(|s| s.to_string())
.or_else(|| config.default_llm_model.clone())
.unwrap_or_else(|| "anthropic/claude-3.5-sonnet".to_string());
let mut provider = OpenRouterProvider::new(api_key).with_model(model);
if let Some(base_url) = &config.base_url {
provider = provider.with_base_url(base_url);
}
let llm_provider = Arc::new(provider);
let embedding: Arc<dyn EmbeddingProvider> = match Self::create_openai() {
Ok((_, embedding)) => embedding,
Err(_) => Arc::new(MockProvider::new()),
};
Ok((llm_provider, embedding))
}
fn create_ollama() -> Result<(Arc<dyn LLMProvider>, Arc<dyn EmbeddingProvider>)> {
let provider = Arc::new(OllamaProvider::from_env()?);
Ok((provider.clone(), provider))
}
fn create_lmstudio() -> Result<(Arc<dyn LLMProvider>, Arc<dyn EmbeddingProvider>)> {
let provider = Arc::new(LMStudioProvider::from_env()?);
Ok((provider.clone(), provider))
}
fn create_xai() -> Result<(Arc<dyn LLMProvider>, Arc<dyn EmbeddingProvider>)> {
let provider = Arc::new(XAIProvider::from_env()?);
let embedding: Arc<dyn EmbeddingProvider> = match Self::create_openai() {
Ok((_, embedding)) => embedding,
Err(_) => Arc::new(MockProvider::new()),
};
Ok((provider, embedding))
}
fn create_huggingface() -> Result<(Arc<dyn LLMProvider>, Arc<dyn EmbeddingProvider>)> {
let provider = Arc::new(HuggingFaceProvider::from_env()?);
let embedding: Arc<dyn EmbeddingProvider> = match Self::create_openai() {
Ok((_, embedding)) => embedding,
Err(_) => Arc::new(MockProvider::new()),
};
Ok((provider, embedding))
}
fn create_vscode_copilot() -> Result<(Arc<dyn LLMProvider>, Arc<dyn EmbeddingProvider>)> {
let model = std::env::var("VSCODE_COPILOT_MODEL").unwrap_or_else(|_| "auto".to_string());
let builder = VsCodeCopilotProvider::new().model(model);
let provider = Arc::new(builder.build().map_err(|e| {
LlmError::ConfigError(format!("Failed to create VSCode Copilot provider: {}", e))
})?);
Ok((provider.clone(), provider))
}
fn create_mock() -> (Arc<dyn LLMProvider>, Arc<dyn EmbeddingProvider>) {
let provider = Arc::new(MockProvider::new());
(provider.clone(), provider)
}
fn create_mistral() -> Result<(Arc<dyn LLMProvider>, Arc<dyn EmbeddingProvider>)> {
let provider = Arc::new(MistralProvider::from_env()?);
Ok((provider.clone(), provider))
}
fn create_mistral_with_model(
model: &str,
) -> Result<(Arc<dyn LLMProvider>, Arc<dyn EmbeddingProvider>)> {
let provider = Arc::new(MistralProvider::from_env()?.with_model(model));
Ok((provider.clone(), provider))
}
fn create_mistral_from_config(
config: &ProviderConfig,
model_name: Option<&str>,
) -> Result<(Arc<dyn LLMProvider>, Arc<dyn EmbeddingProvider>)> {
let mut provider = MistralProvider::from_config(config)?;
if let Some(model) = model_name {
provider = provider.with_model(model);
}
let provider = Arc::new(provider);
Ok((provider.clone(), provider))
}
fn create_nvidia() -> Result<(Arc<dyn LLMProvider>, Arc<dyn EmbeddingProvider>)> {
let provider = Arc::new(NvidiaProvider::from_env()?);
let embedding: Arc<dyn EmbeddingProvider> = match Self::create_openai() {
Ok((_, embedding)) => embedding,
Err(_) => Arc::new(MockProvider::new()),
};
Ok((provider, embedding))
}
fn create_nvidia_with_model(
model: &str,
) -> Result<(Arc<dyn LLMProvider>, Arc<dyn EmbeddingProvider>)> {
let provider = Arc::new(NvidiaProvider::from_env()?.with_model(model));
let embedding: Arc<dyn EmbeddingProvider> = match Self::create_openai() {
Ok((_, embedding)) => embedding,
Err(_) => Arc::new(MockProvider::new()),
};
Ok((provider, embedding))
}
fn create_cohere() -> Result<(Arc<dyn LLMProvider>, Arc<dyn EmbeddingProvider>)> {
let provider = Arc::new(
crate::providers::cohere::CohereProvider::from_env()
.map_err(|e| LlmError::ProviderError(format!("Cohere init failed: {e}")))?,
);
Ok((provider.clone(), provider))
}
fn create_cohere_with_model(
model: &str,
) -> Result<(Arc<dyn LLMProvider>, Arc<dyn EmbeddingProvider>)> {
let provider = Arc::new(
crate::providers::cohere::CohereProvider::from_env()
.map_err(|e| LlmError::ProviderError(format!("Cohere init failed: {e}")))?,
);
let provider = Arc::new(
Arc::try_unwrap(provider)
.unwrap_or_else(|p| (*p).clone())
.with_model(model),
);
Ok((provider.clone(), provider))
}
pub fn create_azure_openai() -> Result<(Arc<dyn LLMProvider>, Arc<dyn EmbeddingProvider>)> {
let provider = Arc::new(AzureOpenAIProvider::from_env_auto()?);
Ok((provider.clone(), provider))
}
fn create_azure_openai_with_deployment(
deployment: &str,
) -> Result<(Arc<dyn LLMProvider>, Arc<dyn EmbeddingProvider>)> {
let provider = Arc::new(AzureOpenAIProvider::from_env_auto()?.with_deployment(deployment));
Ok((provider.clone(), provider))
}
#[cfg(feature = "bedrock")]
fn create_bedrock() -> Result<(Arc<dyn LLMProvider>, Arc<dyn EmbeddingProvider>)> {
use crate::providers::bedrock::BedrockProvider;
let rt = tokio::runtime::Handle::try_current().map_err(|_| {
LlmError::ConfigError(
"Bedrock provider requires a Tokio runtime (use #[tokio::main] or Runtime::new())"
.to_string(),
)
})?;
let provider = Arc::new(tokio::task::block_in_place(|| {
rt.block_on(BedrockProvider::from_env())
})?);
let embedding: Arc<dyn EmbeddingProvider> = provider.clone();
Ok((provider, embedding))
}
#[cfg(feature = "bedrock")]
fn create_bedrock_with_model(
model: &str,
) -> Result<(Arc<dyn LLMProvider>, Arc<dyn EmbeddingProvider>)> {
use crate::providers::bedrock::BedrockProvider;
let rt = tokio::runtime::Handle::try_current().map_err(|_| {
LlmError::ConfigError("Bedrock provider requires a Tokio runtime".to_string())
})?;
let provider = Arc::new(
tokio::task::block_in_place(|| rt.block_on(BedrockProvider::from_env()))?
.with_model(model),
);
let embedding: Arc<dyn EmbeddingProvider> = provider.clone();
Ok((provider, embedding))
}
pub fn embedding_dimension() -> Result<usize> {
let (_, embedding_provider) = Self::from_env()?;
Ok(embedding_provider.dimension())
}
pub fn create_embedding_provider(
provider_name: &str,
model: &str,
_dimension: usize,
) -> Result<Arc<dyn EmbeddingProvider>> {
if provider_name.eq_ignore_ascii_case("jina") {
let api_key = std::env::var("JINA_API_KEY").map_err(|_| {
LlmError::ConfigError("JINA_API_KEY is required for Jina embeddings".to_string())
})?;
let base_url = std::env::var("JINA_BASE_URL")
.unwrap_or_else(|_| "https://api.jina.ai".to_string());
let provider = JinaProvider::builder()
.api_key(api_key)
.base_url(base_url)
.embedding_model(model)
.build()?;
return Ok(Arc::new(provider));
}
let provider_type = ProviderType::from_str(provider_name).ok_or_else(|| {
LlmError::ConfigError(format!(
"Unknown embedding provider: {}. Valid: openai, anthropic, gemini, vertexai, openrouter, xai, huggingface, openai-compatible, ollama, lmstudio, vscode-copilot, mistral, azure, bedrock, jina, mock",
provider_name
))
})?;
match provider_type {
ProviderType::OpenAI => {
let api_key = std::env::var("OPENAI_API_KEY").map_err(|_| {
LlmError::ConfigError(
"OPENAI_API_KEY required for OpenAI embedding provider".to_string(),
)
})?;
let provider = OpenAIProvider::new(api_key).with_embedding_model(model);
Ok(Arc::new(provider))
}
ProviderType::Anthropic => {
warn!("Anthropic doesn't support embeddings, using mock provider");
Ok(Arc::new(MockProvider::new()))
}
ProviderType::OpenRouter => {
warn!("OpenRouter doesn't support embeddings, using mock provider");
Ok(Arc::new(MockProvider::new()))
}
ProviderType::XAI => {
warn!("xAI doesn't support embeddings, using mock provider");
Ok(Arc::new(MockProvider::new()))
}
ProviderType::OpenAICompatible => {
let (_, embedding) = Self::create_openai_compatible_from_env_with_model(model)?;
Ok(embedding)
}
ProviderType::HuggingFace => {
warn!("HuggingFace LLM provider doesn't support embeddings, using mock provider");
Ok(Arc::new(MockProvider::new()))
}
ProviderType::Gemini => {
match GeminiProvider::from_env() {
Ok(provider) => Ok(Arc::new(provider.with_embedding_model(model))),
Err(e) => {
warn!(
"Gemini credentials unavailable ({}), falling back to mock embedding provider",
e
);
Ok(Arc::new(MockProvider::new()))
}
}
}
ProviderType::VertexAI => {
match GeminiProvider::from_env_vertex_ai() {
Ok(provider) => Ok(Arc::new(provider.with_embedding_model(model))),
Err(e) => {
warn!(
"Vertex AI credentials unavailable ({}), falling back to mock embedding provider",
e
);
Ok(Arc::new(MockProvider::new()))
}
}
}
ProviderType::Ollama => {
let provider = OllamaProvider::from_env()?.with_embedding_model(model);
Ok(Arc::new(provider))
}
ProviderType::LMStudio => {
let host = std::env::var("LMSTUDIO_HOST")
.unwrap_or_else(|_| "http://localhost:1234".to_string());
let provider = LMStudioProvider::builder()
.host(&host)
.embedding_model(model)
.build()?;
Ok(Arc::new(provider))
}
ProviderType::Mock => {
Ok(Arc::new(MockProvider::new()))
}
ProviderType::VsCodeCopilot => {
let provider = VsCodeCopilotProvider::new()
.embedding_model(model)
.build()
.map_err(|e| LlmError::ApiError(e.to_string()))?;
Ok(Arc::new(provider))
}
ProviderType::Mistral => {
let provider = MistralProvider::from_env()?.with_embedding_model(model);
Ok(Arc::new(provider))
}
ProviderType::AzureOpenAI => {
let provider =
AzureOpenAIProvider::from_env_auto()?.with_embedding_deployment(model);
Ok(Arc::new(provider))
}
#[cfg(feature = "bedrock")]
ProviderType::Bedrock => {
let (_, embedding) = Self::create_bedrock_with_model(model)?;
Ok(embedding)
}
ProviderType::Nvidia => {
warn!("NVIDIA NIM does not support embeddings via NvidiaProvider, using mock provider");
Ok(Arc::new(MockProvider::new()))
}
ProviderType::Cohere => {
let provider = crate::providers::cohere::CohereProvider::from_env()
.map_err(|e| LlmError::ProviderError(format!("Cohere init failed: {e}")))?;
Ok(Arc::new(provider))
}
}
}
pub fn create_llm_provider(provider_name: &str, model: &str) -> Result<Arc<dyn LLMProvider>> {
let provider_type = ProviderType::from_str(provider_name).ok_or_else(|| {
LlmError::ConfigError(format!(
"Unknown LLM provider: {}. Valid: openai, anthropic, gemini, vertexai, openrouter, xai, huggingface, openai-compatible, ollama, lmstudio, vscode-copilot, mistral, azure, bedrock, mock",
provider_name
))
})?;
match provider_type {
ProviderType::OpenAI => {
let api_key = std::env::var("OPENAI_API_KEY").map_err(|_| {
LlmError::ConfigError(
"OPENAI_API_KEY required for OpenAI LLM provider".to_string(),
)
})?;
let provider = OpenAIProvider::new(api_key).with_model(model);
Ok(Arc::new(provider))
}
ProviderType::Anthropic => {
let (provider, _) = Self::create_anthropic_with_model(model)?;
Ok(provider)
}
ProviderType::OpenRouter => {
let api_key = std::env::var("OPENROUTER_API_KEY").map_err(|_| {
LlmError::ConfigError(
"OPENROUTER_API_KEY required for OpenRouter LLM provider".to_string(),
)
})?;
let provider = OpenRouterProvider::new(api_key).with_model(model);
Ok(Arc::new(provider))
}
ProviderType::XAI => {
let provider = XAIProvider::from_env()?.with_model(model);
Ok(Arc::new(provider))
}
ProviderType::OpenAICompatible => {
let (provider, _) = Self::create_openai_compatible_from_env_with_model(model)?;
Ok(provider)
}
ProviderType::HuggingFace => {
let provider = HuggingFaceProvider::from_env()?.with_model(model);
Ok(Arc::new(provider))
}
ProviderType::Gemini => {
if model.starts_with("vertexai:") {
let actual_model = model.strip_prefix("vertexai:").unwrap_or(model);
let provider = GeminiProvider::from_env_vertex_ai()?.with_model(actual_model);
Ok(Arc::new(provider))
} else {
let provider = GeminiProvider::from_env()?.with_model(model);
Ok(Arc::new(provider))
}
}
ProviderType::VertexAI => {
let actual_model = model.strip_prefix("vertexai:").unwrap_or(model);
let provider = GeminiProvider::from_env_vertex_ai()?.with_model(actual_model);
Ok(Arc::new(provider))
}
ProviderType::Ollama => {
let provider = OllamaProvider::from_env_with_model(model)?;
Ok(Arc::new(provider))
}
ProviderType::LMStudio => {
let host = std::env::var("LMSTUDIO_HOST")
.unwrap_or_else(|_| "http://localhost:1234".to_string());
let provider = LMStudioProvider::builder()
.host(&host)
.model(model)
.build()?;
Ok(Arc::new(provider))
}
ProviderType::Mock => {
Ok(Arc::new(MockProvider::new()))
}
ProviderType::VsCodeCopilot => {
let proxy_url = std::env::var("VSCODE_COPILOT_PROXY_URL")
.unwrap_or_else(|_| "http://localhost:4141".to_string());
let provider = VsCodeCopilotProvider::new()
.proxy_url(&proxy_url)
.model(model)
.build()?;
Ok(Arc::new(provider))
}
ProviderType::Mistral => {
let provider = MistralProvider::from_env()?.with_model(model);
Ok(Arc::new(provider))
}
ProviderType::AzureOpenAI => {
let provider = AzureOpenAIProvider::from_env_auto()?.with_deployment(model);
Ok(Arc::new(provider))
}
#[cfg(feature = "bedrock")]
ProviderType::Bedrock => {
use crate::providers::bedrock::BedrockProvider;
let handle = tokio::runtime::Handle::try_current().map_err(|_| {
LlmError::ConfigError("Bedrock provider requires a Tokio runtime".to_string())
})?;
let provider =
tokio::task::block_in_place(|| handle.block_on(BedrockProvider::from_env()))
.map_err(|e| {
LlmError::ConfigError(format!(
"Failed to initialize Bedrock provider: {e}"
))
})?;
let provider = provider.with_model(model);
Ok(Arc::new(provider))
}
ProviderType::Nvidia => {
let provider = NvidiaProvider::from_env()?.with_model(model);
Ok(Arc::new(provider))
}
ProviderType::Cohere => {
let provider = crate::providers::cohere::CohereProvider::from_env()
.map_err(|e| LlmError::ProviderError(format!("Cohere init failed: {e}")))?
.with_model(model);
Ok(Arc::new(provider))
}
}
}
pub fn create_llm_provider_with_context(
provider_name: &str,
model: &str,
ctx: ApplicationContext,
) -> Result<Arc<dyn LLMProvider>> {
Self::create_llm_provider_with_context_policy(
provider_name,
model,
ctx,
AttributionPolicy::BestEffort,
)
}
pub fn create_llm_provider_with_context_policy(
provider_name: &str,
model: &str,
ctx: ApplicationContext,
policy: AttributionPolicy,
) -> Result<Arc<dyn LLMProvider>> {
if ctx.is_empty() || policy == AttributionPolicy::Disabled {
return Self::create_llm_provider(provider_name, model);
}
let provider_type = ProviderType::from_str(provider_name).ok_or_else(|| {
LlmError::ConfigError(format!(
"Unknown LLM provider: {}. Valid: openai, anthropic, gemini, vertexai, \
openrouter, xai, huggingface, openai-compatible, ollama, lmstudio, \
vscode-copilot, mistral, azure, bedrock, mock",
provider_name
))
})?;
let kind = AttributionProviderKind::from(provider_type);
if policy == AttributionPolicy::RequireAppId
&& ctx.has_app_attribution()
&& !supports_attribution(kind)
{
return Err(LlmError::AttributionError(format!(
"Provider '{}' does not support application attribution propagation",
provider_name
)));
}
if ctx.has_app_attribution() && !supports_attribution(kind) {
warn!(
provider = provider_name,
attribution_dropped = true,
"application attribution not supported for this provider"
);
}
match provider_type {
ProviderType::OpenAI => {
let provider = OpenAIProvider::from_env()?
.with_model(model)
.with_application_context(ctx);
Ok(Arc::new(provider))
}
ProviderType::Anthropic => {
let api_key = AnthropicProvider::resolve_api_key_from_env()?;
let mut provider = AnthropicProvider::new(&api_key);
if let Ok(base_url) = std::env::var("ANTHROPIC_BASE_URL") {
provider = provider.with_base_url(&base_url);
}
Ok(Arc::new(
provider.with_model(model).with_application_context(ctx),
))
}
ProviderType::OpenRouter => {
let provider = OpenRouterProvider::from_env()?
.with_model(model)
.with_application_context(ctx);
Ok(Arc::new(provider))
}
ProviderType::OpenAICompatible => {
let base_url = std::env::var("OPENAI_COMPATIBLE_BASE_URL").map_err(|_| {
LlmError::ConfigError(
"OPENAI_COMPATIBLE_BASE_URL not set for OpenAI-compatible provider"
.to_string(),
)
})?;
let mut config = ProviderConfig {
name: "openai-compatible".to_string(),
display_name: "OpenAI Compatible".to_string(),
provider_type: ConfigProviderType::OpenAICompatible,
base_url: Some(base_url),
default_llm_model: Some(model.to_string()),
default_embedding_model: std::env::var("OPENAI_COMPATIBLE_EMBEDDING_MODEL")
.ok(),
..Default::default()
};
if let Ok(api_key) = std::env::var("OPENAI_COMPATIBLE_API_KEY") {
if !api_key.is_empty() {
config.api_key = Some(api_key);
}
}
Ok(Arc::new(
OpenAICompatibleProvider::from_config(config)?
.with_model(model)
.with_application_context(ctx),
))
}
ProviderType::Gemini => {
if model.starts_with("vertexai:") {
let actual_model = model.strip_prefix("vertexai:").unwrap_or(model);
Ok(Arc::new(
GeminiProvider::from_env_vertex_ai()?
.with_model(actual_model)
.with_application_context(ctx),
))
} else {
Ok(Arc::new(
GeminiProvider::from_env()?
.with_model(model)
.with_application_context(ctx),
))
}
}
ProviderType::VertexAI => {
let actual_model = model.strip_prefix("vertexai:").unwrap_or(model);
Ok(Arc::new(
GeminiProvider::from_env_vertex_ai()?
.with_model(actual_model)
.with_application_context(ctx),
))
}
ProviderType::Mistral => Ok(Arc::new(
MistralProvider::from_env()?
.with_model(model)
.with_application_context(ctx),
)),
ProviderType::Nvidia => Ok(Arc::new(
NvidiaProvider::from_env()?
.with_model(model)
.with_application_context(ctx),
)),
ProviderType::XAI => Ok(Arc::new(
XAIProvider::from_env()?
.with_model(model)
.with_application_context(ctx),
)),
ProviderType::HuggingFace => Ok(Arc::new(
HuggingFaceProvider::from_env()?
.with_model(model)
.with_application_context(ctx),
)),
ProviderType::LMStudio => {
let host = std::env::var("LMSTUDIO_HOST")
.unwrap_or_else(|_| "http://localhost:1234".to_string());
Ok(Arc::new(
LMStudioProvider::builder()
.host(host)
.model(model)
.build()?
.with_application_context(ctx),
))
}
ProviderType::Cohere => Ok(Arc::new(
crate::providers::cohere::CohereProvider::from_env()
.map_err(|e| LlmError::ProviderError(format!("Cohere init failed: {e}")))?
.with_model(model)
.with_application_context(ctx),
)),
ProviderType::AzureOpenAI => Ok(Arc::new(
AzureOpenAIProvider::from_env_auto()?
.with_deployment(model)
.with_application_context(ctx),
)),
ProviderType::Ollama => Ok(Arc::new(
OllamaProvider::from_env_with_model(model)?.with_application_context(ctx),
)),
ProviderType::Mock | ProviderType::VsCodeCopilot => {
Self::create_llm_provider(provider_name, model)
}
#[cfg(feature = "bedrock")]
ProviderType::Bedrock => {
use crate::providers::bedrock::BedrockProvider;
let handle = tokio::runtime::Handle::try_current().map_err(|_| {
LlmError::ConfigError("Bedrock provider requires a Tokio runtime".to_string())
})?;
let provider =
tokio::task::block_in_place(|| handle.block_on(BedrockProvider::from_env()))
.map_err(|e| {
LlmError::ConfigError(format!(
"Failed to initialize Bedrock provider: {e}"
))
})?;
Ok(Arc::new(
provider.with_model(model).with_application_context(ctx),
))
}
}
}
pub fn create_llm_provider_with_headers(
provider_name: &str,
model: &str,
headers: impl IntoIterator<Item = (String, String)>,
) -> Result<Arc<dyn LLMProvider>> {
let headers_vec: Vec<(String, String)> = headers.into_iter().collect();
if headers_vec.is_empty() {
return Self::create_llm_provider(provider_name, model);
}
let ctx = ApplicationContext {
extra_headers: headers_vec.into_iter().collect(),
..Default::default()
};
Self::create_llm_provider_with_context(provider_name, model, ctx)
}
pub fn create_with_context(
provider: ProviderType,
model: Option<&str>,
ctx: ApplicationContext,
) -> Result<Arc<dyn LLMProvider>> {
let model_name = model.unwrap_or("default");
Self::create_llm_provider_with_context(provider.canonical_id(), model_name, ctx)
}
}
#[cfg(test)]
mod tests {
use super::*;
use serial_test::serial;
#[test]
fn test_provider_type_parsing() {
assert_eq!(ProviderType::from_str("openai"), Some(ProviderType::OpenAI));
assert_eq!(ProviderType::from_str("OLLAMA"), Some(ProviderType::Ollama));
assert_eq!(
ProviderType::from_str("lmstudio"),
Some(ProviderType::LMStudio)
);
assert_eq!(
ProviderType::from_str("lm-studio"),
Some(ProviderType::LMStudio)
);
assert_eq!(
ProviderType::from_str("lm_studio"),
Some(ProviderType::LMStudio)
);
assert_eq!(ProviderType::from_str("mock"), Some(ProviderType::Mock));
assert_eq!(ProviderType::from_str("gemini"), Some(ProviderType::Gemini));
assert_eq!(ProviderType::from_str("google"), Some(ProviderType::Gemini));
assert_eq!(
ProviderType::from_str("vertex"),
Some(ProviderType::VertexAI)
);
assert_eq!(
ProviderType::from_str("vertexai"),
Some(ProviderType::VertexAI)
);
assert_eq!(
ProviderType::from_str("openrouter"),
Some(ProviderType::OpenRouter)
);
assert_eq!(ProviderType::from_str("xai"), Some(ProviderType::XAI));
assert_eq!(ProviderType::from_str("grok"), Some(ProviderType::XAI));
assert_eq!(
ProviderType::from_str("huggingface"),
Some(ProviderType::HuggingFace)
);
assert_eq!(
ProviderType::from_str("hf"),
Some(ProviderType::HuggingFace)
);
assert_eq!(
ProviderType::from_str("hugging-face"),
Some(ProviderType::HuggingFace)
);
assert_eq!(
ProviderType::from_str("azure"),
Some(ProviderType::AzureOpenAI)
);
assert_eq!(
ProviderType::from_str("azure-openai"),
Some(ProviderType::AzureOpenAI)
);
assert_eq!(
ProviderType::from_str("azure_openai"),
Some(ProviderType::AzureOpenAI)
);
assert_eq!(
ProviderType::from_str("AZURE"),
Some(ProviderType::AzureOpenAI)
);
assert_eq!(ProviderType::from_str("invalid"), None);
assert_eq!(ProviderType::from_str(""), None);
}
#[test]
fn test_mock_creation() {
let (llm, embedding) = ProviderFactory::create_mock();
assert_eq!(llm.name(), "mock");
assert_eq!(embedding.name(), "mock");
assert_eq!(embedding.dimension(), 1536);
}
#[test]
fn test_explicit_mock_creation() {
let (llm, embedding) = ProviderFactory::create(ProviderType::Mock).unwrap();
assert_eq!(llm.name(), "mock");
assert_eq!(embedding.dimension(), 1536);
}
#[test]
#[serial]
fn test_from_env_fallback_to_mock() {
std::env::remove_var("EDGEQUAKE_LLM_PROVIDER");
std::env::remove_var("OPENAI_API_KEY");
std::env::remove_var("XAI_API_KEY");
std::env::remove_var("GOOGLE_API_KEY");
std::env::remove_var("GEMINI_API_KEY");
std::env::remove_var("OPENROUTER_API_KEY");
std::env::remove_var("ANTHROPIC_API_KEY");
std::env::remove_var("AZURE_OPENAI_API_KEY");
std::env::remove_var("AZURE_OPENAI_ENDPOINT");
std::env::remove_var("AZURE_OPENAI_DEPLOYMENT_NAME");
std::env::remove_var("AZURE_OPENAI_CONTENTGEN_API_KEY");
std::env::remove_var("AZURE_OPENAI_CONTENTGEN_API_ENDPOINT");
std::env::remove_var("AZURE_OPENAI_CONTENTGEN_MODEL_DEPLOYMENT");
std::env::remove_var("HUGGINGFACE_API_KEY"); std::env::remove_var("HF_TOKEN"); std::env::remove_var("HUGGINGFACE_TOKEN"); std::env::remove_var("OLLAMA_HOST");
std::env::remove_var("OLLAMA_MODEL");
std::env::remove_var("LMSTUDIO_HOST");
std::env::remove_var("LMSTUDIO_MODEL");
std::env::remove_var("MISTRAL_API_KEY");
std::env::remove_var("NVIDIA_API_KEY");
let (llm, _) = ProviderFactory::from_env().unwrap();
assert_eq!(llm.name(), "mock");
}
#[test]
#[serial]
fn test_anthropic_auto_detection_uses_auth_token_when_api_key_empty() {
std::env::remove_var("EDGEQUAKE_LLM_PROVIDER");
std::env::remove_var("OLLAMA_HOST");
std::env::remove_var("OLLAMA_MODEL");
std::env::remove_var("LMSTUDIO_HOST");
std::env::remove_var("LMSTUDIO_MODEL");
std::env::remove_var("OPENAI_API_KEY");
std::env::remove_var("XAI_API_KEY");
std::env::remove_var("GOOGLE_API_KEY");
std::env::remove_var("GEMINI_API_KEY");
std::env::remove_var("OPENROUTER_API_KEY");
std::env::remove_var("AZURE_OPENAI_API_KEY");
std::env::set_var("ANTHROPIC_API_KEY", "");
std::env::set_var("ANTHROPIC_AUTH_TOKEN", "poe-token");
let (llm, _) = ProviderFactory::from_env().unwrap();
assert_eq!(llm.name(), "anthropic");
std::env::remove_var("ANTHROPIC_API_KEY");
std::env::remove_var("ANTHROPIC_AUTH_TOKEN");
}
#[test]
#[serial]
fn test_create_anthropic_from_config_uses_auth_token_fallback() {
std::env::remove_var("ANTHROPIC_API_KEY");
std::env::remove_var("ANTHROPIC_AUTH_TOKEN");
std::env::set_var("ANTHROPIC_API_KEY", "");
std::env::set_var("ANTHROPIC_AUTH_TOKEN", "poe-token");
let config = ProviderConfig {
provider_type: ConfigProviderType::Anthropic,
api_key_env: Some("ANTHROPIC_API_KEY".to_string()),
base_url: Some("https://api.poe.com".to_string()),
default_llm_model: Some("claude-sonnet-4-6".to_string()),
..Default::default()
};
let (llm, _) = ProviderFactory::from_config(&config).unwrap();
assert_eq!(llm.name(), "anthropic");
assert_eq!(llm.model(), "claude-sonnet-4-6");
std::env::remove_var("ANTHROPIC_API_KEY");
std::env::remove_var("ANTHROPIC_AUTH_TOKEN");
}
#[test]
#[serial]
fn test_create_llm_provider_anthropic_uses_auth_token_fallback() {
std::env::remove_var("ANTHROPIC_API_KEY");
std::env::remove_var("ANTHROPIC_AUTH_TOKEN");
std::env::set_var("ANTHROPIC_API_KEY", "");
std::env::set_var("ANTHROPIC_AUTH_TOKEN", "poe-token");
let provider = ProviderFactory::create_llm_provider("anthropic", "claude-haiku-4-5")
.expect("anthropic LLM provider should use auth token fallback");
assert_eq!(provider.name(), "anthropic");
assert_eq!(provider.model(), "claude-haiku-4-5");
std::env::remove_var("ANTHROPIC_API_KEY");
std::env::remove_var("ANTHROPIC_AUTH_TOKEN");
}
#[test]
#[serial]
fn test_explicit_provider_env() {
std::env::remove_var("EDGEQUAKE_LLM_PROVIDER");
std::env::remove_var("OLLAMA_HOST");
std::env::remove_var("OPENAI_API_KEY");
std::env::remove_var("XAI_API_KEY");
std::env::remove_var("GOOGLE_API_KEY");
std::env::remove_var("GEMINI_API_KEY");
std::env::remove_var("OPENROUTER_API_KEY");
std::env::remove_var("ANTHROPIC_API_KEY");
std::env::remove_var("AZURE_OPENAI_API_KEY");
std::env::remove_var("LMSTUDIO_HOST");
std::env::set_var("EDGEQUAKE_LLM_PROVIDER", "mock");
let (llm, _) = ProviderFactory::from_env().unwrap();
assert_eq!(llm.name(), "mock");
std::env::remove_var("EDGEQUAKE_LLM_PROVIDER");
}
#[test]
#[serial]
fn test_lmstudio_auto_detection() {
std::env::remove_var("EDGEQUAKE_LLM_PROVIDER");
std::env::remove_var("OLLAMA_HOST");
std::env::remove_var("OLLAMA_MODEL");
std::env::remove_var("OPENAI_API_KEY");
std::env::remove_var("XAI_API_KEY");
std::env::remove_var("GOOGLE_API_KEY");
std::env::remove_var("GEMINI_API_KEY");
std::env::remove_var("OPENROUTER_API_KEY");
std::env::remove_var("ANTHROPIC_API_KEY");
std::env::remove_var("AZURE_OPENAI_API_KEY");
std::env::set_var("LMSTUDIO_HOST", "http://localhost:1234");
let (llm, embedding) = ProviderFactory::from_env().unwrap();
assert_eq!(llm.name(), "lmstudio");
assert_eq!(embedding.name(), "lmstudio");
std::env::remove_var("LMSTUDIO_HOST");
}
#[test]
#[serial]
fn test_lmstudio_model_detection() {
std::env::remove_var("EDGEQUAKE_LLM_PROVIDER");
std::env::remove_var("OLLAMA_HOST");
std::env::remove_var("OLLAMA_MODEL");
std::env::remove_var("OPENAI_API_KEY");
std::env::remove_var("XAI_API_KEY");
std::env::remove_var("GOOGLE_API_KEY");
std::env::remove_var("GEMINI_API_KEY");
std::env::remove_var("OPENROUTER_API_KEY");
std::env::remove_var("ANTHROPIC_API_KEY");
std::env::remove_var("AZURE_OPENAI_API_KEY");
std::env::remove_var("LMSTUDIO_HOST");
std::env::set_var("LMSTUDIO_MODEL", "mistral-7b");
let (llm, _) = ProviderFactory::from_env().unwrap();
assert_eq!(llm.name(), "lmstudio");
std::env::remove_var("LMSTUDIO_MODEL");
}
#[test]
fn test_explicit_lmstudio_creation() {
let (llm, embedding) = ProviderFactory::create(ProviderType::LMStudio).unwrap();
assert_eq!(llm.name(), "lmstudio");
assert_eq!(embedding.name(), "lmstudio");
assert_eq!(embedding.dimension(), 768);
}
#[test]
#[serial]
fn test_invalid_provider_env() {
std::env::remove_var("EDGEQUAKE_LLM_PROVIDER");
std::env::remove_var("OLLAMA_HOST");
std::env::remove_var("OPENAI_API_KEY");
std::env::remove_var("XAI_API_KEY");
std::env::remove_var("GOOGLE_API_KEY");
std::env::remove_var("GEMINI_API_KEY");
std::env::remove_var("OPENROUTER_API_KEY");
std::env::remove_var("ANTHROPIC_API_KEY");
std::env::remove_var("AZURE_OPENAI_API_KEY");
std::env::remove_var("LMSTUDIO_HOST");
std::env::set_var("EDGEQUAKE_LLM_PROVIDER", "invalid_provider");
let result = ProviderFactory::from_env();
assert!(result.is_err());
if let Err(e) = result {
assert!(e.to_string().contains("Unknown provider type"));
}
std::env::remove_var("EDGEQUAKE_LLM_PROVIDER");
}
#[test]
#[serial]
fn test_openai_creation_requires_api_key() {
std::env::remove_var("OPENAI_API_KEY");
std::env::remove_var("EDGEQUAKE_LLM_PROVIDER");
std::env::remove_var("OLLAMA_HOST");
std::env::remove_var("LMSTUDIO_HOST");
let result = ProviderFactory::create(ProviderType::OpenAI);
assert!(result.is_err());
if let Err(e) = result {
assert!(e.to_string().contains("OPENAI_API_KEY not set"));
}
}
#[test]
#[serial]
fn test_embedding_dimension_detection() {
std::env::remove_var("EDGEQUAKE_LLM_PROVIDER");
std::env::remove_var("OLLAMA_HOST");
std::env::remove_var("OPENAI_API_KEY");
std::env::remove_var("LMSTUDIO_HOST");
std::env::set_var("EDGEQUAKE_LLM_PROVIDER", "mock");
let dim = ProviderFactory::embedding_dimension().unwrap();
assert_eq!(dim, 1536);
std::env::remove_var("EDGEQUAKE_LLM_PROVIDER");
}
#[test]
#[serial]
fn test_provider_priority_ollama_over_lmstudio() {
std::env::remove_var("EDGEQUAKE_LLM_PROVIDER");
std::env::remove_var("OPENAI_API_KEY");
std::env::remove_var("LMSTUDIO_HOST");
std::env::remove_var("LMSTUDIO_MODEL");
std::env::set_var("OLLAMA_HOST", "http://localhost:11434");
std::env::set_var("LMSTUDIO_HOST", "http://localhost:1234");
let (llm, _) = ProviderFactory::from_env().unwrap();
assert_eq!(llm.name(), "ollama");
std::env::remove_var("OLLAMA_HOST");
std::env::remove_var("LMSTUDIO_HOST");
}
#[test]
fn test_create_with_model_mock_none() {
let (llm, emb) = ProviderFactory::create_with_model(ProviderType::Mock, None).unwrap();
assert_eq!(llm.name(), "mock");
assert_eq!(emb.name(), "mock");
}
#[test]
fn test_create_with_model_mock_some() {
let (llm, _) =
ProviderFactory::create_with_model(ProviderType::Mock, Some("any-model")).unwrap();
assert_eq!(llm.name(), "mock");
}
#[test]
fn test_create_embedding_provider_mock() {
let provider =
ProviderFactory::create_embedding_provider("mock", "mock-model", 1536).unwrap();
assert_eq!(provider.name(), "mock");
assert_eq!(provider.dimension(), 1536);
}
#[test]
fn test_create_embedding_provider_unknown() {
let result = ProviderFactory::create_embedding_provider("unknown", "model", 1536);
match result {
Err(e) => assert!(e.to_string().contains("Unknown embedding provider")),
Ok(_) => panic!("Expected error for unknown provider"),
}
}
#[test]
fn test_create_llm_provider_mock() {
let provider = ProviderFactory::create_llm_provider("mock", "mock-model").unwrap();
assert_eq!(provider.name(), "mock");
}
#[test]
fn test_create_llm_provider_unknown() {
let result = ProviderFactory::create_llm_provider("unknown", "model");
match result {
Err(e) => assert!(e.to_string().contains("Unknown LLM provider")),
Ok(_) => panic!("Expected error for unknown provider"),
}
}
#[test]
fn test_provider_type_debug() {
let pt = ProviderType::Mock;
let debug = format!("{:?}", pt);
assert_eq!(debug, "Mock");
}
#[test]
fn test_provider_type_clone_eq() {
let pt1 = ProviderType::OpenAI;
let pt2 = pt1;
assert_eq!(pt1, pt2);
assert_ne!(pt1, ProviderType::Ollama);
}
#[test]
#[serial]
fn test_from_config_azure_no_creds() {
use crate::model_config::{ProviderConfig, ProviderType as ConfigProviderType};
std::env::set_var("AZURE_OPENAI_CONTENTGEN_API_KEY", "");
std::env::set_var("AZURE_OPENAI_CONTENTGEN_API_ENDPOINT", "");
std::env::set_var("AZURE_OPENAI_CONTENTGEN_MODEL_DEPLOYMENT", "");
std::env::set_var("AZURE_OPENAI_API_KEY", "");
std::env::set_var("AZURE_OPENAI_ENDPOINT", "");
std::env::set_var("AZURE_OPENAI_DEPLOYMENT_NAME", "");
let config = ProviderConfig {
provider_type: ConfigProviderType::Azure,
..ProviderConfig::default()
};
let result = ProviderFactory::from_config(&config);
std::env::remove_var("AZURE_OPENAI_CONTENTGEN_API_KEY");
std::env::remove_var("AZURE_OPENAI_CONTENTGEN_API_ENDPOINT");
std::env::remove_var("AZURE_OPENAI_CONTENTGEN_MODEL_DEPLOYMENT");
std::env::remove_var("AZURE_OPENAI_API_KEY");
std::env::remove_var("AZURE_OPENAI_ENDPOINT");
std::env::remove_var("AZURE_OPENAI_DEPLOYMENT_NAME");
assert!(
result.is_err(),
"Expected error when Azure credentials are not set"
);
}
#[test]
fn test_from_config_mock() {
use crate::model_config::{ProviderConfig, ProviderType as ConfigProviderType};
let config = ProviderConfig {
provider_type: ConfigProviderType::Mock,
..ProviderConfig::default()
};
let (llm, emb) = ProviderFactory::from_config(&config).unwrap();
assert_eq!(llm.name(), "mock");
assert_eq!(emb.name(), "mock");
}
#[test]
fn test_vscode_copilot_parsing() {
assert_eq!(
ProviderType::from_str("vscode"),
Some(ProviderType::VsCodeCopilot)
);
assert_eq!(
ProviderType::from_str("copilot"),
Some(ProviderType::VsCodeCopilot)
);
assert_eq!(
ProviderType::from_str("vscode-copilot"),
Some(ProviderType::VsCodeCopilot)
);
}
#[test]
fn test_vscode_copilot_default_model_is_auto() {
std::env::remove_var("VSCODE_COPILOT_MODEL");
let (llm, embedding) = ProviderFactory::create(ProviderType::VsCodeCopilot).unwrap();
assert_eq!(llm.model(), "auto");
assert_eq!(embedding.name(), "vscode-copilot");
}
#[test]
fn test_create_embedding_provider_anthropic_fallback() {
let provider =
ProviderFactory::create_embedding_provider("anthropic", "any-model", 1536).unwrap();
assert_eq!(provider.name(), "mock");
}
#[test]
fn test_create_embedding_provider_openrouter_fallback() {
let provider =
ProviderFactory::create_embedding_provider("openrouter", "any-model", 1536).unwrap();
assert_eq!(provider.name(), "mock");
}
#[test]
fn test_create_embedding_provider_xai_fallback() {
let provider =
ProviderFactory::create_embedding_provider("xai", "any-model", 1536).unwrap();
assert_eq!(provider.name(), "mock");
}
#[test]
fn test_create_embedding_provider_huggingface_fallback() {
let provider =
ProviderFactory::create_embedding_provider("huggingface", "any-model", 768).unwrap();
assert_eq!(provider.name(), "mock");
}
#[test]
fn test_create_embedding_provider_gemini_fallback() {
let provider =
ProviderFactory::create_embedding_provider("gemini", "any-model", 768).unwrap();
let name = provider.name();
assert!(
name == "gemini" || name == "mock",
"Expected 'gemini' (with API key) or 'mock' (without), got '{}'",
name
);
}
#[test]
fn test_create_embedding_provider_ollama() {
let provider =
ProviderFactory::create_embedding_provider("ollama", "nomic-embed-text", 768).unwrap();
assert_eq!(provider.name(), "ollama");
}
#[test]
fn test_create_embedding_provider_lmstudio() {
let provider =
ProviderFactory::create_embedding_provider("lmstudio", "nomic-embed-text-v1.5", 768)
.unwrap();
assert_eq!(provider.name(), "lmstudio");
}
#[test]
fn test_create_embedding_provider_vscode_copilot() {
let provider = ProviderFactory::create_embedding_provider(
"vscode-copilot",
"text-embedding-3-small",
1536,
)
.unwrap();
assert_eq!(provider.name(), "vscode-copilot");
}
#[test]
fn test_from_config_ollama() {
use crate::model_config::{ProviderConfig, ProviderType as ConfigProviderType};
let config = ProviderConfig {
provider_type: ConfigProviderType::Ollama,
..ProviderConfig::default()
};
let (llm, emb) = ProviderFactory::from_config(&config).unwrap();
assert_eq!(llm.name(), "ollama");
assert_eq!(emb.name(), "ollama");
}
#[test]
fn test_from_config_lmstudio() {
use crate::model_config::{ProviderConfig, ProviderType as ConfigProviderType};
let config = ProviderConfig {
provider_type: ConfigProviderType::LMStudio,
..ProviderConfig::default()
};
let (llm, emb) = ProviderFactory::from_config(&config).unwrap();
assert_eq!(llm.name(), "lmstudio");
assert_eq!(emb.name(), "lmstudio");
}
#[test]
#[serial]
fn test_from_config_openai_requires_api_key() {
use crate::model_config::{ProviderConfig, ProviderType as ConfigProviderType};
std::env::remove_var("OPENAI_API_KEY");
let config = ProviderConfig {
provider_type: ConfigProviderType::OpenAI,
..ProviderConfig::default()
};
let result = ProviderFactory::from_config(&config);
assert!(result.is_err());
if let Err(e) = result {
assert!(e.to_string().contains("OPENAI_API_KEY"));
}
}
#[test]
fn test_create_with_model_ollama() {
let (llm, _) =
ProviderFactory::create_with_model(ProviderType::Ollama, Some("llama3:8b")).unwrap();
assert_eq!(llm.name(), "ollama");
assert_eq!(llm.model(), "llama3:8b");
}
#[test]
fn test_create_with_model_lmstudio() {
let (llm, _) =
ProviderFactory::create_with_model(ProviderType::LMStudio, Some("mistral-7b")).unwrap();
assert_eq!(llm.name(), "lmstudio");
assert_eq!(llm.model(), "mistral-7b");
}
#[test]
#[serial]
fn test_provider_type_parsing_azure() {
assert_eq!(
ProviderType::from_str("azure"),
Some(ProviderType::AzureOpenAI)
);
assert_eq!(
ProviderType::from_str("azure-openai"),
Some(ProviderType::AzureOpenAI)
);
assert_eq!(
ProviderType::from_str("AZURE"),
Some(ProviderType::AzureOpenAI)
);
}
#[test]
#[serial]
fn test_create_azure_openai_fails_without_env() {
std::env::set_var("AZURE_OPENAI_CONTENTGEN_API_KEY", "");
std::env::set_var("AZURE_OPENAI_CONTENTGEN_API_ENDPOINT", "");
std::env::set_var("AZURE_OPENAI_CONTENTGEN_MODEL_DEPLOYMENT", "");
std::env::set_var("AZURE_OPENAI_API_KEY", "");
std::env::set_var("AZURE_OPENAI_ENDPOINT", "");
std::env::set_var("AZURE_OPENAI_DEPLOYMENT_NAME", "");
let result = ProviderFactory::create(ProviderType::AzureOpenAI);
std::env::remove_var("AZURE_OPENAI_CONTENTGEN_API_KEY");
std::env::remove_var("AZURE_OPENAI_CONTENTGEN_API_ENDPOINT");
std::env::remove_var("AZURE_OPENAI_CONTENTGEN_MODEL_DEPLOYMENT");
std::env::remove_var("AZURE_OPENAI_API_KEY");
std::env::remove_var("AZURE_OPENAI_ENDPOINT");
std::env::remove_var("AZURE_OPENAI_DEPLOYMENT_NAME");
assert!(
result.is_err(),
"Azure provider should fail when env vars are empty"
);
}
#[test]
#[serial]
fn test_from_env_auto_detects_azure_with_contentgen_vars() {
std::env::remove_var("EDGEQUAKE_LLM_PROVIDER");
std::env::remove_var("OLLAMA_HOST");
std::env::remove_var("OLLAMA_MODEL");
std::env::remove_var("LMSTUDIO_HOST");
std::env::remove_var("LMSTUDIO_MODEL");
std::env::remove_var("ANTHROPIC_API_KEY");
std::env::remove_var("GEMINI_API_KEY");
std::env::remove_var("GOOGLE_API_KEY");
std::env::remove_var("MISTRAL_API_KEY");
std::env::set_var("AZURE_OPENAI_CONTENTGEN_API_KEY", "test-azure-key");
std::env::set_var(
"AZURE_OPENAI_CONTENTGEN_API_ENDPOINT",
"https://test.openai.azure.com",
);
std::env::set_var("AZURE_OPENAI_CONTENTGEN_MODEL_DEPLOYMENT", "gpt-4o");
let result = ProviderFactory::from_env();
std::env::remove_var("AZURE_OPENAI_CONTENTGEN_API_KEY");
std::env::remove_var("AZURE_OPENAI_CONTENTGEN_API_ENDPOINT");
std::env::remove_var("AZURE_OPENAI_CONTENTGEN_MODEL_DEPLOYMENT");
let (llm, _) = result.expect("Should detect Azure from CONTENTGEN vars");
assert_eq!(llm.name(), "azure-openai");
}
#[test]
#[serial]
fn test_from_env_auto_detects_azure_with_standard_vars() {
std::env::remove_var("EDGEQUAKE_LLM_PROVIDER");
std::env::remove_var("OLLAMA_HOST");
std::env::remove_var("OLLAMA_MODEL");
std::env::remove_var("LMSTUDIO_HOST");
std::env::remove_var("LMSTUDIO_MODEL");
std::env::remove_var("ANTHROPIC_API_KEY");
std::env::remove_var("GEMINI_API_KEY");
std::env::remove_var("GOOGLE_API_KEY");
std::env::remove_var("MISTRAL_API_KEY");
std::env::remove_var("AZURE_OPENAI_CONTENTGEN_API_KEY");
std::env::set_var("AZURE_OPENAI_API_KEY", "test-azure-key");
std::env::set_var("AZURE_OPENAI_ENDPOINT", "https://test.openai.azure.com");
std::env::set_var("AZURE_OPENAI_DEPLOYMENT_NAME", "gpt-4o");
let result = ProviderFactory::from_env();
std::env::remove_var("AZURE_OPENAI_API_KEY");
std::env::remove_var("AZURE_OPENAI_ENDPOINT");
std::env::remove_var("AZURE_OPENAI_DEPLOYMENT_NAME");
let (llm, _) = result.expect("Should detect Azure from standard vars");
assert_eq!(llm.name(), "azure-openai");
}
#[test]
#[serial]
fn test_explicit_azure_provider_selection() {
std::env::remove_var("EDGEQUAKE_LLM_PROVIDER");
std::env::remove_var("OLLAMA_HOST");
std::env::remove_var("LMSTUDIO_HOST");
std::env::remove_var("OPENAI_API_KEY");
std::env::set_var("AZURE_OPENAI_CONTENTGEN_API_KEY", "test-key");
std::env::set_var(
"AZURE_OPENAI_CONTENTGEN_API_ENDPOINT",
"https://test.openai.azure.com",
);
std::env::set_var("AZURE_OPENAI_CONTENTGEN_MODEL_DEPLOYMENT", "gpt-4.1-mini");
std::env::set_var("EDGEQUAKE_LLM_PROVIDER", "azure");
let result = ProviderFactory::from_env();
std::env::remove_var("EDGEQUAKE_LLM_PROVIDER");
std::env::remove_var("AZURE_OPENAI_CONTENTGEN_API_KEY");
std::env::remove_var("AZURE_OPENAI_CONTENTGEN_API_ENDPOINT");
std::env::remove_var("AZURE_OPENAI_CONTENTGEN_MODEL_DEPLOYMENT");
let (llm, _) = result.expect("Explicit azure provider selection should succeed");
assert_eq!(llm.name(), "azure-openai");
assert_eq!(llm.model(), "gpt-4.1-mini");
}
#[test]
#[serial]
fn test_create_with_model_azure() {
std::env::set_var("AZURE_OPENAI_CONTENTGEN_API_KEY", "test-key");
std::env::set_var(
"AZURE_OPENAI_CONTENTGEN_API_ENDPOINT",
"https://test.openai.azure.com",
);
std::env::set_var("AZURE_OPENAI_CONTENTGEN_MODEL_DEPLOYMENT", "gpt-4o");
let result = ProviderFactory::create_with_model(
ProviderType::AzureOpenAI,
Some("my-custom-deployment"),
);
std::env::remove_var("AZURE_OPENAI_CONTENTGEN_API_KEY");
std::env::remove_var("AZURE_OPENAI_CONTENTGEN_API_ENDPOINT");
std::env::remove_var("AZURE_OPENAI_CONTENTGEN_MODEL_DEPLOYMENT");
let (llm, _) = result.expect("Azure create_with_model should succeed");
assert_eq!(llm.name(), "azure-openai");
assert_eq!(llm.model(), "my-custom-deployment");
}
#[test]
fn test_vertex_provider_type_is_distinct_from_gemini() {
assert_eq!(
ProviderType::from_str("vertex"),
Some(ProviderType::VertexAI)
);
assert_eq!(
ProviderType::from_str("vertexai"),
Some(ProviderType::VertexAI)
);
assert_eq!(
ProviderType::from_str("VERTEXAI"),
Some(ProviderType::VertexAI)
);
assert_eq!(ProviderType::from_str("gemini"), Some(ProviderType::Gemini));
assert_eq!(ProviderType::from_str("google"), Some(ProviderType::Gemini));
assert_ne!(ProviderType::VertexAI, ProviderType::Gemini);
}
#[test]
#[serial]
fn test_create_llm_provider_vertexai_uses_vertex_auth() {
std::env::remove_var("GOOGLE_CLOUD_PROJECT");
std::env::remove_var("GOOGLE_ACCESS_TOKEN");
let result = ProviderFactory::create_llm_provider("vertexai", "gemini-2.5-flash");
let msg = result
.err()
.expect("Expected error without Vertex AI credentials")
.to_string();
assert!(
msg.contains("GOOGLE_CLOUD_PROJECT"),
"VertexAI error must mention GOOGLE_CLOUD_PROJECT, got: {msg}"
);
}
#[test]
#[serial]
fn test_create_llm_provider_vertex_alias_same_as_vertexai() {
std::env::remove_var("GOOGLE_CLOUD_PROJECT");
std::env::remove_var("GOOGLE_ACCESS_TOKEN");
let result = ProviderFactory::create_llm_provider("vertex", "gemini-2.5-flash");
let msg = result
.err()
.expect("Expected error from vertex alias")
.to_string();
assert!(
msg.contains("GOOGLE_CLOUD_PROJECT"),
"\"vertex\" alias must route to VertexAI arm: {msg}"
);
}
#[test]
#[serial]
fn test_vertexai_prefix_stripped_on_model_name() {
std::env::remove_var("GOOGLE_CLOUD_PROJECT");
std::env::remove_var("GOOGLE_ACCESS_TOKEN");
let result = ProviderFactory::create_llm_provider("vertexai", "vertexai:gemini-2.5-flash");
let msg = result
.err()
.expect("Expected error for prefixed model")
.to_string();
assert!(
msg.contains("GOOGLE_CLOUD_PROJECT"),
"vertexai: prefix must be stripped before auth check: {msg}"
);
}
#[test]
#[serial]
fn test_vertexai_ignores_gemini_api_key() {
std::env::set_var("GEMINI_API_KEY", "fake-key-should-not-satisfy-vertexai");
std::env::remove_var("GOOGLE_CLOUD_PROJECT");
std::env::remove_var("GOOGLE_ACCESS_TOKEN");
let result = ProviderFactory::create_llm_provider("vertexai", "gemini-2.5-flash");
std::env::remove_var("GEMINI_API_KEY");
let msg = result
.err()
.expect("VertexAI must fail: GEMINI_API_KEY alone is not sufficient")
.to_string();
assert!(
msg.contains("GOOGLE_CLOUD_PROJECT"),
"VertexAI must require GOOGLE_CLOUD_PROJECT even when GEMINI_API_KEY is set: {msg}"
);
}
#[test]
fn test_require_app_id_rejects_unsupported_provider() {
let ctx = crate::application_context::ApplicationContextBuilder::new()
.app_id("my-backend")
.build()
.unwrap();
let err = ProviderFactory::create_llm_provider_with_context_policy(
"vscode-copilot",
"default",
ctx,
AttributionPolicy::RequireAppId,
);
assert!(matches!(err, Err(LlmError::AttributionError(_))));
}
}