rlm_rs/embedding/
fastembed_impl.rs1use crate::Result;
7use crate::embedding::{DEFAULT_DIMENSIONS, Embedder};
8use crate::error::StorageError;
9use std::panic::{AssertUnwindSafe, catch_unwind};
10use std::sync::OnceLock;
11
12static EMBEDDING_MODEL: OnceLock<std::sync::Mutex<fastembed::TextEmbedding>> = OnceLock::new();
15
16pub struct FastEmbedEmbedder {
36 model_name: &'static str,
38}
39
40impl FastEmbedEmbedder {
41 #[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 fn get_model() -> Result<&'static std::sync::Mutex<fastembed::TextEmbedding>> {
66 if let Some(model) = EMBEDDING_MODEL.get() {
68 return Ok(model);
69 }
70
71 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 let _ = EMBEDDING_MODEL.set(std::sync::Mutex::new(model));
79
80 EMBEDDING_MODEL.get().ok_or_else(|| {
82 StorageError::Embedding("Model initialization race condition".to_string()).into()
83 })
84 }
85
86 #[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 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 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 #[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}