1use crate::api::{ModelAliasSpec, ModelTask};
2use crate::error::{Result, RuntimeError};
3use crate::provider::remote_common::{
4 RemoteProviderBase, check_http_status, parse_openai_embeddings_response, resolve_api_key,
5};
6use crate::traits::{
7 EmbedResult, EmbeddingModel, GenerationOptions, GenerationResult, GeneratorModel,
8 LoadedModelHandle, Message, MessageRole, ModelProvider, ProviderCapabilities, ProviderHealth,
9 TokenUsage,
10};
11use async_trait::async_trait;
12use reqwest::Client;
13use serde_json::json;
14use std::sync::Arc;
15
16pub struct RemoteOpenAIProvider {
27 base: RemoteProviderBase,
28}
29
30impl Default for RemoteOpenAIProvider {
31 fn default() -> Self {
32 Self {
33 base: RemoteProviderBase::new(),
34 }
35 }
36}
37
38impl RemoteOpenAIProvider {
39 pub fn new() -> Self {
40 Self::default()
41 }
42
43 #[cfg(test)]
44 fn insert_test_breaker(&self, key: crate::api::ModelRuntimeKey, age: std::time::Duration) {
45 self.base.insert_test_breaker(key, age);
46 }
47
48 #[cfg(test)]
49 fn breaker_count(&self) -> usize {
50 self.base.breaker_count()
51 }
52
53 #[cfg(test)]
54 fn force_cleanup_now_for_test(&self) {
55 self.base.force_cleanup_now_for_test();
56 }
57}
58
59#[async_trait]
60impl ModelProvider for RemoteOpenAIProvider {
61 fn provider_id(&self) -> &'static str {
62 "remote/openai"
63 }
64
65 fn capabilities(&self) -> ProviderCapabilities {
66 ProviderCapabilities {
67 supported_tasks: vec![ModelTask::Embed, ModelTask::Generate],
68 }
69 }
70
71 async fn load(&self, spec: &ModelAliasSpec) -> Result<LoadedModelHandle> {
72 match spec.task {
73 ModelTask::Embed => {
74 let api_key = resolve_api_key(&spec.options, "api_key_env", "OPENAI_API_KEY")?;
75 let base_url = resolve_base_url(&spec.options);
76 let embedding_dimensions = spec
77 .options
78 .get("embedding_dimensions")
79 .and_then(|v| v.as_u64())
80 .map(|v| v as u32);
81 let default_dims = match spec.model_id.as_str() {
82 "text-embedding-3-large" => 3072,
83 _ => 1536,
84 };
85 let model = OpenAIEmbeddingModel {
86 client: self.base.client.clone(),
87 cb: self.base.circuit_breaker_for(spec),
88 model_id: spec.model_id.clone(),
89 api_key,
90 base_url,
91 dimensions: embedding_dimensions.unwrap_or(default_dims),
92 };
93 let handle: Arc<dyn EmbeddingModel> = Arc::new(model);
94 Ok(Arc::new(handle) as LoadedModelHandle)
95 }
96 ModelTask::Generate => {
97 let api_key = resolve_api_key(&spec.options, "api_key_env", "OPENAI_API_KEY")?;
98 let base_url = resolve_base_url(&spec.options);
99 let model = OpenAIGeneratorModel {
100 client: self.base.client.clone(),
101 cb: self.base.circuit_breaker_for(spec),
102 model_id: spec.model_id.clone(),
103 api_key,
104 base_url,
105 };
106 let handle: Arc<dyn GeneratorModel> = Arc::new(model);
107 Ok(Arc::new(handle) as LoadedModelHandle)
108 }
109 ModelTask::Raw => Err(RuntimeError::CapabilityMismatch(
110 "OpenAI provider does not support task Raw".to_string(),
111 )),
112 _ => Err(RuntimeError::CapabilityMismatch(format!(
113 "OpenAI provider does not support task {:?}",
114 spec.task
115 ))),
116 }
117 }
118
119 async fn health(&self) -> ProviderHealth {
120 ProviderHealth::Healthy
121 }
122}
123
124const DEFAULT_BASE_URL: &str = "https://api.openai.com/v1";
125
126fn resolve_base_url(options: &serde_json::Value) -> String {
129 let raw = options
130 .get("base_url")
131 .and_then(|v| v.as_str())
132 .unwrap_or(DEFAULT_BASE_URL);
133 raw.trim_end_matches('/').to_string()
134}
135
136pub struct OpenAIEmbeddingModel {
138 client: Client,
139 cb: crate::reliability::CircuitBreakerWrapper,
140 model_id: String,
141 api_key: String,
142 base_url: String,
143 dimensions: u32,
144}
145
146#[async_trait]
147impl EmbeddingModel for OpenAIEmbeddingModel {
148 async fn embed(&self, texts: &[&str]) -> Result<EmbedResult> {
149 let texts: Vec<String> = texts.iter().map(|s| s.to_string()).collect();
150
151 self.cb
152 .call(move || async move {
153 let response = self
154 .client
155 .post(format!("{}/embeddings", self.base_url))
156 .header("Authorization", format!("Bearer {}", self.api_key))
157 .json(&json!({
158 "model": self.model_id,
159 "input": texts
160 }))
161 .send()
162 .await
163 .map_err(|e| RuntimeError::ApiError(e.to_string()))?;
164
165 let body: serde_json::Value = check_http_status("OpenAI", response)?
166 .json()
167 .await
168 .map_err(|e| RuntimeError::ApiError(e.to_string()))?;
169
170 parse_openai_embeddings_response("OpenAI", &body, None)
173 })
174 .await
175 }
176
177 fn dimensions(&self) -> u32 {
178 self.dimensions
179 }
180}
181
182impl crate::traits::ModelInfo for OpenAIEmbeddingModel {
183 fn model_id(&self) -> &str {
184 &self.model_id
185 }
186}
187
188struct OpenAIGeneratorModel {
193 client: Client,
194 cb: crate::reliability::CircuitBreakerWrapper,
195 model_id: String,
196 api_key: String,
197 base_url: String,
198}
199
200impl crate::traits::ModelInfo for OpenAIGeneratorModel {
201 fn model_id(&self) -> &str {
202 &self.model_id
203 }
204}
205
206#[async_trait]
207impl GeneratorModel for OpenAIGeneratorModel {
208 async fn generate(
209 &self,
210 messages: &[Message],
211 options: GenerationOptions,
212 ) -> Result<GenerationResult> {
213 let messages: Vec<serde_json::Value> = messages
214 .iter()
215 .map(|msg| {
216 let role = match msg.role {
217 MessageRole::System => "system",
218 MessageRole::User => "user",
219 MessageRole::Assistant => "assistant",
220 };
221 json!({ "role": role, "content": msg.text() })
222 })
223 .collect();
224
225 self.cb
226 .call(move || async move {
227 let mut body = json!({
228 "model": self.model_id,
229 "messages": messages,
230 });
231
232 if let Some(max_tokens) = options.max_tokens {
233 body["max_completion_tokens"] = json!(max_tokens);
234 }
235 if let Some(temperature) = options.temperature {
236 body["temperature"] = json!(temperature);
237 }
238 if let Some(top_p) = options.top_p {
239 body["top_p"] = json!(top_p);
240 }
241
242 let response = self
243 .client
244 .post(format!("{}/chat/completions", self.base_url))
245 .header("Authorization", format!("Bearer {}", self.api_key))
246 .json(&body)
247 .send()
248 .await
249 .map_err(|e| RuntimeError::ApiError(e.to_string()))?;
250
251 let body: serde_json::Value = check_http_status("OpenAI", response)?
252 .json()
253 .await
254 .map_err(|e| RuntimeError::ApiError(e.to_string()))?;
255
256 let text = body["choices"][0]["message"]["content"]
257 .as_str()
258 .unwrap_or("")
259 .to_string();
260
261 let usage = body.get("usage").map(|u| TokenUsage {
262 prompt_tokens: u["prompt_tokens"].as_u64().unwrap_or(0) as usize,
263 completion_tokens: u["completion_tokens"].as_u64().unwrap_or(0) as usize,
264 total_tokens: u["total_tokens"].as_u64().unwrap_or(0) as usize,
265 });
266
267 Ok(GenerationResult {
268 text,
269 usage,
270 images: vec![],
271 audio: None,
272 })
273 })
274 .await
275 }
276}
277
278#[cfg(test)]
279mod tests {
280 use super::*;
281 use crate::api::ModelRuntimeKey;
282 use crate::provider::remote_common::RemoteProviderBase;
283 use crate::traits::ModelProvider;
284 use std::time::Duration;
285
286 static ENV_LOCK: tokio::sync::Mutex<()> = tokio::sync::Mutex::const_new(());
287
288 fn spec(alias: &str, task: ModelTask, model_id: &str) -> ModelAliasSpec {
289 ModelAliasSpec {
290 alias: alias.to_string(),
291 task,
292 provider_id: "remote/openai".to_string(),
293 model_id: model_id.to_string(),
294 revision: None,
295 warmup: crate::api::WarmupPolicy::Lazy,
296 required: false,
297 timeout: None,
298 load_timeout: None,
299 retry: None,
300 options: serde_json::Value::Null,
301 }
302 }
303
304 #[tokio::test]
305 async fn breaker_reused_for_same_runtime_key() {
306 let _lock = ENV_LOCK.lock().await;
307 unsafe { std::env::set_var("OPENAI_API_KEY", "test-key") };
309
310 let provider = RemoteOpenAIProvider::new();
311 let s1 = spec("embed/a", ModelTask::Embed, "text-embedding-3-small");
312 let s2 = spec("embed/b", ModelTask::Embed, "text-embedding-3-small");
313
314 let _ = provider.load(&s1).await.unwrap();
315 let _ = provider.load(&s2).await.unwrap();
316
317 assert_eq!(provider.breaker_count(), 1);
318
319 unsafe { std::env::remove_var("OPENAI_API_KEY") };
321 }
322
323 #[tokio::test]
324 async fn breaker_isolated_by_task_and_model() {
325 let _lock = ENV_LOCK.lock().await;
326 unsafe { std::env::set_var("OPENAI_API_KEY", "test-key") };
328
329 let provider = RemoteOpenAIProvider::new();
330 let embed = spec("embed/a", ModelTask::Embed, "text-embedding-3-small");
331 let gen_spec = spec("chat/a", ModelTask::Generate, "gpt-4o-mini");
332
333 let _ = provider.load(&embed).await.unwrap();
334 let _ = provider.load(&gen_spec).await.unwrap();
335
336 assert_eq!(provider.breaker_count(), 2);
337
338 unsafe { std::env::remove_var("OPENAI_API_KEY") };
340 }
341
342 #[tokio::test]
343 async fn breaker_cleanup_evicts_stale_entries() {
344 let _lock = ENV_LOCK.lock().await;
345 unsafe { std::env::set_var("OPENAI_API_KEY", "test-key") };
347
348 let provider = RemoteOpenAIProvider::new();
349 let stale = spec("embed/stale", ModelTask::Embed, "text-embedding-3-small");
350 let fresh = spec("embed/fresh", ModelTask::Embed, "text-embedding-3-large");
351 provider.insert_test_breaker(
352 ModelRuntimeKey::new(&stale),
353 RemoteProviderBase::BREAKER_TTL + Duration::from_secs(5),
354 );
355 provider.insert_test_breaker(ModelRuntimeKey::new(&fresh), Duration::from_secs(1));
356 assert_eq!(provider.breaker_count(), 2);
357
358 provider.force_cleanup_now_for_test();
359 let _ = provider.load(&fresh).await.unwrap();
360
361 assert_eq!(provider.breaker_count(), 1);
362
363 unsafe { std::env::remove_var("OPENAI_API_KEY") };
365 }
366
367 #[tokio::test]
368 async fn default_embedding_dimensions() {
369 let _lock = ENV_LOCK.lock().await;
370 unsafe { std::env::set_var("OPENAI_API_KEY", "test-key") };
372
373 let provider = RemoteOpenAIProvider::new();
374 let s = spec("embed/dim", ModelTask::Embed, "text-embedding-3-small");
375
376 let handle = provider.load(&s).await.unwrap();
377 let model = handle.downcast_ref::<Arc<dyn EmbeddingModel>>().unwrap();
378 assert_eq!(model.dimensions(), 1536);
379
380 unsafe { std::env::remove_var("OPENAI_API_KEY") };
382 }
383
384 #[tokio::test]
385 async fn custom_embedding_dimensions() {
386 let _lock = ENV_LOCK.lock().await;
387 unsafe { std::env::set_var("OPENAI_API_KEY", "test-key") };
389
390 let provider = RemoteOpenAIProvider::new();
391 let mut s = spec(
392 "embed/dim-custom",
393 ModelTask::Embed,
394 "text-embedding-3-small",
395 );
396 s.options = serde_json::json!({"embedding_dimensions": 256});
397
398 let handle = provider.load(&s).await.unwrap();
399 let model = handle.downcast_ref::<Arc<dyn EmbeddingModel>>().unwrap();
400 assert_eq!(model.dimensions(), 256);
401
402 unsafe { std::env::remove_var("OPENAI_API_KEY") };
404 }
405
406 #[test]
407 fn resolve_base_url_defaults_to_openai() {
408 assert_eq!(
409 resolve_base_url(&serde_json::Value::Null),
410 "https://api.openai.com/v1"
411 );
412 assert_eq!(
413 resolve_base_url(&serde_json::json!({})),
414 "https://api.openai.com/v1"
415 );
416 }
417
418 #[test]
419 fn resolve_base_url_uses_custom_value() {
420 assert_eq!(
421 resolve_base_url(&serde_json::json!({"base_url": "http://localhost:8000/v1"})),
422 "http://localhost:8000/v1"
423 );
424 }
425
426 #[test]
427 fn resolve_base_url_strips_trailing_slash() {
428 assert_eq!(
429 resolve_base_url(&serde_json::json!({"base_url": "http://localhost:8000/v1/"})),
430 "http://localhost:8000/v1"
431 );
432 }
433
434 #[tokio::test]
435 async fn load_accepts_custom_base_url() {
436 let _lock = ENV_LOCK.lock().await;
437 unsafe { std::env::set_var("OPENAI_API_KEY", "test-key") };
439
440 let provider = RemoteOpenAIProvider::new();
441 let mut s = spec("embed/local", ModelTask::Embed, "text-embedding-3-small");
442 s.options = serde_json::json!({"base_url": "http://localhost:8000/v1"});
443
444 assert!(provider.load(&s).await.is_ok());
445
446 unsafe { std::env::remove_var("OPENAI_API_KEY") };
448 }
449}