use std::env;
use reqwest::header::HeaderMap;
use crate::request_options::RequestOptions;
const DEFAULT_BASE_URL: &str = "https://api.openai.com/v1";
const DEFAULT_TIMEOUT_SECS: u64 = 600;
const DEFAULT_MAX_RETRIES: u32 = 2;
pub trait Config: Send + Sync + std::fmt::Debug {
fn base_url(&self) -> &str;
fn api_key(&self) -> &str;
fn build_request(&self, request: reqwest::RequestBuilder) -> reqwest::RequestBuilder;
fn organization(&self) -> Option<&str> {
None
}
fn project(&self) -> Option<&str> {
None
}
fn timeout_secs(&self) -> u64 {
DEFAULT_TIMEOUT_SECS
}
fn max_retries(&self) -> u32 {
DEFAULT_MAX_RETRIES
}
fn default_headers(&self) -> Option<&HeaderMap> {
None
}
fn default_query(&self) -> Option<&[(String, String)]> {
None
}
fn initial_options(&self) -> RequestOptions {
let mut opts = RequestOptions::new();
if let Some(h) = self.default_headers() {
opts.headers = Some(h.clone());
}
if let Some(q) = self.default_query() {
opts.query = Some(q.to_vec());
}
opts
}
}
#[derive(Debug, Clone)]
pub struct ClientConfig {
pub api_key: String,
pub base_url: String,
pub organization: Option<String>,
pub project: Option<String>,
pub timeout_secs: u64,
pub max_retries: u32,
pub default_headers: Option<HeaderMap>,
pub default_query: Option<Vec<(String, String)>>,
pub(crate) use_azure_api_key_header: bool,
}
impl ClientConfig {
pub fn new(api_key: impl Into<String>) -> Self {
Self {
api_key: api_key.into(),
base_url: DEFAULT_BASE_URL.to_string(),
organization: None,
project: None,
timeout_secs: DEFAULT_TIMEOUT_SECS,
max_retries: DEFAULT_MAX_RETRIES,
default_headers: None,
default_query: None,
use_azure_api_key_header: false,
}
}
pub fn from_env() -> Result<Self, crate::error::OpenAIError> {
Self::from_env_with(|key| env::var(key).ok())
}
fn from_env_with(
lookup: impl Fn(&str) -> Option<String>,
) -> Result<Self, crate::error::OpenAIError> {
let api_key = lookup("OPENAI_API_KEY").ok_or_else(|| {
crate::error::OpenAIError::InvalidArgument(
"OPENAI_API_KEY environment variable not set".to_string(),
)
})?;
let mut config = Self::new(api_key);
let set = |key: &str| {
lookup(key)
.map(|v| v.trim().to_string())
.filter(|v| !v.is_empty())
};
if let Some(base_url) = set("OPENAI_BASE_URL") {
config = config.base_url(base_url);
}
if let Some(organization) = set("OPENAI_ORG_ID") {
config = config.organization(organization);
}
if let Some(project) = set("OPENAI_PROJECT_ID") {
config = config.project(project);
}
Ok(config)
}
pub fn base_url(mut self, url: impl Into<String>) -> Self {
let url = url.into();
self.base_url = url.trim_end_matches('/').to_string();
self
}
pub fn organization(mut self, org: impl Into<String>) -> Self {
self.organization = Some(org.into());
self
}
pub fn project(mut self, project: impl Into<String>) -> Self {
self.project = Some(project.into());
self
}
pub fn timeout_secs(mut self, secs: u64) -> Self {
self.timeout_secs = secs;
self
}
pub fn max_retries(mut self, retries: u32) -> Self {
self.max_retries = retries;
self
}
pub fn default_headers(mut self, headers: HeaderMap) -> Self {
self.default_headers = Some(headers);
self
}
pub fn default_query(mut self, query: Vec<(String, String)>) -> Self {
self.default_query = Some(query);
self
}
pub(crate) fn use_azure_api_key_header(mut self, enabled: bool) -> Self {
self.use_azure_api_key_header = enabled;
self
}
}
impl Config for ClientConfig {
fn base_url(&self) -> &str {
&self.base_url
}
fn api_key(&self) -> &str {
&self.api_key
}
fn build_request(&self, mut req: reqwest::RequestBuilder) -> reqwest::RequestBuilder {
if self.use_azure_api_key_header {
req = req.header("api-key", &self.api_key);
} else {
req = req.bearer_auth(&self.api_key);
}
if let Some(ref org) = self.organization {
req = req.header("OpenAI-Organization", org);
}
if let Some(ref project) = self.project {
req = req.header("OpenAI-Project", project);
}
req
}
fn organization(&self) -> Option<&str> {
self.organization.as_deref()
}
fn project(&self) -> Option<&str> {
self.project.as_deref()
}
fn timeout_secs(&self) -> u64 {
self.timeout_secs
}
fn max_retries(&self) -> u32 {
self.max_retries
}
fn default_headers(&self) -> Option<&HeaderMap> {
self.default_headers.as_ref()
}
fn default_query(&self) -> Option<&[(String, String)]> {
self.default_query.as_deref()
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::collections::HashMap;
fn env_of(pairs: &[(&str, &str)]) -> impl Fn(&str) -> Option<String> + use<> {
let map: HashMap<String, String> = pairs
.iter()
.map(|(k, v)| ((*k).to_string(), (*v).to_string()))
.collect();
move |key| map.get(key).cloned()
}
#[test]
fn api_key_only_keeps_the_default_base_url() {
let cfg = ClientConfig::from_env_with(env_of(&[("OPENAI_API_KEY", "sk-test")])).unwrap();
assert_eq!(cfg.api_key, "sk-test");
assert_eq!(cfg.base_url, DEFAULT_BASE_URL);
assert_eq!(cfg.organization, None);
assert_eq!(cfg.project, None);
}
#[test]
fn base_url_org_and_project_come_from_the_environment() {
let cfg = ClientConfig::from_env_with(env_of(&[
("OPENAI_API_KEY", "sk-test"),
("OPENAI_BASE_URL", "https://proxy.internal/v1"),
("OPENAI_ORG_ID", "org-123"),
("OPENAI_PROJECT_ID", "proj-456"),
]))
.unwrap();
assert_eq!(cfg.base_url, "https://proxy.internal/v1");
assert_eq!(cfg.organization.as_deref(), Some("org-123"));
assert_eq!(cfg.project.as_deref(), Some("proj-456"));
}
#[test]
fn blank_and_padded_values_are_ignored_or_trimmed() {
let cfg = ClientConfig::from_env_with(env_of(&[
("OPENAI_API_KEY", "sk-test"),
("OPENAI_BASE_URL", " "),
("OPENAI_ORG_ID", " org-pad "),
]))
.unwrap();
assert_eq!(cfg.base_url, DEFAULT_BASE_URL);
assert_eq!(cfg.organization.as_deref(), Some("org-pad"));
}
#[test]
fn a_trailing_slash_would_double_up_in_request_paths() {
let cfg = ClientConfig::from_env_with(env_of(&[
("OPENAI_API_KEY", "sk-test"),
("OPENAI_BASE_URL", "https://proxy.internal/v1/"),
]))
.unwrap();
assert_eq!(cfg.base_url, "https://proxy.internal/v1");
assert_eq!(
ClientConfig::new("k").base_url("https://x/v1//").base_url,
"https://x/v1"
);
}
#[test]
fn an_empty_api_key_is_accepted_for_keyless_proxies() {
let cfg = ClientConfig::from_env_with(env_of(&[
("OPENAI_API_KEY", ""),
("OPENAI_BASE_URL", "http://localhost:11434/v1"),
]))
.unwrap();
assert_eq!(cfg.api_key, "");
assert_eq!(cfg.base_url, "http://localhost:11434/v1");
}
#[test]
fn a_missing_api_key_is_an_error() {
let err = ClientConfig::from_env_with(env_of(&[("OPENAI_BASE_URL", "https://x")]));
assert!(err.is_err());
}
}