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