pub mod base;
#[cfg(feature = "providers-extended")]
#[cfg_attr(
not(test),
deprecated(since = "0.6.0", note = "use catalog amazon_nova before 0.7")
)]
pub mod amazon_nova;
pub mod anthropic;
#[cfg(feature = "providers-extra")]
pub mod azure;
#[cfg(feature = "providers-extra")]
pub mod azure_ai;
pub mod bedrock;
pub mod cloudflare;
#[cfg(feature = "providers-extended")]
pub mod cohere;
#[cfg(feature = "providers-extended")]
#[cfg_attr(
not(test),
deprecated(
since = "0.6.0",
note = "use a catalog or typed provider before 0.7.0 removal"
)
)]
pub mod custom_api;
#[cfg(feature = "providers-extended")]
pub mod fal_ai;
#[cfg(any(feature = "providers-extended", feature = "providers-extra"))]
pub mod gemini;
#[cfg(feature = "providers-extended")]
pub mod github;
#[cfg(feature = "providers-extended")]
pub mod github_copilot;
pub(crate) mod google_tool_loop;
#[cfg(feature = "providers-extra")]
pub mod meta_llama;
pub mod mistral;
#[cfg(feature = "providers-extended")]
pub mod ollama;
pub mod openai;
pub mod openai_like;
#[cfg(feature = "providers-extended")]
pub mod replicate;
#[cfg(feature = "providers-extra")]
pub mod v0;
#[cfg(feature = "providers-extra")]
pub mod vertex_ai;
pub mod macros; pub mod shared; pub mod thinking; pub mod provider_type;
pub use provider_type::ProviderType;
pub mod factory;
pub use factory::{create_provider, is_provider_selector_supported};
pub mod contextual_error;
pub mod failure;
pub mod provider_error_conversions;
pub mod provider_registry;
pub mod registry; pub mod unified_provider;
#[cfg(test)]
mod unified_provider_tests;
pub use crate::core::traits::provider::llm_provider::trait_definition::LLMProvider;
use crate::core::types::responses::{
ChatChunk, ChatResponse, EmbeddingResponse, ImageGenerationResponse,
};
use crate::core::types::{
chat::ChatRequest, embedding::EmbeddingRequest, image::ImageGenerationRequest,
};
use crate::core::types::{context::RequestContext, model::ProviderCapability};
pub use contextual_error::ContextualError;
pub use failure::{ProviderFailureFacts, ProviderRetryHint};
pub use provider_registry::ProviderRegistry;
pub use unified_provider::ProviderError;
#[derive(Debug, Clone)]
pub(crate) struct GeminiNativeRequest {
pub(crate) api_version: String,
pub(crate) model: String,
pub(crate) method: &'static str,
pub(crate) stream: bool,
pub(crate) body: serde_json::Value,
}
pub(crate) fn gemini_native_url(
base_url: &str,
api_key: &str,
request: &GeminiNativeRequest,
) -> Result<reqwest::Url, ProviderError> {
if !matches!(request.api_version.as_str(), "v1" | "v1beta")
|| !matches!(request.method, "generateContent" | "streamGenerateContent")
|| request.model.is_empty()
|| !request.model.chars().all(|character| {
character.is_ascii_alphanumeric() || matches!(character, '-' | '_' | '.')
})
{
return Err(ProviderError::invalid_request(
"gemini_proxy",
"invalid Gemini native route segment",
));
}
let mut url = reqwest::Url::parse(&format!(
"{}/{}/models/{}:{}",
base_url.trim_end_matches('/'),
request.api_version,
request.model,
request.method
))
.map_err(|_| ProviderError::configuration("gemini_proxy", "invalid Gemini API base URL"))?;
let mut query = url.query_pairs_mut();
if request.stream {
query.append_pair("alt", "sse");
}
query.append_pair("key", api_key);
drop(query);
Ok(url)
}
pub(crate) async fn gemini_response_or_provider_error(
response: reqwest::Response,
api_key: &str,
) -> Result<reqwest::Response, ProviderError> {
if response.status().is_success() {
return Ok(response);
}
let status = response.status().as_u16();
let header_retry = response
.headers()
.get(reqwest::header::RETRY_AFTER)
.and_then(|value| value.to_str().ok())
.and_then(|value| value.trim().parse::<u64>().ok());
let body = base::read_streaming_error_body(response)
.await
.map_err(|error| error.into_provider_error("gemini_proxy"))?;
let body = redact_gemini_key(&body, api_key);
let message = if body.trim().is_empty() {
format!("Gemini upstream returned HTTP {status}")
} else {
format!("Gemini upstream returned HTTP {status}: {body}")
};
Err(if status == 429 {
ProviderError::rate_limit_with_retry(
"gemini_proxy",
message,
header_retry.or_else(|| shared::parse_retry_after_from_body(&body)),
)
} else {
ProviderError::api_error("gemini_proxy", status, message)
})
}
fn redact_gemini_key(body: &str, api_key: &str) -> String {
if api_key.is_empty() {
return body.to_string();
}
let encoded: String = url::form_urlencoded::byte_serialize(api_key.as_bytes()).collect();
body.replace(api_key, "[REDACTED]")
.replace(&encoded, "[REDACTED]")
}
pub(crate) fn gemini_transport_error(is_timeout: bool) -> ProviderError {
let message = "Gemini upstream request failed";
if is_timeout {
return ProviderError::timeout("gemini_proxy", message);
}
ProviderError::network("gemini_proxy", message)
}
macro_rules! dispatch_provider {
(sync, $self:expr, $method:ident) => {
dispatch_provider!(@expand sync, $self, $method,)
};
(sync, $self:expr, $method:ident, $($arg:expr),+ $(,)?) => {
dispatch_provider!(@expand sync, $self, $method, $($arg),+)
};
(async_err, $self:expr, $method:ident $(, $arg:expr)* $(,)?) => {
dispatch_provider!(@expand async_err, $self, $method, $($arg),*)
};
(value, $self:expr, $method:ident) => {
dispatch_provider!(@expand value, $self, $method,)
};
(value, $self:expr, $method:ident, $($arg:expr),+ $(,)?) => {
dispatch_provider!(@expand value, $self, $method, $($arg),+)
};
(async_direct, $self:expr, $method:ident) => {
dispatch_provider!(@expand async_direct, $self, $method,)
};
(@expand sync, $self:expr, $method:ident, $($arg:expr),*) => {
match $self {
Provider::OpenAI(p) => p.$method($($arg),*),
Provider::Anthropic(p) => p.$method($($arg),*),
Provider::Bedrock(p) => p.$method($($arg),*),
#[cfg(feature = "providers-extra")]
Provider::Azure(p) => p.$method($($arg),*),
#[cfg(feature = "providers-extra")]
Provider::AzureAI(p) => p.$method($($arg),*),
#[cfg(feature = "providers-extra")]
Provider::VertexAI(p) => p.$method($($arg),*),
#[cfg(feature = "providers-extended")]
Provider::Gemini(p) => p.$method($($arg),*),
#[cfg(feature = "providers-extended")]
Provider::GitHubCopilot(p) => p.$method($($arg),*),
#[cfg(feature = "providers-extended")]
Provider::Ollama(p) => p.$method($($arg),*),
#[cfg(feature = "providers-extended")]
Provider::FalAI(p) => p.$method($($arg),*),
Provider::Mistral(p) => p.$method($($arg),*),
Provider::Cloudflare(p) => p.$method($($arg),*),
#[cfg(feature = "providers-extended")]
Provider::Cohere(p) => p.$method($($arg),*),
#[cfg(feature = "providers-extended")]
Provider::Replicate(p) => p.$method($($arg),*),
Provider::OpenAILike(p) => p.$method($($arg),*),
}
};
(@expand async_err, $self:expr, $method:ident, $($arg:expr),*) => {
match $self {
Provider::OpenAI(p) => LLMProvider::$method(p, $($arg),*).await.map_err(ProviderError::from),
Provider::Anthropic(p) => LLMProvider::$method(p, $($arg),*).await.map_err(ProviderError::from),
Provider::Bedrock(p) => LLMProvider::$method(p, $($arg),*).await.map_err(ProviderError::from),
#[cfg(feature = "providers-extra")]
Provider::Azure(p) => LLMProvider::$method(p, $($arg),*).await.map_err(ProviderError::from),
#[cfg(feature = "providers-extra")]
Provider::AzureAI(p) => LLMProvider::$method(p, $($arg),*).await.map_err(ProviderError::from),
#[cfg(feature = "providers-extra")]
Provider::VertexAI(p) => LLMProvider::$method(p, $($arg),*).await.map_err(ProviderError::from),
#[cfg(feature = "providers-extended")]
Provider::Gemini(p) => LLMProvider::$method(p.as_ref(), $($arg),*).await.map_err(ProviderError::from),
#[cfg(feature = "providers-extended")]
Provider::GitHubCopilot(p) => LLMProvider::$method(p, $($arg),*).await.map_err(ProviderError::from),
#[cfg(feature = "providers-extended")]
Provider::Ollama(p) => LLMProvider::$method(p, $($arg),*).await.map_err(ProviderError::from),
#[cfg(feature = "providers-extended")]
Provider::FalAI(p) => LLMProvider::$method(p, $($arg),*).await.map_err(ProviderError::from),
Provider::Mistral(p) => LLMProvider::$method(p, $($arg),*).await.map_err(ProviderError::from),
Provider::Cloudflare(p) => LLMProvider::$method(p, $($arg),*).await.map_err(ProviderError::from),
#[cfg(feature = "providers-extended")]
Provider::Cohere(p) => LLMProvider::$method(p, $($arg),*).await.map_err(ProviderError::from),
#[cfg(feature = "providers-extended")]
Provider::Replicate(p) => LLMProvider::$method(p, $($arg),*).await.map_err(ProviderError::from),
Provider::OpenAILike(p) => LLMProvider::$method(p, $($arg),*).await.map_err(ProviderError::from),
}
};
(@expand value, $self:expr, $method:ident, $($arg:expr),*) => {
match $self {
Provider::OpenAI(p) => LLMProvider::$method(p, $($arg),*),
Provider::Anthropic(p) => LLMProvider::$method(p, $($arg),*),
Provider::Bedrock(p) => LLMProvider::$method(p, $($arg),*),
#[cfg(feature = "providers-extra")]
Provider::Azure(p) => LLMProvider::$method(p, $($arg),*),
#[cfg(feature = "providers-extra")]
Provider::AzureAI(p) => LLMProvider::$method(p, $($arg),*),
#[cfg(feature = "providers-extra")]
Provider::VertexAI(p) => LLMProvider::$method(p, $($arg),*),
#[cfg(feature = "providers-extended")]
Provider::Gemini(p) => LLMProvider::$method(p.as_ref(), $($arg),*),
#[cfg(feature = "providers-extended")]
Provider::GitHubCopilot(p) => LLMProvider::$method(p, $($arg),*),
#[cfg(feature = "providers-extended")]
Provider::Ollama(p) => LLMProvider::$method(p, $($arg),*),
#[cfg(feature = "providers-extended")]
Provider::FalAI(p) => LLMProvider::$method(p, $($arg),*),
Provider::Mistral(p) => LLMProvider::$method(p, $($arg),*),
Provider::Cloudflare(p) => LLMProvider::$method(p, $($arg),*),
#[cfg(feature = "providers-extended")]
Provider::Cohere(p) => LLMProvider::$method(p, $($arg),*),
#[cfg(feature = "providers-extended")]
Provider::Replicate(p) => LLMProvider::$method(p, $($arg),*),
Provider::OpenAILike(p) => LLMProvider::$method(p, $($arg),*),
}
};
(@expand async_direct, $self:expr, $method:ident, $($arg:expr),*) => {
match $self {
Provider::OpenAI(p) => LLMProvider::$method(p).await,
Provider::Anthropic(p) => LLMProvider::$method(p).await,
Provider::Bedrock(p) => LLMProvider::$method(p).await,
#[cfg(feature = "providers-extra")]
Provider::Azure(p) => LLMProvider::$method(p).await,
#[cfg(feature = "providers-extra")]
Provider::AzureAI(p) => LLMProvider::$method(p).await,
#[cfg(feature = "providers-extra")]
Provider::VertexAI(p) => LLMProvider::$method(p).await,
#[cfg(feature = "providers-extended")]
Provider::Gemini(p) => LLMProvider::$method(p.as_ref()).await,
#[cfg(feature = "providers-extended")]
Provider::GitHubCopilot(p) => LLMProvider::$method(p).await,
#[cfg(feature = "providers-extended")]
Provider::Ollama(p) => LLMProvider::$method(p).await,
#[cfg(feature = "providers-extended")]
Provider::FalAI(p) => LLMProvider::$method(p).await,
Provider::Mistral(p) => LLMProvider::$method(p).await,
Provider::Cloudflare(p) => LLMProvider::$method(p).await,
#[cfg(feature = "providers-extended")]
Provider::Cohere(p) => LLMProvider::$method(p).await,
#[cfg(feature = "providers-extended")]
Provider::Replicate(p) => LLMProvider::$method(p).await,
Provider::OpenAILike(p) => LLMProvider::$method(p).await,
}
};
}
#[allow(unused_macros)]
macro_rules! dispatch_provider_selective {
($self:expr, $method:ident, { $($provider:ident),+ }, $default:expr) => {
match $self {
$(Provider::$provider(p) => p.$method()),+,
_ => $default,
}
};
($self:expr, $method:ident($($arg:expr),+), { $($provider:ident),+ }, $default:expr) => {
match $self {
$(Provider::$provider(p) => p.$method($($arg),+)),+,
_ => $default,
}
};
}
mod audio_dispatch;
mod capability_dispatch;
#[derive(Debug, Clone)]
pub enum Provider {
OpenAI(openai::OpenAIProvider),
Anthropic(anthropic::AnthropicProvider),
Bedrock(bedrock::BedrockProvider),
#[cfg(feature = "providers-extra")]
Azure(azure::AzureOpenAIProvider),
#[cfg(feature = "providers-extra")]
AzureAI(azure_ai::AzureAIProvider),
#[cfg(feature = "providers-extra")]
VertexAI(vertex_ai::VertexAIProvider),
#[cfg(feature = "providers-extended")]
Gemini(std::sync::Arc<gemini::GeminiProvider>),
#[cfg(feature = "providers-extended")]
GitHubCopilot(github_copilot::GitHubCopilotProvider),
#[cfg(feature = "providers-extended")]
Ollama(ollama::OllamaProvider),
#[cfg(feature = "providers-extended")]
FalAI(fal_ai::FalAIProvider),
Mistral(mistral::MistralProvider),
Cloudflare(cloudflare::CloudflareProvider),
#[cfg(feature = "providers-extended")]
Cohere(cohere::CohereProvider),
#[cfg(feature = "providers-extended")]
Replicate(replicate::ReplicateProvider),
OpenAILike(openai_like::OpenAILikeProvider),
}
impl Provider {
pub(crate) async fn gemini_generate_content(
&self,
request: GeminiNativeRequest,
) -> Result<reqwest::Response, ProviderError> {
match self {
#[cfg(feature = "providers-extended")]
Provider::Gemini(provider) => provider.gemini_generate_content(request).await,
Provider::OpenAILike(provider) => provider.gemini_generate_content(request).await,
_ => Err(ProviderError::not_supported(
"provider",
"Gemini native generateContent",
)),
}
}
pub fn name(&self) -> &str {
match self {
Provider::OpenAI(p) => {
use crate::core::traits::provider::llm_provider::trait_definition::LLMProvider;
p.name()
}
Provider::Anthropic(_) => "anthropic",
Provider::Bedrock(_) => "bedrock",
#[cfg(feature = "providers-extra")]
Provider::Azure(_) => "azure",
#[cfg(feature = "providers-extra")]
Provider::AzureAI(_) => "azure_ai",
#[cfg(feature = "providers-extra")]
Provider::VertexAI(_) => "vertex_ai",
#[cfg(feature = "providers-extended")]
Provider::Gemini(_) => "gemini",
#[cfg(feature = "providers-extended")]
Provider::GitHubCopilot(_) => "github_copilot",
#[cfg(feature = "providers-extended")]
Provider::Ollama(_) => "ollama",
#[cfg(feature = "providers-extended")]
Provider::FalAI(_) => "fal_ai",
Provider::Mistral(_) => "mistral",
Provider::Cloudflare(_) => "cloudflare",
#[cfg(feature = "providers-extended")]
Provider::Cohere(_) => "cohere",
#[cfg(feature = "providers-extended")]
Provider::Replicate(_) => "replicate",
Provider::OpenAILike(p) => {
use crate::core::traits::provider::llm_provider::trait_definition::LLMProvider;
p.name()
}
}
}
pub fn provider_type(&self) -> ProviderType {
match self {
Provider::OpenAI(_) => ProviderType::OpenAI,
Provider::Anthropic(_) => ProviderType::Anthropic,
Provider::Bedrock(_) => ProviderType::Bedrock,
#[cfg(feature = "providers-extra")]
Provider::Azure(_) => ProviderType::Azure,
#[cfg(feature = "providers-extra")]
Provider::AzureAI(_) => ProviderType::AzureAI,
#[cfg(feature = "providers-extra")]
Provider::VertexAI(_) => ProviderType::VertexAI,
#[cfg(feature = "providers-extended")]
Provider::Gemini(_) => ProviderType::Gemini,
#[cfg(feature = "providers-extended")]
Provider::GitHubCopilot(_) => ProviderType::GitHubCopilot,
#[cfg(feature = "providers-extended")]
Provider::Ollama(_) => ProviderType::Ollama,
#[cfg(feature = "providers-extended")]
Provider::FalAI(_) => ProviderType::FalAI,
Provider::Mistral(_) => ProviderType::Mistral,
Provider::Cloudflare(_) => ProviderType::Cloudflare,
#[cfg(feature = "providers-extended")]
Provider::Cohere(_) => ProviderType::Cohere,
#[cfg(feature = "providers-extended")]
Provider::Replicate(_) => ProviderType::Replicate,
Provider::OpenAILike(_) => ProviderType::OpenAICompatible,
}
}
pub fn factory_supported_provider_types() -> &'static [ProviderType] {
registry::dispatchable_provider_types_slice()
}
pub fn supports_model(&self, model: &str) -> bool {
use crate::core::traits::provider::llm_provider::trait_definition::LLMProvider;
dispatch_provider!(value, self, supports_model, model)
}
pub fn capabilities(&self) -> &'static [ProviderCapability] {
dispatch_provider!(sync, self, capabilities)
}
pub fn supports_capability(&self, capability: &ProviderCapability) -> bool {
use crate::core::traits::provider::llm_provider::trait_definition::LLMProvider;
dispatch_provider!(value, self, supports_capability, capability)
}
pub async fn chat_completion(
&self,
request: ChatRequest,
context: RequestContext,
) -> Result<ChatResponse, ProviderError> {
use crate::core::traits::provider::llm_provider::trait_definition::LLMProvider;
dispatch_provider!(async_err, self, chat_completion, request, context)
}
pub async fn health_check(&self) -> crate::core::types::health::HealthStatus {
use crate::core::traits::provider::llm_provider::trait_definition::LLMProvider;
dispatch_provider!(async_direct, self, health_check)
}
pub fn list_models(&self) -> &[crate::core::types::model::ModelInfo] {
use crate::core::traits::provider::llm_provider::trait_definition::LLMProvider;
dispatch_provider!(value, self, models)
}
pub async fn calculate_cost(
&self,
model: &str,
input_tokens: u32,
output_tokens: u32,
) -> Result<f64, ProviderError> {
use crate::core::traits::provider::llm_provider::trait_definition::LLMProvider;
let model = self.strip_provider_prefix(model);
dispatch_provider!(
async_err,
self,
calculate_cost,
model,
input_tokens,
output_tokens
)
}
fn strip_provider_prefix<'a>(&self, model: &'a str) -> &'a str {
model
.strip_prefix(self.name())
.and_then(|model| model.strip_prefix('/'))
.unwrap_or(model)
}
pub async fn chat_completion_stream(
&self,
request: ChatRequest,
context: RequestContext,
) -> Result<
std::pin::Pin<
Box<dyn futures::Stream<Item = Result<ChatChunk, ProviderError>> + Send + 'static>,
>,
ProviderError,
> {
use crate::core::traits::provider::llm_provider::trait_definition::LLMProvider;
dispatch_provider!(async_err, self, chat_completion_stream, request, context)
}
pub async fn create_embeddings(
&self,
request: EmbeddingRequest,
context: RequestContext,
) -> Result<EmbeddingResponse, ProviderError> {
use crate::core::traits::provider::llm_provider::trait_definition::LLMProvider;
dispatch_provider!(async_err, self, embeddings, request, context)
}
pub async fn create_images(
&self,
request: ImageGenerationRequest,
context: RequestContext,
) -> Result<ImageGenerationResponse, ProviderError> {
use crate::core::traits::provider::llm_provider::trait_definition::LLMProvider;
dispatch_provider!(async_err, self, image_generation, request, context)
}
pub async fn get_model(
&self,
model_id: &str,
) -> Result<Option<crate::core::types::model::ModelInfo>, ProviderError> {
let models = self.list_models();
for model in models {
if model.id == model_id || model.name == model_id {
return Ok(Some(model.clone()));
}
}
Ok(None)
}
}
#[cfg(test)]
#[path = "provider_tests.rs"]
mod tests;