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