From 48c36baf6984dab349667df973e4251ba8ee5860 Mon Sep 17 00:00:00 2001 From: Andreas Zwinkau Date: Mon, 13 Apr 2026 22:22:08 +0200 Subject: [PATCH 1/4] Provide API to add HTTP headers Just for OpenAI-compatible providers. Partially fixes #119 --- src/backends/cohere.rs | 2 ++ src/backends/groq.rs | 2 ++ src/backends/huggingface.rs | 2 ++ src/backends/mistral.rs | 2 ++ src/backends/openai.rs | 23 +++++++++++-------- src/backends/openrouter.rs | 2 ++ src/builder/build/backends/openai.rs | 1 + .../backends/openai_compatible/cohere.rs | 1 + .../build/backends/openai_compatible/groq.rs | 1 + .../backends/openai_compatible/huggingface.rs | 1 + .../backends/openai_compatible/mistral.rs | 1 + .../backends/openai_compatible/openrouter.rs | 1 + src/builder/llm_builder.rs | 6 +++++ src/builder/state.rs | 1 + src/providers/openai_compatible.rs | 17 ++++++++++++++ 15 files changed, 53 insertions(+), 10 deletions(-) diff --git a/src/backends/cohere.rs b/src/backends/cohere.rs index a8c11b9c..822ce71f 100644 --- a/src/backends/cohere.rs +++ b/src/backends/cohere.rs @@ -55,6 +55,7 @@ impl Cohere { json_schema: Option, parallel_tool_calls: Option, normalize_response: Option, + headers: Vec<(String, String)>, ) -> Self { >::new( api_key, @@ -76,6 +77,7 @@ impl Cohere { normalize_response, embedding_encoding_format, embedding_dimensions, + headers, ) } } diff --git a/src/backends/groq.rs b/src/backends/groq.rs index 8fd04fce..eccaff50 100644 --- a/src/backends/groq.rs +++ b/src/backends/groq.rs @@ -66,6 +66,7 @@ impl Groq { json_schema: Option, parallel_tool_calls: Option, normalize_response: Option, + headers: Vec<(String, String)>, ) -> Self { OpenAICompatibleProvider::::new( api_key, @@ -87,6 +88,7 @@ impl Groq { normalize_response, None, // embedding_encoding_format - not supported by Groq None, // embedding_dimensions - not supported by Groq + headers, ) } } diff --git a/src/backends/huggingface.rs b/src/backends/huggingface.rs index 0faa0490..4843265b 100644 --- a/src/backends/huggingface.rs +++ b/src/backends/huggingface.rs @@ -52,6 +52,7 @@ impl HuggingFace { json_schema: Option, parallel_tool_calls: Option, normalize_response: Option, + headers: Vec<(String, String)>, ) -> Self { OpenAICompatibleProvider::::new( api_key, @@ -73,6 +74,7 @@ impl HuggingFace { normalize_response, None, // embedding_encoding_format None, // embedding_dimensions + headers, ) } } diff --git a/src/backends/mistral.rs b/src/backends/mistral.rs index 862d9fb2..43e421c4 100644 --- a/src/backends/mistral.rs +++ b/src/backends/mistral.rs @@ -55,6 +55,7 @@ impl Mistral { json_schema: Option, parallel_tool_calls: Option, normalize_response: Option, + headers: Vec<(String, String)>, ) -> Self { >::new( api_key, @@ -76,6 +77,7 @@ impl Mistral { normalize_response, embedding_encoding_format, embedding_dimensions, + headers, ) } } diff --git a/src/backends/openai.rs b/src/backends/openai.rs index a2c3ecde..9cd4b378 100644 --- a/src/backends/openai.rs +++ b/src/backends/openai.rs @@ -220,6 +220,7 @@ impl OpenAI { json_schema: Option, voice: Option, extra_body: Option, + headers: Vec<(String, String)>, enable_web_search: Option, web_search_context_size: Option, web_search_user_location_type: Option, @@ -252,6 +253,7 @@ impl OpenAI { normalize_response, embedding_encoding_format, embedding_dimensions, + headers, ), enable_web_search: enable_web_search.unwrap_or(false), web_search_context_size, @@ -446,6 +448,7 @@ impl SpeechToTextProvider for OpenAI { .post(url) .bearer_auth(&self.provider.config.api_key) .multipart(form); + req = self.provider.apply_headers(req); if let Some(t) = self.provider.config.timeout_seconds { req = req.timeout(Duration::from_secs(t)); @@ -480,6 +483,7 @@ impl SpeechToTextProvider for OpenAI { .post(url) .bearer_auth(&self.provider.config.api_key) .multipart(form); + req = self.provider.apply_headers(req); if let Some(t) = self.provider.config.timeout_seconds { req = req.timeout(Duration::from_secs(t)); @@ -526,15 +530,14 @@ impl EmbeddingProvider for OpenAI { .join("embeddings") .map_err(|e| LLMError::HttpError(e.to_string()))?; - let resp = self + let mut req = self .provider .client .post(url) .bearer_auth(&self.provider.config.api_key) - .json(&body) - .send() - .await? - .error_for_status()?; + .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(); @@ -555,14 +558,13 @@ impl ModelsProvider for OpenAI { .join("models") .map_err(|e| LLMError::HttpError(e.to_string()))?; - let resp = self + let mut req = self .provider .client .get(url) - .bearer_auth(&self.provider.config.api_key) - .send() - .await? - .error_for_status()?; + .bearer_auth(&self.provider.config.api_key); + req = self.provider.apply_headers(req); + let resp = req.send().await?.error_for_status()?; let result = StandardModelListResponse { inner: resp.json().await?, @@ -611,6 +613,7 @@ impl OpenAI { .post(url) .bearer_auth(&self.provider.config.api_key) .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) diff --git a/src/backends/openrouter.rs b/src/backends/openrouter.rs index 633400b5..0b3a0652 100644 --- a/src/backends/openrouter.rs +++ b/src/backends/openrouter.rs @@ -52,6 +52,7 @@ impl OpenRouter { json_schema: Option, parallel_tool_calls: Option, normalize_response: Option, + headers: Vec<(String, String)>, ) -> Self { OpenAICompatibleProvider::::new( api_key, @@ -73,6 +74,7 @@ impl OpenRouter { normalize_response, None, // embedding_encoding_format - not supported by OpenRouter None, // embedding_dimensions - not supported by OpenRouter + headers, ) } } diff --git a/src/builder/build/backends/openai.rs b/src/builder/build/backends/openai.rs index dfc49644..2e7cef96 100644 --- a/src/builder/build/backends/openai.rs +++ b/src/builder/build/backends/openai.rs @@ -35,6 +35,7 @@ pub(super) fn build_openai( state.json_schema.take(), state.voice.take(), state.extra_body.take(), + std::mem::take(&mut state.headers), state.openai_enable_web_search, state.openai_web_search_context_size.take(), state.openai_web_search_user_location_type.take(), diff --git a/src/builder/build/backends/openai_compatible/cohere.rs b/src/builder/build/backends/openai_compatible/cohere.rs index 1108d5d0..fddc878a 100644 --- a/src/builder/build/backends/openai_compatible/cohere.rs +++ b/src/builder/build/backends/openai_compatible/cohere.rs @@ -35,6 +35,7 @@ pub(super) fn build_cohere( state.normalize_response, state.embedding_encoding_format.take(), state.embedding_dimensions, + std::mem::take(&mut state.headers), ); Ok(Box::new(provider)) } diff --git a/src/builder/build/backends/openai_compatible/groq.rs b/src/builder/build/backends/openai_compatible/groq.rs index 51893d08..d539a436 100644 --- a/src/builder/build/backends/openai_compatible/groq.rs +++ b/src/builder/build/backends/openai_compatible/groq.rs @@ -34,6 +34,7 @@ pub(super) fn build_groq( state.json_schema.take(), state.enable_parallel_tool_use, state.normalize_response, + std::mem::take(&mut state.headers), ); Ok(Box::new(provider)) } diff --git a/src/builder/build/backends/openai_compatible/huggingface.rs b/src/builder/build/backends/openai_compatible/huggingface.rs index 3b16f700..bbea91c5 100644 --- a/src/builder/build/backends/openai_compatible/huggingface.rs +++ b/src/builder/build/backends/openai_compatible/huggingface.rs @@ -34,6 +34,7 @@ pub(super) fn build_huggingface( state.json_schema.take(), state.enable_parallel_tool_use, state.normalize_response, + std::mem::take(&mut state.headers), ); Ok(Box::new(provider)) } diff --git a/src/builder/build/backends/openai_compatible/mistral.rs b/src/builder/build/backends/openai_compatible/mistral.rs index ff9fcd8d..2d25307e 100644 --- a/src/builder/build/backends/openai_compatible/mistral.rs +++ b/src/builder/build/backends/openai_compatible/mistral.rs @@ -34,6 +34,7 @@ pub(super) fn build_mistral( state.json_schema.take(), state.enable_parallel_tool_use, state.normalize_response, + std::mem::take(&mut state.headers), ); Ok(Box::new(provider)) } diff --git a/src/builder/build/backends/openai_compatible/openrouter.rs b/src/builder/build/backends/openai_compatible/openrouter.rs index 19d47e7b..5a849b4c 100644 --- a/src/builder/build/backends/openai_compatible/openrouter.rs +++ b/src/builder/build/backends/openai_compatible/openrouter.rs @@ -34,6 +34,7 @@ pub(super) fn build_openrouter( state.json_schema.take(), state.enable_parallel_tool_use, state.normalize_response, + std::mem::take(&mut state.headers), ); Ok(Box::new(provider)) } diff --git a/src/builder/llm_builder.rs b/src/builder/llm_builder.rs index b51a4a73..b9010be4 100644 --- a/src/builder/llm_builder.rs +++ b/src/builder/llm_builder.rs @@ -111,4 +111,10 @@ impl LLMBuilder { self.state.top_k = Some(top_k); self } + + /// Adds a custom HTTP header to every request. + pub fn header(mut self, key: impl Into, value: impl Into) -> Self { + self.state.headers.push((key.into(), value.into())); + self + } } diff --git a/src/builder/state.rs b/src/builder/state.rs index fc009137..b528c8cf 100644 --- a/src/builder/state.rs +++ b/src/builder/state.rs @@ -21,6 +21,7 @@ pub(crate) struct BuilderState { pub(crate) timeout_seconds: Option, pub(crate) top_p: Option, pub(crate) top_k: Option, + pub(crate) headers: Vec<(String, String)>, pub(crate) embedding_encoding_format: Option, pub(crate) embedding_dimensions: Option, pub(crate) validator: Option>, diff --git a/src/providers/openai_compatible.rs b/src/providers/openai_compatible.rs index 5923becd..9b585de2 100644 --- a/src/providers/openai_compatible.rs +++ b/src/providers/openai_compatible.rs @@ -67,6 +67,8 @@ pub struct OpenAICompatibleProviderConfig { pub embedding_dimensions: Option, /// Whether to normalize streaming responses. pub normalize_response: bool, + /// User-supplied custom headers to attach to every request. + pub headers: Vec<(String, String)>, } /// Generic OpenAI-compatible provider @@ -347,6 +349,7 @@ impl OpenAICompatibleProvider { normalize_response: Option, embedding_encoding_format: Option, embedding_dimensions: Option, + headers: Vec<(String, String)>, ) -> Self { let mut builder = Client::builder(); if let Some(sec) = timeout_seconds { @@ -374,6 +377,7 @@ impl OpenAICompatibleProvider { normalize_response, embedding_encoding_format, embedding_dimensions, + headers, ) } @@ -400,6 +404,7 @@ impl OpenAICompatibleProvider { normalize_response: Option, embedding_encoding_format: Option, embedding_dimensions: Option, + headers: Vec<(String, String)>, ) -> Self { let extra_body = match extra_body { Some(serde_json::Value::Object(map)) => map, @@ -431,6 +436,7 @@ impl OpenAICompatibleProvider { normalize_response: normalize_response.unwrap_or(true), embedding_encoding_format, embedding_dimensions, + headers, }; Self { config: Arc::new(config), @@ -519,6 +525,14 @@ impl OpenAICompatibleProvider { &self.client } + /// Attaches all user-supplied custom headers to a request builder. + pub fn apply_headers(&self, mut request: reqwest::RequestBuilder) -> reqwest::RequestBuilder { + for (key, value) in &self.config.headers { + request = request.header(key.as_str(), value.as_str()); + } + request + } + pub fn prepare_messages(&self, messages: &[ChatMessage]) -> Vec> { let mut openai_msgs: Vec = messages .iter() @@ -626,6 +640,7 @@ impl ChatProvider for OpenAICompatibleProvider { .post(url) .bearer_auth(&self.config.api_key) .json(&body); + request = self.apply_headers(request); // Add custom headers if provider specifies them if let Some(headers) = T::custom_headers() { for (key, value) in headers { @@ -748,6 +763,7 @@ impl ChatProvider for OpenAICompatibleProvider { .post(url) .bearer_auth(&self.config.api_key) .json(&body); + request = self.apply_headers(request); if let Some(headers) = T::custom_headers() { for (key, value) in headers { request = request.header(key, value); @@ -852,6 +868,7 @@ impl ChatProvider for OpenAICompatibleProvider { .bearer_auth(&self.config.api_key) .json(&body); + request = self.apply_headers(request); if let Some(headers) = T::custom_headers() { for (key, value) in headers { request = request.header(key, value); From 703074d115134b3add8c1132dad87fc9269ea98f Mon Sep 17 00:00:00 2001 From: Andreas Zwinkau Date: Mon, 13 Apr 2026 22:56:30 +0200 Subject: [PATCH 2/4] Provide API to request bearer token dynamically Just for OpenAI-compatible providers. Partially fixes #119 --- src/backends/cohere.rs | 6 +- src/backends/groq.rs | 6 +- src/backends/huggingface.rs | 6 +- src/backends/mistral.rs | 6 +- src/backends/openai.rs | 29 +++++--- src/backends/openrouter.rs | 6 +- src/builder/build/backends/openai.rs | 3 +- .../backends/openai_compatible/cohere.rs | 3 +- .../build/backends/openai_compatible/groq.rs | 3 +- .../backends/openai_compatible/huggingface.rs | 3 +- .../backends/openai_compatible/mistral.rs | 3 +- .../backends/openai_compatible/openrouter.rs | 3 +- src/builder/build/helpers.rs | 16 ++++ src/builder/llm_builder.rs | 27 ++++++- src/builder/state.rs | 2 + src/providers/openai_compatible.rs | 73 +++++++++++++------ 16 files changed, 150 insertions(+), 45 deletions(-) diff --git a/src/backends/cohere.rs b/src/backends/cohere.rs index 822ce71f..329977aa 100644 --- a/src/backends/cohere.rs +++ b/src/backends/cohere.rs @@ -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}, @@ -56,6 +58,7 @@ impl Cohere { parallel_tool_calls: Option, normalize_response: Option, headers: Vec<(String, String)>, + token_provider: Option, ) -> Self { >::new( api_key, @@ -78,6 +81,7 @@ impl Cohere { embedding_encoding_format, embedding_dimensions, headers, + token_provider, ) } } diff --git a/src/backends/groq.rs b/src/backends/groq.rs index eccaff50..229ad015 100644 --- a/src/backends/groq.rs +++ b/src/backends/groq.rs @@ -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, @@ -67,6 +69,7 @@ impl Groq { parallel_tool_calls: Option, normalize_response: Option, headers: Vec<(String, String)>, + token_provider: Option, ) -> Self { OpenAICompatibleProvider::::new( api_key, @@ -89,6 +92,7 @@ impl Groq { None, // embedding_encoding_format - not supported by Groq None, // embedding_dimensions - not supported by Groq headers, + token_provider, ) } } diff --git a/src/backends/huggingface.rs b/src/backends/huggingface.rs index 4843265b..2a15c679 100644 --- a/src/backends/huggingface.rs +++ b/src/backends/huggingface.rs @@ -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, @@ -53,6 +55,7 @@ impl HuggingFace { parallel_tool_calls: Option, normalize_response: Option, headers: Vec<(String, String)>, + token_provider: Option, ) -> Self { OpenAICompatibleProvider::::new( api_key, @@ -75,6 +78,7 @@ impl HuggingFace { None, // embedding_encoding_format None, // embedding_dimensions headers, + token_provider, ) } } diff --git a/src/backends/mistral.rs b/src/backends/mistral.rs index 43e421c4..1f69c9bd 100644 --- a/src/backends/mistral.rs +++ b/src/backends/mistral.rs @@ -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}, @@ -56,6 +58,7 @@ impl Mistral { parallel_tool_calls: Option, normalize_response: Option, headers: Vec<(String, String)>, + token_provider: Option, ) -> Self { >::new( api_key, @@ -78,6 +81,7 @@ impl Mistral { embedding_encoding_format, embedding_dimensions, headers, + token_provider, ) } } diff --git a/src/backends/openai.rs b/src/backends/openai.rs index 9cd4b378..cd484b00 100644 --- a/src/backends/openai.rs +++ b/src/backends/openai.rs @@ -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::{ @@ -221,6 +221,7 @@ impl OpenAI { voice: Option, extra_body: Option, headers: Vec<(String, String)>, + token_provider: Option, enable_web_search: Option, web_search_context_size: Option, web_search_user_location_type: Option, @@ -229,8 +230,10 @@ impl OpenAI { web_search_user_location_approximate_region: Option, ) -> Result { 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: >::new( @@ -254,6 +257,7 @@ impl OpenAI { embedding_encoding_format, embedding_dimensions, headers, + token_provider, ), enable_web_search: enable_web_search.unwrap_or(false), web_search_context_size, @@ -442,11 +446,12 @@ 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); @@ -477,11 +482,12 @@ 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); @@ -530,11 +536,12 @@ impl EmbeddingProvider for OpenAI { .join("embeddings") .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) .json(&body); req = self.provider.apply_headers(req); let resp = req.send().await?.error_for_status()?; @@ -558,11 +565,8 @@ impl ModelsProvider for OpenAI { .join("models") .map_err(|e| LLMError::HttpError(e.to_string()))?; - let mut req = self - .provider - .client - .get(url) - .bearer_auth(&self.provider.config.api_key); + 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()?; @@ -607,11 +611,12 @@ impl OpenAI { label: &str, ) -> Result { 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); diff --git a/src/backends/openrouter.rs b/src/backends/openrouter.rs index 0b3a0652..104fd729 100644 --- a/src/backends/openrouter.rs +++ b/src/backends/openrouter.rs @@ -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, @@ -53,6 +55,7 @@ impl OpenRouter { parallel_tool_calls: Option, normalize_response: Option, headers: Vec<(String, String)>, + token_provider: Option, ) -> Self { OpenAICompatibleProvider::::new( api_key, @@ -75,6 +78,7 @@ impl OpenRouter { None, // embedding_encoding_format - not supported by OpenRouter None, // embedding_dimensions - not supported by OpenRouter headers, + token_provider, ) } } diff --git a/src/builder/build/backends/openai.rs b/src/builder/build/backends/openai.rs index 2e7cef96..2b25c6e5 100644 --- a/src/builder/build/backends/openai.rs +++ b/src/builder/build/backends/openai.rs @@ -13,7 +13,7 @@ pub(super) fn build_openai( tools: Option>, tool_choice: Option, ) -> Result, 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( @@ -36,6 +36,7 @@ pub(super) fn build_openai( 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(), diff --git a/src/builder/build/backends/openai_compatible/cohere.rs b/src/builder/build/backends/openai_compatible/cohere.rs index fddc878a..b652f006 100644 --- a/src/builder/build/backends/openai_compatible/cohere.rs +++ b/src/builder/build/backends/openai_compatible/cohere.rs @@ -13,7 +13,7 @@ pub(super) fn build_cohere( tools: Option>, tool_choice: Option, ) -> Result, 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, @@ -36,6 +36,7 @@ pub(super) fn build_cohere( state.embedding_encoding_format.take(), state.embedding_dimensions, std::mem::take(&mut state.headers), + state.token_provider.take(), ); Ok(Box::new(provider)) } diff --git a/src/builder/build/backends/openai_compatible/groq.rs b/src/builder/build/backends/openai_compatible/groq.rs index d539a436..f4a8cd4b 100644 --- a/src/builder/build/backends/openai_compatible/groq.rs +++ b/src/builder/build/backends/openai_compatible/groq.rs @@ -13,7 +13,7 @@ pub(super) fn build_groq( tools: Option>, tool_choice: Option, ) -> Result, 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, @@ -35,6 +35,7 @@ pub(super) fn build_groq( state.enable_parallel_tool_use, state.normalize_response, std::mem::take(&mut state.headers), + state.token_provider.take(), ); Ok(Box::new(provider)) } diff --git a/src/builder/build/backends/openai_compatible/huggingface.rs b/src/builder/build/backends/openai_compatible/huggingface.rs index bbea91c5..9198bdc6 100644 --- a/src/builder/build/backends/openai_compatible/huggingface.rs +++ b/src/builder/build/backends/openai_compatible/huggingface.rs @@ -13,7 +13,7 @@ pub(super) fn build_huggingface( tools: Option>, tool_choice: Option, ) -> Result, 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, @@ -35,6 +35,7 @@ pub(super) fn build_huggingface( state.enable_parallel_tool_use, state.normalize_response, std::mem::take(&mut state.headers), + state.token_provider.take(), ); Ok(Box::new(provider)) } diff --git a/src/builder/build/backends/openai_compatible/mistral.rs b/src/builder/build/backends/openai_compatible/mistral.rs index 2d25307e..3d376559 100644 --- a/src/builder/build/backends/openai_compatible/mistral.rs +++ b/src/builder/build/backends/openai_compatible/mistral.rs @@ -13,7 +13,7 @@ pub(super) fn build_mistral( tools: Option>, tool_choice: Option, ) -> Result, 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, @@ -35,6 +35,7 @@ pub(super) fn build_mistral( state.enable_parallel_tool_use, state.normalize_response, std::mem::take(&mut state.headers), + state.token_provider.take(), ); Ok(Box::new(provider)) } diff --git a/src/builder/build/backends/openai_compatible/openrouter.rs b/src/builder/build/backends/openai_compatible/openrouter.rs index 5a849b4c..9ab1e4fe 100644 --- a/src/builder/build/backends/openai_compatible/openrouter.rs +++ b/src/builder/build/backends/openai_compatible/openrouter.rs @@ -13,7 +13,7 @@ pub(super) fn build_openrouter( tools: Option>, tool_choice: Option, ) -> Result, 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, @@ -35,6 +35,7 @@ pub(super) fn build_openrouter( state.enable_parallel_tool_use, state.normalize_response, std::mem::take(&mut state.headers), + state.token_provider.take(), ); Ok(Box::new(provider)) } diff --git a/src/builder/build/helpers.rs b/src/builder/build/helpers.rs index a850d22e..d217487c 100644 --- a/src/builder/build/helpers.rs +++ b/src/builder/build/helpers.rs @@ -62,6 +62,22 @@ pub(super) fn require_api_key( Ok(key.expose_secret().to_string()) } +/// Like [`require_api_key`], but succeeds with an empty string when a +/// [`TokenProviderFn`] has been configured on the builder state. +/// +/// Use this in every OpenAI-compatible backend builder so that callers who +/// supply `.auth_provider(...)` don't also have to supply a static API key. +pub(super) fn require_api_key_or_token( + state: &mut BuilderState, + provider: &str, +) -> Result { + if state.token_provider.is_some() { + Ok(optional_api_key(state).unwrap_or_default()) + } else { + require_api_key(state, provider) + } +} + pub(super) fn optional_api_key(state: &mut BuilderState) -> Option { state .api_key diff --git a/src/builder/llm_builder.rs b/src/builder/llm_builder.rs index b9010be4..0db03070 100644 --- a/src/builder/llm_builder.rs +++ b/src/builder/llm_builder.rs @@ -1,6 +1,11 @@ +use std::future::Future; +use std::pin::Pin; +use std::sync::Arc; + use secrecy::SecretString; use crate::chat::ReasoningEffort; +use crate::error::LLMError; use super::{backend::LLMBackend, state::BuilderState}; @@ -112,9 +117,29 @@ impl LLMBuilder { self } - /// Adds a custom HTTP header to every request. + /// Adds a custom HTTP header to every request for OpenAI-compatible backends. + /// + /// Can be called multiple times to add multiple headers. pub fn header(mut self, key: impl Into, value: impl Into) -> Self { self.state.headers.push((key.into(), value.into())); self } + + /// Sets an async callback that is invoked before every request to obtain + /// (or refresh) the bearer token for OpenAI-compatible backends. + /// + /// When set, the callback replaces the static API key supplied via + /// [`.api_key()`](Self::api_key) — you may omit `.api_key()` entirely. + /// The callback is called once per request, so it can perform token + /// exchange or refresh on every call without any extra coordination. + pub fn auth_provider(mut self, f: F) -> Self + where + F: Fn() -> Fut + Send + Sync + 'static, + Fut: Future> + Send + 'static, + { + self.state.token_provider = Some(Arc::new(move || { + Box::pin(f()) as Pin> + Send>> + })); + self + } } diff --git a/src/builder/state.rs b/src/builder/state.rs index b528c8cf..8e071f7a 100644 --- a/src/builder/state.rs +++ b/src/builder/state.rs @@ -1,3 +1,4 @@ +use crate::providers::openai_compatible::TokenProviderFn; use secrecy::SecretString; use crate::{ @@ -37,6 +38,7 @@ pub(crate) struct BuilderState { pub(crate) api_version: Option, pub(crate) deployment_id: Option, pub(crate) voice: Option, + pub(crate) token_provider: Option, pub(crate) extra_body: Option, pub(crate) xai_search_mode: Option, pub(crate) xai_search_source_type: Option, diff --git a/src/providers/openai_compatible.rs b/src/providers/openai_compatible.rs index 9b585de2..d3179312 100644 --- a/src/providers/openai_compatible.rs +++ b/src/providers/openai_compatible.rs @@ -20,10 +20,31 @@ use futures::{stream::Stream, StreamExt}; use reqwest::{Client, Url}; use serde::{Deserialize, Serialize}; use std::collections::HashMap; +use std::future::Future; use std::marker::PhantomData; use std::pin::Pin; use std::sync::Arc; +/// Type alias for a dynamic async function that returns a bearer token. +/// +/// Used by [`LLMBuilder::auth_provider`](crate::builder::LLMBuilder::auth_provider) to supply a +/// callback that is invoked before every request to obtain (or refresh) the bearer token. +pub type TokenProviderFn = + Arc Pin> + Send>> + Send + Sync>; + +/// Newtype wrapper around [`TokenProviderFn`] that implements [`Debug`] and [`Clone`]. +/// +/// Stored inside [`OpenAICompatibleProviderConfig`] so the config can continue to +/// derive `Debug`. +#[derive(Clone)] +pub struct TokenProvider(pub TokenProviderFn); + +impl std::fmt::Debug for TokenProvider { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.write_str("TokenProvider()") + } +} + const AUDIO_UNSUPPORTED: &str = "Audio messages are not supported for this provider"; /// Configuration for OpenAI-compatible providers. @@ -69,6 +90,9 @@ pub struct OpenAICompatibleProviderConfig { pub normalize_response: bool, /// User-supplied custom headers to attach to every request. pub headers: Vec<(String, String)>, + /// Optional async callback invoked before every request to obtain the bearer token. + /// When set, its return value is used instead of the static `api_key`. + pub token_provider: Option, } /// Generic OpenAI-compatible provider @@ -350,6 +374,7 @@ impl OpenAICompatibleProvider { embedding_encoding_format: Option, embedding_dimensions: Option, headers: Vec<(String, String)>, + token_provider: Option, ) -> Self { let mut builder = Client::builder(); if let Some(sec) = timeout_seconds { @@ -378,6 +403,7 @@ impl OpenAICompatibleProvider { embedding_encoding_format, embedding_dimensions, headers, + token_provider, ) } @@ -405,6 +431,7 @@ impl OpenAICompatibleProvider { embedding_encoding_format: Option, embedding_dimensions: Option, headers: Vec<(String, String)>, + token_provider: Option, ) -> Self { let extra_body = match extra_body { Some(serde_json::Value::Object(map)) => map, @@ -437,6 +464,7 @@ impl OpenAICompatibleProvider { embedding_encoding_format, embedding_dimensions, headers, + token_provider: token_provider.map(TokenProvider), }; Self { config: Arc::new(config), @@ -533,6 +561,18 @@ impl OpenAICompatibleProvider { request } + /// Returns the bearer token for the current request. + /// + /// If a [`TokenProvider`] callback was configured via + /// [`LLMBuilder::auth_provider`](crate::builder::LLMBuilder::auth_provider), it is invoked and + /// its result is used. Otherwise the static `api_key` is returned. + pub async fn get_bearer_token(&self) -> Result { + match &self.config.token_provider { + Some(provider) => (provider.0)().await, + None => Ok(self.config.api_key.clone()), + } + } + pub fn prepare_messages(&self, messages: &[ChatMessage]) -> Vec> { let mut openai_msgs: Vec = messages .iter() @@ -584,9 +624,9 @@ impl ChatProvider for OpenAICompatibleProvider { tools: Option<&[Tool]>, ) -> Result, LLMError> { crate::chat::ensure_no_audio(messages, AUDIO_UNSUPPORTED)?; - if self.config.api_key.is_empty() { + if self.config.api_key.is_empty() && self.config.token_provider.is_none() { return Err(LLMError::AuthError(format!( - "Missing {} API key", + "Missing {} credentials: provide an API key via `.api_key()` or a dynamic token provider via `.auth_provider()`", T::PROVIDER_NAME ))); } @@ -635,11 +675,8 @@ impl ChatProvider for OpenAICompatibleProvider { .base_url .join(T::CHAT_ENDPOINT) .map_err(|e| LLMError::HttpError(e.to_string()))?; - let mut request = self - .client - .post(url) - .bearer_auth(&self.config.api_key) - .json(&body); + let token = self.get_bearer_token().await?; + let mut request = self.client.post(url).bearer_auth(&token).json(&body); request = self.apply_headers(request); // Add custom headers if provider specifies them if let Some(headers) = T::custom_headers() { @@ -716,9 +753,9 @@ impl ChatProvider for OpenAICompatibleProvider { LLMError, > { crate::chat::ensure_no_audio(messages, AUDIO_UNSUPPORTED)?; - if self.config.api_key.is_empty() { + if self.config.api_key.is_empty() && self.config.token_provider.is_none() { return Err(LLMError::AuthError(format!( - "Missing {} API key", + "Missing {} credentials: provide an API key via `.api_key()` or a dynamic token provider via `.auth_provider()`", T::PROVIDER_NAME ))); } @@ -758,11 +795,8 @@ impl ChatProvider for OpenAICompatibleProvider { .base_url .join(T::CHAT_ENDPOINT) .map_err(|e| LLMError::HttpError(e.to_string()))?; - let mut request = self - .client - .post(url) - .bearer_auth(&self.config.api_key) - .json(&body); + let token = self.get_bearer_token().await?; + let mut request = self.client.post(url).bearer_auth(&token).json(&body); request = self.apply_headers(request); if let Some(headers) = T::custom_headers() { for (key, value) in headers { @@ -811,9 +845,9 @@ impl ChatProvider for OpenAICompatibleProvider { ) -> Result> + Send>>, LLMError> { crate::chat::ensure_no_audio(messages, AUDIO_UNSUPPORTED)?; - if self.config.api_key.is_empty() { + if self.config.api_key.is_empty() && self.config.token_provider.is_none() { return Err(LLMError::AuthError(format!( - "Missing {} API key", + "Missing {} credentials: provide an API key via `.api_key()` or a dynamic token provider via `.auth_provider()`", T::PROVIDER_NAME ))); } @@ -862,11 +896,8 @@ impl ChatProvider for OpenAICompatibleProvider { .join(T::CHAT_ENDPOINT) .map_err(|e| LLMError::HttpError(e.to_string()))?; - let mut request = self - .client - .post(url) - .bearer_auth(&self.config.api_key) - .json(&body); + let token = self.get_bearer_token().await?; + let mut request = self.client.post(url).bearer_auth(&token).json(&body); request = self.apply_headers(request); if let Some(headers) = T::custom_headers() { From 337ad7dd800507fc4173f8fcab19d266c94168a8 Mon Sep 17 00:00:00 2001 From: Andreas Zwinkau Date: Mon, 13 Apr 2026 23:20:21 +0200 Subject: [PATCH 3/4] Adapt test to API change --- src/backends/openai/responses/request/tests.rs | 2 ++ 1 file changed, 2 insertions(+) diff --git a/src/backends/openai/responses/request/tests.rs b/src/backends/openai/responses/request/tests.rs index 91b8a2a9..8e2dc079 100644 --- a/src/backends/openai/responses/request/tests.rs +++ b/src/backends/openai/responses/request/tests.rs @@ -27,6 +27,8 @@ fn base_config() -> OpenAICompatibleProviderConfig { embedding_encoding_format: None, embedding_dimensions: None, normalize_response: false, + headers: Vec::new(), + token_provider: None, } } From 4e20890ffb487f2b233e79797d27e8d7ffdf309a Mon Sep 17 00:00:00 2001 From: Andreas Zwinkau Date: Sat, 18 Apr 2026 16:46:06 +0200 Subject: [PATCH 4/4] Fix: missed some places where a bearer token was necessary too --- src/backends/cohere.rs | 9 ++++++--- src/backends/groq.rs | 9 ++++++--- src/backends/huggingface.rs | 7 ++++--- src/backends/mistral.rs | 20 ++++++++++++++------ src/backends/openrouter.rs | 7 ++++--- 5 files changed, 34 insertions(+), 18 deletions(-) diff --git a/src/backends/cohere.rs b/src/backends/cohere.rs index 329977aa..0581316c 100644 --- a/src/backends/cohere.rs +++ b/src/backends/cohere.rs @@ -142,8 +142,10 @@ impl SpeechToTextProvider for Cohere { #[async_trait] impl EmbeddingProvider for Cohere { async fn embed(&self, input: Vec) -> Result>, 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 { @@ -163,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? diff --git a/src/backends/groq.rs b/src/backends/groq.rs index 229ad015..9070e6f2 100644 --- a/src/backends/groq.rs +++ b/src/backends/groq.rs @@ -139,16 +139,19 @@ impl ModelsProvider for Groq { &self, _request: Option<&ModelListRequest>, ) -> Result, 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()?; diff --git a/src/backends/huggingface.rs b/src/backends/huggingface.rs index 2a15c679..6196073d 100644 --- a/src/backends/huggingface.rs +++ b/src/backends/huggingface.rs @@ -125,18 +125,19 @@ impl ModelsProvider for HuggingFace { &self, _request: Option<&ModelListRequest>, ) -> Result, 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()?; diff --git a/src/backends/mistral.rs b/src/backends/mistral.rs index 1f69c9bd..1a9dc2e4 100644 --- a/src/backends/mistral.rs +++ b/src/backends/mistral.rs @@ -142,8 +142,10 @@ impl SpeechToTextProvider for Mistral { #[async_trait] impl EmbeddingProvider for Mistral { async fn embed(&self, input: Vec) -> Result>, 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 { @@ -163,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? @@ -184,14 +187,18 @@ impl ModelsProvider for Mistral { &self, _request: Option<&ModelListRequest>, ) -> Result, 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()?; @@ -212,3 +219,4 @@ impl TextToSpeechProvider for Mistral { )) } } + diff --git a/src/backends/openrouter.rs b/src/backends/openrouter.rs index 104fd729..453967a0 100644 --- a/src/backends/openrouter.rs +++ b/src/backends/openrouter.rs @@ -125,18 +125,19 @@ impl ModelsProvider for OpenRouter { &self, _request: Option<&ModelListRequest>, ) -> Result, 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 OpenRouter API key".to_string(), + "Missing OpenRouter credentials: provide an API key via `.api_key()` or a dynamic token provider via `.auth_provider()`".to_string(), )); } let url = format!("{}models", OpenRouterConfig::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()?;