Skip to main content

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}