Skip to main content

lc_embeddings/
openai.rs

1// lc-embeddings/src/openai.rs
2//! OpenAI Embeddings implementation
3//!
4//! Uses OpenAI's text-embedding-ada-002 or other embedding models.
5
6use crate::{EmbeddingError, Embeddings};
7use async_trait::async_trait;
8use serde::Deserialize;
9
10/// OpenAI Embeddings configuration
11#[derive(Debug, Clone)]
12pub struct OpenAIEmbeddingsConfig {
13    /// API key
14    pub api_key: String,
15
16    /// API base URL
17    pub base_url: String,
18
19    /// Model name (default: text-embedding-ada-002)
20    pub model: String,
21
22    /// Batch size (default: 2048)
23    pub batch_size: usize,
24}
25
26impl Default for OpenAIEmbeddingsConfig {
27    fn default() -> Self {
28        Self {
29            api_key: std::env::var("OPENAI_API_KEY").unwrap_or_default(),
30            base_url: "https://api.openai.com/v1".to_string(),
31            model: "text-embedding-ada-002".to_string(),
32            batch_size: 2048,
33        }
34    }
35}
36
37impl OpenAIEmbeddingsConfig {
38    /// Create a new configuration
39    pub fn new(api_key: impl Into<String>) -> Self {
40        Self {
41            api_key: api_key.into(),
42            ..Default::default()
43        }
44    }
45
46    /// Set the model
47    pub fn with_model(mut self, model: impl Into<String>) -> Self {
48        self.model = model.into();
49        self
50    }
51
52    /// Set the base URL
53    pub fn with_base_url(mut self, url: impl Into<String>) -> Self {
54        self.base_url = url.into();
55        self
56    }
57}
58
59/// OpenAI Embeddings client
60pub struct OpenAIEmbeddings {
61    config: OpenAIEmbeddingsConfig,
62    client: reqwest::Client,
63    dimension: usize,
64}
65
66impl OpenAIEmbeddings {
67    /// Create a new OpenAI Embeddings client
68    pub fn new(config: OpenAIEmbeddingsConfig) -> Self {
69        // Determine dimension based on model
70        let dimension = match config.model.as_str() {
71            "text-embedding-ada-002" => 1536,
72            "text-embedding-3-small" => 1536,
73            "text-embedding-3-large" => 3072,
74            _ => 1536, // Default dimension
75        };
76
77        Self {
78            config,
79            client: reqwest::Client::new(),
80            dimension,
81        }
82    }
83
84    /// Creates OpenAIEmbeddings from environment variables.
85    #[deprecated(
86        since = "0.7.0",
87        note = "Use from_env_result() which returns Result<Self, String>"
88    )]
89    #[allow(deprecated)]
90    pub fn from_env() -> Self {
91        Self::from_env_result().unwrap_or_else(|_| Self::new(OpenAIEmbeddingsConfig::default()))
92    }
93
94    /// Creates OpenAIEmbeddings from environment variables, returning a Result.
95    ///
96    /// Environment variables:
97    /// - `OPENAI_API_KEY`: API key (required)
98    /// - `OPENAI_BASE_URL`: API endpoint (optional)
99    /// - `OPENAI_EMBED_MODEL`: Model name (optional)
100    pub fn from_env_result() -> Result<Self, String> {
101        let api_key = std::env::var("OPENAI_API_KEY")
102            .map_err(|_| "OPENAI_API_KEY environment variable not set".to_string())?;
103        let base_url = std::env::var("OPENAI_BASE_URL")
104            .unwrap_or_else(|_| "https://api.openai.com/v1".to_string());
105        let model = std::env::var("OPENAI_EMBED_MODEL")
106            .unwrap_or_else(|_| "text-embedding-ada-002".to_string());
107        Ok(Self::new(OpenAIEmbeddingsConfig {
108            api_key,
109            base_url,
110            model,
111            batch_size: 2048,
112        }))
113    }
114}
115
116#[async_trait]
117impl Embeddings for OpenAIEmbeddings {
118    async fn embed_query(&self, text: &str) -> Result<Vec<f32>, EmbeddingError> {
119        if text.is_empty() {
120            return Err(EmbeddingError::EmptyInput);
121        }
122
123        let url = format!("{}/embeddings", self.config.base_url);
124
125        let body = serde_json::json!({
126            "model": self.config.model,
127            "input": text,
128        });
129
130        let response = self
131            .client
132            .post(&url)
133            .header("Authorization", format!("Bearer {}", self.config.api_key))
134            .header("Content-Type", "application/json")
135            .json(&body)
136            .send()
137            .await
138            .map_err(|e| EmbeddingError::HttpError(e.to_string()))?;
139
140        let status = response.status();
141        if !status.is_success() {
142            let error_text = response.text().await.unwrap_or_default();
143            return Err(EmbeddingError::ApiError(format!(
144                "HTTP {}: {}",
145                status, error_text
146            )));
147        }
148
149        let embedding_response: OpenAIEmbeddingResponse = response
150            .json()
151            .await
152            .map_err(|e| EmbeddingError::ParseError(e.to_string()))?;
153
154        Ok(embedding_response
155            .data
156            .first()
157            .ok_or_else(|| EmbeddingError::ApiError("No embedding data in response".to_string()))?
158            .embedding
159            .clone())
160    }
161
162    async fn embed_documents(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>, EmbeddingError> {
163        if texts.is_empty() {
164            return Ok(Vec::new());
165        }
166
167        let url = format!("{}/embeddings", self.config.base_url);
168        let batch_size = self.config.batch_size.max(1);
169        let mut all_results = vec![Vec::new(); texts.len()];
170        let mut offset = 0;
171
172        // Call API in batches of batch_size
173        for chunk in texts.chunks(batch_size) {
174            let body = serde_json::json!({
175                "model": self.config.model,
176                "input": chunk,
177            });
178
179            let response = self
180                .client
181                .post(&url)
182                .header("Authorization", format!("Bearer {}", self.config.api_key))
183                .header("Content-Type", "application/json")
184                .json(&body)
185                .send()
186                .await
187                .map_err(|e| EmbeddingError::HttpError(e.to_string()))?;
188
189            let status = response.status();
190            if !status.is_success() {
191                let error_text = response.text().await.unwrap_or_default();
192                return Err(EmbeddingError::ApiError(format!(
193                    "HTTP {}: {}",
194                    status, error_text
195                )));
196            }
197
198            let embedding_response: OpenAIEmbeddingResponse = response
199                .json()
200                .await
201                .map_err(|e| EmbeddingError::ParseError(e.to_string()))?;
202
203            for item in embedding_response.data {
204                let global_index = offset + item.index as usize;
205                if global_index < all_results.len() {
206                    all_results[global_index] = item.embedding;
207                }
208            }
209            offset += chunk.len();
210        }
211
212        Ok(all_results)
213    }
214
215    fn dimension(&self) -> usize {
216        self.dimension
217    }
218
219    fn model_name(&self) -> &str {
220        &self.config.model
221    }
222}
223
224/// OpenAI Embedding API response
225#[derive(Debug, Deserialize)]
226#[allow(dead_code)]
227struct OpenAIEmbeddingResponse {
228    data: Vec<OpenAIEmbeddingData>,
229    model: String,
230    usage: OpenAIEmbeddingUsage,
231}
232
233#[derive(Debug, Deserialize)]
234#[allow(dead_code)]
235struct OpenAIEmbeddingData {
236    embedding: Vec<f32>,
237    index: i32,
238    object: String,
239}
240
241#[derive(Debug, Deserialize)]
242#[allow(dead_code)]
243struct OpenAIEmbeddingUsage {
244    prompt_tokens: usize,
245    total_tokens: usize,
246}
247
248#[cfg(test)]
249mod tests_env {
250    use super::*;
251    use std::env;
252
253    fn save_and_set(key: &str, value: &str) -> Option<String> {
254        let old = env::var(key).ok();
255        env::set_var(key, value);
256        old
257    }
258
259    fn restore(key: &str, old: Option<String>) {
260        match old {
261            Some(v) => env::set_var(key, v),
262            None => env::remove_var(key),
263        }
264    }
265
266    #[test]
267    fn test_from_env_result_ok_when_key_set() {
268        let _lock = crate::ENV_TEST_LOCK
269            .lock()
270            .unwrap_or_else(|e| e.into_inner());
271        let old = save_and_set("OPENAI_API_KEY", "test-key-123");
272        let result = OpenAIEmbeddings::from_env_result();
273        assert!(result.is_ok());
274        restore("OPENAI_API_KEY", old);
275    }
276
277    #[test]
278    fn test_from_env_result_err_when_key_missing() {
279        let _lock = crate::ENV_TEST_LOCK
280            .lock()
281            .unwrap_or_else(|e| e.into_inner());
282        let old = env::var("OPENAI_API_KEY").ok();
283        env::remove_var("OPENAI_API_KEY");
284        let result = OpenAIEmbeddings::from_env_result();
285        match result {
286            Err(msg) => assert!(msg.contains("OPENAI_API_KEY")),
287            Ok(_) => panic!("expected error when OPENAI_API_KEY is missing"),
288        }
289        restore("OPENAI_API_KEY", old);
290    }
291
292    #[test]
293    fn test_from_env_result_uses_optional_vars() {
294        let _lock = crate::ENV_TEST_LOCK
295            .lock()
296            .unwrap_or_else(|e| e.into_inner());
297        let old_key = save_and_set("OPENAI_API_KEY", "key");
298        let old_url = save_and_set("OPENAI_BASE_URL", "https://custom.api.com/v1");
299        let old_model = save_and_set("OPENAI_EMBED_MODEL", "text-embedding-3-small");
300        let embeddings = OpenAIEmbeddings::from_env_result().unwrap();
301        assert_eq!(embeddings.model_name(), "text-embedding-3-small");
302        restore("OPENAI_API_KEY", old_key);
303        restore("OPENAI_BASE_URL", old_url);
304        restore("OPENAI_EMBED_MODEL", old_model);
305    }
306
307    #[test]
308    fn test_from_env_result_uses_defaults_for_optional_vars() {
309        let _lock = crate::ENV_TEST_LOCK
310            .lock()
311            .unwrap_or_else(|e| e.into_inner());
312        let old_key = save_and_set("OPENAI_API_KEY", "key");
313        let old_url = env::var("OPENAI_BASE_URL").ok();
314        env::remove_var("OPENAI_BASE_URL");
315        let old_model = env::var("OPENAI_EMBED_MODEL").ok();
316        env::remove_var("OPENAI_EMBED_MODEL");
317        let embeddings = OpenAIEmbeddings::from_env_result().unwrap();
318        assert_eq!(embeddings.model_name(), "text-embedding-ada-002");
319        restore("OPENAI_API_KEY", old_key);
320        restore("OPENAI_BASE_URL", old_url);
321        restore("OPENAI_EMBED_MODEL", old_model);
322    }
323}
324
325#[cfg(test)]
326mod tests {
327    use super::*;
328
329    #[test]
330    fn test_config_default() {
331        let config = OpenAIEmbeddingsConfig::default();
332        assert_eq!(config.model, "text-embedding-ada-002");
333        assert_eq!(config.batch_size, 2048);
334    }
335
336    #[test]
337    fn test_config_builder() {
338        let config = OpenAIEmbeddingsConfig::new("test-key")
339            .with_model("text-embedding-3-large")
340            .with_base_url("https://custom.api.com/v1");
341
342        assert_eq!(config.api_key, "test-key");
343        assert_eq!(config.model, "text-embedding-3-large");
344        assert_eq!(config.base_url, "https://custom.api.com/v1");
345    }
346
347    #[tokio::test]
348    #[ignore = "requires real API call"]
349    async fn test_real_embedding() {
350        let config = OpenAIEmbeddingsConfig {
351            api_key: "sk-6eb65fcf5d17491ca10b984efe1f43e7".to_string(),
352            base_url:
353                "https://llm-8xo1b7o30z27y2xc.cn-beijing.maas.aliyuncs.com/compatible-mode/v1"
354                    .to_string(),
355            model: "text-embedding-ada-002".to_string(),
356            batch_size: 2048,
357        };
358
359        let embeddings = OpenAIEmbeddings::new(config);
360
361        let result = embeddings.embed_query("Hello, world!").await;
362        assert!(result.is_ok());
363
364        let embedding = result.unwrap();
365        assert_eq!(embedding.len(), 1536);
366    }
367}