diff --git a/Cargo.lock b/Cargo.lock index 03720dca1101fb..f43ed8dc7706fd 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2117,6 +2117,7 @@ dependencies = [ "aws-sdk-bedrockruntime", "aws-smithy-types", "futures 0.3.32", + "http_client", "schemars 1.0.4", "serde", "serde_json", diff --git a/crates/agent_ui/src/agent_configuration/add_llm_provider_modal.rs b/crates/agent_ui/src/agent_configuration/add_llm_provider_modal.rs index 8eeda6447e878d..99413e10638c2b 100644 --- a/crates/agent_ui/src/agent_configuration/add_llm_provider_modal.rs +++ b/crates/agent_ui/src/agent_configuration/add_llm_provider_modal.rs @@ -279,6 +279,7 @@ fn save_provider_to_settings( OpenAiCompatibleSettingsContent { api_url, available_models: models, + custom_headers: None, }, ); }); diff --git a/crates/anthropic/src/anthropic.rs b/crates/anthropic/src/anthropic.rs index e18c7b9e9619c7..b8d4253d859fe2 100644 --- a/crates/anthropic/src/anthropic.rs +++ b/crates/anthropic/src/anthropic.rs @@ -6,7 +6,10 @@ use anyhow::{Context as _, Result}; use chrono::{DateTime, Utc}; use futures::{AsyncBufReadExt, AsyncReadExt, StreamExt, io::BufReader, stream::BoxStream}; use http_client::http::{self, HeaderMap, HeaderValue}; -use http_client::{AsyncBody, HttpClient, Method, Request as HttpRequest, StatusCode}; +use http_client::{ + AsyncBody, CustomHeaders, HttpClient, Method, Request as HttpRequest, RequestBuilderExt, + StatusCode, +}; use serde::{Deserialize, Serialize}; use strum::EnumString; use thiserror::Error; @@ -202,10 +205,18 @@ pub async fn stream_completion( api_key: &str, request: Request, beta_headers: Option, + extra_headers: &CustomHeaders, ) -> Result>, AnthropicError> { - stream_completion_with_rate_limit_info(client, api_url, api_key, request, beta_headers) - .await - .map(|output| output.0) + stream_completion_with_rate_limit_info( + client, + api_url, + api_key, + request, + beta_headers, + extra_headers, + ) + .await + .map(|output| output.0) } /// A raw model entry returned by the Anthropic models listing endpoint. @@ -233,6 +244,7 @@ pub async fn list_models( client: &dyn HttpClient, api_url: &str, api_key: &str, + extra_headers: &CustomHeaders, ) -> Result> { let uri = format!("{api_url}/v1/models?limit=1000"); @@ -242,6 +254,7 @@ pub async fn list_models( .header("Anthropic-Version", "2023-06-01") .header("X-Api-Key", api_key.trim()) .header("Accept", "application/json") + .extra_headers(extra_headers) .body(AsyncBody::default()) .context("failed to build Anthropic models list request")?; @@ -282,9 +295,17 @@ pub async fn non_streaming_completion( api_key: &str, request: Request, beta_headers: Option, + extra_headers: &CustomHeaders, ) -> Result { - let (mut response, rate_limits) = - send_request(client, api_url, api_key, &request, beta_headers).await?; + let (mut response, rate_limits) = send_request( + client, + api_url, + api_key, + &request, + beta_headers, + extra_headers, + ) + .await?; if response.status().is_success() { let mut body = String::new(); @@ -306,6 +327,7 @@ async fn send_request( api_key: &str, request: impl Serialize, beta_headers: Option, + extra_headers: &CustomHeaders, ) -> Result<(http::Response, RateLimitInfo), AnthropicError> { let uri = format!("{api_url}/v1/messages"); @@ -323,6 +345,7 @@ async fn send_request( let serialized_request = serde_json::to_string(&request).map_err(AnthropicError::SerializeRequest)?; let request = request_builder + .extra_headers(extra_headers) .body(AsyncBody::from(serialized_request)) .map_err(AnthropicError::BuildRequestBody)?; @@ -462,6 +485,7 @@ pub async fn stream_completion_with_rate_limit_info( api_key: &str, request: Request, beta_headers: Option, + extra_headers: &CustomHeaders, ) -> Result< ( BoxStream<'static, Result>, @@ -474,8 +498,15 @@ pub async fn stream_completion_with_rate_limit_info( stream: true, }; - let (response, rate_limits) = - send_request(client, api_url, api_key, &request, beta_headers).await?; + let (response, rate_limits) = send_request( + client, + api_url, + api_key, + &request, + beta_headers, + extra_headers, + ) + .await?; if response.status().is_success() { let reader = BufReader::new(response.into_body()); diff --git a/crates/bedrock/Cargo.toml b/crates/bedrock/Cargo.toml index f8f6fa46017309..17f69cde6b45d6 100644 --- a/crates/bedrock/Cargo.toml +++ b/crates/bedrock/Cargo.toml @@ -20,6 +20,7 @@ anyhow.workspace = true aws-sdk-bedrockruntime = { workspace = true, features = ["behavior-version-latest"] } aws-smithy-types = {workspace = true} futures.workspace = true +http_client.workspace = true schemars = { workspace = true, optional = true } serde.workspace = true serde_json.workspace = true diff --git a/crates/bedrock/src/bedrock.rs b/crates/bedrock/src/bedrock.rs index f3102b2d39c9fe..e8ce37e8c8f423 100644 --- a/crates/bedrock/src/bedrock.rs +++ b/crates/bedrock/src/bedrock.rs @@ -35,6 +35,7 @@ pub use crate::models::*; pub async fn stream_completion( client: bedrock::Client, request: Request, + extra_headers: http_client::CustomHeaders, ) -> Result>, BedrockError> { let mut response = bedrock::Client::converse_stream(&client) .model_id(request.model.clone()) @@ -99,30 +100,45 @@ pub async fn stream_completion( ); } - let output = response.send().await.map_err(|err| match err { - bedrock::error::SdkError::ServiceError(ctx) => { - use bedrock::operation::converse_stream::ConverseStreamError; - let err = ctx.into_err(); - match &err { - ConverseStreamError::ValidationException(e) => { - BedrockError::Validation(e.message().unwrap_or("validation error").to_string()) - } - ConverseStreamError::ThrottlingException(_) => BedrockError::RateLimited, - ConverseStreamError::ServiceUnavailableException(_) - | ConverseStreamError::ModelNotReadyException(_) => { - BedrockError::ServiceUnavailable - } - ConverseStreamError::AccessDeniedException(e) => { - BedrockError::AccessDenied(e.message().unwrap_or("access denied").to_string()) + let output = response + .customize() + .mutate_request(move |http_request| { + let headers = http_request.headers_mut(); + for (name, value) in extra_headers.iter() { + headers.insert( + name.as_str().to_owned(), + value.to_str().unwrap_or("").to_owned(), + ); + } + }) + .send() + .await + .map_err(|err| match err { + bedrock::error::SdkError::ServiceError(ctx) => { + use bedrock::operation::converse_stream::ConverseStreamError; + let err = ctx.into_err(); + match &err { + ConverseStreamError::ValidationException(e) => BedrockError::Validation( + e.message().unwrap_or("validation error").to_string(), + ), + ConverseStreamError::ThrottlingException(_) => BedrockError::RateLimited, + ConverseStreamError::ServiceUnavailableException(_) + | ConverseStreamError::ModelNotReadyException(_) => { + BedrockError::ServiceUnavailable + } + ConverseStreamError::AccessDeniedException(e) => BedrockError::AccessDenied( + e.message().unwrap_or("access denied").to_string(), + ), + ConverseStreamError::InternalServerException(e) => { + BedrockError::InternalServer( + e.message().unwrap_or("internal server error").to_string(), + ) + } + _ => BedrockError::Other(err.into()), } - ConverseStreamError::InternalServerException(e) => BedrockError::InternalServer( - e.message().unwrap_or("internal server error").to_string(), - ), - _ => BedrockError::Other(err.into()), } - } - other => BedrockError::Other(other.into()), - }); + other => BedrockError::Other(other.into()), + }); let stream = Box::pin(stream::unfold( output?.stream, diff --git a/crates/deepseek/src/deepseek.rs b/crates/deepseek/src/deepseek.rs index 56609a47492cc6..4ec7e918045b97 100644 --- a/crates/deepseek/src/deepseek.rs +++ b/crates/deepseek/src/deepseek.rs @@ -4,7 +4,9 @@ use futures::{ io::BufReader, stream::{BoxStream, StreamExt}, }; -use http_client::{AsyncBody, HttpClient, Method, Request as HttpRequest}; +use http_client::{ + AsyncBody, CustomHeaders, HttpClient, Method, Request as HttpRequest, RequestBuilderExt, +}; use serde::{Deserialize, Serialize}; use serde_json::Value; use std::convert::TryFrom; @@ -297,15 +299,16 @@ pub async fn stream_completion( api_url: &str, api_key: &str, request: Request, + extra_headers: &CustomHeaders, ) -> Result>> { let uri = format!("{api_url}/chat/completions"); - let request_builder = HttpRequest::builder() + let request = HttpRequest::builder() .method(Method::POST) .uri(uri) .header("Content-Type", "application/json") - .header("Authorization", format!("Bearer {}", api_key.trim())); - - let request = request_builder.body(AsyncBody::from(serde_json::to_string(&request)?))?; + .header("Authorization", format!("Bearer {}", api_key.trim())) + .extra_headers(extra_headers) + .body(AsyncBody::from(serde_json::to_string(&request)?))?; let mut response = client.send(request).await?; if response.status().is_success() { diff --git a/crates/edit_prediction_cli/src/anthropic_client.rs b/crates/edit_prediction_cli/src/anthropic_client.rs index c1706d6161b5f7..56c8f00a6ac51e 100644 --- a/crates/edit_prediction_cli/src/anthropic_client.rs +++ b/crates/edit_prediction_cli/src/anthropic_client.rs @@ -63,6 +63,7 @@ impl PlainLlmClient { &self.api_key, request, None, + &http_client::CustomHeaders::default(), ) .await .map_err(|e| anyhow::anyhow!("{:?}", e))?; @@ -104,6 +105,7 @@ impl PlainLlmClient { &self.api_key, request, None, + &http_client::CustomHeaders::default(), ) .await .map_err(|e| anyhow::anyhow!("{:?}", e))?; diff --git a/crates/google_ai/src/google_ai.rs b/crates/google_ai/src/google_ai.rs index 56ec48c83a866c..fd2acd40f89eeb 100644 --- a/crates/google_ai/src/google_ai.rs +++ b/crates/google_ai/src/google_ai.rs @@ -2,7 +2,9 @@ use std::mem; use anyhow::{Result, anyhow, bail}; use futures::{AsyncBufReadExt, AsyncReadExt, StreamExt, io::BufReader, stream::BoxStream}; -use http_client::{AsyncBody, HttpClient, Method, Request as HttpRequest}; +use http_client::{ + AsyncBody, CustomHeaders, HttpClient, Method, Request as HttpRequest, RequestBuilderExt, +}; pub use language_model_core::ModelMode as GoogleModelMode; use serde::{Deserialize, Deserializer, Serialize, Serializer}; pub mod completion; @@ -14,6 +16,7 @@ pub async fn stream_generate_content( api_url: &str, api_key: &str, mut request: GenerateContentRequest, + extra_headers: &CustomHeaders, ) -> Result>> { let api_key = api_key.trim(); validate_generate_content_request(&request)?; @@ -24,12 +27,12 @@ pub async fn stream_generate_content( let uri = format!("{api_url}/v1beta/models/{model_id}:streamGenerateContent?alt=sse&key={api_key}",); - let request_builder = HttpRequest::builder() + let request = HttpRequest::builder() .method(Method::POST) .uri(uri) - .header("Content-Type", "application/json"); - - let request = request_builder.body(AsyncBody::from(serde_json::to_string(&request)?))?; + .header("Content-Type", "application/json") + .extra_headers(extra_headers) + .body(AsyncBody::from(serde_json::to_string(&request)?))?; let mut response = client.send(request).await?; if response.status().is_success() { let reader = BufReader::new(response.into_body()); diff --git a/crates/http_client/src/http_client.rs b/crates/http_client/src/http_client.rs index bbbe3b1a832332..1c3e1859b7f6bc 100644 --- a/crates/http_client/src/http_client.rs +++ b/crates/http_client/src/http_client.rs @@ -7,8 +7,8 @@ pub mod github_download; pub use anyhow::{Result, anyhow}; pub use async_body::{AsyncBody, Inner, Json}; use derive_more::Deref; -use http::HeaderValue; pub use http::{self, Method, Request, Response, StatusCode, Uri, request::Builder}; +use http::{HeaderName, HeaderValue}; use futures::future::BoxFuture; use parking_lot::Mutex; @@ -57,6 +57,58 @@ impl HttpRequestExt for http::request::Builder { } } +/// A set of pre-validated user-supplied HTTP headers. +/// +/// Construction (and the per-name validation that goes with it) happens once +/// at settings load time. Cloning is `Arc`-cheap, so providers can hand a copy +/// to each outgoing request without re-parsing or re-allocating. +#[derive(Default, Clone, Debug)] +pub struct CustomHeaders(Arc<[(HeaderName, HeaderValue)]>); + +impl CustomHeaders { + pub fn new(headers: Vec<(HeaderName, HeaderValue)>) -> Self { + Self(headers.into()) + } + + pub fn is_empty(&self) -> bool { + self.0.is_empty() + } + + pub fn iter(&self) -> impl ExactSizeIterator { + self.0.iter().map(|(n, v)| (n, v)) + } +} + +impl PartialEq for CustomHeaders { + fn eq(&self, other: &Self) -> bool { + self.0.len() == other.0.len() + && self + .0 + .iter() + .zip(other.0.iter()) + .all(|(a, b)| a.0 == b.0 && a.1 == b.1) + } +} + +pub trait RequestBuilderExt { + /// Append every header in `headers` to the request being built. + fn extra_headers(self, headers: &CustomHeaders) -> Self; +} + +impl RequestBuilderExt for http::request::Builder { + fn extra_headers(mut self, headers: &CustomHeaders) -> Self { + if headers.is_empty() { + return self; + } + if let Some(map) = self.headers_mut() { + for (name, value) in headers.iter() { + map.append(name.clone(), value.clone()); + } + } + self + } +} + pub trait HttpClient: 'static + Send + Sync { fn user_agent(&self) -> Option<&HeaderValue>; diff --git a/crates/language_models/src/provider.rs b/crates/language_models/src/provider.rs index 8c0b4fbfd6c8e9..51b323e6c7babf 100644 --- a/crates/language_models/src/provider.rs +++ b/crates/language_models/src/provider.rs @@ -1,3 +1,7 @@ +use collections::HashMap; +use http_client::CustomHeaders; +use http_client::http::{HeaderName, HeaderValue}; + pub mod anthropic; pub mod bedrock; pub mod cloud; @@ -15,3 +19,118 @@ pub mod opencode; pub mod vercel_ai_gateway; pub mod x_ai; + +const COMMON_RESERVED_HEADER_NAMES: &[&str] = &["Authorization", "Content-Type", "Accept"]; + +/// Validate the user-supplied custom-headers map once at settings load time, +/// dropping reserved or malformed entries (each with a `log::warn!`) and +/// returning a typed `CustomHeaders` ready to be appended to outgoing requests. +pub(crate) fn resolve_custom_headers( + provider_name: &str, + settings: &HashMap, + reserved_header_names: &[&str], +) -> CustomHeaders { + let headers = settings + .iter() + .filter_map(|(name, value)| { + if COMMON_RESERVED_HEADER_NAMES + .iter() + .chain(reserved_header_names) + .any(|reserved| reserved.eq_ignore_ascii_case(name)) + { + log::warn!( + "ignoring custom {provider_name} header `{name}`: managed by Zed and cannot be overridden" + ); + return None; + } + let header_name = match name.parse::() { + Ok(header_name) => header_name, + Err(err) => { + log::warn!("ignoring custom {provider_name} header `{name}`: invalid header name ({err})"); + return None; + } + }; + let header_value = match HeaderValue::from_str(value) { + Ok(header_value) => header_value, + Err(err) => { + log::warn!( + "ignoring custom {provider_name} header `{name}`: invalid header value ({err})" + ); + return None; + } + }; + Some((header_name, header_value)) + }) + .collect(); + CustomHeaders::new(headers) +} + +#[cfg(test)] +mod tests { + use super::*; + + fn map(pairs: &[(&str, &str)]) -> HashMap { + pairs + .iter() + .map(|(key, value)| ((*key).to_string(), (*value).to_string())) + .collect() + } + + fn names(headers: &CustomHeaders) -> Vec { + let mut names: Vec = headers + .iter() + .map(|(name, _)| name.as_str().to_owned()) + .collect(); + names.sort(); + names + } + + #[test] + fn drops_common_and_provider_reserved_headers() { + let settings = map(&[ + ("Authorization", "Bearer leak"), + ("Content-Type", "text/plain"), + ("Accept", "text/plain"), + ("X-Api-Key", "leak"), + ("X-Allowed", "yes"), + ]); + let merged = resolve_custom_headers("Test", &settings, &["X-Api-Key"]); + assert_eq!(names(&merged), vec!["x-allowed".to_string()]); + } + + #[test] + fn reserved_header_match_is_case_insensitive() { + let settings = map(&[ + ("authorization", "Bearer leak"), + ("CONTENT-TYPE", "text/plain"), + ("x-api-key", "leak"), + ("X-Allowed", "yes"), + ]); + let merged = resolve_custom_headers("Test", &settings, &["X-Api-Key"]); + assert_eq!(names(&merged), vec!["x-allowed".to_string()]); + } + + #[test] + fn headers_with_reserved_prefix_are_kept() { + let settings = map(&[("Authorization-Forwarded", "ok"), ("X-Api-Key-Trace", "ok")]); + let merged = resolve_custom_headers("Test", &settings, &["X-Api-Key"]); + assert_eq!( + names(&merged), + vec![ + "authorization-forwarded".to_string(), + "x-api-key-trace".to_string(), + ] + ); + } + + #[test] + fn drops_invalid_header_name_and_value() { + let settings = map(&[ + ("Bad Name", "ok"), + ("X-Bad-Value", "line1\nline2"), + ("X-Allowed", "yes"), + ]); + let merged = resolve_custom_headers("Test", &settings, &[]); + assert_eq!(names(&merged), vec!["x-allowed".to_string()]); + } +} diff --git a/crates/language_models/src/provider/anthropic.rs b/crates/language_models/src/provider/anthropic.rs index b919e2b61092d5..cb7f8b7aa114fb 100644 --- a/crates/language_models/src/provider/anthropic.rs +++ b/crates/language_models/src/provider/anthropic.rs @@ -6,7 +6,7 @@ use collections::BTreeMap; use credentials_provider::CredentialsProvider; use futures::{FutureExt, StreamExt, future::BoxFuture, stream::BoxStream}; use gpui::{AnyView, App, AsyncApp, Context, Entity, Task, TaskExt}; -use http_client::HttpClient; +use http_client::{CustomHeaders, HttpClient}; use language_model::{ ANTHROPIC_PROVIDER_ID, ANTHROPIC_PROVIDER_NAME, ApiKeyState, AuthenticateError, ConfigurationViewTargetAgent, EnvVar, FastModeConfirmation, IconOrSvg, LanguageModel, @@ -32,6 +32,8 @@ pub struct AnthropicSettings { pub api_url: String, /// Extend Zed's list of Anthropic models. pub available_models: Vec, + /// User-configured headers added to every Anthropic request. + pub custom_headers: CustomHeaders, } pub struct AnthropicLanguageModelProvider { @@ -42,6 +44,9 @@ pub struct AnthropicLanguageModelProvider { const API_KEY_ENV_VAR_NAME: &str = "ANTHROPIC_API_KEY"; static API_KEY_ENV_VAR: LazyLock = env_var!(API_KEY_ENV_VAR_NAME); +pub(crate) const RESERVED_HEADER_NAMES: &[&str] = + &["X-Api-Key", "Anthropic-Version", "Anthropic-Beta"]; + pub struct State { api_key_state: ApiKeyState, credentials_provider: Arc, @@ -105,10 +110,18 @@ impl State { "cannot fetch Anthropic models without an API key" ))); }; + let extra_headers = AnthropicLanguageModelProvider::settings(cx) + .custom_headers + .clone(); cx.spawn(async move |this, cx| { - let models = - anthropic::list_models(http_client.as_ref(), &api_url, api_key.as_ref()).await?; + let models = anthropic::list_models( + http_client.as_ref(), + &api_url, + api_key.as_ref(), + &extra_headers, + ) + .await?; this.update(cx, |this, cx| { this.fetched_models = models; @@ -375,9 +388,12 @@ impl AnthropicModel { > { let http_client = self.http_client.clone(); - let (api_key, api_url) = self.state.read_with(cx, |state, cx| { + let (api_key, api_url, extra_headers) = self.state.read_with(cx, |state, cx| { let api_url = AnthropicLanguageModelProvider::api_url(cx); - (state.api_key_state.key(&api_url), api_url) + let extra_headers = AnthropicLanguageModelProvider::settings(cx) + .custom_headers + .clone(); + (state.api_key_state.key(&api_url), api_url, extra_headers) }); let beta_headers = self.model.beta_headers(); @@ -394,6 +410,7 @@ impl AnthropicModel { &api_key, request, beta_headers, + &extra_headers, ); request.await.map_err(Into::into) } diff --git a/crates/language_models/src/provider/bedrock.rs b/crates/language_models/src/provider/bedrock.rs index 9d595b516ad253..7024a0970244e9 100644 --- a/crates/language_models/src/provider/bedrock.rs +++ b/crates/language_models/src/provider/bedrock.rs @@ -49,12 +49,21 @@ use ui_input::InputField; use util::ResultExt; use crate::AllLanguageModelSettings; +use http_client::CustomHeaders; use language_model::util::{fix_streamed_json, parse_tool_arguments}; actions!(bedrock, [Tab, TabPrev]); const PROVIDER_ID: LanguageModelProviderId = LanguageModelProviderId::new("amazon-bedrock"); const PROVIDER_NAME: LanguageModelProviderName = LanguageModelProviderName::new("Amazon Bedrock"); +pub(crate) const RESERVED_HEADER_NAMES: &[&str] = &[ + "host", + "x-amz-date", + "x-amz-security-token", + "x-amz-content-sha256", + "amz-sdk-invocation-id", + "amz-sdk-request", +]; /// Credentials stored in the keychain for static authentication. /// Region is handled separately since it's orthogonal to auth method. @@ -107,6 +116,7 @@ impl BedrockCredentials { #[derive(Default, Clone, Debug, PartialEq)] pub struct AmazonBedrockSettings { pub available_models: Vec, + pub custom_headers: CustomHeaders, pub region: Option, pub endpoint: Option, pub profile_name: Option, @@ -617,8 +627,17 @@ impl BedrockModel { return futures::future::ready(Err(BedrockError::Other(anyhow!("App state dropped")))) .boxed(); }; + let extra_headers = self.state.read_with(cx, |_, cx| { + AllLanguageModelSettings::get_global(cx) + .bedrock + .custom_headers + .clone() + }); - let task = Tokio::spawn(cx, bedrock::stream_completion(runtime_client, request)); + let task = Tokio::spawn( + cx, + bedrock::stream_completion(runtime_client, request, extra_headers), + ); async move { task.await.map_err(|e| BedrockError::Other(e.into()))? }.boxed() } } diff --git a/crates/language_models/src/provider/deepseek.rs b/crates/language_models/src/provider/deepseek.rs index 1e2869bd7b1013..a7da87f355c02c 100644 --- a/crates/language_models/src/provider/deepseek.rs +++ b/crates/language_models/src/provider/deepseek.rs @@ -6,7 +6,7 @@ use deepseek::DEEPSEEK_API_URL; use futures::Stream; use futures::{FutureExt, StreamExt, future::BoxFuture, stream::BoxStream}; use gpui::{AnyView, App, AsyncApp, Context, Entity, SharedString, Task, TaskExt, Window}; -use http_client::HttpClient; +use http_client::{CustomHeaders, HttpClient}; use language_model::{ ApiKeyState, AuthenticateError, EnvVar, IconOrSvg, LanguageModel, LanguageModelCompletionError, LanguageModelCompletionEvent, LanguageModelEffortLevel, LanguageModelId, LanguageModelName, @@ -43,6 +43,7 @@ struct RawToolCall { pub struct DeepSeekSettings { pub api_url: String, pub available_models: Vec, + pub custom_headers: CustomHeaders, } pub struct DeepSeekLanguageModelProvider { http_client: Arc, @@ -228,9 +229,12 @@ impl DeepSeekLanguageModel { ) -> BoxFuture<'static, Result>>> { let http_client = self.http_client.clone(); - let (api_key, api_url) = self.state.read_with(cx, |state, cx| { + let (api_key, api_url, extra_headers) = self.state.read_with(cx, |state, cx| { let api_url = DeepSeekLanguageModelProvider::api_url(cx); - (state.api_key_state.key(&api_url), api_url) + let extra_headers = DeepSeekLanguageModelProvider::settings(cx) + .custom_headers + .clone(); + (state.api_key_state.key(&api_url), api_url, extra_headers) }); let future = self.request_limiter.stream(async move { @@ -239,8 +243,13 @@ impl DeepSeekLanguageModel { provider: PROVIDER_NAME, }); }; - let request = - deepseek::stream_completion(http_client.as_ref(), &api_url, &api_key, request); + let request = deepseek::stream_completion( + http_client.as_ref(), + &api_url, + &api_key, + request, + &extra_headers, + ); let response = request.await?; Ok(response) }); diff --git a/crates/language_models/src/provider/google.rs b/crates/language_models/src/provider/google.rs index 774de74a6b2a72..392a81454bbe11 100644 --- a/crates/language_models/src/provider/google.rs +++ b/crates/language_models/src/provider/google.rs @@ -5,7 +5,7 @@ use futures::{FutureExt, StreamExt, future::BoxFuture}; use google_ai::GenerateContentResponse; pub use google_ai::completion::{GoogleEventMapper, into_google}; use gpui::{AnyView, App, AsyncApp, Context, Entity, SharedString, Task, TaskExt, Window}; -use http_client::HttpClient; +use http_client::{CustomHeaders, HttpClient}; use language_model::{ AuthenticateError, ConfigurationViewTargetAgent, EnvVar, LanguageModelCompletionError, LanguageModelCompletionEvent, LanguageModelToolChoice, LanguageModelToolSchemaFormat, @@ -34,6 +34,7 @@ const PROVIDER_NAME: LanguageModelProviderName = GOOGLE_PROVIDER_NAME; pub struct GoogleSettings { pub api_url: String, pub available_models: Vec, + pub custom_headers: CustomHeaders, } #[derive(Clone, Copy, Debug, Default, PartialEq, Serialize, Deserialize, JsonSchema)] @@ -255,9 +256,12 @@ impl GoogleLanguageModel { > { let http_client = self.http_client.clone(); - let (api_key, api_url) = self.state.read_with(cx, |state, cx| { + let (api_key, api_url, extra_headers) = self.state.read_with(cx, |state, cx| { let api_url = GoogleLanguageModelProvider::api_url(cx); - (state.api_key_state.key(&api_url), api_url) + let extra_headers = GoogleLanguageModelProvider::settings(cx) + .custom_headers + .clone(); + (state.api_key_state.key(&api_url), api_url, extra_headers) }); async move { @@ -267,6 +271,7 @@ impl GoogleLanguageModel { &api_url, &api_key, request, + &extra_headers, ); request.await.context("failed to stream completion") } diff --git a/crates/language_models/src/provider/lmstudio.rs b/crates/language_models/src/provider/lmstudio.rs index ea19c265e9c5d2..1e09a3b370469a 100644 --- a/crates/language_models/src/provider/lmstudio.rs +++ b/crates/language_models/src/provider/lmstudio.rs @@ -1,11 +1,10 @@ use anyhow::{Result, anyhow}; -use collections::HashMap; use credentials_provider::CredentialsProvider; use fs::Fs; use futures::Stream; use futures::{FutureExt, StreamExt, future::BoxFuture, stream::BoxStream}; use gpui::{AnyView, App, AsyncApp, Context, CursorStyle, Entity, Subscription, Task, TaskExt}; -use http_client::HttpClient; +use http_client::{CustomHeaders, HttpClient}; use language_model::{ ApiKeyState, AuthenticateError, EnvVar, IconOrSvg, LanguageModel, LanguageModelCompletionError, LanguageModelCompletionEvent, LanguageModelToolChoice, LanguageModelToolResultContent, @@ -21,7 +20,10 @@ pub use settings::LmStudioAvailableModel as AvailableModel; use settings::{Settings, SettingsStore, update_settings_file}; use std::pin::Pin; use std::sync::LazyLock; -use std::{collections::BTreeMap, sync::Arc}; +use std::{ + collections::{BTreeMap, HashMap}, + sync::Arc, +}; use ui::{ ButtonLike, ConfiguredApiCard, ElevationIndex, List, ListBulletItem, Tooltip, prelude::*, }; @@ -44,6 +46,7 @@ static API_KEY_ENV_VAR: LazyLock = env_var!(API_KEY_ENV_VAR_NAME); pub struct LmStudioSettings { pub api_url: String, pub available_models: Vec, + pub custom_headers: CustomHeaders, } pub struct LmStudioLanguageModelProvider { @@ -84,11 +87,18 @@ impl State { let http_client = self.http_client.clone(); let api_url = settings.api_url.clone(); let api_key = self.api_key_state.key(&api_url); + let extra_headers = settings.custom_headers.clone(); // As a proxy for the server being "authenticated", we'll check if its up by fetching the models cx.spawn(async move |this, cx| { - let models = - get_models(http_client.as_ref(), &api_url, api_key.as_deref(), None).await?; + let models = get_models( + http_client.as_ref(), + &api_url, + api_key.as_deref(), + None, + &extra_headers, + ) + .await?; let mut models: Vec = models .into_iter() @@ -447,9 +457,13 @@ impl LmStudioLanguageModel { Result>>, > { let http_client = self.http_client.clone(); - let (api_key, api_url) = self.state.read_with(cx, |state, cx| { + let (api_key, api_url, extra_headers) = self.state.read_with(cx, |state, cx| { let api_url = LmStudioLanguageModelProvider::api_url(cx); - (state.api_key_state.key(&api_url), api_url) + let extra_headers = AllLanguageModelSettings::get_global(cx) + .lmstudio + .custom_headers + .clone(); + (state.api_key_state.key(&api_url), api_url, extra_headers) }); let future = self.request_limiter.stream(async move { @@ -458,6 +472,7 @@ impl LmStudioLanguageModel { &api_url, api_key.as_deref(), request, + &extra_headers, ) .await?; Ok(stream) diff --git a/crates/language_models/src/provider/mistral.rs b/crates/language_models/src/provider/mistral.rs index 30e3836e8bba5b..92c83342f3e94d 100644 --- a/crates/language_models/src/provider/mistral.rs +++ b/crates/language_models/src/provider/mistral.rs @@ -1,10 +1,10 @@ use anyhow::{Result, anyhow}; -use collections::BTreeMap; +use collections::{BTreeMap, HashMap}; use credentials_provider::CredentialsProvider; use futures::{FutureExt, Stream, StreamExt, future::BoxFuture, stream::BoxStream}; use gpui::{AnyView, App, AsyncApp, Context, Entity, Global, SharedString, Task, TaskExt, Window}; -use http_client::HttpClient; +use http_client::{CustomHeaders, HttpClient}; use language_model::{ ApiKeyState, AuthenticateError, EnvVar, IconOrSvg, LanguageModel, LanguageModelCompletionError, LanguageModelCompletionEvent, LanguageModelId, LanguageModelName, LanguageModelProvider, @@ -15,7 +15,6 @@ use language_model::{ pub use mistral::{MISTRAL_API_URL, StreamResponse}; pub use settings::MistralAvailableModel as AvailableModel; use settings::{Settings, SettingsStore}; -use std::collections::HashMap; use std::pin::Pin; use std::sync::{Arc, LazyLock}; use strum::IntoEnumIterator; @@ -30,11 +29,13 @@ const PROVIDER_NAME: LanguageModelProviderName = LanguageModelProviderName::new( const API_KEY_ENV_VAR_NAME: &str = "MISTRAL_API_KEY"; static API_KEY_ENV_VAR: LazyLock = env_var!(API_KEY_ENV_VAR_NAME); +pub(crate) const RESERVED_HEADER_NAMES: &[&str] = &["x-affinity"]; #[derive(Default, Clone, Debug, PartialEq)] pub struct MistralSettings { pub api_url: String, pub available_models: Vec, + pub custom_headers: CustomHeaders, } pub struct MistralLanguageModelProvider { @@ -256,9 +257,12 @@ impl MistralLanguageModel { > { let http_client = self.http_client.clone(); - let (api_key, api_url) = self.state.read_with(cx, |state, cx| { + let (api_key, api_url, extra_headers) = self.state.read_with(cx, |state, cx| { let api_url = MistralLanguageModelProvider::api_url(cx); - (state.api_key_state.key(&api_url), api_url) + let extra_headers = MistralLanguageModelProvider::settings(cx) + .custom_headers + .clone(); + (state.api_key_state.key(&api_url), api_url, extra_headers) }); let future = self.request_limiter.stream(async move { @@ -273,6 +277,7 @@ impl MistralLanguageModel { &api_key, request, affinity, + &extra_headers, ); let response = request.await?; Ok(response) diff --git a/crates/language_models/src/provider/ollama.rs b/crates/language_models/src/provider/ollama.rs index ea3a6f65035200..ad866cb63fa4cd 100644 --- a/crates/language_models/src/provider/ollama.rs +++ b/crates/language_models/src/provider/ollama.rs @@ -1,10 +1,11 @@ use anyhow::{Result, anyhow}; +use collections::HashMap; use credentials_provider::CredentialsProvider; use fs::Fs; use futures::{FutureExt, StreamExt, future::BoxFuture, stream::BoxStream}; use futures::{Stream, TryFutureExt, stream}; use gpui::{AnyView, App, AsyncApp, Context, CursorStyle, Entity, Task, TaskExt}; -use http_client::HttpClient; +use http_client::{CustomHeaders, HttpClient}; use language_model::{ ApiKeyState, AuthenticateError, EnvVar, IconOrSvg, LanguageModel, LanguageModelCompletionError, LanguageModelCompletionEvent, LanguageModelId, LanguageModelName, LanguageModelProvider, @@ -20,8 +21,8 @@ use ollama::{ pub use settings::OllamaAvailableModel as AvailableModel; use settings::{Settings, SettingsStore, update_settings_file}; use std::pin::Pin; +use std::sync::Arc; use std::sync::LazyLock; -use std::{collections::HashMap, sync::Arc}; use ui::{ ButtonLike, ButtonLink, ConfiguredApiCard, ElevationIndex, List, ListBulletItem, Tooltip, prelude::*, @@ -46,6 +47,7 @@ pub struct OllamaSettings { pub auto_discover: bool, pub available_models: Vec, pub context_window: Option, + pub custom_headers: CustomHeaders, } pub struct OllamaLanguageModelProvider { @@ -109,12 +111,20 @@ impl State { fn fetch_models(&mut self, cx: &mut Context) -> Task> { let http_client = Arc::clone(&self.http_client); + let settings = OllamaLanguageModelProvider::settings(cx); let api_url = OllamaLanguageModelProvider::api_url(cx); let api_key = self.api_key_state.key(&api_url); + let extra_headers = settings.custom_headers.clone(); // As a proxy for the server being "authenticated", we'll check if its up by fetching the models cx.spawn(async move |this, cx| { - let models = get_models(http_client.as_ref(), &api_url, api_key.as_deref()).await?; + let models = get_models( + http_client.as_ref(), + &api_url, + api_key.as_deref(), + &extra_headers, + ) + .await?; let tasks = models .into_iter() @@ -126,11 +136,17 @@ impl State { let http_client = Arc::clone(&http_client); let api_url = api_url.clone(); let api_key = api_key.clone(); + let extra_headers = extra_headers.clone(); async move { let name = model.name.as_str(); - let model = - show_model(http_client.as_ref(), &api_url, api_key.as_deref(), name) - .await?; + let model = show_model( + http_client.as_ref(), + &api_url, + api_key.as_deref(), + name, + &extra_headers, + ) + .await?; let ollama_model = ollama::Model::new( name, None, @@ -266,7 +282,7 @@ impl LanguageModelProvider for OllamaLanguageModelProvider { } fn provided_models(&self, cx: &App) -> Vec> { - let mut models: HashMap = HashMap::new(); + let mut models: HashMap = HashMap::default(); let settings = OllamaLanguageModelProvider::settings(cx); if settings.auto_discover { @@ -512,15 +528,23 @@ impl LanguageModel for OllamaLanguageModel { let request = self.to_ollama_request(request); let http_client = self.http_client.clone(); - let (api_key, api_url) = self.state.read_with(cx, |state, cx| { + let (api_key, api_url, extra_headers) = self.state.read_with(cx, |state, cx| { let api_url = OllamaLanguageModelProvider::api_url(cx); - (state.api_key_state.key(&api_url), api_url) + let extra_headers = OllamaLanguageModelProvider::settings(cx) + .custom_headers + .clone(); + (state.api_key_state.key(&api_url), api_url, extra_headers) }); let future = self.request_limiter.stream(async move { - let stream = - stream_chat_completion(http_client.as_ref(), &api_url, api_key.as_deref(), request) - .await?; + let stream = stream_chat_completion( + http_client.as_ref(), + &api_url, + api_key.as_deref(), + request, + &extra_headers, + ) + .await?; let stream = map_to_language_model_completion_events(stream); Ok(stream) }); @@ -1091,7 +1115,7 @@ mod tests { // When multiple models share the same base name (e.g., qwen2.5-coder:1.5b and qwen2.5-coder:3b), // each model should get its own display_name from settings, not a random one. - let mut models: HashMap = HashMap::new(); + let mut models: HashMap = HashMap::default(); models.insert( "qwen2.5-coder:1.5b".to_string(), ollama::Model { diff --git a/crates/language_models/src/provider/open_ai.rs b/crates/language_models/src/provider/open_ai.rs index 1ef535f39a6bd7..42200e349648e9 100644 --- a/crates/language_models/src/provider/open_ai.rs +++ b/crates/language_models/src/provider/open_ai.rs @@ -3,7 +3,7 @@ use collections::BTreeMap; use credentials_provider::CredentialsProvider; use futures::{FutureExt, StreamExt, future::BoxFuture}; use gpui::{AnyView, App, AsyncApp, Context, Entity, SharedString, Task, TaskExt, Window}; -use http_client::HttpClient; +use http_client::{CustomHeaders, HttpClient}; use language_model::{ ApiKeyState, AuthenticateError, EnvVar, FastModeConfirmation, IconOrSvg, LanguageModel, LanguageModelCompletionError, LanguageModelCompletionEvent, LanguageModelEffortLevel, @@ -38,6 +38,7 @@ static API_KEY_ENV_VAR: LazyLock = env_var!(API_KEY_ENV_VAR_NAME); pub struct OpenAiSettings { pub api_url: String, pub available_models: Vec, + pub custom_headers: CustomHeaders, } pub struct OpenAiLanguageModelProvider { @@ -345,9 +346,12 @@ impl OpenAiLanguageModel { { let http_client = self.http_client.clone(); - let (api_key, api_url) = self.state.read_with(cx, |state, cx| { + let (api_key, api_url, extra_headers) = self.state.read_with(cx, |state, cx| { let api_url = OpenAiLanguageModelProvider::api_url(cx); - (state.api_key_state.key(&api_url), api_url) + let extra_headers = OpenAiLanguageModelProvider::settings(cx) + .custom_headers + .clone(); + (state.api_key_state.key(&api_url), api_url, extra_headers) }); let future = self.request_limiter.stream(async move { @@ -361,6 +365,7 @@ impl OpenAiLanguageModel { &api_url, &api_key, request, + &extra_headers, ); let response = request.await?; Ok(response) @@ -377,9 +382,12 @@ impl OpenAiLanguageModel { { let http_client = self.http_client.clone(); - let (api_key, api_url) = self.state.read_with(cx, |state, cx| { + let (api_key, api_url, extra_headers) = self.state.read_with(cx, |state, cx| { let api_url = OpenAiLanguageModelProvider::api_url(cx); - (state.api_key_state.key(&api_url), api_url) + let extra_headers = OpenAiLanguageModelProvider::settings(cx) + .custom_headers + .clone(); + (state.api_key_state.key(&api_url), api_url, extra_headers) }); let provider = PROVIDER_NAME; @@ -393,7 +401,7 @@ impl OpenAiLanguageModel { &api_url, &api_key, request, - vec![], + &extra_headers, ); let response = request.await?; Ok(response) diff --git a/crates/language_models/src/provider/open_ai_compatible.rs b/crates/language_models/src/provider/open_ai_compatible.rs index d4e1f21f20a057..51f277be33f47e 100644 --- a/crates/language_models/src/provider/open_ai_compatible.rs +++ b/crates/language_models/src/provider/open_ai_compatible.rs @@ -3,7 +3,7 @@ use convert_case::{Case, Casing}; use credentials_provider::CredentialsProvider; use futures::{FutureExt, StreamExt, future::BoxFuture}; use gpui::{AnyView, App, AsyncApp, Context, Entity, SharedString, Task, TaskExt, Window}; -use http_client::HttpClient; +use http_client::{CustomHeaders, HttpClient}; use language_model::{ ApiKeyState, AuthenticateError, EnvVar, IconOrSvg, LanguageModel, LanguageModelCompletionError, LanguageModelCompletionEvent, LanguageModelId, LanguageModelName, LanguageModelProvider, @@ -32,6 +32,7 @@ pub use settings::OpenAiCompatibleModelCapabilities as ModelCapabilities; pub struct OpenAiCompatibleSettings { pub api_url: String, pub available_models: Vec, + pub custom_headers: CustomHeaders, } pub struct OpenAiCompatibleLanguageModelProvider { @@ -235,11 +236,12 @@ impl OpenAiCompatibleLanguageModel { > { let http_client = self.http_client.clone(); - let (api_key, api_url) = self.state.read_with(cx, |state, _cx| { + let (api_key, api_url, extra_headers) = self.state.read_with(cx, |state, _cx| { let api_url = &state.settings.api_url; ( state.api_key_state.key(api_url), state.settings.api_url.clone(), + state.settings.custom_headers.clone(), ) }); @@ -254,6 +256,7 @@ impl OpenAiCompatibleLanguageModel { &api_url, &api_key, request, + &extra_headers, ); let response = request.await?; Ok(response) @@ -270,11 +273,12 @@ impl OpenAiCompatibleLanguageModel { { let http_client = self.http_client.clone(); - let (api_key, api_url) = self.state.read_with(cx, |state, _cx| { + let (api_key, api_url, extra_headers) = self.state.read_with(cx, |state, _cx| { let api_url = &state.settings.api_url; ( state.api_key_state.key(api_url), state.settings.api_url.clone(), + state.settings.custom_headers.clone(), ) }); @@ -289,7 +293,7 @@ impl OpenAiCompatibleLanguageModel { &api_url, &api_key, request, - vec![], + &extra_headers, ); let response = request.await?; Ok(response) diff --git a/crates/language_models/src/provider/open_router.rs b/crates/language_models/src/provider/open_router.rs index d9632c402f03d0..ef434eed859992 100644 --- a/crates/language_models/src/provider/open_router.rs +++ b/crates/language_models/src/provider/open_router.rs @@ -3,7 +3,7 @@ use collections::HashMap; use credentials_provider::CredentialsProvider; use futures::{FutureExt, Stream, StreamExt, future::BoxFuture}; use gpui::{AnyView, App, AsyncApp, Context, Entity, SharedString, Task, TaskExt}; -use http_client::HttpClient; +use http_client::{CustomHeaders, HttpClient}; use language_model::{ ApiKeyState, AuthenticateError, EnvVar, IconOrSvg, LanguageModel, LanguageModelCompletionError, LanguageModelCompletionEvent, LanguageModelId, LanguageModelName, LanguageModelProvider, @@ -29,11 +29,13 @@ const PROVIDER_NAME: LanguageModelProviderName = LanguageModelProviderName::new( const API_KEY_ENV_VAR_NAME: &str = "OPENROUTER_API_KEY"; static API_KEY_ENV_VAR: LazyLock = env_var!(API_KEY_ENV_VAR_NAME); +pub(crate) const RESERVED_HEADER_NAMES: &[&str] = &["HTTP-Referer", "X-Title"]; #[derive(Default, Clone, Debug, PartialEq)] pub struct OpenRouterSettings { pub api_url: String, pub available_models: Vec, + pub custom_headers: CustomHeaders, } pub struct OpenRouterLanguageModelProvider { @@ -90,13 +92,16 @@ impl State { ) -> Task> { let http_client = self.http_client.clone(); let api_url = OpenRouterLanguageModelProvider::api_url(cx); + let extra_headers = OpenRouterLanguageModelProvider::settings(cx) + .custom_headers + .clone(); let Some(api_key) = self.api_key_state.key(&api_url) else { return Task::ready(Err(LanguageModelCompletionError::NoApiKey { provider: PROVIDER_NAME, })); }; cx.spawn(async move |this, cx| { - let models = list_models(http_client.as_ref(), &api_url, &api_key) + let models = list_models(http_client.as_ref(), &api_url, &api_key, &extra_headers) .await .map_err(|e| { LanguageModelCompletionError::Other(anyhow::anyhow!( @@ -291,9 +296,12 @@ impl OpenRouterLanguageModel { >, > { let http_client = self.http_client.clone(); - let (api_key, api_url) = self.state.read_with(cx, |state, cx| { + let (api_key, api_url, extra_headers) = self.state.read_with(cx, |state, cx| { let api_url = OpenRouterLanguageModelProvider::api_url(cx); - (state.api_key_state.key(&api_url), api_url) + let extra_headers = OpenRouterLanguageModelProvider::settings(cx) + .custom_headers + .clone(); + (state.api_key_state.key(&api_url), api_url, extra_headers) }); async move { @@ -302,8 +310,13 @@ impl OpenRouterLanguageModel { provider: PROVIDER_NAME, }); }; - let request = - open_router::stream_completion(http_client.as_ref(), &api_url, &api_key, request); + let request = open_router::stream_completion( + http_client.as_ref(), + &api_url, + &api_key, + request, + &extra_headers, + ); request.await.map_err(Into::into) } .boxed() diff --git a/crates/language_models/src/provider/openai_subscribed.rs b/crates/language_models/src/provider/openai_subscribed.rs index 66716ebdadb881..6d40e8e465f36f 100644 --- a/crates/language_models/src/provider/openai_subscribed.rs +++ b/crates/language_models/src/provider/openai_subscribed.rs @@ -4,7 +4,10 @@ use base64::engine::general_purpose::URL_SAFE_NO_PAD; use credentials_provider::CredentialsProvider; use futures::{FutureExt, StreamExt, future::BoxFuture, future::Shared}; use gpui::{AnyView, App, AsyncApp, Context, Entity, SharedString, Task, Window}; -use http_client::{AsyncBody, HttpClient, Method, Request as HttpRequest}; +use http_client::{ + AsyncBody, CustomHeaders, HttpClient, Method, Request as HttpRequest, + http::{HeaderName, HeaderValue}, +}; use language_model::{ AuthenticateError, FastModeConfirmation, IconOrSvg, LanguageModel, LanguageModelCompletionError, LanguageModelCompletionEvent, LanguageModelEffortLevel, @@ -510,15 +513,24 @@ impl LanguageModel for OpenAiSubscribedLanguageModel { let future = cx.spawn(async move |cx| { let creds = get_fresh_credentials(&state, &http_client, cx).await?; - let mut extra_headers: Vec<(String, String)> = vec![ - ("originator".into(), "zed".into()), - ("OpenAI-Beta".into(), "responses=experimental".into()), + let mut header_pairs: Vec<(HeaderName, HeaderValue)> = vec![ + ( + HeaderName::from_static("originator"), + HeaderValue::from_static("zed"), + ), + ( + HeaderName::from_static("openai-beta"), + HeaderValue::from_static("responses=experimental"), + ), ]; if let Some(ref id) = creds.account_id { if !id.is_empty() { - extra_headers.push(("ChatGPT-Account-Id".into(), id.clone())); + if let Ok(value) = HeaderValue::from_str(id) { + header_pairs.push((HeaderName::from_static("chatgpt-account-id"), value)); + } } } + let extra_headers = CustomHeaders::new(header_pairs); let access_token = creds.access_token.clone(); request_limiter @@ -529,7 +541,7 @@ impl LanguageModel for OpenAiSubscribedLanguageModel { CODEX_BASE_URL, &access_token, responses_request, - extra_headers, + &extra_headers, ) .await .map_err(LanguageModelCompletionError::from) diff --git a/crates/language_models/src/provider/opencode.rs b/crates/language_models/src/provider/opencode.rs index a5c98151feb0ea..28501560c28b94 100644 --- a/crates/language_models/src/provider/opencode.rs +++ b/crates/language_models/src/provider/opencode.rs @@ -4,7 +4,7 @@ use credentials_provider::CredentialsProvider; use fs::Fs; use futures::{FutureExt, StreamExt, future::BoxFuture}; use gpui::{AnyView, App, AsyncApp, Context, Entity, SharedString, Task, TaskExt, Window}; -use http_client::{AsyncBody, HttpClient, http}; +use http_client::{AsyncBody, CustomHeaders, HttpClient, http}; use language_model::{ ApiKeyState, AuthenticateError, EnvVar, IconOrSvg, LanguageModel, LanguageModelCompletionError, LanguageModelCompletionEvent, LanguageModelEffortLevel, LanguageModelId, LanguageModelName, @@ -58,11 +58,13 @@ const PROVIDER_NAME: LanguageModelProviderName = LanguageModelProviderName::new( const API_KEY_ENV_VAR_NAME: &str = "OPENCODE_API_KEY"; static API_KEY_ENV_VAR: LazyLock = env_var!(API_KEY_ENV_VAR_NAME); +pub(crate) const RESERVED_HEADER_NAMES: &[&str] = &["x-opencode-session"]; #[derive(Default, Clone, Debug, PartialEq)] pub struct OpenCodeSettings { pub api_url: String, pub available_models: Vec, + pub custom_headers: CustomHeaders, pub show_zen_models: bool, pub show_go_models: bool, pub show_free_models: bool, @@ -378,10 +380,19 @@ impl OpenCodeLanguageModel { }) } + fn custom_headers(&self, cx: &AsyncApp) -> CustomHeaders { + self.state.read_with(cx, |_, cx| { + OpenCodeLanguageModelProvider::settings(cx) + .custom_headers + .clone() + }) + } + fn stream_anthropic( &self, request: anthropic::Request, http_client: Arc, + extra_headers: CustomHeaders, cx: &AsyncApp, ) -> BoxFuture< 'static, @@ -409,6 +420,7 @@ impl OpenCodeLanguageModel { &api_key, request, None, + &extra_headers, ); let response = request.await?; Ok(response) @@ -421,6 +433,7 @@ impl OpenCodeLanguageModel { &self, request: open_ai::Request, http_client: Arc, + extra_headers: CustomHeaders, cx: &AsyncApp, ) -> BoxFuture< 'static, @@ -444,6 +457,7 @@ impl OpenCodeLanguageModel { &api_url, &api_key, request, + &extra_headers, ); let response = request.await?; Ok(response) @@ -456,6 +470,7 @@ impl OpenCodeLanguageModel { &self, request: open_ai::responses::Request, http_client: Arc, + extra_headers: CustomHeaders, cx: &AsyncApp, ) -> BoxFuture< 'static, @@ -479,7 +494,7 @@ impl OpenCodeLanguageModel { &api_url, &api_key, request, - vec![], + &extra_headers, ); let response = request.await?; Ok(response) @@ -492,6 +507,7 @@ impl OpenCodeLanguageModel { &self, request: google_ai::GenerateContentRequest, http_client: Arc, + extra_headers: CustomHeaders, cx: &AsyncApp, ) -> BoxFuture< 'static, @@ -511,6 +527,7 @@ impl OpenCodeLanguageModel { &api_url, &api_key, request, + &extra_headers, ); let response = request.await?; Ok(response) @@ -634,6 +651,7 @@ impl LanguageModel for OpenCodeLanguageModel { } else { self.http_client.clone() }; + let extra_headers = self.custom_headers(cx); match self.model.protocol(self.subscription) { ApiProtocol::Anthropic => { @@ -652,7 +670,8 @@ impl LanguageModel for OpenCodeLanguageModel { mode, anthropic::completion::AnthropicPromptCacheMode::Automatic, ); - let stream = self.stream_anthropic(anthropic_request, http_client, cx); + let stream = + self.stream_anthropic(anthropic_request, http_client, extra_headers, cx); async move { let mapper = AnthropicEventMapper::new(); Ok(mapper.map_stream(stream.await?).boxed()) @@ -677,7 +696,8 @@ impl LanguageModel for OpenCodeLanguageModel { reasoning_effort, self.model.interleaved_reasoning(), ); - let stream = self.stream_openai_chat(openai_request, http_client, cx); + let stream = + self.stream_openai_chat(openai_request, http_client, extra_headers, cx); async move { let mapper = OpenAiEventMapper::new(); Ok(mapper.map_stream(stream.await?).boxed()) @@ -698,7 +718,8 @@ impl LanguageModel for OpenCodeLanguageModel { None, supports_none_reasoning_effort, ); - let stream = self.stream_openai_response(response_request, http_client, cx); + let stream = + self.stream_openai_response(response_request, http_client, extra_headers, cx); async move { let mapper = OpenAiResponseEventMapper::new(); Ok(mapper.map_stream(stream.await?).boxed()) @@ -711,7 +732,7 @@ impl LanguageModel for OpenCodeLanguageModel { self.model.id().to_string(), google_ai::GoogleModelMode::Default, ); - let stream = self.stream_google(google_request, http_client, cx); + let stream = self.stream_google(google_request, http_client, extra_headers, cx); async move { let mapper = GoogleEventMapper::new(); Ok(mapper.map_stream(stream.await?.boxed()).boxed()) diff --git a/crates/language_models/src/provider/vercel_ai_gateway.rs b/crates/language_models/src/provider/vercel_ai_gateway.rs index 312cdee5a6605b..694972e48ae25a 100644 --- a/crates/language_models/src/provider/vercel_ai_gateway.rs +++ b/crates/language_models/src/provider/vercel_ai_gateway.rs @@ -3,7 +3,9 @@ use collections::BTreeMap; use credentials_provider::CredentialsProvider; use futures::{AsyncReadExt, FutureExt, StreamExt, future::BoxFuture}; use gpui::{AnyView, App, AsyncApp, Context, Entity, SharedString, Task, TaskExt, Window}; -use http_client::{AsyncBody, HttpClient, Method, Request as HttpRequest, http}; +use http_client::{ + AsyncBody, CustomHeaders, HttpClient, Method, Request as HttpRequest, RequestBuilderExt, http, +}; use language_model::{ ApiKeyState, AuthenticateError, EnvVar, IconOrSvg, LanguageModel, LanguageModelCompletionError, LanguageModelCompletionEvent, LanguageModelId, LanguageModelName, LanguageModelProvider, @@ -33,6 +35,7 @@ static API_KEY_ENV_VAR: LazyLock = env_var!(API_KEY_ENV_VAR_NAME); pub struct VercelAiGatewaySettings { pub api_url: String, pub available_models: Vec, + pub custom_headers: CustomHeaders, } pub struct VercelAiGatewayLanguageModelProvider { @@ -90,8 +93,17 @@ impl State { let http_client = self.http_client.clone(); let api_url = VercelAiGatewayLanguageModelProvider::api_url(cx); let api_key = self.api_key_state.key(&api_url); + let extra_headers = VercelAiGatewayLanguageModelProvider::settings(cx) + .custom_headers + .clone(); cx.spawn(async move |this, cx| { - let models = list_models(http_client.as_ref(), &api_url, api_key.as_deref()).await?; + let models = list_models( + http_client.as_ref(), + &api_url, + api_key.as_deref(), + &extra_headers, + ) + .await?; this.update(cx, |this, cx| { this.available_models = models; cx.notify(); @@ -271,9 +283,12 @@ impl VercelAiGatewayLanguageModel { >, > { let http_client = self.http_client.clone(); - let (api_key, api_url) = self.state.read_with(cx, |state, cx| { + let (api_key, api_url, extra_headers) = self.state.read_with(cx, |state, cx| { let api_url = VercelAiGatewayLanguageModelProvider::api_url(cx); - (state.api_key_state.key(&api_url), api_url) + let extra_headers = VercelAiGatewayLanguageModelProvider::settings(cx) + .custom_headers + .clone(); + (state.api_key_state.key(&api_url), api_url, extra_headers) }); let future = self.request_limiter.stream(async move { @@ -287,6 +302,7 @@ impl VercelAiGatewayLanguageModel { &api_url, &api_key, request, + &extra_headers, ); let response = request.await.map_err(map_open_ai_error)?; Ok(response) @@ -484,6 +500,7 @@ async fn list_models( client: &dyn HttpClient, api_url: &str, api_key: Option<&str>, + extra_headers: &CustomHeaders, ) -> Result, LanguageModelCompletionError> { let uri = format!("{api_url}/models?include_mappings=true"); let mut request_builder = HttpRequest::builder() @@ -494,6 +511,7 @@ async fn list_models( request_builder = request_builder.header("Authorization", format!("Bearer {}", api_key)); } let request = request_builder + .extra_headers(extra_headers) .body(AsyncBody::default()) .map_err(|error| LanguageModelCompletionError::BuildRequestBody { provider: PROVIDER_NAME, diff --git a/crates/language_models/src/provider/x_ai.rs b/crates/language_models/src/provider/x_ai.rs index 51eb2e3b81d533..8ac48213c0eb20 100644 --- a/crates/language_models/src/provider/x_ai.rs +++ b/crates/language_models/src/provider/x_ai.rs @@ -3,7 +3,7 @@ use collections::BTreeMap; use credentials_provider::CredentialsProvider; use futures::{FutureExt, StreamExt, future::BoxFuture}; use gpui::{AnyView, App, AsyncApp, Context, Entity, Task, TaskExt, Window}; -use http_client::HttpClient; +use http_client::{CustomHeaders, HttpClient}; use language_model::{ ApiKeyState, AuthenticateError, EnvVar, IconOrSvg, LanguageModel, LanguageModelCompletionError, LanguageModelCompletionEvent, LanguageModelEffortLevel, LanguageModelId, LanguageModelName, @@ -31,6 +31,7 @@ static API_KEY_ENV_VAR: LazyLock = env_var!(API_KEY_ENV_VAR_NAME); pub struct XAiSettings { pub api_url: String, pub available_models: Vec, + pub custom_headers: CustomHeaders, } pub struct XAiLanguageModelProvider { @@ -230,9 +231,12 @@ impl XAiLanguageModel { > { let http_client = self.http_client.clone(); - let (api_key, api_url) = self.state.read_with(cx, |state, cx| { + let (api_key, api_url, extra_headers) = self.state.read_with(cx, |state, cx| { let api_url = XAiLanguageModelProvider::api_url(cx); - (state.api_key_state.key(&api_url), api_url) + let extra_headers = XAiLanguageModelProvider::settings(cx) + .custom_headers + .clone(); + (state.api_key_state.key(&api_url), api_url, extra_headers) }); let future = self.request_limiter.stream(async move { @@ -246,6 +250,7 @@ impl XAiLanguageModel { &api_url, &api_key, request, + &extra_headers, ); let response = request.await?; Ok(response) diff --git a/crates/language_models/src/settings.rs b/crates/language_models/src/settings.rs index 53a1e060ec5d71..9e3f6d708bc8e1 100644 --- a/crates/language_models/src/settings.rs +++ b/crates/language_models/src/settings.rs @@ -4,11 +4,12 @@ use collections::HashMap; use settings::RegisterSetting; use crate::provider::{ - anthropic::AnthropicSettings, bedrock::AmazonBedrockSettings, cloud::ZedDotDevSettings, - deepseek::DeepSeekSettings, google::GoogleSettings, lmstudio::LmStudioSettings, - mistral::MistralSettings, ollama::OllamaSettings, open_ai::OpenAiSettings, - open_ai_compatible::OpenAiCompatibleSettings, open_router::OpenRouterSettings, - opencode::OpenCodeSettings, vercel_ai_gateway::VercelAiGatewaySettings, x_ai::XAiSettings, + anthropic, anthropic::AnthropicSettings, bedrock, bedrock::AmazonBedrockSettings, + cloud::ZedDotDevSettings, deepseek::DeepSeekSettings, google::GoogleSettings, + lmstudio::LmStudioSettings, mistral, mistral::MistralSettings, ollama::OllamaSettings, + open_ai::OpenAiSettings, open_ai_compatible::OpenAiCompatibleSettings, open_router, + open_router::OpenRouterSettings, opencode, opencode::OpenCodeSettings, resolve_custom_headers, + vercel_ai_gateway::VercelAiGatewaySettings, x_ai::XAiSettings, }; #[derive(Debug, RegisterSetting)] @@ -29,6 +30,17 @@ pub struct AllLanguageModelSettings { pub zed_dot_dev: ZedDotDevSettings, } +fn custom_headers_from( + provider_name: &str, + raw: Option>, + reserved: &[&str], +) -> http_client::CustomHeaders { + raw.as_ref() + .filter(|map| !map.is_empty()) + .map(|map| resolve_custom_headers(provider_name, map, reserved)) + .unwrap_or_default() +} + impl settings::Settings for AllLanguageModelSettings { const PRESERVED_KEYS: Option<&'static [&'static str]> = Some(&["version"]); @@ -52,9 +64,19 @@ impl settings::Settings for AllLanguageModelSettings { anthropic: AnthropicSettings { api_url: anthropic.api_url.unwrap(), available_models: anthropic.available_models.unwrap_or_default(), + custom_headers: custom_headers_from( + "Anthropic", + anthropic.custom_headers, + anthropic::RESERVED_HEADER_NAMES, + ), }, bedrock: AmazonBedrockSettings { available_models: bedrock.available_models.unwrap_or_default(), + custom_headers: custom_headers_from( + "Amazon Bedrock", + bedrock.custom_headers, + bedrock::RESERVED_HEADER_NAMES, + ), region: bedrock.region, endpoint: bedrock.endpoint_url, // todo(should be api_url) profile_name: bedrock.profile, @@ -67,28 +89,42 @@ impl settings::Settings for AllLanguageModelSettings { deepseek: DeepSeekSettings { api_url: deepseek.api_url.unwrap(), available_models: deepseek.available_models.unwrap_or_default(), + custom_headers: custom_headers_from("DeepSeek", deepseek.custom_headers, &[]), }, google: GoogleSettings { api_url: google.api_url.unwrap(), available_models: google.available_models.unwrap_or_default(), + custom_headers: custom_headers_from("Google AI", google.custom_headers, &[]), }, lmstudio: LmStudioSettings { api_url: lmstudio.api_url.unwrap(), available_models: lmstudio.available_models.unwrap_or_default(), + custom_headers: custom_headers_from("LM Studio", lmstudio.custom_headers, &[]), }, mistral: MistralSettings { api_url: mistral.api_url.unwrap(), available_models: mistral.available_models.unwrap_or_default(), + custom_headers: custom_headers_from( + "Mistral", + mistral.custom_headers, + mistral::RESERVED_HEADER_NAMES, + ), }, ollama: OllamaSettings { api_url: ollama.api_url.unwrap(), auto_discover: ollama.auto_discover.unwrap_or(true), available_models: ollama.available_models.unwrap_or_default(), context_window: ollama.context_window, + custom_headers: custom_headers_from("Ollama", ollama.custom_headers, &[]), }, opencode: OpenCodeSettings { api_url: opencode.api_url.unwrap(), available_models: opencode.available_models.unwrap_or_default(), + custom_headers: custom_headers_from( + "OpenCode", + opencode.custom_headers, + opencode::RESERVED_HEADER_NAMES, + ), show_zen_models: opencode.show_zen_models.unwrap_or(true), show_go_models: opencode.show_go_models.unwrap_or(true), show_free_models: opencode.show_free_models.unwrap_or(true), @@ -96,19 +132,31 @@ impl settings::Settings for AllLanguageModelSettings { open_router: OpenRouterSettings { api_url: open_router.api_url.unwrap(), available_models: open_router.available_models.unwrap_or_default(), + custom_headers: custom_headers_from( + "OpenRouter", + open_router.custom_headers, + open_router::RESERVED_HEADER_NAMES, + ), }, openai: OpenAiSettings { api_url: openai.api_url.unwrap(), available_models: openai.available_models.unwrap_or_default(), + custom_headers: custom_headers_from("OpenAI", openai.custom_headers, &[]), }, openai_compatible: openai_compatible .into_iter() .map(|(key, value)| { + let provider_label = format!("OpenAI Compatible ({key})"); ( key, OpenAiCompatibleSettings { api_url: value.api_url, available_models: value.available_models, + custom_headers: custom_headers_from( + &provider_label, + value.custom_headers, + &[], + ), }, ) }) @@ -116,10 +164,16 @@ impl settings::Settings for AllLanguageModelSettings { vercel_ai_gateway: VercelAiGatewaySettings { api_url: vercel_ai_gateway.api_url.unwrap(), available_models: vercel_ai_gateway.available_models.unwrap_or_default(), + custom_headers: custom_headers_from( + "Vercel AI Gateway", + vercel_ai_gateway.custom_headers, + &[], + ), }, x_ai: XAiSettings { api_url: x_ai.api_url.unwrap(), available_models: x_ai.available_models.unwrap_or_default(), + custom_headers: custom_headers_from("xAI", x_ai.custom_headers, &[]), }, zed_dot_dev: ZedDotDevSettings { available_models: zed_dot_dev.available_models.unwrap_or_default(), diff --git a/crates/lmstudio/src/lmstudio.rs b/crates/lmstudio/src/lmstudio.rs index 8a44b7fdefe526..57963bbb040c07 100644 --- a/crates/lmstudio/src/lmstudio.rs +++ b/crates/lmstudio/src/lmstudio.rs @@ -1,6 +1,8 @@ use anyhow::{Context as _, Result, anyhow}; use futures::{AsyncBufReadExt, AsyncReadExt, StreamExt, io::BufReader, stream::BoxStream}; -use http_client::{AsyncBody, HttpClient, Method, Request as HttpRequest, http}; +use http_client::{ + AsyncBody, CustomHeaders, HttpClient, Method, Request as HttpRequest, RequestBuilderExt, http, +}; use serde::{Deserialize, Serialize}; use serde_json::Value; use std::{convert::TryFrom, time::Duration}; @@ -356,6 +358,7 @@ pub async fn complete( api_url: &str, api_key: Option<&str>, request: ChatCompletionRequest, + extra_headers: &CustomHeaders, ) -> Result { let uri = format!("{api_url}/chat/completions"); let mut request_builder = HttpRequest::builder() @@ -368,7 +371,9 @@ pub async fn complete( } let serialized_request = serde_json::to_string(&request)?; - let request = request_builder.body(AsyncBody::from(serialized_request))?; + let request = request_builder + .extra_headers(extra_headers) + .body(AsyncBody::from(serialized_request))?; let mut response = client.send(request).await?; if response.status().is_success() { @@ -393,6 +398,7 @@ pub async fn stream_chat_completion( api_url: &str, api_key: Option<&str>, request: ChatCompletionRequest, + extra_headers: &CustomHeaders, ) -> Result>> { let uri = format!("{api_url}/chat/completions"); let mut request_builder = http::Request::builder() @@ -404,7 +410,9 @@ pub async fn stream_chat_completion( request_builder = request_builder.header("Authorization", format!("Bearer {}", api_key)); } - let request = request_builder.body(AsyncBody::from(serde_json::to_string(&request)?))?; + let request = request_builder + .extra_headers(extra_headers) + .body(AsyncBody::from(serde_json::to_string(&request)?))?; let mut response = client.send(request).await?; if response.status().is_success() { let reader = BufReader::new(response.into_body()); @@ -446,6 +454,7 @@ pub async fn get_models( api_url: &str, api_key: Option<&str>, _: Option, + extra_headers: &CustomHeaders, ) -> Result> { let uri = format!("{api_url}/models"); let mut request_builder = HttpRequest::builder() @@ -457,7 +466,9 @@ pub async fn get_models( request_builder = request_builder.header("Authorization", format!("Bearer {}", api_key)); } - let request = request_builder.body(AsyncBody::default())?; + let request = request_builder + .extra_headers(extra_headers) + .body(AsyncBody::default())?; let mut response = client.send(request).await?; diff --git a/crates/mistral/src/mistral.rs b/crates/mistral/src/mistral.rs index e7a3ee421eb7c3..c0bcfcf4e49345 100644 --- a/crates/mistral/src/mistral.rs +++ b/crates/mistral/src/mistral.rs @@ -1,6 +1,9 @@ use anyhow::{Result, anyhow}; use futures::{AsyncBufReadExt, AsyncReadExt, StreamExt, io::BufReader, stream::BoxStream}; -use http_client::{AsyncBody, HttpClient, HttpRequestExt, Method, Request as HttpRequest}; +use http_client::{ + AsyncBody, CustomHeaders, HttpClient, HttpRequestExt, Method, Request as HttpRequest, + RequestBuilderExt, +}; use serde::{Deserialize, Serialize}; use serde_json::Value; use std::convert::TryFrom; @@ -425,6 +428,7 @@ pub async fn stream_completion( api_key: &str, request: Request, affinity: Option, + extra_headers: &CustomHeaders, ) -> Result>> { let uri = format!("{api_url}/chat/completions"); let request_builder = HttpRequest::builder() @@ -434,7 +438,8 @@ pub async fn stream_completion( .header("Authorization", format!("Bearer {}", api_key.trim())) .when_some(affinity, |this, affinity| { this.header("x-affinity", affinity) - }); + }) + .extra_headers(extra_headers); let request = request_builder.body(AsyncBody::from(serde_json::to_string(&request)?))?; let mut response = client.send(request).await?; diff --git a/crates/ollama/src/ollama.rs b/crates/ollama/src/ollama.rs index b8666c47e9e041..356bce5d74de9c 100644 --- a/crates/ollama/src/ollama.rs +++ b/crates/ollama/src/ollama.rs @@ -1,7 +1,10 @@ use anyhow::{Context, Result}; use futures::{AsyncBufReadExt, AsyncReadExt, StreamExt, io::BufReader, stream::BoxStream}; -use http_client::{AsyncBody, HttpClient, HttpRequestExt, Method, Request as HttpRequest}; +use http_client::{ + AsyncBody, CustomHeaders, HttpClient, HttpRequestExt, Method, Request as HttpRequest, + RequestBuilderExt, +}; use serde::{Deserialize, Serialize}; use serde_json::Value; pub use settings::KeepAlive; @@ -281,6 +284,7 @@ pub async fn stream_chat_completion( api_url: &str, api_key: Option<&str>, request: ChatRequest, + extra_headers: &CustomHeaders, ) -> Result>> { let uri = format!("{api_url}/api/chat"); let request = HttpRequest::builder() @@ -290,6 +294,7 @@ pub async fn stream_chat_completion( .when_some(api_key, |builder, api_key| { builder.header("Authorization", format!("Bearer {api_key}")) }) + .extra_headers(extra_headers) .body(AsyncBody::from(serde_json::to_string(&request)?))?; let mut response = client.send(request).await?; @@ -318,6 +323,7 @@ pub async fn get_models( client: &dyn HttpClient, api_url: &str, api_key: Option<&str>, + extra_headers: &CustomHeaders, ) -> Result> { let uri = format!("{api_url}/api/tags"); let request = HttpRequest::builder() @@ -327,6 +333,7 @@ pub async fn get_models( .when_some(api_key, |builder, api_key| { builder.header("Authorization", format!("Bearer {api_key}")) }) + .extra_headers(extra_headers) .body(AsyncBody::default())?; let mut response = client.send(request).await?; @@ -351,6 +358,7 @@ pub async fn show_model( api_url: &str, api_key: Option<&str>, model: &str, + extra_headers: &CustomHeaders, ) -> Result { let uri = format!("{api_url}/api/show"); let request = HttpRequest::builder() @@ -360,6 +368,7 @@ pub async fn show_model( .when_some(api_key, |builder, api_key| { builder.header("Authorization", format!("Bearer {api_key}")) }) + .extra_headers(extra_headers) .body(AsyncBody::from( serde_json::json!({ "model": model }).to_string(), ))?; diff --git a/crates/open_ai/src/open_ai.rs b/crates/open_ai/src/open_ai.rs index 0ff1308d52a0ec..f50948221ac53e 100644 --- a/crates/open_ai/src/open_ai.rs +++ b/crates/open_ai/src/open_ai.rs @@ -5,7 +5,8 @@ pub mod responses; use anyhow::{Context as _, Result, anyhow}; use futures::{AsyncBufReadExt, AsyncReadExt, StreamExt, io::BufReader, stream::BoxStream}; use http_client::{ - AsyncBody, HttpClient, Method, Request as HttpRequest, StatusCode, + AsyncBody, CustomHeaders, HttpClient, Method, Request as HttpRequest, RequestBuilderExt, + StatusCode, http::{HeaderMap, HeaderValue}, }; pub use language_model_core::ReasoningEffort; @@ -758,15 +759,15 @@ pub async fn stream_completion( api_url: &str, api_key: &str, request: Request, + extra_headers: &CustomHeaders, ) -> Result>, RequestError> { let uri = format!("{api_url}/chat/completions"); - let request_builder = HttpRequest::builder() + let request = HttpRequest::builder() .method(Method::POST) .uri(uri) .header("Content-Type", "application/json") - .header("Authorization", format!("Bearer {}", api_key.trim())); - - let request = request_builder + .header("Authorization", format!("Bearer {}", api_key.trim())) + .extra_headers(extra_headers) .body(AsyncBody::from( serde_json::to_string(&request).map_err(|e| RequestError::Other(e.into()))?, )) diff --git a/crates/open_ai/src/responses.rs b/crates/open_ai/src/responses.rs index 6cc05699254e17..fb9b44b1978e6b 100644 --- a/crates/open_ai/src/responses.rs +++ b/crates/open_ai/src/responses.rs @@ -1,6 +1,8 @@ use anyhow::{Result, anyhow}; use futures::{AsyncBufReadExt, AsyncReadExt, StreamExt, io::BufReader, stream::BoxStream}; -use http_client::{AsyncBody, HttpClient, Method, Request as HttpRequest}; +use http_client::{ + AsyncBody, CustomHeaders, HttpClient, Method, Request as HttpRequest, RequestBuilderExt, +}; use serde::{Deserialize, Serialize}; use serde_json::Value; @@ -439,20 +441,16 @@ pub async fn stream_response( api_url: &str, api_key: &str, request: Request, - extra_headers: Vec<(String, String)>, + extra_headers: &CustomHeaders, ) -> Result>, RequestError> { let uri = format!("{api_url}/responses"); - let mut request_builder = HttpRequest::builder() + let is_streaming = request.stream; + let request = HttpRequest::builder() .method(Method::POST) .uri(uri) .header("Content-Type", "application/json") - .header("Authorization", format!("Bearer {}", api_key.trim())); - for (name, value) in &extra_headers { - request_builder = request_builder.header(name.as_str(), value.as_str()); - } - - let is_streaming = request.stream; - let request = request_builder + .header("Authorization", format!("Bearer {}", api_key.trim())) + .extra_headers(extra_headers) .body(AsyncBody::from( serde_json::to_string(&request).map_err(|e| RequestError::Other(e.into()))?, )) diff --git a/crates/open_router/src/open_router.rs b/crates/open_router/src/open_router.rs index 83f744e101e4c8..01dbe69ed7df9e 100644 --- a/crates/open_router/src/open_router.rs +++ b/crates/open_router/src/open_router.rs @@ -1,6 +1,8 @@ use anyhow::{Result, anyhow}; use futures::{AsyncBufReadExt, AsyncReadExt, StreamExt, io::BufReader, stream::BoxStream}; -use http_client::{AsyncBody, HttpClient, Method, Request as HttpRequest, http}; +use http_client::{ + AsyncBody, CustomHeaders, HttpClient, Method, Request as HttpRequest, RequestBuilderExt, http, +}; use serde::{Deserialize, Serialize}; use serde_json::Value; pub use settings::DataCollection; @@ -440,17 +442,17 @@ pub async fn stream_completion( api_url: &str, api_key: &str, request: Request, + extra_headers: &CustomHeaders, ) -> Result>, OpenRouterError> { let uri = format!("{api_url}/chat/completions"); - let request_builder = HttpRequest::builder() + let request = HttpRequest::builder() .method(Method::POST) .uri(uri) .header("Content-Type", "application/json") .header("Authorization", format!("Bearer {}", api_key)) .header("HTTP-Referer", "https://zed.dev") - .header("X-Title", "Zed Editor"); - - let request = request_builder + .header("X-Title", "Zed Editor") + .extra_headers(extra_headers) .body(AsyncBody::from( serde_json::to_string(&request).map_err(OpenRouterError::SerializeRequest)?, )) @@ -533,17 +535,17 @@ pub async fn list_models( client: &dyn HttpClient, api_url: &str, api_key: &str, + extra_headers: &CustomHeaders, ) -> Result, OpenRouterError> { let uri = format!("{api_url}/models/user"); - let request_builder = HttpRequest::builder() + let request = HttpRequest::builder() .method(Method::GET) .uri(uri) .header("Accept", "application/json") .header("Authorization", format!("Bearer {}", api_key)) .header("HTTP-Referer", "https://zed.dev") - .header("X-Title", "Zed Editor"); - - let request = request_builder + .header("X-Title", "Zed Editor") + .extra_headers(extra_headers) .body(AsyncBody::default()) .map_err(OpenRouterError::BuildRequestBody)?; let mut response = client diff --git a/crates/opencode/src/opencode.rs b/crates/opencode/src/opencode.rs index 4dfc30c62b9957..91a89840bbc704 100644 --- a/crates/opencode/src/opencode.rs +++ b/crates/opencode/src/opencode.rs @@ -1,6 +1,8 @@ use anyhow::{Result, anyhow}; use futures::{AsyncBufReadExt, AsyncReadExt, StreamExt, io::BufReader, stream::BoxStream}; -use http_client::{AsyncBody, HttpClient, Method, Request as HttpRequest}; +use http_client::{ + AsyncBody, CustomHeaders, HttpClient, Method, Request as HttpRequest, RequestBuilderExt, +}; use language_model_core::ReasoningEffort; use serde::{Deserialize, Serialize}; use strum::EnumIter; @@ -642,6 +644,7 @@ pub async fn stream_generate_content( api_url: &str, api_key: &str, request: google_ai::GenerateContentRequest, + extra_headers: &CustomHeaders, ) -> Result>> { let api_key = api_key.trim(); @@ -649,13 +652,13 @@ pub async fn stream_generate_content( let uri = format!("{api_url}/v1/models/{model_id}:streamGenerateContent?alt=sse"); - let request_builder = HttpRequest::builder() + let request = HttpRequest::builder() .method(Method::POST) .uri(uri) .header("Content-Type", "application/json") - .header("Authorization", format!("Bearer {api_key}")); - - let request = request_builder.body(AsyncBody::from(serde_json::to_string(&request)?))?; + .header("Authorization", format!("Bearer {api_key}")) + .extra_headers(extra_headers) + .body(AsyncBody::from(serde_json::to_string(&request)?))?; let mut response = client.send(request).await?; if response.status().is_success() { let reader = BufReader::new(response.into_body()); diff --git a/crates/settings_content/src/language_model.rs b/crates/settings_content/src/language_model.rs index a1be3ec89bb08c..259856feade1dd 100644 --- a/crates/settings_content/src/language_model.rs +++ b/crates/settings_content/src/language_model.rs @@ -32,6 +32,7 @@ pub struct AllLanguageModelSettingsContent { pub struct AnthropicSettingsContent { pub api_url: Option, pub available_models: Option>, + pub custom_headers: Option>, } #[with_fallible_options] @@ -60,6 +61,7 @@ pub struct AnthropicAvailableModel { #[derive(Default, Clone, Debug, Serialize, Deserialize, PartialEq, JsonSchema, MergeFrom)] pub struct AmazonBedrockSettingsContent { pub available_models: Option>, + pub custom_headers: Option>, pub endpoint_url: Option, pub region: Option, pub profile: Option, @@ -104,6 +106,7 @@ pub struct OllamaSettingsContent { pub auto_discover: Option, pub available_models: Option>, pub context_window: Option, + pub custom_headers: Option>, } #[with_fallible_options] @@ -152,6 +155,7 @@ impl Default for KeepAlive { pub struct OpenCodeSettingsContent { pub api_url: Option, pub available_models: Option>, + pub custom_headers: Option>, /// Whether to show OpenCode Zen models. Defaults to true. pub show_zen_models: Option, /// Whether to show OpenCode Go models. Defaults to true. @@ -194,6 +198,7 @@ pub struct LmStudioSettingsContent { pub api_url: Option, pub api_key: Option, pub available_models: Option>, + pub custom_headers: Option>, } #[with_fallible_options] @@ -211,6 +216,7 @@ pub struct LmStudioAvailableModel { pub struct DeepseekSettingsContent { pub api_url: Option, pub available_models: Option>, + pub custom_headers: Option>, } #[with_fallible_options] @@ -227,6 +233,7 @@ pub struct DeepseekAvailableModel { pub struct MistralSettingsContent { pub api_url: Option, pub available_models: Option>, + pub custom_headers: Option>, } #[with_fallible_options] @@ -247,6 +254,7 @@ pub struct MistralAvailableModel { pub struct OpenAiSettingsContent { pub api_url: Option, pub available_models: Option>, + pub custom_headers: Option>, } #[with_fallible_options] @@ -275,6 +283,7 @@ impl MergeFrom for OpenAiReasoningEffort { pub struct OpenAiCompatibleSettingsContent { pub api_url: String, pub available_models: Vec, + pub custom_headers: Option>, } #[with_fallible_options] @@ -339,6 +348,7 @@ impl Default for OpenAiCompatibleModelCapabilities { pub struct VercelAiGatewaySettingsContent { pub api_url: Option, pub available_models: Option>, + pub custom_headers: Option>, } #[with_fallible_options] @@ -358,6 +368,7 @@ pub struct VercelAiGatewayAvailableModel { pub struct GoogleSettingsContent { pub api_url: Option, pub available_models: Option>, + pub custom_headers: Option>, } #[with_fallible_options] @@ -374,6 +385,7 @@ pub struct GoogleAvailableModel { pub struct XAiSettingsContent { pub api_url: Option, pub available_models: Option>, + pub custom_headers: Option>, } #[with_fallible_options] @@ -437,6 +449,7 @@ pub enum ZedDotDevAvailableProvider { pub struct OpenRouterSettingsContent { pub api_url: Option, pub available_models: Option>, + pub custom_headers: Option>, } #[with_fallible_options] diff --git a/docs/src/ai/llm-providers.md b/docs/src/ai/llm-providers.md index 0c8645648c830a..2185e5113bd93f 100644 --- a/docs/src/ai/llm-providers.md +++ b/docs/src/ai/llm-providers.md @@ -37,6 +37,27 @@ Zed supports these providers with your own API keys: - [Vercel AI Gateway](#vercel-ai-gateway) - [xAI](#xai) +### Custom Headers {#custom-headers} + +You can attach extra HTTP headers to every request Zed makes to supported HTTP-based providers. This is useful in corporate environments, or for observability tooling. + +Configure them via the `language_models..custom_headers` settings key: + +```json [settings] +{ + "language_models": { + "openai": { + "custom_headers": { + "Fancy-Auth": "Bearer ", + "X-My-Tag": "zed" + } + } + } +} +``` + +`custom_headers` is supported by Amazon Bedrock, Anthropic, DeepSeek, Google AI, LM Studio, Mistral, Ollama, OpenAI, OpenAI API Compatible providers, OpenCode, OpenRouter, Vercel AI Gateway, and xAI. Headers managed by Zed for each provider, such as `Authorization`, `Content-Type`, `Accept`, and provider-specific authentication headers, are ignored with a warning if you try to override them. + ### Amazon Bedrock {#amazon-bedrock} > Supports tool use with models that support streaming tool use.