Skip to main content

cera_client/
provider.rs

1//! Provider definitions and URL / header resolution for remote endpoints.
2
3use reqwest::Url;
4use reqwest::header::{AUTHORIZATION, HeaderMap, HeaderName, HeaderValue};
5
6use crate::error::ClientError;
7
8/// Default base URL for the OpenAI API.
9pub const OPENAI_DEFAULT_BASE_URL: &str = "https://api.openai.com/v1";
10
11/// Default base URL for the OpenRouter API.
12pub const OPENROUTER_DEFAULT_BASE_URL: &str = "https://openrouter.ai/api/v1";
13
14/// Environment variable for OpenAI API key.
15pub const OPENAI_API_KEY_ENV: &str = "OPENAI_API_KEY";
16
17/// Environment variable for OpenRouter API key.
18pub const OPENROUTER_API_KEY_ENV: &str = "OPENROUTER_API_KEY";
19
20/// Environment variable for custom OpenAI base URL.
21pub const OPENAI_BASE_URL_ENV: &str = "OPENAI_BASE_URL";
22
23/// Endpoint provider target.
24#[derive(Debug, Clone, PartialEq, Eq)]
25pub enum Provider {
26    /// Official OpenAI API endpoint (`https://api.openai.com/v1`).
27    OpenAi,
28
29    /// OpenRouter API endpoint (`https://openrouter.ai/api/v1`) with optional attribution headers.
30    OpenRouter {
31        /// Optional site URL sent as `HTTP-Referer` for OpenRouter rankings.
32        app_url: Option<String>,
33        /// Optional site name sent as `X-Title` for OpenRouter rankings.
34        app_name: Option<String>,
35    },
36
37    /// Any custom OpenAI-compatible server (e.g. vLLM, Ollama, Groq, Together, DeepSeek).
38    Custom {
39        /// Base URL for the API (e.g. `http://localhost:8000/v1`).
40        base_url: String,
41    },
42}
43
44impl Provider {
45    /// Create an OpenRouter provider with default settings.
46    pub fn openrouter() -> Self {
47        Self::OpenRouter {
48            app_url: None,
49            app_name: None,
50        }
51    }
52
53    /// Create an OpenRouter provider with site URL and application name for OpenRouter rankings.
54    pub fn openrouter_with_attribution(
55        app_url: impl Into<String>,
56        app_name: impl Into<String>,
57    ) -> Self {
58        Self::OpenRouter {
59            app_url: Some(app_url.into()),
60            app_name: Some(app_name.into()),
61        }
62    }
63
64    /// Create a custom OpenAI-compatible provider with a specific base URL.
65    pub fn custom(base_url: impl Into<String>) -> Self {
66        Self::Custom {
67            base_url: base_url.into(),
68        }
69    }
70
71    /// Returns the base URL string for this provider.
72    pub fn base_url(&self) -> &str {
73        match self {
74            Self::OpenAi => OPENAI_DEFAULT_BASE_URL,
75            Self::OpenRouter { .. } => OPENROUTER_DEFAULT_BASE_URL,
76            Self::Custom { base_url } => base_url.as_str(),
77        }
78    }
79
80    /// Returns the canonical environment variable name for this provider's API key.
81    pub fn default_env_var(&self) -> Option<&'static str> {
82        match self {
83            Self::OpenAi => Some(OPENAI_API_KEY_ENV),
84            Self::OpenRouter { .. } => Some(OPENROUTER_API_KEY_ENV),
85            Self::Custom { .. } => None,
86        }
87    }
88
89    /// Constructs a full URL for a given relative path (e.g. `/chat/completions`).
90    pub fn endpoint_url(&self, path: &str) -> Result<Url, ClientError> {
91        let base = self.base_url().trim_end_matches('/');
92        let subpath = path.trim_start_matches('/');
93        let full = format!("{base}/{subpath}");
94        Url::parse(&full).map_err(|e| ClientError::InvalidUrl(format!("{full}: {e}")))
95    }
96
97    /// Injects authentication and provider-specific headers into a request header map.
98    pub fn apply_headers(
99        &self,
100        headers: &mut HeaderMap,
101        api_key: Option<&str>,
102    ) -> Result<(), ClientError> {
103        if let Some(key) = api_key {
104            let auth_value = format!("Bearer {key}");
105            let mut val = HeaderValue::from_str(&auth_value)
106                .map_err(|e| ClientError::InvalidHeader(format!("Invalid auth header: {e}")))?;
107            val.set_sensitive(true);
108            headers.insert(AUTHORIZATION, val);
109        }
110
111        if let Self::OpenRouter { app_url, app_name } = self {
112            if let Some(url) = app_url {
113                match HeaderValue::from_str(url) {
114                    Ok(val) => {
115                        headers.insert(HeaderName::from_static("http-referer"), val);
116                    }
117                    Err(e) => {
118                        tracing::warn!(target: "cera_client", error = %e, "Invalid http-referer header value; omitting");
119                    }
120                }
121            }
122            if let Some(name) = app_name {
123                match HeaderValue::from_str(name) {
124                    Ok(val) => {
125                        headers.insert(HeaderName::from_static("x-title"), val);
126                    }
127                    Err(e) => {
128                        tracing::warn!(target: "cera_client", error = %e, "Invalid x-title header value; omitting");
129                    }
130                }
131            }
132        }
133
134        Ok(())
135    }
136}
137
138#[cfg(test)]
139mod tests {
140    use super::*;
141
142    #[test]
143    fn test_openai_endpoints_and_headers() {
144        let provider = Provider::OpenAi;
145        assert_eq!(provider.base_url(), "https://api.openai.com/v1");
146        assert_eq!(provider.default_env_var(), Some("OPENAI_API_KEY"));
147
148        let url = provider.endpoint_url("chat/completions").unwrap();
149        assert_eq!(url.as_str(), "https://api.openai.com/v1/chat/completions");
150
151        let mut headers = HeaderMap::new();
152        provider
153            .apply_headers(&mut headers, Some("sk-test123"))
154            .unwrap();
155        assert_eq!(
156            headers.get(AUTHORIZATION).unwrap().to_str().unwrap(),
157            "Bearer sk-test123"
158        );
159        assert!(!headers.contains_key("http-referer"));
160    }
161
162    #[test]
163    fn test_openrouter_endpoints_and_attribution_headers() {
164        let provider = Provider::openrouter_with_attribution("https://myapp.example.com", "MyApp");
165        assert_eq!(provider.base_url(), "https://openrouter.ai/api/v1");
166        assert_eq!(provider.default_env_var(), Some("OPENROUTER_API_KEY"));
167
168        let url = provider.endpoint_url("/embeddings").unwrap();
169        assert_eq!(url.as_str(), "https://openrouter.ai/api/v1/embeddings");
170
171        let mut headers = HeaderMap::new();
172        provider
173            .apply_headers(&mut headers, Some("or-key-abc"))
174            .unwrap();
175        assert_eq!(
176            headers.get(AUTHORIZATION).unwrap().to_str().unwrap(),
177            "Bearer or-key-abc"
178        );
179        assert_eq!(
180            headers.get("http-referer").unwrap().to_str().unwrap(),
181            "https://myapp.example.com"
182        );
183        assert_eq!(headers.get("x-title").unwrap().to_str().unwrap(), "MyApp");
184    }
185
186    #[test]
187    fn test_custom_provider() {
188        let provider = Provider::custom("http://127.0.0.1:8000/v1/");
189        assert_eq!(provider.base_url(), "http://127.0.0.1:8000/v1/");
190        assert_eq!(provider.default_env_var(), None);
191
192        let url = provider.endpoint_url("models").unwrap();
193        assert_eq!(url.as_str(), "http://127.0.0.1:8000/v1/models");
194
195        let mut headers = HeaderMap::new();
196        provider.apply_headers(&mut headers, None).unwrap();
197        assert!(!headers.contains_key(AUTHORIZATION));
198    }
199
200    #[test]
201    fn test_invalid_auth_header() {
202        let provider = Provider::OpenAi;
203        let mut headers = HeaderMap::new();
204        let err = provider
205            .apply_headers(&mut headers, Some("key_with_\ninvalid_newline"))
206            .unwrap_err();
207        match err {
208            ClientError::InvalidHeader(msg) => assert!(msg.contains("Invalid auth header")),
209            other => panic!("expected InvalidHeader error, got {other:?}"),
210        }
211    }
212}