Skip to main content

rlm_rs/embedding/
fastembed_impl.rs

1//! `FastEmbed`-based semantic embedder.
2//!
3//! Provides real semantic embeddings using the BGE-M3 model via fastembed-rs.
4//! Only available when the `fastembed-embeddings` feature is enabled.
5
6use crate::Result;
7use crate::embedding::{DEFAULT_DIMENSIONS, Embedder};
8use crate::error::StorageError;
9use std::panic::{AssertUnwindSafe, catch_unwind};
10use std::sync::OnceLock;
11
12/// Thread-safe singleton for the embedding model.
13/// Uses `OnceLock` for lazy initialization on first use.
14static EMBEDDING_MODEL: OnceLock<std::sync::Mutex<fastembed::TextEmbedding>> = OnceLock::new();
15
16/// `FastEmbed` embedder using BGE-M3.
17///
18/// Uses the fastembed-rs library for real semantic embeddings.
19/// The model is lazily loaded on first embed call to preserve cold start time.
20///
21/// BGE-M3 provides:
22/// - 1024 dimensions (vs 384 for `MiniLM`)
23/// - 8192 token context (vs ~512 for `MiniLM`)
24/// - Better multilingual support
25///
26/// # Examples
27///
28/// ```ignore
29/// use rlm_rs::embedding::FastEmbedEmbedder;
30///
31/// let embedder = FastEmbedEmbedder::new()?;
32/// let embedding = embedder.embed("Hello, world!")?;
33/// assert_eq!(embedding.len(), 1024);
34/// ```
35pub struct FastEmbedEmbedder {
36    /// Model name for debugging.
37    model_name: &'static str,
38}
39
40impl FastEmbedEmbedder {
41    /// Creates a new `FastEmbed` embedder.
42    ///
43    /// Note: Model is lazily loaded on first `embed()` call.
44    ///
45    /// # Errors
46    ///
47    /// Returns an error if model initialization fails.
48    #[allow(clippy::missing_const_for_fn)]
49    pub fn new() -> Result<Self> {
50        Ok(Self {
51            model_name: "BGE-M3",
52        })
53    }
54
55    fn init_options() -> fastembed::InitOptions {
56        fastembed::InitOptions::new(fastembed::EmbeddingModel::BGEM3)
57            .with_max_length(8192)
58            .with_show_download_progress(false)
59    }
60
61    /// Gets or initializes the embedding model (thread-safe).
62    ///
63    /// The model is loaded lazily on first use to preserve cold start time.
64    /// Subsequent calls return the cached instance.
65    fn get_model() -> Result<&'static std::sync::Mutex<fastembed::TextEmbedding>> {
66        // Check if already initialized
67        if let Some(model) = EMBEDDING_MODEL.get() {
68            return Ok(model);
69        }
70
71        // Initialize the model
72        let options = Self::init_options();
73
74        let model = fastembed::TextEmbedding::try_new(options)
75            .map_err(|e| StorageError::Embedding(format!("Failed to load embedding model: {e}")))?;
76
77        // Store the model, ignoring if another thread beat us to it
78        let _ = EMBEDDING_MODEL.set(std::sync::Mutex::new(model));
79
80        // Return the (possibly other thread's) model
81        EMBEDDING_MODEL.get().ok_or_else(|| {
82            StorageError::Embedding("Model initialization race condition".to_string()).into()
83        })
84    }
85
86    /// Returns the model name.
87    #[must_use]
88    pub const fn model_name(&self) -> &'static str {
89        self.model_name
90    }
91}
92
93impl Embedder for FastEmbedEmbedder {
94    fn dimensions(&self) -> usize {
95        DEFAULT_DIMENSIONS
96    }
97
98    fn model_name(&self) -> &'static str {
99        self.model_name
100    }
101
102    fn embed(&self, text: &str) -> Result<Vec<f32>> {
103        if text.is_empty() {
104            return Err(crate::Error::Chunking(
105                crate::error::ChunkingError::InvalidConfig {
106                    reason: "Cannot embed empty text".to_string(),
107                },
108            ));
109        }
110
111        let model = Self::get_model()?;
112        let mut model = model
113            .lock()
114            .map_err(|e| StorageError::Embedding(format!("Failed to lock embedding model: {e}")))?;
115
116        let texts = [text];
117
118        // Wrap ONNX runtime call in catch_unwind for graceful degradation.
119        // ONNX runtime can panic on malformed inputs or internal errors.
120        let result = catch_unwind(AssertUnwindSafe(|| model.embed(texts, None)));
121
122        let embeddings = result
123            .map_err(|panic_info| {
124                let panic_msg = panic_info
125                    .downcast_ref::<&str>()
126                    .map(|s| (*s).to_string())
127                    .or_else(|| panic_info.downcast_ref::<String>().cloned())
128                    .unwrap_or_else(|| "unknown panic".to_string());
129                StorageError::Embedding(format!("ONNX runtime panic: {panic_msg}"))
130            })?
131            .map_err(|e| StorageError::Embedding(format!("Embedding failed: {e}")))?;
132
133        embeddings.into_iter().next().ok_or_else(|| {
134            StorageError::Embedding("No embedding returned from model".to_string()).into()
135        })
136    }
137
138    fn embed_batch(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>> {
139        if texts.is_empty() {
140            return Ok(Vec::new());
141        }
142
143        if texts.iter().any(|t| t.is_empty()) {
144            return Err(crate::Error::Chunking(
145                crate::error::ChunkingError::InvalidConfig {
146                    reason: "Cannot embed empty text".to_string(),
147                },
148            ));
149        }
150
151        let model = Self::get_model()?;
152        let mut model = model
153            .lock()
154            .map_err(|e| StorageError::Embedding(format!("Failed to lock embedding model: {e}")))?;
155
156        // Wrap ONNX runtime call in catch_unwind for graceful degradation.
157        let result = catch_unwind(AssertUnwindSafe(|| model.embed(texts, None)));
158
159        result
160            .map_err(|panic_info| {
161                let panic_msg = panic_info
162                    .downcast_ref::<&str>()
163                    .map(|s| (*s).to_string())
164                    .or_else(|| panic_info.downcast_ref::<String>().cloned())
165                    .unwrap_or_else(|| "unknown panic".to_string());
166                crate::Error::Storage(StorageError::Embedding(format!(
167                    "ONNX runtime panic: {panic_msg}"
168                )))
169            })?
170            .map_err(|e| {
171                crate::Error::Storage(StorageError::Embedding(format!(
172                    "Batch embedding failed: {e}"
173                )))
174            })
175    }
176}
177
178#[cfg(test)]
179mod tests {
180    use super::*;
181
182    #[test]
183    fn test_embedder_creation() {
184        let embedder = FastEmbedEmbedder::new();
185        assert!(embedder.is_ok());
186        assert_eq!(embedder.unwrap().dimensions(), DEFAULT_DIMENSIONS);
187    }
188
189    #[test]
190    fn test_model_name() {
191        let embedder = FastEmbedEmbedder::new().unwrap();
192        assert_eq!(embedder.model_name(), "BGE-M3");
193    }
194
195    #[test]
196    fn test_bge_m3_max_length() {
197        assert_eq!(FastEmbedEmbedder::init_options().max_length, 8192);
198    }
199
200    // Integration tests that require model download are marked #[ignore]
201    // Run with: cargo test --features fastembed-embeddings -- --ignored
202
203    #[test]
204    #[ignore = "requires fastembed model download"]
205    fn test_embed_success() {
206        let embedder = FastEmbedEmbedder::new().unwrap();
207        let result = embedder.embed("Hello, world!");
208        assert!(result.is_ok());
209        assert_eq!(result.unwrap().len(), DEFAULT_DIMENSIONS);
210    }
211
212    #[test]
213    #[ignore = "requires fastembed model download"]
214    fn test_embed_batch_success() {
215        let embedder = FastEmbedEmbedder::new().unwrap();
216        let texts = vec!["Hello", "World"];
217        let result = embedder.embed_batch(&texts);
218        assert!(result.is_ok());
219        let embeddings = result.unwrap();
220        assert_eq!(embeddings.len(), 2);
221        assert_eq!(embeddings[0].len(), DEFAULT_DIMENSIONS);
222    }
223
224    #[test]
225    fn test_embed_empty_fails() {
226        let embedder = FastEmbedEmbedder::new().unwrap();
227        let result = embedder.embed("");
228        assert!(result.is_err());
229    }
230
231    #[test]
232    fn test_embed_batch_empty_list() {
233        let embedder = FastEmbedEmbedder::new().unwrap();
234        let result = embedder.embed_batch(&[]);
235        assert!(result.is_ok());
236        assert!(result.unwrap().is_empty());
237    }
238
239    #[test]
240    fn test_embed_batch_with_empty_fails() {
241        let embedder = FastEmbedEmbedder::new().unwrap();
242        let texts = vec!["Valid", "", "Also valid"];
243        let result = embedder.embed_batch(&texts);
244        assert!(result.is_err());
245    }
246}