litellm-rs 0.6.0

A high-performance AI Gateway written in Rust, providing OpenAI-compatible APIs with intelligent routing, load balancing, and enterprise features
Documentation
//! OpenAI Provider Configuration
//!
//! Unified configuration system following the base provider pattern

use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::time::Duration;

use crate::core::net::ProviderEndpointAccess;
use crate::core::providers::base::BaseConfig;
use crate::core::traits::provider::ProviderConfig;

/// OpenAI provider configuration
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct OpenAIConfig {
    /// Base configuration shared across all providers
    #[serde(flatten)]
    pub base: BaseConfig,

    /// Gateway provider name for router deployment identity.
    #[serde(default = "default_provider_name")]
    pub provider_name: String,

    /// OpenAI-specific configuration
    /// Organization ID (optional)
    pub organization: Option<String>,

    /// Project ID (optional)  
    pub project: Option<String>,

    /// Custom model mappings
    pub model_mappings: HashMap<String, String>,

    /// Feature flags
    pub features: OpenAIFeatures,
}

/// OpenAI feature configuration
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct OpenAIFeatures {
    /// Enable O-series model optimizations
    pub o_series_optimizations: bool,

    /// Enable GPT-5 specific features
    pub gpt5_features: bool,

    /// Enable audio model support
    pub audio_models: bool,

    /// Enable DALL-E image generation
    pub image_generation: bool,

    /// Enable Whisper transcription
    pub audio_transcription: bool,

    /// Enable fine-tuning capabilities
    pub fine_tuning: bool,

    /// Enable vector store integration
    pub vector_stores: bool,

    /// Enable real-time audio (beta)
    pub realtime_audio: bool,
}

fn default_provider_name() -> String {
    "openai".to_string()
}

pub(crate) fn is_official_openai_endpoint(raw_url: &str) -> bool {
    url::Url::parse(raw_url).is_ok_and(|url| {
        url.host_str().is_some_and(|host| {
            host.trim_end_matches('.')
                .eq_ignore_ascii_case("api.openai.com")
        })
    })
}

pub(crate) fn validate_private_official_openai_endpoint(
    access: ProviderEndpointAccess,
    endpoint: Option<&str>,
) -> Result<(), &'static str> {
    if access == ProviderEndpointAccess::PrivateNetwork
        && endpoint.is_some_and(is_official_openai_endpoint)
    {
        return Err("private_network access cannot target the official OpenAI endpoint");
    }
    Ok(())
}

impl Default for OpenAIFeatures {
    fn default() -> Self {
        Self {
            o_series_optimizations: true,
            gpt5_features: false, // Beta feature
            audio_models: true,
            image_generation: true,
            audio_transcription: true,
            fine_tuning: false,    // Enterprise feature
            vector_stores: false,  // Enterprise feature
            realtime_audio: false, // Beta feature
        }
    }
}

impl Default for OpenAIConfig {
    fn default() -> Self {
        Self {
            base: BaseConfig {
                api_base: Some("https://api.openai.com/v1".to_string()),
                ..Default::default()
            },
            provider_name: default_provider_name(),
            organization: None,
            project: None,
            model_mappings: HashMap::new(),
            features: OpenAIFeatures::default(),
        }
    }
}

impl OpenAIConfig {
    /// Create configuration from environment variables
    pub fn from_env() -> Self {
        let mut config = Self::default();

        // API Key
        if let Ok(api_key) = std::env::var("OPENAI_API_KEY") {
            config.base.api_key = Some(api_key);
        }

        // Organization
        if let Ok(org) = std::env::var("OPENAI_ORG_ID") {
            config.organization = Some(org);
        }

        // Project
        if let Ok(project) = std::env::var("OPENAI_PROJECT_ID") {
            config.project = Some(project);
        }

        // Base URL
        if let Ok(base_url) = std::env::var("OPENAI_API_BASE") {
            config.base.api_base = Some(base_url);
        }

        // Timeout
        if let Ok(timeout_str) = std::env::var("OPENAI_TIMEOUT")
            && let Ok(timeout) = timeout_str.parse::<u64>()
        {
            config.base.timeout = timeout;
        }

        config
    }

    /// Validate the configuration
    pub fn validate(&self) -> Result<(), String> {
        // Validate base config
        self.base.validate("openai")?;
        validate_private_official_openai_endpoint(
            self.base.endpoint_access,
            self.base.api_base.as_deref(),
        )
        .map_err(str::to_owned)?;

        // OpenAI specific validations
        if let Some(ref api_key) = self.base.api_key
            && !api_key.starts_with("sk-")
            && !api_key.starts_with("sk-proj-")
        {
            return Err("OpenAI API key must start with 'sk-' or 'sk-proj-'".to_string());
        }

        if let Some(ref org) = self.organization
            && org.is_empty()
        {
            return Err("Organization ID cannot be empty".to_string());
        }

        if let Some(ref project) = self.project
            && project.is_empty()
        {
            return Err("Project ID cannot be empty".to_string());
        }

        Ok(())
    }

    /// Get the effective API base URL
    pub fn get_api_base(&self) -> String {
        self.base
            .api_base
            .as_ref()
            .unwrap_or(&"https://api.openai.com/v1".to_string())
            .clone()
    }

    /// Check if a feature is enabled
    pub fn is_feature_enabled(&self, feature: OpenAIFeature) -> bool {
        match feature {
            OpenAIFeature::OSeriesOptimizations => self.features.o_series_optimizations,
            OpenAIFeature::GPT5Features => self.features.gpt5_features,
            OpenAIFeature::AudioModels => self.features.audio_models,
            OpenAIFeature::ImageGeneration => self.features.image_generation,
            OpenAIFeature::AudioTranscription => self.features.audio_transcription,
            OpenAIFeature::FineTuning => self.features.fine_tuning,
            OpenAIFeature::VectorStores => self.features.vector_stores,
            OpenAIFeature::RealtimeAudio => self.features.realtime_audio,
        }
    }

    /// Get model mapping for a given model name
    pub fn get_model_mapping(&self, model: &str) -> String {
        self.model_mappings
            .get(model)
            .unwrap_or(&model.to_string())
            .clone()
    }
}

/// OpenAI feature enumeration
#[derive(Debug, Clone, PartialEq)]
pub enum OpenAIFeature {
    OSeriesOptimizations,
    GPT5Features,
    AudioModels,
    ImageGeneration,
    AudioTranscription,
    FineTuning,
    VectorStores,
    RealtimeAudio,
}

impl ProviderConfig for OpenAIConfig {
    fn validate(&self) -> Result<(), String> {
        self.validate()
    }

    fn api_key(&self) -> Option<&str> {
        self.base.api_key.as_deref()
    }

    fn api_base(&self) -> Option<&str> {
        self.base.api_base.as_deref()
    }

    fn timeout(&self) -> Duration {
        Duration::from_secs(self.base.timeout)
    }

    fn max_retries(&self) -> u32 {
        self.base.max_retries
    }

    fn endpoint_access(&self) -> ProviderEndpointAccess {
        self.base.endpoint_access
    }
}

#[cfg(test)]
pub(crate) fn test_openai_config(
    api_base: impl Into<String>,
    api_key: impl Into<String>,
) -> OpenAIConfig {
    let mut config = OpenAIConfig::default();
    config.base.api_base = Some(api_base.into());
    config.base.api_key = Some(api_key.into());
    config.base.endpoint_access = ProviderEndpointAccess::PrivateNetwork;
    config
}

#[cfg(test)]
mod tests {
    use super::*;

    #[test]
    fn test_default_config() {
        let config = OpenAIConfig::default();
        // Remove provider_name check as it doesn't exist in BaseConfig
        assert_eq!(config.get_api_base(), "https://api.openai.com/v1");
        assert!(config.features.image_generation);
        assert_eq!(
            ProviderConfig::endpoint_access(&config),
            ProviderEndpointAccess::PublicOnly
        );
    }

    #[test]
    fn test_endpoint_access_flattened_serde_round_trip() {
        let mut value = serde_json::to_value(OpenAIConfig::default()).unwrap();
        value["endpoint_access"] = serde_json::json!("private_network");
        let config: OpenAIConfig = serde_json::from_value(value).unwrap();
        assert_eq!(
            ProviderConfig::endpoint_access(&config),
            ProviderEndpointAccess::PrivateNetwork
        );
        assert_eq!(
            serde_json::to_value(config).unwrap()["endpoint_access"],
            "private_network"
        );
    }

    #[test]
    fn private_access_rejects_the_official_openai_authority() {
        let mut config = OpenAIConfig::default();
        config.base.api_key = Some("sk-test".to_string());
        config.base.endpoint_access = ProviderEndpointAccess::PrivateNetwork;
        let error = config
            .validate()
            .expect_err("official OpenAI must stay public");
        assert!(error.contains("official OpenAI"), "{error}");

        config.base.api_base = Some("http://127.0.0.1:18080/v1".to_string());
        assert!(config.validate().is_ok());

        config.base.api_base = Some("https://api.openai.com/v1".to_string());
        config.base.endpoint_access = ProviderEndpointAccess::PublicOnly;
        assert!(config.validate().is_ok());
    }

    #[test]
    fn test_config_validation() {
        let mut config = OpenAIConfig::default();

        // Should fail without API key
        assert!(config.validate().is_err());

        // Should pass with valid API key
        config.base.api_key = Some("sk-test123".to_string());
        assert!(config.validate().is_ok());

        // Should fail with invalid API key
        config.base.api_key = Some("invalid-key".to_string());
        assert!(config.validate().is_err());
    }

    #[test]
    fn test_feature_flags() {
        let config = OpenAIConfig::default();

        assert!(config.is_feature_enabled(OpenAIFeature::ImageGeneration));
        assert!(!config.is_feature_enabled(OpenAIFeature::GPT5Features));
        assert!(!config.is_feature_enabled(OpenAIFeature::RealtimeAudio));
    }

    #[test]
    fn test_model_mapping() {
        let mut config = OpenAIConfig::default();
        config
            .model_mappings
            .insert("gpt-4".to_string(), "gpt-4-0613".to_string());

        assert_eq!(config.get_model_mapping("gpt-4"), "gpt-4-0613");
        assert_eq!(config.get_model_mapping("gpt-3.5-turbo"), "gpt-3.5-turbo");
    }
}