1use reqwest::Url;
4use reqwest::header::{AUTHORIZATION, HeaderMap, HeaderName, HeaderValue};
5
6use crate::error::ClientError;
7
8pub const OPENAI_DEFAULT_BASE_URL: &str = "https://api.openai.com/v1";
10
11pub const OPENROUTER_DEFAULT_BASE_URL: &str = "https://openrouter.ai/api/v1";
13
14pub const OPENAI_API_KEY_ENV: &str = "OPENAI_API_KEY";
16
17pub const OPENROUTER_API_KEY_ENV: &str = "OPENROUTER_API_KEY";
19
20pub const OPENAI_BASE_URL_ENV: &str = "OPENAI_BASE_URL";
22
23#[derive(Debug, Clone, PartialEq, Eq)]
25pub enum Provider {
26 OpenAi,
28
29 OpenRouter {
31 app_url: Option<String>,
33 app_name: Option<String>,
35 },
36
37 Custom {
39 base_url: String,
41 },
42}
43
44impl Provider {
45 pub fn openrouter() -> Self {
47 Self::OpenRouter {
48 app_url: None,
49 app_name: None,
50 }
51 }
52
53 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 pub fn custom(base_url: impl Into<String>) -> Self {
66 Self::Custom {
67 base_url: base_url.into(),
68 }
69 }
70
71 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 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 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 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}