Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
17 changes: 13 additions & 4 deletions src/backends/cohere.rs
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,9 @@
//!
//! This module provides integration with Cohere's LLM models through their API.

use crate::providers::openai_compatible::{OpenAICompatibleProvider, OpenAIProviderConfig};
use crate::providers::openai_compatible::{
OpenAICompatibleProvider, OpenAIProviderConfig, TokenProviderFn,
};
use crate::{
chat::{StructuredOutputFormat, Tool, ToolChoice},
completion::{CompletionProvider, CompletionRequest, CompletionResponse},
Expand Down Expand Up @@ -55,6 +57,8 @@ impl Cohere {
json_schema: Option<StructuredOutputFormat>,
parallel_tool_calls: Option<bool>,
normalize_response: Option<bool>,
headers: Vec<(String, String)>,
token_provider: Option<TokenProviderFn>,
) -> Self {
<OpenAICompatibleProvider<CohereConfig>>::new(
api_key,
Expand All @@ -76,6 +80,8 @@ impl Cohere {
normalize_response,
embedding_encoding_format,
embedding_dimensions,
headers,
token_provider,
)
}
}
Expand Down Expand Up @@ -136,8 +142,10 @@ impl SpeechToTextProvider for Cohere {
#[async_trait]
impl EmbeddingProvider for Cohere {
async fn embed(&self, input: Vec<String>) -> Result<Vec<Vec<f32>>, LLMError> {
if self.config.api_key.is_empty() {
return Err(LLMError::AuthError("Missing Cohere API key".into()));
if self.config.api_key.is_empty() && self.config.token_provider.is_none() {
return Err(LLMError::AuthError(
"Missing Cohere credentials: provide an API key via `.api_key()` or a dynamic token provider via `.auth_provider()`".into(),
));
}

let body = CohereEmbeddingRequest {
Expand All @@ -157,10 +165,11 @@ impl EmbeddingProvider for Cohere {
.join("embeddings")
.map_err(|e| LLMError::HttpError(e.to_string()))?;

let token = self.get_bearer_token().await?;
let resp = self
.client
.post(url)
.bearer_auth(&self.config.api_key)
.bearer_auth(&token)
.json(&body)
.send()
.await?
Expand Down
17 changes: 13 additions & 4 deletions src/backends/groq.rs
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,9 @@ use crate::{
embedding::EmbeddingProvider,
error::LLMError,
models::{ModelListRequest, ModelListResponse, ModelsProvider, StandardModelListResponse},
providers::openai_compatible::{OpenAICompatibleProvider, OpenAIProviderConfig},
providers::openai_compatible::{
OpenAICompatibleProvider, OpenAIProviderConfig, TokenProviderFn,
},
stt::SpeechToTextProvider,
tts::TextToSpeechProvider,
LLMProvider,
Expand Down Expand Up @@ -66,6 +68,8 @@ impl Groq {
json_schema: Option<StructuredOutputFormat>,
parallel_tool_calls: Option<bool>,
normalize_response: Option<bool>,
headers: Vec<(String, String)>,
token_provider: Option<TokenProviderFn>,
) -> Self {
OpenAICompatibleProvider::<GroqConfig>::new(
api_key,
Expand All @@ -87,6 +91,8 @@ impl Groq {
normalize_response,
None, // embedding_encoding_format - not supported by Groq
None, // embedding_dimensions - not supported by Groq
headers,
token_provider,
)
}
}
Expand Down Expand Up @@ -133,16 +139,19 @@ impl ModelsProvider for Groq {
&self,
_request: Option<&ModelListRequest>,
) -> Result<Box<dyn ModelListResponse>, LLMError> {
if self.config.api_key.is_empty() {
return Err(LLMError::AuthError("Missing Groq API key".to_string()));
if self.config.api_key.is_empty() && self.config.token_provider.is_none() {
return Err(LLMError::AuthError(
"Missing Groq credentials: provide an API key via `.api_key()` or a dynamic token provider via `.auth_provider()`".to_string(),
));
}

let url = format!("{}/models", GroqConfig::DEFAULT_BASE_URL);

let token = self.get_bearer_token().await?;
let resp = self
.client
.get(&url)
.bearer_auth(&self.config.api_key)
.bearer_auth(&token)
.send()
.await?
.error_for_status()?;
Expand Down
15 changes: 11 additions & 4 deletions src/backends/huggingface.rs
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,9 @@ use crate::{
embedding::EmbeddingProvider,
error::LLMError,
models::{ModelListRequest, ModelListResponse, ModelsProvider, StandardModelListResponse},
providers::openai_compatible::{OpenAICompatibleProvider, OpenAIProviderConfig},
providers::openai_compatible::{
OpenAICompatibleProvider, OpenAIProviderConfig, TokenProviderFn,
},
stt::SpeechToTextProvider,
tts::TextToSpeechProvider,
LLMProvider,
Expand Down Expand Up @@ -52,6 +54,8 @@ impl HuggingFace {
json_schema: Option<StructuredOutputFormat>,
parallel_tool_calls: Option<bool>,
normalize_response: Option<bool>,
headers: Vec<(String, String)>,
token_provider: Option<TokenProviderFn>,
) -> Self {
OpenAICompatibleProvider::<HuggingFaceConfig>::new(
api_key,
Expand All @@ -73,6 +77,8 @@ impl HuggingFace {
normalize_response,
None, // embedding_encoding_format
None, // embedding_dimensions
headers,
token_provider,
)
}
}
Expand Down Expand Up @@ -119,18 +125,19 @@ impl ModelsProvider for HuggingFace {
&self,
_request: Option<&ModelListRequest>,
) -> Result<Box<dyn ModelListResponse>, LLMError> {
if self.config.api_key.is_empty() {
if self.config.api_key.is_empty() && self.config.token_provider.is_none() {
return Err(LLMError::AuthError(
"Missing HuggingFace API key".to_string(),
"Missing HuggingFace credentials: provide an API key via `.api_key()` or a dynamic token provider via `.auth_provider()`".to_string(),
));
}

let url = format!("{}/models", HuggingFaceConfig::DEFAULT_BASE_URL);

let token = self.get_bearer_token().await?;
let resp = self
.client
.get(&url)
.bearer_auth(&self.config.api_key)
.bearer_auth(&token)
.send()
.await?
.error_for_status()?;
Expand Down
28 changes: 21 additions & 7 deletions src/backends/mistral.rs
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,9 @@

use crate::builder::LLMBackend;
use crate::models::{ModelListRequest, ModelListResponse, StandardModelListResponse};
use crate::providers::openai_compatible::{OpenAICompatibleProvider, OpenAIProviderConfig};
use crate::providers::openai_compatible::{
OpenAICompatibleProvider, OpenAIProviderConfig, TokenProviderFn,
};
use crate::{
chat::{StructuredOutputFormat, Tool, ToolChoice},
completion::{CompletionProvider, CompletionRequest, CompletionResponse},
Expand Down Expand Up @@ -55,6 +57,8 @@ impl Mistral {
json_schema: Option<StructuredOutputFormat>,
parallel_tool_calls: Option<bool>,
normalize_response: Option<bool>,
headers: Vec<(String, String)>,
token_provider: Option<TokenProviderFn>,
) -> Self {
<OpenAICompatibleProvider<MistralConfig>>::new(
api_key,
Expand All @@ -76,6 +80,8 @@ impl Mistral {
normalize_response,
embedding_encoding_format,
embedding_dimensions,
headers,
token_provider,
)
}
}
Expand Down Expand Up @@ -136,8 +142,10 @@ impl SpeechToTextProvider for Mistral {
#[async_trait]
impl EmbeddingProvider for Mistral {
async fn embed(&self, input: Vec<String>) -> Result<Vec<Vec<f32>>, LLMError> {
if self.config.api_key.is_empty() {
return Err(LLMError::AuthError("Missing Mistral API key".into()));
if self.config.api_key.is_empty() && self.config.token_provider.is_none() {
return Err(LLMError::AuthError(
"Missing Mistral credentials: provide an API key via `.api_key()` or a dynamic token provider via `.auth_provider()`".into(),
));
}

let body = MistralEmbeddingRequest {
Expand All @@ -157,10 +165,11 @@ impl EmbeddingProvider for Mistral {
.join("embeddings")
.map_err(|e| LLMError::HttpError(e.to_string()))?;

let token = self.get_bearer_token().await?;
let resp = self
.client
.post(url)
.bearer_auth(&self.config.api_key)
.bearer_auth(&token)
.json(&body)
.send()
.await?
Expand All @@ -178,14 +187,18 @@ impl ModelsProvider for Mistral {
&self,
_request: Option<&ModelListRequest>,
) -> Result<Box<dyn ModelListResponse>, LLMError> {
if self.config.api_key.is_empty() {
return Err(LLMError::AuthError("Missing Mistral API key".to_string()));
if self.config.api_key.is_empty() && self.config.token_provider.is_none() {
return Err(LLMError::AuthError(
"Missing Mistral credentials: provide an API key via `.api_key()` or a dynamic token provider via `.auth_provider()`".to_string(),
));
}
let url = format!("{}models", MistralConfig::DEFAULT_BASE_URL);

let token = self.get_bearer_token().await?;
let resp = self
.client
.get(&url)
.bearer_auth(&self.config.api_key)
.bearer_auth(&token)
.send()
.await?
.error_for_status()?;
Expand All @@ -206,3 +219,4 @@ impl TextToSpeechProvider for Mistral {
))
}
}

48 changes: 28 additions & 20 deletions src/backends/openai.rs
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,7 @@ use crate::builder::LLMBackend;
use crate::chat::Usage;
use crate::providers::openai_compatible::{
OpenAIChatMessage, OpenAICompatibleProvider, OpenAIProviderConfig, OpenAIResponseFormat,
OpenAIStreamOptions,
OpenAIStreamOptions, TokenProviderFn,
};
use crate::{
chat::{
Expand Down Expand Up @@ -220,6 +220,8 @@ impl OpenAI {
json_schema: Option<StructuredOutputFormat>,
voice: Option<String>,
extra_body: Option<serde_json::Value>,
headers: Vec<(String, String)>,
token_provider: Option<TokenProviderFn>,
enable_web_search: Option<bool>,
web_search_context_size: Option<String>,
web_search_user_location_type: Option<String>,
Expand All @@ -228,8 +230,10 @@ impl OpenAI {
web_search_user_location_approximate_region: Option<String>,
) -> Result<Self, LLMError> {
let api_key_str = api_key.into();
if api_key_str.is_empty() {
return Err(LLMError::AuthError("Missing OpenAI API key".to_string()));
if api_key_str.is_empty() && token_provider.is_none() {
return Err(LLMError::AuthError(
"Missing OpenAI credentials: provide an API key via `.api_key()` or a dynamic token provider via `.auth_provider()`".to_string(),
));
}
Ok(OpenAI {
provider: <OpenAICompatibleProvider<OpenAIConfig>>::new(
Expand All @@ -252,6 +256,8 @@ impl OpenAI {
normalize_response,
embedding_encoding_format,
embedding_dimensions,
headers,
token_provider,
),
enable_web_search: enable_web_search.unwrap_or(false),
web_search_context_size,
Expand Down Expand Up @@ -440,12 +446,14 @@ impl SpeechToTextProvider for OpenAI {
.text("response_format", RESPONSE_FORMAT)
.part("file", part);

let token = self.provider.get_bearer_token().await?;
let mut req = self
.provider
.client
.post(url)
.bearer_auth(&self.provider.config.api_key)
.bearer_auth(&token)
.multipart(form);
req = self.provider.apply_headers(req);

if let Some(t) = self.provider.config.timeout_seconds {
req = req.timeout(Duration::from_secs(t));
Expand Down Expand Up @@ -474,12 +482,14 @@ impl SpeechToTextProvider for OpenAI {
.await
.map_err(|e| LLMError::HttpError(e.to_string()))?;

let token = self.provider.get_bearer_token().await?;
let mut req = self
.provider
.client
.post(url)
.bearer_auth(&self.provider.config.api_key)
.bearer_auth(&token)
.multipart(form);
req = self.provider.apply_headers(req);

if let Some(t) = self.provider.config.timeout_seconds {
req = req.timeout(Duration::from_secs(t));
Expand Down Expand Up @@ -526,15 +536,15 @@ impl EmbeddingProvider for OpenAI {
.join("embeddings")
.map_err(|e| LLMError::HttpError(e.to_string()))?;

let resp = self
let token = self.provider.get_bearer_token().await?;
let mut req = self
.provider
.client
.post(url)
.bearer_auth(&self.provider.config.api_key)
.json(&body)
.send()
.await?
.error_for_status()?;
.bearer_auth(&token)
.json(&body);
req = self.provider.apply_headers(req);
let resp = req.send().await?.error_for_status()?;

let json_resp: OpenAIEmbeddingResponse = resp.json().await?;
let embeddings = json_resp.data.into_iter().map(|d| d.embedding).collect();
Expand All @@ -555,14 +565,10 @@ impl ModelsProvider for OpenAI {
.join("models")
.map_err(|e| LLMError::HttpError(e.to_string()))?;

let resp = self
.provider
.client
.get(url)
.bearer_auth(&self.provider.config.api_key)
.send()
.await?
.error_for_status()?;
let token = self.provider.get_bearer_token().await?;
let mut req = self.provider.client.get(url).bearer_auth(&token);
req = self.provider.apply_headers(req);
let resp = req.send().await?.error_for_status()?;

let result = StandardModelListResponse {
inner: resp.json().await?,
Expand Down Expand Up @@ -605,12 +611,14 @@ impl OpenAI {
label: &str,
) -> Result<reqwest::Response, LLMError> {
let url = self.responses_url()?;
let token = self.provider.get_bearer_token().await?;
let mut request = self
.provider
.client
.post(url)
.bearer_auth(&self.provider.config.api_key)
.bearer_auth(&token)
.json(body);
request = self.provider.apply_headers(request);
self.log_request_payload(label, body);
request = self.apply_timeout(request);
request.send().await.map_err(LLMError::from)
Expand Down
2 changes: 2 additions & 0 deletions src/backends/openai/responses/request/tests.rs
Original file line number Diff line number Diff line change
Expand Up @@ -27,6 +27,8 @@ fn base_config() -> OpenAICompatibleProviderConfig {
embedding_encoding_format: None,
embedding_dimensions: None,
normalize_response: false,
headers: Vec::new(),
token_provider: None,
}
}

Expand Down
Loading