Support custom headers for openai provider (#1801)

This commit is contained in:
Tarun Boddupalli
2025-03-26 14:22:44 -04:00
committed by GitHub
parent 4f2e193554
commit 5bbd4c4691
2 changed files with 36 additions and 1 deletions
+25
View File
@@ -2,6 +2,7 @@ use anyhow::Result;
use async_trait::async_trait;
use reqwest::Client;
use serde_json::Value;
use std::collections::HashMap;
use std::time::Duration;
use super::base::{ConfigKey, Provider, ProviderMetadata, ProviderUsage, Usage};
@@ -33,6 +34,7 @@ pub struct OpenAiProvider {
organization: Option<String>,
project: Option<String>,
model: ModelConfig,
custom_headers: Option<HashMap<String, String>>,
}
impl Default for OpenAiProvider {
@@ -54,6 +56,10 @@ impl OpenAiProvider {
.unwrap_or_else(|_| "v1/chat/completions".to_string());
let organization: Option<String> = config.get_param("OPENAI_ORGANIZATION").ok();
let project: Option<String> = config.get_param("OPENAI_PROJECT").ok();
let custom_headers: Option<HashMap<String, String>> = config
.get_secret("OPENAI_CUSTOM_HEADERS")
.ok()
.map(parse_custom_headers);
let timeout_secs: u64 = config.get_param("OPENAI_TIMEOUT").unwrap_or(600);
let client = Client::builder()
.timeout(Duration::from_secs(timeout_secs))
@@ -67,6 +73,7 @@ impl OpenAiProvider {
organization,
project,
model,
custom_headers,
})
}
@@ -92,6 +99,12 @@ impl OpenAiProvider {
request = request.header("OpenAI-Project", project);
}
if let Some(custom_headers) = &self.custom_headers {
for (key, value) in custom_headers {
request = request.header(key, value);
}
}
let response = request.json(&payload).send().await?;
handle_response_openai_compat(response).await
@@ -117,6 +130,7 @@ impl Provider for OpenAiProvider {
ConfigKey::new("OPENAI_BASE_PATH", true, false, Some("v1/chat/completions")),
ConfigKey::new("OPENAI_ORGANIZATION", false, false, None),
ConfigKey::new("OPENAI_PROJECT", false, false, None),
ConfigKey::new("OPENAI_CUSTOM_HEADERS", false, true, None),
ConfigKey::new("OPENAI_TIMEOUT", false, false, Some("600")),
],
)
@@ -156,3 +170,14 @@ impl Provider for OpenAiProvider {
Ok((message, ProviderUsage::new(model, usage)))
}
}
fn parse_custom_headers(s: String) -> HashMap<String, String> {
s.split(',')
.filter_map(|header| {
let mut parts = header.splitn(2, '=');
let key = parts.next().map(|s| s.trim().to_string())?;
let value = parts.next().map(|s| s.trim().to_string())?;
Some((key, value))
})
.collect()
}