1use crate::{EmbeddingError, Embeddings};
7use async_trait::async_trait;
8use serde::Deserialize;
9
10#[derive(Debug, Clone)]
12pub struct OpenAIEmbeddingsConfig {
13 pub api_key: String,
15
16 pub base_url: String,
18
19 pub model: String,
21
22 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 pub fn new(api_key: impl Into<String>) -> Self {
40 Self {
41 api_key: api_key.into(),
42 ..Default::default()
43 }
44 }
45
46 pub fn with_model(mut self, model: impl Into<String>) -> Self {
48 self.model = model.into();
49 self
50 }
51
52 pub fn with_base_url(mut self, url: impl Into<String>) -> Self {
54 self.base_url = url.into();
55 self
56 }
57}
58
59pub struct OpenAIEmbeddings {
61 config: OpenAIEmbeddingsConfig,
62 client: reqwest::Client,
63 dimension: usize,
64}
65
66impl OpenAIEmbeddings {
67 pub fn new(config: OpenAIEmbeddingsConfig) -> Self {
69 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, };
76
77 Self {
78 config,
79 client: reqwest::Client::new(),
80 dimension,
81 }
82 }
83
84 #[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 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 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#[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}