lc_embeddings/lib.rs
1#![warn(missing_docs)]
2// lc-embeddings/src/lib.rs
3//! Embedding model implementations for LangChainRust.
4//!
5//! Provides embedding generation via multiple backends:
6//! - OpenAI (`text-embedding-ada-002`, `text-embedding-3-small/large`)
7//! - DeepSeek
8//! - Qwen (Alibaba Cloud / DashScope)
9//! - Local: `BagOfWordsEmbeddings` (always available) and `LocalEmbeddings` (ONNX, feature-gated)
10//! - `MockEmbeddings` for testing
11
12mod cohere;
13mod deepseek;
14mod local;
15mod mock;
16mod openai;
17pub mod openai_compat;
18mod qwen;
19mod retry;
20
21#[cfg(test)]
22mod test_support;
23
24#[cfg(feature = "fastembed")]
25mod fastembed_emb;
26
27pub use cohere::{
28 CohereEmbedInputType, CohereEmbeddings, CohereEmbeddingsConfig, COHERE_EMBED_BASE_URL,
29 COHERE_EMBED_MODEL,
30};
31pub use deepseek::{DeepSeekEmbeddings, DeepSeekEmbeddingsConfig, DEEPSEEK_EMBED_MODEL};
32pub use local::BagOfWordsEmbeddings;
33// 1.0:无 `local-embeddings` feature 时 `LocalEmbeddings` 名字不可用(原降级别名
34// 已移除),需显式选边——`BagOfWordsEmbeddings` 或开启 feature 用 ONNX 版。
35#[cfg(feature = "local-embeddings")]
36pub use local::LocalEmbeddings;
37pub use mock::MockEmbeddings;
38pub use openai::{OpenAIEmbeddings, OpenAIEmbeddingsConfig};
39pub use qwen::{QwenEmbeddings, QwenEmbeddingsConfig, QWEN_EMBED_MODEL};
40
41#[cfg(feature = "fastembed")]
42pub use fastembed_emb::FastEmbedEmbeddings;
43
44use async_trait::async_trait;
45
46/// Embedding error type
47#[derive(Debug, thiserror::Error)]
48#[non_exhaustive]
49pub enum EmbeddingError {
50 /// HTTP request error
51 #[error("HTTP error: {0}")]
52 HttpError(String),
53
54 /// API error
55 #[error("API error: {0}")]
56 ApiError(String),
57
58 /// Parse error
59 #[error("Parse error: {0}")]
60 ParseError(String),
61
62 /// 配置错误(如 API key 为空、模型维度未知)——构造期 fail fast。
63 #[error("Configuration error: {0}")]
64 Config(String),
65
66 /// Empty input
67 #[error("Input is empty")]
68 EmptyInput,
69
70 /// 批量 embedding 数量错位:请求 N 条文本,服务端返回的向量数量或 index
71 /// 超出了预期范围(某 chunk 少返回/乱序导致)。
72 ///
73 /// P0-1: 拒绝静默错数据——绝不把缺失向量当成"不相似"。
74 #[error("Embedding batch mismatch: expected {expected} vectors, got position {actual}")]
75 BatchMismatch {
76 /// 期望返回的向量数量
77 expected: usize,
78 /// 出现错位的位置索引
79 actual: usize,
80 },
81
82 /// 批量 embedding 中某条文本未取到向量(服务端返回量 < 请求量)。
83 ///
84 /// P0-1: 拒绝静默空向量——缺失即显式报错,而非留下零向量。
85 #[error("Embedding batch contains an empty vector (provider returned fewer embeddings than requested)")]
86 EmptyVectorInBatch,
87}
88
89/// Embedding model trait
90///
91/// Defines the interface for generating text embedding vectors.
92///
93/// # 归一化契约(P2-8)
94///
95/// 所有返回的向量都是 **L2 归一化**的单位向量(零向量除外),与 provider 内部
96/// 是否已归一化无关。这保证下游无论用 cosine、点积还是 L2 距离,结果都不随
97/// provider 漂移。HTTP provider 在返回前统一调用 [`l2_normalize`]。
98#[async_trait]
99pub trait Embeddings: Send + Sync {
100 /// Generate an embedding vector for a single text.
101 ///
102 /// # Arguments
103 /// * `text` - Input text
104 ///
105 /// # Returns
106 /// Embedding vector (typically 1536 dimensions or higher)
107 async fn embed_query(&self, text: &str) -> Result<Vec<f32>, EmbeddingError>;
108
109 /// Generate embedding vectors for multiple documents.
110 ///
111 /// # Arguments
112 /// * `texts` - List of input texts
113 ///
114 /// # Returns
115 /// List of embedding vectors
116 ///
117 /// # 语义约定(P1-1)
118 ///
119 /// - 任一文本为空或全空白(`trim().is_empty()`)→ `Err(EmbeddingError::EmptyInput)`;
120 /// - 空切片 `&[]` 视为"没有要嵌入的文本"→ `Ok(vec![])`(无事可做不算错误)。
121 ///
122 /// 默认实现循环调用 [`Self::embed_query`],并前置统一判空;各 provider
123 /// 覆写时须遵循同样的契约,不得再各自分裂。
124 async fn embed_documents(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>, EmbeddingError> {
125 if texts.iter().any(|t| t.trim().is_empty()) {
126 return Err(EmbeddingError::EmptyInput);
127 }
128 let mut embeddings = Vec::new();
129 for text in texts {
130 embeddings.push(self.embed_query(text).await?);
131 }
132 Ok(embeddings)
133 }
134
135 /// Get the embedding vector dimension.
136 fn dimension(&self) -> usize;
137
138 /// Get the model name.
139 fn model_name(&self) -> &str;
140}
141
142/// Compute cosine similarity between two vectors.
143///
144/// Re-exported from [`lc_core::math::cosine_similarity`].
145pub use lc_core::math::cosine_similarity;
146
147/// In-place L2 normalization: scale `vec` to unit length.
148///
149/// P2-8: 各 provider 返回向量的归一化口径不一(OpenAI 已归一化、BOW 自归一化、
150/// Cohere 等远程 provider 可能不),下游一旦用点积/L2 距离而非 cosine,结果会因
151/// provider 漂移。本函数是**唯一**的归一化实现,HTTP provider 在返回前统一调用,
152/// 保证 `Embeddings` 产出的向量恒为单位长度。
153///
154/// 零向量保持零向量(不产生 NaN)。
155pub fn l2_normalize(vec: &mut [f32]) {
156 let norm: f32 = vec.iter().map(|x| x * x).sum::<f32>().sqrt();
157 if norm > 0.0 {
158 for v in vec.iter_mut() {
159 *v /= norm;
160 }
161 }
162}
163
164/// Mutex for synchronizing environment-variable mutations in tests.
165///
166/// Tests that set/remove env vars must acquire this lock to avoid data races
167/// when running tests in parallel (the default for `cargo test`).
168#[cfg(test)]
169static ENV_TEST_LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(());
170
171#[cfg(test)]
172mod tests {
173 use super::*;
174
175 #[test]
176 fn test_cosine_similarity() {
177 // Identical vectors
178 let a = vec![1.0, 0.0, 0.0];
179 let b = vec![1.0, 0.0, 0.0];
180 assert!((cosine_similarity(&a, &b).unwrap() - 1.0).abs() < 0.0001);
181
182 // Orthogonal vectors
183 let a = vec![1.0, 0.0, 0.0];
184 let b = vec![0.0, 1.0, 0.0];
185 assert!((cosine_similarity(&a, &b).unwrap() - 0.0).abs() < 0.0001);
186
187 // Opposite vectors
188 let a = vec![1.0, 0.0, 0.0];
189 let b = vec![-1.0, 0.0, 0.0];
190 assert!((cosine_similarity(&a, &b).unwrap() - (-1.0)).abs() < 0.0001);
191 }
192
193 #[test]
194 fn test_cosine_similarity_different_lengths() {
195 let a = vec![1.0, 0.0];
196 let b = vec![1.0, 0.0, 0.0];
197 assert!(cosine_similarity(&a, &b).is_err());
198 }
199}