Skip to main content

aither_core/
embedding.rs

1//! # Embedding Module
2//!
3//! This module provides types and traits for working with text embeddings.
4//!
5//! ## What are Embeddings?
6//!
7//! Embeddings are dense vector representations of text that capture semantic meaning.
8//! They transform human-readable text into numerical vectors that machine learning
9//! models can process effectively. Similar texts produce similar embedding vectors,
10//! making them useful for:
11//!
12//! - **Semantic search**: Finding relevant documents based on meaning rather than exact keywords
13//! - **Text similarity**: Measuring how similar two pieces of text are
14//! - **Classification**: Categorizing text based on content
15//! - **Clustering**: Grouping similar texts together
16//! - **Recommendation systems**: Finding related content
17//!
18//! ## Embedding Models
19//!
20//! An embedding model is a neural network that has been trained to convert text into
21//! meaningful vector representations. Different models have different characteristics:
22//!
23//! - **Dimension**: The length of the embedding vector (e.g., 768, 1536)
24//! - **Domain**: Some models are optimized for specific types of content
25//! - **Performance**: Trade-offs between speed, accuracy, and resource usage
26//!
27//! Popular embedding models include:
28//! - OpenAI's `text-embedding-ada-002` (1536 dimensions)
29//! - Sentence Transformers like `all-MiniLM-L6-v2` (384 dimensions)
30//! - Cohere's embedding models
31//!
32//! ## Usage
33//!
34//! This module provides the [`EmbeddingModel`] trait that abstracts over different
35//! embedding implementations, allowing you to switch between providers while
36//! maintaining the same interface.
37//!
38//! ```rust,ignore
39//! use aither::EmbeddingModel;
40//!
41//! async fn example<T: EmbeddingModel>(model: &mut T) -> aither::Result<()> {
42//!     // Get the embedding dimension
43//!     let dim = model.dim();
44//!     println!("Model produces {}-dimensional embeddings", dim);
45//!
46//!     // Convert text to embedding
47//!     let embedding = model.embed("Hello, world!").await?;
48//!     assert_eq!(embedding.len(), dim);
49//!
50//!     Ok(())
51//! }
52//! ```
53
54use alloc::vec::Vec;
55use core::future::Future;
56
57/// A type alias for an embedding vector of 32-bit floats.
58///
59/// Embeddings are dense vector representations where each dimension captures
60/// different semantic features of the input text. The vector length is determined
61/// by the embedding model's architecture.
62pub type Embedding = Vec<f32>;
63
64/// Converts text to vector representations.
65///
66/// This trait provides a unified interface for different embedding model implementations,
67/// allowing you to switch between providers (`OpenAI`, `Cohere`, `Hugging Face`, etc.) while
68/// maintaining the same API.
69///
70/// See the [module documentation](crate::embedding) for more details on embeddings and their use cases.
71///
72/// # Implementation Requirements
73///
74/// - The [`embed`](EmbeddingModel::embed) method must return vectors with length equal to [`dim`](EmbeddingModel::dim)
75/// - Embeddings should be normalized if the underlying model requires it
76/// - The implementation should handle errors gracefully (network issues, API limits, etc.)
77///
78/// # Example
79///
80/// ```rust,ignore
81/// use aither::EmbeddingModel;
82///
83/// struct MyEmbedding {
84///     api_key: String,
85/// }
86///
87/// impl EmbeddingModel for MyEmbedding {
88///     fn dim(&self) -> usize {
89///         1536 // OpenAI text-embedding-ada-002 dimension
90///     }
91///     
92///     async fn embed(&mut self, text: &str) -> aither::Result<Vec<f32>> {
93///         // In a real implementation, this would call the embedding API
94///         Ok(vec![0.0; self.dim()])
95///     }
96/// }
97///
98/// # async {
99/// let mut model = MyEmbedding { api_key: "sk-...".to_string() };
100/// let embedding = model.embed("The quick brown fox").await.unwrap();
101/// assert_eq!(embedding.len(), 1536);
102/// # };
103/// ```
104///
105/// # Performance Considerations
106///
107/// - Batch multiple texts when possible to reduce API calls
108/// - Consider caching embeddings for frequently used texts
109/// - Be aware of rate limits when using cloud-based embedding services
110pub trait EmbeddingModel: Send + Sized + Sync {
111    /// Returns the embedding vector dimension.
112    ///
113    /// This value determines the length of vectors returned by [`embed`](EmbeddingModel::embed).
114    /// Common dimensions include:
115    /// - 384 (`Sentence Transformers MiniLM`)
116    /// - 768 (`BERT-base`)
117    /// - 1536 (`OpenAI text-embedding-ada-002`)
118    /// - 3072 (`OpenAI text-embedding-3-large`)
119    fn dim(&self) -> usize;
120
121    /// Converts text to an embedding vector.
122    ///
123    /// # Arguments
124    ///
125    /// * `text` - The input text to embed. Can be a word, sentence, paragraph, or document.
126    ///
127    /// # Returns
128    ///
129    /// A [`Vec<f32>`] with length equal to [`Self::dim`](EmbeddingModel::dim).
130    /// The vector represents the semantic meaning of the input text in high-dimensional space.
131    ///
132    /// Implementations that need mutable state should use interior mutability.
133    fn embed(&self, text: &str) -> impl Future<Output = crate::Result<Vec<f32>>> + Send;
134}
135
136#[cfg(test)]
137mod tests {
138    use super::*;
139    use alloc::vec;
140
141    struct MockEmbeddingModel {
142        dimension: usize,
143    }
144
145    impl EmbeddingModel for MockEmbeddingModel {
146        fn dim(&self) -> usize {
147            self.dimension
148        }
149
150        #[allow(clippy::cast_precision_loss)]
151        fn embed(&self, text: &str) -> impl Future<Output = crate::Result<Vec<f32>>> + Send {
152            // Create a simple mock embedding based on text length
153            let mut embedding = vec![0.0; self.dimension];
154            let text_len = text.len();
155
156            for (i, value) in embedding.iter_mut().enumerate() {
157                *value = (text_len + i) as f32 * 0.01;
158            }
159
160            core::future::ready(Ok(embedding))
161        }
162    }
163
164    #[tokio::test]
165    async fn embedding_model_dimension() {
166        let model = MockEmbeddingModel { dimension: 768 };
167        assert_eq!(model.dim(), 768);
168    }
169
170    #[tokio::test]
171    async fn embedding_generation() {
172        let model = MockEmbeddingModel { dimension: 4 };
173        let embedding = model.embed("test").await.unwrap();
174
175        assert_eq!(embedding.len(), 4);
176        assert!((embedding[0] - 0.04).abs() < f32::EPSILON); // text length 4 + index 0 = 4 * 0.01
177        assert!((embedding[1] - 0.05).abs() < f32::EPSILON); // text length 4 + index 1 = 5 * 0.01
178        assert!((embedding[2] - 0.06).abs() < f32::EPSILON); // text length 4 + index 2 = 6 * 0.01
179        assert!((embedding[3] - 0.07).abs() < f32::EPSILON); // text length 4 + index 3 = 7 * 0.01
180    }
181
182    #[tokio::test]
183    #[allow(clippy::float_cmp)]
184    async fn embedding_different_texts() {
185        let model = MockEmbeddingModel { dimension: 2 };
186
187        let embedding1 = model.embed("a").await.unwrap();
188        let embedding2 = model.embed("ab").await.unwrap();
189
190        assert_eq!(embedding1.len(), 2);
191        assert_eq!(embedding2.len(), 2);
192
193        // Different text lengths should produce different embeddings
194        assert_ne!(embedding1[0], embedding2[0]);
195        assert_ne!(embedding1[1], embedding2[1]);
196    }
197
198    #[tokio::test]
199    #[allow(clippy::float_cmp)]
200    async fn embedding_empty_text() {
201        let model = MockEmbeddingModel { dimension: 3 };
202        let embedding = model.embed("").await.unwrap();
203
204        assert_eq!(embedding.len(), 3);
205        assert_eq!(embedding[0], 0.00); // length 0 + index 0 = 0 * 0.01
206        assert_eq!(embedding[1], 0.01); // length 0 + index 1 = 1 * 0.01
207        assert_eq!(embedding[2], 0.02); // length 0 + index 2 = 2 * 0.01
208    }
209
210    #[tokio::test]
211    async fn embedding_large_dimension() {
212        let model = MockEmbeddingModel { dimension: 1536 }; // Common OpenAI dimension
213        let embedding = model.embed("test text").await.unwrap();
214
215        assert_eq!(embedding.len(), 1536);
216        assert!((embedding[0] - 0.09).abs() < f32::EPSILON); // text length 9 + index 0 = 9 * 0.01
217        assert!((embedding[1535] - 15.44).abs() < 0.01); // text length 9 + index 1535 = 1544 * 0.01
218    }
219}