Skip to content
Open
Show file tree
Hide file tree
Changes from 3 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
8 changes: 7 additions & 1 deletion 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
8 changes: 7 additions & 1 deletion 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
8 changes: 7 additions & 1 deletion 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
8 changes: 7 additions & 1 deletion 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
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
8 changes: 7 additions & 1 deletion src/backends/openrouter.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 OpenRouter {
json_schema: Option<StructuredOutputFormat>,
parallel_tool_calls: Option<bool>,
normalize_response: Option<bool>,
headers: Vec<(String, String)>,
token_provider: Option<TokenProviderFn>,
) -> Self {
OpenAICompatibleProvider::<OpenRouterConfig>::new(
api_key,
Expand All @@ -73,6 +77,8 @@ impl OpenRouter {
normalize_response,
None, // embedding_encoding_format - not supported by OpenRouter
None, // embedding_dimensions - not supported by OpenRouter
headers,
token_provider,
)
}
}
Expand Down
4 changes: 3 additions & 1 deletion src/builder/build/backends/openai.rs
Original file line number Diff line number Diff line change
Expand Up @@ -13,7 +13,7 @@ pub(super) fn build_openai(
tools: Option<Vec<Tool>>,
tool_choice: Option<ToolChoice>,
) -> Result<Box<dyn LLMProvider>, LLMError> {
let key = helpers::require_api_key(state, "OpenAI")?;
let key = helpers::require_api_key_or_token(state, "OpenAI")?;
let timeout = helpers::timeout_or_default(state);

let provider = crate::backends::openai::OpenAI::new(
Expand All @@ -35,6 +35,8 @@ pub(super) fn build_openai(
state.json_schema.take(),
state.voice.take(),
state.extra_body.take(),
std::mem::take(&mut state.headers),
state.token_provider.take(),
state.openai_enable_web_search,
state.openai_web_search_context_size.take(),
state.openai_web_search_user_location_type.take(),
Expand Down
4 changes: 3 additions & 1 deletion src/builder/build/backends/openai_compatible/cohere.rs
Original file line number Diff line number Diff line change
Expand Up @@ -13,7 +13,7 @@ pub(super) fn build_cohere(
tools: Option<Vec<Tool>>,
tool_choice: Option<ToolChoice>,
) -> Result<Box<dyn LLMProvider>, LLMError> {
let api_key = helpers::require_api_key(state, "Cohere")?;
let api_key = helpers::require_api_key_or_token(state, "Cohere")?;
let timeout = helpers::timeout_or_default(state);
let provider = crate::backends::cohere::Cohere::new(
api_key,
Expand All @@ -35,6 +35,8 @@ pub(super) fn build_cohere(
state.normalize_response,
state.embedding_encoding_format.take(),
state.embedding_dimensions,
std::mem::take(&mut state.headers),
state.token_provider.take(),
);
Ok(Box::new(provider))
}
Expand Down
4 changes: 3 additions & 1 deletion src/builder/build/backends/openai_compatible/groq.rs
Original file line number Diff line number Diff line change
Expand Up @@ -13,7 +13,7 @@ pub(super) fn build_groq(
tools: Option<Vec<Tool>>,
tool_choice: Option<ToolChoice>,
) -> Result<Box<dyn LLMProvider>, LLMError> {
let api_key = helpers::require_api_key(state, "Groq")?;
let api_key = helpers::require_api_key_or_token(state, "Groq")?;
let timeout = helpers::timeout_or_default(state);
let provider = crate::backends::groq::Groq::with_config(
api_key,
Expand All @@ -34,6 +34,8 @@ pub(super) fn build_groq(
state.json_schema.take(),
state.enable_parallel_tool_use,
state.normalize_response,
std::mem::take(&mut state.headers),
state.token_provider.take(),
);
Ok(Box::new(provider))
}
Expand Down
4 changes: 3 additions & 1 deletion src/builder/build/backends/openai_compatible/huggingface.rs
Original file line number Diff line number Diff line change
Expand Up @@ -13,7 +13,7 @@ pub(super) fn build_huggingface(
tools: Option<Vec<Tool>>,
tool_choice: Option<ToolChoice>,
) -> Result<Box<dyn LLMProvider>, LLMError> {
let api_key = helpers::require_api_key(state, "HuggingFace Inference Providers")?;
let api_key = helpers::require_api_key_or_token(state, "HuggingFace Inference Providers")?;
let timeout = helpers::timeout_or_default(state);
let provider = crate::backends::huggingface::HuggingFace::with_config(
api_key,
Expand All @@ -34,6 +34,8 @@ pub(super) fn build_huggingface(
state.json_schema.take(),
state.enable_parallel_tool_use,
state.normalize_response,
std::mem::take(&mut state.headers),
state.token_provider.take(),
);
Ok(Box::new(provider))
}
Expand Down
4 changes: 3 additions & 1 deletion src/builder/build/backends/openai_compatible/mistral.rs
Original file line number Diff line number Diff line change
Expand Up @@ -13,7 +13,7 @@ pub(super) fn build_mistral(
tools: Option<Vec<Tool>>,
tool_choice: Option<ToolChoice>,
) -> Result<Box<dyn LLMProvider>, LLMError> {
let api_key = helpers::require_api_key(state, "Mistral")?;
let api_key = helpers::require_api_key_or_token(state, "Mistral")?;
let timeout = helpers::timeout_or_default(state);
let provider = crate::backends::mistral::Mistral::with_config(
api_key,
Expand All @@ -34,6 +34,8 @@ pub(super) fn build_mistral(
state.json_schema.take(),
state.enable_parallel_tool_use,
state.normalize_response,
std::mem::take(&mut state.headers),
state.token_provider.take(),
);
Ok(Box::new(provider))
}
Expand Down
4 changes: 3 additions & 1 deletion src/builder/build/backends/openai_compatible/openrouter.rs
Original file line number Diff line number Diff line change
Expand Up @@ -13,7 +13,7 @@ pub(super) fn build_openrouter(
tools: Option<Vec<Tool>>,
tool_choice: Option<ToolChoice>,
) -> Result<Box<dyn LLMProvider>, LLMError> {
let api_key = helpers::require_api_key(state, "OpenRouter")?;
let api_key = helpers::require_api_key_or_token(state, "OpenRouter")?;
let timeout = helpers::timeout_or_default(state);
let provider = crate::backends::openrouter::OpenRouter::with_config(
api_key,
Expand All @@ -34,6 +34,8 @@ pub(super) fn build_openrouter(
state.json_schema.take(),
state.enable_parallel_tool_use,
state.normalize_response,
std::mem::take(&mut state.headers),
state.token_provider.take(),
);
Ok(Box::new(provider))
}
Expand Down
Loading