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 futures_util::StreamExt;
9use serde::Deserialize;
10
11/// P2-6: 批量切块后的并发上限——避免一次性打爆 provider 限流。
12const MAX_CONCURRENT_CHUNKS: usize = 8;
13
14/// OpenAI Embeddings configuration
15#[derive(Debug, Clone)]
16pub struct OpenAIEmbeddingsConfig {
17    /// API key
18    pub api_key: String,
19
20    /// API base URL
21    pub base_url: String,
22
23    /// Model name (default: text-embedding-ada-002)
24    pub model: String,
25
26    /// Batch size (default: 2048)
27    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    /// Create a new configuration
43    pub fn new(api_key: impl Into<String>) -> Self {
44        Self {
45            api_key: api_key.into(),
46            ..Default::default()
47        }
48    }
49
50    /// Set the model
51    pub fn with_model(mut self, model: impl Into<String>) -> Self {
52        self.model = model.into();
53        self
54    }
55
56    /// Set the base URL
57    pub fn with_base_url(mut self, url: impl Into<String>) -> Self {
58        self.base_url = url.into();
59        self
60    }
61}
62
63/// OpenAI Embeddings client
64pub struct OpenAIEmbeddings {
65    config: OpenAIEmbeddingsConfig,
66    client: reqwest::Client,
67    dimension: usize,
68}
69
70impl OpenAIEmbeddings {
71    /// Create a new OpenAI Embeddings client.
72    ///
73    /// 构造时 fail fast(P1-3):API key 为空立即报错,而不是拖到发请求才 401。
74    /// 模型维度已知才构造(P1-2):未知模型返回 `Err`,不得回落默认 1536 撒谎。
75    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    /// 已知模型的 embedding 维度表;未知模型返回 `Err`(P1-2)。
91    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    /// Creates OpenAIEmbeddings from environment variables.
105    #[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    /// Creates OpenAIEmbeddings from environment variables, returning a Result.
118    ///
119    /// Environment variables:
120    /// - `OPENAI_API_KEY`: API key (required)
121    /// - `OPENAI_BASE_URL`: API endpoint (optional)
122    /// - `OPENAI_EMBED_MODEL`: Model name (optional)
123    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        // P2-5: 429/5xx 指数退避重试,瞬时故障不再一次失败即抛错。
155        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            // P1-4: 读失败的错误体也要报错,不能 unwrap_or_default() 吞掉。
168            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        // P2-8: 统一 L2 归一化,保证单位长度,消除 provider 漂移。
189        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        // P1-1: 任一空/全空白文本都报错,与 trait 默认契约一致。
198        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        // P2-6: 各 chunk 并发请求(buffer_unordered + 并发上限),而非串行 await,
205        // 提升大批量吞吐;并发上限避免一次性打爆 provider 限流。
206        let concurrency = texts.len().div_ceil(batch_size).min(MAX_CONCURRENT_CHUNKS);
207
208        // 每个 future 返回 (chunk_idx, data);并发完成顺序不定,由收集端按 chunk_idx
209        // 放回全局槽位。P0-1: 任一槽位空缺即显式报错,绝不把缺失向量当成"不相似"。
210        // 参考 faithfulness.rs 的并发写法:先把各 chunk 转成 owned Vec<String> 再
211        // stream::iter,map 闭包输入无生命周期 → 闭包天然可泛化;async move 只捕获
212        // owned chunk + Copy 引用(client/api_key/model/url),map(FnMut) 才编译得过。
213        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                // P2-5: 429/5xx 指数退避重试。
230                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                    // P1-4: 读失败的错误体也要报错,不能 unwrap_or_default() 吞掉。
243                    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                    // 服务端 index 超出请求范围 = 批次错位,直接报错。
269                    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        // 展开为 Result:任一槽位空缺即显式报错,而非留下零向量;并统一 L2 归一化(P2-8)。
279        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/// OpenAI Embedding API response
299#[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    /// P0-1: 跨 chunk 边界批量对齐——顺序与输入一致,无空向量、无错位。
406    #[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        // 5 条文本 → 3 个 chunk(2/2/1),验证跨 chunk 顺序正确。
418        // stub 以文本字节和编码向量,因此每条文本必须落到自己的槽位,
419        // 任何错位/重复都会让对应槽位的向量与文本不匹配。
420        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            // stub 返回 [sum, 1.0];P2-8 归一化后与逐条归一化的期望值一致,
428            // 且每条向量仍互不相同,可验证跨 chunk 顺序对齐。
429            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    /// P0-1: 服务端某 chunk 少返回 → 显式 `EmptyVectorInBatch`,而非静默空向量。
437    #[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    /// P0-1: 服务端返回 index 超出请求范围 → 显式 `BatchMismatch`。
457    #[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    /// P2-5: `embed_query` 接线重试——429 两次后 200,请求共 3 次且返回成功。
477    #[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        // P2-8: 返回向量应已归一化。
497        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    /// P2-5: `embed_documents` 的每个 chunk 同样走重试——429 后成功。
503    #[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    /// P2-6: 多 chunk 并发请求——并发度受 `MAX_CONCURRENT_CHUNKS` 约束,
530    /// 且确实并行(最大在途请求数 > 1),而非串行。
531    #[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                    // 50ms 让各 chunk 有重叠窗口,验证确实并行。
589                    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, // 每条文本一个 chunk → 5 个并发 future
628        };
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}