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