Skip to main content

uni_xervo/provider/
openai.rs

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
16/// Remote provider that calls the [OpenAI API](https://platform.openai.com/docs/api-reference)
17/// for embedding (`/embeddings`) and text generation (`/chat/completions`).
18///
19/// Requires the `OPENAI_API_KEY` environment variable (or a custom env var name
20/// via the `api_key_env` option).
21///
22/// Set the `base_url` option to target an OpenAI-compatible server (OpenRouter,
23/// vLLM, LM Studio, Ollama, internal proxies). The value should include the
24/// version path segment (e.g. `http://localhost:8000/v1`). Defaults to
25/// `https://api.openai.com/v1` when unset.
26pub 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
126/// Resolve `base_url` from options, falling back to the OpenAI default and
127/// stripping a trailing `/` so callers can append `/embeddings` etc. directly.
128fn 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
136/// Embedding model backed by the OpenAI embeddings API.
137pub 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                // Lenient decode (historical behaviour): see
171                // `remote_common::parse_openai_embeddings_response`.
172                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
188// ---------------------------------------------------------------------------
189// Generator
190// ---------------------------------------------------------------------------
191
192struct 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        // SAFETY: protected by ENV_LOCK
308        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        // SAFETY: protected by ENV_LOCK
320        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        // SAFETY: protected by ENV_LOCK
327        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        // SAFETY: protected by ENV_LOCK
339        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        // SAFETY: protected by ENV_LOCK
346        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        // SAFETY: protected by ENV_LOCK
364        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        // SAFETY: protected by ENV_LOCK
371        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        // SAFETY: protected by ENV_LOCK
381        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        // SAFETY: protected by ENV_LOCK
388        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        // SAFETY: protected by ENV_LOCK
403        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        // SAFETY: protected by ENV_LOCK
438        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        // SAFETY: protected by ENV_LOCK
447        unsafe { std::env::remove_var("OPENAI_API_KEY") };
448    }
449}