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, returning a Result.
105    ///
106    /// Environment variables:
107    /// - `OPENAI_API_KEY`: API key (required)
108    /// - `OPENAI_BASE_URL`: API endpoint (optional)
109    /// - `OPENAI_EMBED_MODEL`: Model name (optional)
110    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        // P2-5: 429/5xx 指数退避重试,瞬时故障不再一次失败即抛错。
142        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            // P1-4: 读失败的错误体也要报错,不能 unwrap_or_default() 吞掉。
155            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        // P2-8: 统一 L2 归一化,保证单位长度,消除 provider 漂移。
176        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        // P1-1: 任一空/全空白文本都报错,与 trait 默认契约一致。
185        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        // P2-6: 各 chunk 并发请求(buffer_unordered + 并发上限),而非串行 await,
192        // 提升大批量吞吐;并发上限避免一次性打爆 provider 限流。
193        let concurrency = texts.len().div_ceil(batch_size).min(MAX_CONCURRENT_CHUNKS);
194
195        // 每个 future 返回 (chunk_idx, data);并发完成顺序不定,由收集端按 chunk_idx
196        // 放回全局槽位。P0-1: 任一槽位空缺即显式报错,绝不把缺失向量当成"不相似"。
197        // 参考 faithfulness.rs 的并发写法:先把各 chunk 转成 owned Vec<String> 再
198        // stream::iter,map 闭包输入无生命周期 → 闭包天然可泛化;async move 只捕获
199        // owned chunk + Copy 引用(client/api_key/model/url),map(FnMut) 才编译得过。
200        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                // P2-5: 429/5xx 指数退避重试。
217                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                    // P1-4: 读失败的错误体也要报错,不能 unwrap_or_default() 吞掉。
230                    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                    // 服务端 index 超出请求范围 = 批次错位,直接报错。
256                    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        // 展开为 Result:任一槽位空缺即显式报错,而非留下零向量;并统一 L2 归一化(P2-8)。
266        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/// OpenAI Embedding API response
286#[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    /// P0-1: 跨 chunk 边界批量对齐——顺序与输入一致,无空向量、无错位。
393    #[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        // 5 条文本 → 3 个 chunk(2/2/1),验证跨 chunk 顺序正确。
405        // stub 以文本字节和编码向量,因此每条文本必须落到自己的槽位,
406        // 任何错位/重复都会让对应槽位的向量与文本不匹配。
407        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            // stub 返回 [sum, 1.0];P2-8 归一化后与逐条归一化的期望值一致,
415            // 且每条向量仍互不相同,可验证跨 chunk 顺序对齐。
416            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    /// P0-1: 服务端某 chunk 少返回 → 显式 `EmptyVectorInBatch`,而非静默空向量。
424    #[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    /// P0-1: 服务端返回 index 超出请求范围 → 显式 `BatchMismatch`。
444    #[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    /// P2-5: `embed_query` 接线重试——429 两次后 200,请求共 3 次且返回成功。
464    #[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        // P2-8: 返回向量应已归一化。
484        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    /// P2-5: `embed_documents` 的每个 chunk 同样走重试——429 后成功。
490    #[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    /// P2-6: 多 chunk 并发请求——并发度受 `MAX_CONCURRENT_CHUNKS` 约束,
517    /// 且确实并行(最大在途请求数 > 1),而非串行。
518    #[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                    // 50ms 让各 chunk 有重叠窗口,验证确实并行。
576                    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, // 每条文本一个 chunk → 5 个并发 future
615        };
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}