1use crate::{EmbeddingError, Embeddings};
7use async_trait::async_trait;
8use futures_util::StreamExt;
9use serde::Deserialize;
10
11const MAX_CONCURRENT_CHUNKS: usize = 8;
13
14#[derive(Debug, Clone)]
16pub struct OpenAIEmbeddingsConfig {
17 pub api_key: String,
19
20 pub base_url: String,
22
23 pub model: String,
25
26 pub batch_size: usize,
28}
29
30impl Default for OpenAIEmbeddingsConfig {
31 fn default() -> Self {
32 Self {
33 api_key: std::env::var("OPENAI_API_KEY").unwrap_or_default(),
34 base_url: "https://api.openai.com/v1".to_string(),
35 model: "text-embedding-ada-002".to_string(),
36 batch_size: 2048,
37 }
38 }
39}
40
41impl OpenAIEmbeddingsConfig {
42 pub fn new(api_key: impl Into<String>) -> Self {
44 Self {
45 api_key: api_key.into(),
46 ..Default::default()
47 }
48 }
49
50 pub fn with_model(mut self, model: impl Into<String>) -> Self {
52 self.model = model.into();
53 self
54 }
55
56 pub fn with_base_url(mut self, url: impl Into<String>) -> Self {
58 self.base_url = url.into();
59 self
60 }
61}
62
63pub struct OpenAIEmbeddings {
65 config: OpenAIEmbeddingsConfig,
66 client: reqwest::Client,
67 dimension: usize,
68}
69
70impl OpenAIEmbeddings {
71 pub fn new(config: OpenAIEmbeddingsConfig) -> Result<Self, EmbeddingError> {
77 if config.api_key.trim().is_empty() {
78 return Err(EmbeddingError::Config(
79 "OPENAI_API_KEY is empty".to_string(),
80 ));
81 }
82 let dimension = Self::dimension_for(&config.model)?;
83
84 Ok(Self {
85 config,
86 client: reqwest::Client::new(),
87 dimension,
88 })
89 }
90
91 fn dimension_for(model: &str) -> Result<usize, EmbeddingError> {
93 match model {
94 "text-embedding-ada-002" => Ok(1536),
95 "text-embedding-3-small" => Ok(1536),
96 "text-embedding-3-large" => Ok(3072),
97 other => Err(EmbeddingError::Config(format!(
98 "unknown embedding dimension for OpenAI model '{other}' \
99 (supported: 'text-embedding-ada-002', 'text-embedding-3-small', \
100 'text-embedding-3-large')"
101 ))),
102 }
103 }
104
105 pub fn from_env_result() -> Result<Self, EmbeddingError> {
112 let api_key = std::env::var("OPENAI_API_KEY").map_err(|_| {
113 EmbeddingError::Config("OPENAI_API_KEY environment variable not set".to_string())
114 })?;
115 let base_url = std::env::var("OPENAI_BASE_URL")
116 .unwrap_or_else(|_| "https://api.openai.com/v1".to_string());
117 let model = std::env::var("OPENAI_EMBED_MODEL")
118 .unwrap_or_else(|_| "text-embedding-ada-002".to_string());
119 Self::new(OpenAIEmbeddingsConfig {
120 api_key,
121 base_url,
122 model,
123 batch_size: 2048,
124 })
125 }
126}
127
128#[async_trait]
129impl Embeddings for OpenAIEmbeddings {
130 async fn embed_query(&self, text: &str) -> Result<Vec<f32>, EmbeddingError> {
131 if text.trim().is_empty() {
132 return Err(EmbeddingError::EmptyInput);
133 }
134
135 let url = format!("{}/embeddings", self.config.base_url);
136
137 let body = serde_json::json!({
138 "model": self.config.model,
139 "input": text,
140 });
141
142 let response = crate::retry::post_json_with_retry(
144 &self.client,
145 &url,
146 &self.config.api_key,
147 &body,
148 &crate::retry::DEFAULT_RETRY,
149 )
150 .await
151 .map_err(|e| EmbeddingError::HttpError(e.to_string()))?;
152
153 let status = response.status();
154 if !status.is_success() {
155 let error_text = response.text().await.map_err(|e| {
157 EmbeddingError::HttpError(format!("failed to read error response body: {e}"))
158 })?;
159 return Err(EmbeddingError::ApiError(format!(
160 "HTTP {}: {}",
161 status, error_text
162 )));
163 }
164
165 let embedding_response: OpenAIEmbeddingResponse = response
166 .json()
167 .await
168 .map_err(|e| EmbeddingError::ParseError(e.to_string()))?;
169
170 let mut embedding = embedding_response
171 .data
172 .first()
173 .ok_or_else(|| EmbeddingError::ApiError("No embedding data in response".to_string()))?
174 .embedding
175 .clone();
176 crate::l2_normalize(&mut embedding);
178 Ok(embedding)
179 }
180
181 async fn embed_documents(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>, EmbeddingError> {
182 if texts.is_empty() {
183 return Ok(Vec::new());
184 }
185 if texts.iter().any(|t| t.trim().is_empty()) {
187 return Err(EmbeddingError::EmptyInput);
188 }
189
190 let url = format!("{}/embeddings", self.config.base_url);
191 let batch_size = self.config.batch_size.max(1);
192 let concurrency = texts.len().div_ceil(batch_size).min(MAX_CONCURRENT_CHUNKS);
196
197 let chunks: Vec<(usize, Vec<String>)> = texts
205 .chunks(batch_size)
206 .enumerate()
207 .map(|(i, chunk)| (i, chunk.iter().map(|s| s.to_string()).collect()))
208 .collect();
209 let client = &self.client;
210 let api_key = self.config.api_key.as_str();
211 let model = self.config.model.as_str();
212 let url = url.as_str();
213 let mut all_results: Vec<Option<Vec<f32>>> = vec![None; texts.len()];
214 let mut stream = futures_util::stream::iter(chunks)
215 .map(|(chunk_idx, chunk)| async move {
216 let body = serde_json::json!({
217 "model": model,
218 "input": chunk,
219 });
220 let response = crate::retry::post_json_with_retry(
222 client,
223 url,
224 api_key,
225 &body,
226 &crate::retry::DEFAULT_RETRY,
227 )
228 .await
229 .map_err(|e| EmbeddingError::HttpError(e.to_string()))?;
230
231 let status = response.status();
232 if !status.is_success() {
233 let error_text = response.text().await.map_err(|e| {
235 EmbeddingError::HttpError(format!(
236 "failed to read error response body: {e}"
237 ))
238 })?;
239 return Err(EmbeddingError::ApiError(format!(
240 "HTTP {}: {}",
241 status, error_text
242 )));
243 }
244
245 let embedding_response: OpenAIEmbeddingResponse = response
246 .json()
247 .await
248 .map_err(|e| EmbeddingError::ParseError(e.to_string()))?;
249
250 Ok::<_, EmbeddingError>((chunk_idx, embedding_response.data))
251 })
252 .buffer_unordered(concurrency);
253 while let Some(result) = stream.next().await {
254 let (chunk_idx, data) = result?;
255 let base = chunk_idx * batch_size;
256 for item in data {
257 let global_index = base + item.index as usize;
258 if global_index >= all_results.len() {
259 return Err(EmbeddingError::BatchMismatch {
261 expected: all_results.len(),
262 actual: global_index + 1,
263 });
264 }
265 all_results[global_index] = Some(item.embedding);
266 }
267 }
268
269 all_results
271 .into_iter()
272 .map(|opt| {
273 let mut v = opt.ok_or(EmbeddingError::EmptyVectorInBatch)?;
274 crate::l2_normalize(&mut v);
275 Ok(v)
276 })
277 .collect()
278 }
279
280 fn dimension(&self) -> usize {
281 self.dimension
282 }
283
284 fn model_name(&self) -> &str {
285 &self.config.model
286 }
287}
288
289#[derive(Debug, Deserialize)]
291#[allow(dead_code)]
292struct OpenAIEmbeddingResponse {
293 data: Vec<OpenAIEmbeddingData>,
294 model: String,
295 usage: OpenAIEmbeddingUsage,
296}
297
298#[derive(Debug, Deserialize)]
299#[allow(dead_code)]
300struct OpenAIEmbeddingData {
301 embedding: Vec<f32>,
302 index: i32,
303 object: String,
304}
305
306#[derive(Debug, Deserialize)]
307#[allow(dead_code)]
308struct OpenAIEmbeddingUsage {
309 prompt_tokens: usize,
310 total_tokens: usize,
311}
312
313#[cfg(test)]
314mod tests_env {
315 use super::*;
316 use std::env;
317
318 fn save_and_set(key: &str, value: &str) -> Option<String> {
319 let old = env::var(key).ok();
320 env::set_var(key, value);
321 old
322 }
323
324 fn restore(key: &str, old: Option<String>) {
325 match old {
326 Some(v) => env::set_var(key, v),
327 None => env::remove_var(key),
328 }
329 }
330
331 #[test]
332 fn test_from_env_result_ok_when_key_set() {
333 let _lock = crate::ENV_TEST_LOCK
334 .lock()
335 .unwrap_or_else(|e| e.into_inner());
336 let old = save_and_set("OPENAI_API_KEY", "test-key-123");
337 let result = OpenAIEmbeddings::from_env_result();
338 assert!(result.is_ok());
339 restore("OPENAI_API_KEY", old);
340 }
341
342 #[test]
343 fn test_from_env_result_err_when_key_missing() {
344 let _lock = crate::ENV_TEST_LOCK
345 .lock()
346 .unwrap_or_else(|e| e.into_inner());
347 let old = env::var("OPENAI_API_KEY").ok();
348 env::remove_var("OPENAI_API_KEY");
349 let result = OpenAIEmbeddings::from_env_result();
350 match result {
351 Err(msg) => assert!(msg.to_string().contains("OPENAI_API_KEY")),
352 Ok(_) => panic!("expected error when OPENAI_API_KEY is missing"),
353 }
354 restore("OPENAI_API_KEY", old);
355 }
356
357 #[test]
358 fn test_from_env_result_uses_optional_vars() {
359 let _lock = crate::ENV_TEST_LOCK
360 .lock()
361 .unwrap_or_else(|e| e.into_inner());
362 let old_key = save_and_set("OPENAI_API_KEY", "key");
363 let old_url = save_and_set("OPENAI_BASE_URL", "https://custom.api.com/v1");
364 let old_model = save_and_set("OPENAI_EMBED_MODEL", "text-embedding-3-small");
365 let embeddings = OpenAIEmbeddings::from_env_result().unwrap();
366 assert_eq!(embeddings.model_name(), "text-embedding-3-small");
367 restore("OPENAI_API_KEY", old_key);
368 restore("OPENAI_BASE_URL", old_url);
369 restore("OPENAI_EMBED_MODEL", old_model);
370 }
371
372 #[test]
373 fn test_from_env_result_uses_defaults_for_optional_vars() {
374 let _lock = crate::ENV_TEST_LOCK
375 .lock()
376 .unwrap_or_else(|e| e.into_inner());
377 let old_key = save_and_set("OPENAI_API_KEY", "key");
378 let old_url = env::var("OPENAI_BASE_URL").ok();
379 env::remove_var("OPENAI_BASE_URL");
380 let old_model = env::var("OPENAI_EMBED_MODEL").ok();
381 env::remove_var("OPENAI_EMBED_MODEL");
382 let embeddings = OpenAIEmbeddings::from_env_result().unwrap();
383 assert_eq!(embeddings.model_name(), "text-embedding-ada-002");
384 restore("OPENAI_API_KEY", old_key);
385 restore("OPENAI_BASE_URL", old_url);
386 restore("OPENAI_EMBED_MODEL", old_model);
387 }
388}
389
390#[cfg(test)]
391mod tests {
392 use super::*;
393 use crate::test_support::{spawn_embeddings_stub, spawn_status_stub};
394 use std::sync::Arc;
395
396 #[tokio::test]
398 async fn test_embed_documents_batch_alignment() {
399 let base_url = spawn_embeddings_stub(Arc::new(|n| n)).await;
400 let config = OpenAIEmbeddingsConfig {
401 api_key: "test-key".into(),
402 base_url,
403 model: "text-embedding-ada-002".into(),
404 batch_size: 2,
405 };
406 let embeddings = OpenAIEmbeddings::new(config).unwrap();
407
408 let texts = ["a", "b", "c", "d", "e"];
412 let results = embeddings
413 .embed_documents(&texts)
414 .await
415 .expect("batch embedding should succeed");
416 assert_eq!(results.len(), 5);
417 for (i, text) in texts.iter().enumerate() {
418 let raw = text.bytes().map(|b| b as f32).sum::<f32>();
422 let mut expected = vec![raw, 1.0];
423 crate::l2_normalize(&mut expected);
424 assert_eq!(results[i], expected, "text #{} out of alignment", i);
425 }
426 }
427
428 #[tokio::test]
430 async fn test_embed_documents_truncated_response_errors() {
431 let base_url = spawn_embeddings_stub(Arc::new(|n| n.saturating_sub(1))).await;
432 let config = OpenAIEmbeddingsConfig {
433 api_key: "test-key".into(),
434 base_url,
435 model: "text-embedding-ada-002".into(),
436 batch_size: 2,
437 };
438 let embeddings = OpenAIEmbeddings::new(config).unwrap();
439
440 let result = embeddings.embed_documents(&["a", "b"]).await;
441 assert!(
442 matches!(result, Err(EmbeddingError::EmptyVectorInBatch)),
443 "truncated response should report EmptyVectorInBatch, got: {:?}",
444 result
445 );
446 }
447
448 #[tokio::test]
450 async fn test_embed_documents_overrun_returns_batch_mismatch() {
451 let base_url = spawn_embeddings_stub(Arc::new(|_| 100)).await;
452 let config = OpenAIEmbeddingsConfig {
453 api_key: "test-key".into(),
454 base_url,
455 model: "text-embedding-ada-002".into(),
456 batch_size: 2,
457 };
458 let embeddings = OpenAIEmbeddings::new(config).unwrap();
459
460 let result = embeddings.embed_documents(&["a", "b"]).await;
461 assert!(
462 matches!(result, Err(EmbeddingError::BatchMismatch { .. })),
463 "out-of-range index should report BatchMismatch, got: {:?}",
464 result
465 );
466 }
467
468 #[tokio::test]
470 async fn test_embed_query_retries_on_429() {
471 use std::sync::atomic::Ordering;
472
473 let success_body = r#"{"data":[{"object":"embedding","index":0,"embedding":[0.6,0.8]}],"model":"stub","usage":{"prompt_tokens":0,"total_tokens":0}}"#;
474 let (base_url, requests) = spawn_status_stub(429, 2, 200, success_body).await;
475 let config = OpenAIEmbeddingsConfig {
476 api_key: "test-key".into(),
477 base_url,
478 model: "text-embedding-ada-002".into(),
479 batch_size: 2048,
480 };
481 let embeddings = OpenAIEmbeddings::new(config).unwrap();
482
483 let v = embeddings
484 .embed_query("hello")
485 .await
486 .expect("should retry successfully after two 429s");
487 assert_eq!(v.len(), 2);
488 let norm: f32 = v.iter().map(|x| x * x).sum::<f32>().sqrt();
490 assert!((norm - 1.0).abs() < 1e-5, "norm = {}", norm);
491 assert_eq!(requests.load(Ordering::SeqCst), 3, "1 initial + 2 retries");
492 }
493
494 #[tokio::test]
496 async fn test_embed_documents_retries_on_429() {
497 use std::sync::atomic::Ordering;
498
499 let success_body = r#"{"data":[{"object":"embedding","index":0,"embedding":[1.0,0.0]},{"object":"embedding","index":1,"embedding":[0.0,1.0]}],"model":"stub","usage":{"prompt_tokens":0,"total_tokens":0}}"#;
500 let (base_url, requests) = spawn_status_stub(429, 2, 200, success_body).await;
501 let config = OpenAIEmbeddingsConfig {
502 api_key: "test-key".into(),
503 base_url,
504 model: "text-embedding-ada-002".into(),
505 batch_size: 2,
506 };
507 let embeddings = OpenAIEmbeddings::new(config).unwrap();
508
509 let results = embeddings
510 .embed_documents(&["a", "b"])
511 .await
512 .expect("should retry successfully after 429");
513 assert_eq!(results.len(), 2);
514 assert_eq!(
515 requests.load(Ordering::SeqCst),
516 3,
517 "single chunk: 1 initial + 2 retries"
518 );
519 }
520
521 #[tokio::test]
524 async fn test_embed_documents_chunks_run_concurrently() {
525 use std::sync::atomic::{AtomicUsize, Ordering};
526 use tokio::io::{AsyncReadExt, AsyncWriteExt};
527
528 let in_flight = Arc::new(AtomicUsize::new(0));
529 let max_in_flight = Arc::new(AtomicUsize::new(0));
530
531 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
532 let addr = listener.local_addr().unwrap();
533 let base_url = format!("http://{addr}");
534
535 let in_flight_server = in_flight.clone();
536 let max_in_flight_server = max_in_flight.clone();
537 tokio::spawn(async move {
538 while let Ok((mut socket, _)) = listener.accept().await {
539 let in_flight = in_flight_server.clone();
540 let max_in_flight = max_in_flight_server.clone();
541 tokio::spawn(async move {
542 let mut header = Vec::new();
543 let mut byte = [0u8; 1];
544 while header.len() < 64 * 1024 {
545 if socket.read_exact(&mut byte).await.is_err() {
546 return;
547 }
548 header.push(byte[0]);
549 if header.ends_with(b"\r\n\r\n") {
550 break;
551 }
552 }
553 let header_str = String::from_utf8_lossy(&header).to_lowercase();
554 let content_length: usize = header_str
555 .lines()
556 .find_map(|l| l.strip_prefix("content-length:"))
557 .and_then(|v| v.trim().parse().ok())
558 .unwrap_or(0);
559 let mut body = vec![0u8; content_length];
560 if content_length > 0 && socket.read_exact(&mut body).await.is_err() {
561 return;
562 }
563 let body_str = String::from_utf8_lossy(&body);
564 let inputs: Vec<String> = serde_json::from_str::<serde_json::Value>(&body_str)
565 .ok()
566 .and_then(|v| v.get("input").cloned())
567 .and_then(|input| match input {
568 serde_json::Value::String(s) => Some(vec![s]),
569 serde_json::Value::Array(a) => Some(
570 a.iter()
571 .filter_map(|x| x.as_str().map(String::from))
572 .collect(),
573 ),
574 _ => None,
575 })
576 .unwrap_or_default();
577
578 let now = in_flight.fetch_add(1, Ordering::SeqCst) + 1;
579 max_in_flight.fetch_max(now, Ordering::SeqCst);
580 tokio::time::sleep(std::time::Duration::from_millis(50)).await;
582 in_flight.fetch_sub(1, Ordering::SeqCst);
583
584 let data: Vec<serde_json::Value> = inputs
585 .iter()
586 .enumerate()
587 .map(|(i, s)| {
588 let raw = s.bytes().map(|b| b as f32).sum::<f32>();
589 let mut v = vec![raw, 1.0];
590 crate::l2_normalize(&mut v);
591 serde_json::json!({
592 "object": "embedding",
593 "index": i,
594 "embedding": v,
595 })
596 })
597 .collect();
598 let json = serde_json::json!({
599 "data": data,
600 "model": "stub",
601 "usage": { "prompt_tokens": 0, "total_tokens": 0 },
602 })
603 .to_string();
604 let response = format!(
605 "HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{}",
606 json.len(),
607 json
608 );
609 let _ = socket.write_all(response.as_bytes()).await;
610 let _ = socket.shutdown().await;
611 });
612 }
613 });
614
615 let config = OpenAIEmbeddingsConfig {
616 api_key: "test-key".into(),
617 base_url,
618 model: "text-embedding-ada-002".into(),
619 batch_size: 1, };
621 let embeddings = OpenAIEmbeddings::new(config).unwrap();
622
623 let results = embeddings
624 .embed_documents(&["a", "b", "c", "d", "e"])
625 .await
626 .expect("concurrent batch should succeed");
627 assert_eq!(results.len(), 5);
628 let peak = max_in_flight.load(Ordering::SeqCst);
629 assert!(
630 peak >= 2,
631 "multiple chunks should run concurrently (max in-flight = {peak}), not serially"
632 );
633 assert!(
634 peak <= super::MAX_CONCURRENT_CHUNKS,
635 "concurrency must not exceed MAX_CONCURRENT_CHUNKS (actual {peak})"
636 );
637 }
638
639 #[test]
640 fn test_config_default() {
641 let config = OpenAIEmbeddingsConfig::default();
642 assert_eq!(config.model, "text-embedding-ada-002");
643 assert_eq!(config.batch_size, 2048);
644 }
645
646 #[test]
647 fn test_config_builder() {
648 let config = OpenAIEmbeddingsConfig::new("test-key")
649 .with_model("text-embedding-3-large")
650 .with_base_url("https://custom.api.com/v1");
651
652 assert_eq!(config.api_key, "test-key");
653 assert_eq!(config.model, "text-embedding-3-large");
654 assert_eq!(config.base_url, "https://custom.api.com/v1");
655 }
656}