1use async_trait::async_trait;
12
13#[cfg(feature = "local-embeddings")]
14use std::path::Path;
15
16use crate::{EmbeddingError, Embeddings};
17
18pub struct BagOfWordsEmbeddings {
30 dim: usize,
31}
32
33impl BagOfWordsEmbeddings {
34 pub fn new(dim: usize) -> Self {
36 Self { dim: dim.max(1) }
37 }
38
39 pub fn default_dim() -> Self {
41 Self::new(256)
42 }
43
44 fn tokenize(text: &str) -> Vec<String> {
46 let mut tokens = Vec::new();
47 let mut current = String::new();
48 for c in text.chars() {
49 if c.is_alphanumeric() {
50 if c.is_ascii() {
51 current.push(c.to_ascii_lowercase());
52 } else {
53 if !current.is_empty() {
55 tokens.push(std::mem::take(&mut current));
56 }
57 tokens.push(c.to_string());
58 }
59 } else if !current.is_empty() {
60 tokens.push(std::mem::take(&mut current));
61 }
62 }
63 if !current.is_empty() {
64 tokens.push(current);
65 }
66 tokens
67 }
68
69 fn hash(s: &str) -> u64 {
71 let mut h: u64 = 0xcbf29ce484222325;
72 for b in s.bytes() {
73 h ^= b as u64;
74 h = h.wrapping_mul(0x100000001b3);
75 }
76 h
77 }
78
79 fn embed(&self, text: &str) -> Vec<f32> {
81 let mut v = vec![0.0f32; self.dim];
82 for token in Self::tokenize(text) {
83 let idx = (Self::hash(&token) as usize) % self.dim;
84 v[idx] += 1.0;
85 }
86 let norm: f32 = v.iter().map(|x| x * x).sum::<f32>().sqrt();
88 if norm > 0.0 {
89 for x in &mut v {
90 *x /= norm;
91 }
92 }
93 v
94 }
95}
96
97impl Default for BagOfWordsEmbeddings {
98 fn default() -> Self {
99 Self::default_dim()
100 }
101}
102
103#[async_trait]
104impl Embeddings for BagOfWordsEmbeddings {
105 async fn embed_query(&self, text: &str) -> Result<Vec<f32>, EmbeddingError> {
106 if text.trim().is_empty() {
107 return Err(EmbeddingError::EmptyInput);
108 }
109 Ok(self.embed(text))
110 }
111
112 fn dimension(&self) -> usize {
113 self.dim
114 }
115
116 fn model_name(&self) -> &str {
117 "local-bow"
118 }
119}
120
121#[cfg(feature = "local-embeddings")]
126mod nn {
127 use super::*;
128 use ort::value::Tensor;
129 use std::sync::RwLock;
130
131 pub struct LocalEmbeddings {
145 session: RwLock<ort::session::Session>,
148 dim: usize,
149 model_name: String,
150 }
151
152 impl LocalEmbeddings {
153 pub fn from_file(model_path: impl AsRef<Path>) -> Result<Self, EmbeddingError> {
161 let path = model_path.as_ref();
162 let model_name = path
163 .file_stem()
164 .and_then(|s| s.to_str())
165 .unwrap_or("unknown")
166 .to_string();
167
168 let session = ort::session::Session::builder()
169 .map_err(|e| {
170 EmbeddingError::ApiError(format!("Failed to create ONNX SessionBuilder: {}", e))
171 })?
172 .commit_from_file(path)
173 .map_err(|e| {
174 EmbeddingError::ApiError(format!(
175 "Failed to load ONNX model ({}): {}",
176 path.display(),
177 e
178 ))
179 })?;
180
181 let dim = Self::infer_dimension(&session)?;
183
184 Ok(Self {
185 session: RwLock::new(session),
186 dim,
187 model_name,
188 })
189 }
190
191 fn infer_dimension(session: &ort::session::Session) -> Result<usize, EmbeddingError> {
193 let outputs = session.outputs();
194 if outputs.is_empty() {
195 return Err(EmbeddingError::ParseError(
196 "ONNX model has no output nodes".to_string(),
197 ));
198 }
199
200 let dtype = outputs[0].dtype();
202 let shape = dtype.tensor_shape().ok_or_else(|| {
203 EmbeddingError::ParseError("Output is not a Tensor type".to_string())
204 })?;
205
206 let dim = shape
209 .iter()
210 .rev()
211 .find_map(|&d| if d > 0 { Some(d as usize) } else { None })
212 .ok_or_else(|| {
213 EmbeddingError::ParseError(format!(
214 "Cannot infer embedding dimension from model output shape: {:?}",
215 *shape
216 ))
217 })?;
218
219 Ok(dim)
220 }
221
222 fn simple_tokenize(text: &str) -> Vec<i64> {
227 text.split_whitespace()
229 .map(|word| {
230 let mut h: u64 = 0xcbf29ce484222325;
231 for b in word.bytes() {
232 h ^= b as u64;
233 h = h.wrapping_mul(0x100000001b3);
234 }
235 (h % 30522) as i64
237 })
238 .collect()
239 }
240
241 fn run_inference(
243 &self,
244 input_ids: &[i64],
245 ) -> Result<(Vec<usize>, Vec<f32>), EmbeddingError> {
246 let seq_len = input_ids.len();
247 if seq_len == 0 {
248 return Err(EmbeddingError::EmptyInput);
249 }
250
251 let input_shape = vec![1i64, seq_len as i64];
253 let input_data = input_ids.to_vec();
254
255 let input_tensor = Tensor::from_array((input_shape, input_data)).map_err(|e| {
256 EmbeddingError::ApiError(format!("Failed to construct input tensor: {}", e))
257 })?;
258
259 let session = self.session.read().map_err(|e| {
261 EmbeddingError::ApiError(format!("Failed to acquire session read lock: {}", e))
262 })?;
263 let input_name = session
264 .inputs()
265 .first()
266 .map(|o| o.name().to_string())
267 .unwrap_or_else(|| "input_ids".to_string());
268
269 drop(session);
271 let mut session = self.session.write().map_err(|e| {
272 EmbeddingError::ApiError(format!("Failed to acquire session write lock: {}", e))
273 })?;
274 let outputs = session
275 .run(ort::inputs![input_name.as_str() => input_tensor]?)
276 .map_err(|e| EmbeddingError::ApiError(format!("ONNX inference failed: {}", e)))?;
277
278 let output_value = outputs.get(0).ok_or_else(|| {
280 EmbeddingError::ParseError("ONNX model has no output".to_string())
281 })?;
282
283 let (shape, data) = output_value.try_extract_tensor::<f32>().map_err(|e| {
285 EmbeddingError::ParseError(format!("Failed to extract output tensor: {}", e))
286 })?;
287
288 let shape_vec: Vec<usize> = shape.iter().map(|&d| d as usize).collect();
289 let data_vec = data.to_vec();
290
291 Ok((shape_vec, data_vec))
292 }
293
294 fn mean_pool(shape: &[usize], data: &[f32]) -> Result<Vec<f32>, EmbeddingError> {
299 match shape.len() {
300 3 => {
301 let dim = shape[2];
302 let seq_len = shape[1];
303 let mut result = vec![0.0f32; dim];
304
305 for s in 0..seq_len {
306 for d in 0..dim {
307 result[d] += data[s * dim + d];
308 }
309 }
310
311 for v in &mut result {
312 *v /= seq_len as f32;
313 }
314
315 Ok(result)
316 }
317 2 => {
318 let dim = shape[1];
319 Ok(data[..dim].to_vec())
320 }
321 _ => Err(EmbeddingError::ParseError(format!(
322 "Unsupported output dimension count: {}",
323 shape.len()
324 ))),
325 }
326 }
327
328 fn l2_normalize(vec: &mut [f32]) {
330 let norm: f32 = vec.iter().map(|x| x * x).sum::<f32>().sqrt();
331 if norm > 0.0 {
332 for v in vec.iter_mut() {
333 *v /= norm;
334 }
335 }
336 }
337
338 fn embed_single(&self, text: &str) -> Result<Vec<f32>, EmbeddingError> {
340 if text.trim().is_empty() {
341 return Err(EmbeddingError::EmptyInput);
342 }
343
344 let input_ids = Self::simple_tokenize(text);
345 if input_ids.is_empty() {
346 return Err(EmbeddingError::EmptyInput);
347 }
348
349 let (shape, raw_data) = self.run_inference(&input_ids)?;
350 let mut pooled = Self::mean_pool(&shape, &raw_data)?;
351 Self::l2_normalize(&mut pooled);
352 Ok(pooled)
353 }
354 }
355
356 #[async_trait]
357 impl Embeddings for LocalEmbeddings {
358 async fn embed_query(&self, text: &str) -> Result<Vec<f32>, EmbeddingError> {
359 let text = text.to_string();
361 tokio::task::spawn_blocking(move || self.embed_single(&text))
362 .await
363 .map_err(|e| EmbeddingError::ApiError(format!("Task execution failed: {}", e)))?
364 }
365
366 async fn embed_documents(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>, EmbeddingError> {
367 if texts.is_empty() {
368 return Ok(Vec::new());
369 }
370
371 let texts: Vec<String> = texts.iter().map(|s| s.to_string()).collect();
373 tokio::task::spawn_blocking(move || {
374 let mut results = Vec::with_capacity(texts.len());
375 for text in &texts {
376 results.push(self.embed_single(text)?);
377 }
378 Ok(results)
379 })
380 .await
381 .map_err(|e| EmbeddingError::ApiError(format!("Task execution failed: {}", e)))?
382 }
383
384 fn dimension(&self) -> usize {
385 self.dim
386 }
387
388 fn model_name(&self) -> &str {
389 &self.model_name
390 }
391 }
392}
393
394#[cfg(feature = "local-embeddings")]
396pub use nn::LocalEmbeddings;
397
398#[cfg(not(feature = "local-embeddings"))]
407pub type LocalEmbeddings = BagOfWordsEmbeddings;
408
409#[cfg(test)]
414mod tests {
415 use super::*;
416 use crate::cosine_similarity;
417
418 #[tokio::test]
421 async fn test_bow_dimension() {
422 let e = BagOfWordsEmbeddings::new(128);
423 let v = e.embed_query("hello world").await.unwrap();
424 assert_eq!(v.len(), 128);
425 assert_eq!(e.dimension(), 128);
426 }
427
428 #[tokio::test]
429 async fn test_bow_same_text_same_vector() {
430 let e = BagOfWordsEmbeddings::new(64);
431 let a = e.embed_query("rust programming").await.unwrap();
432 let b = e.embed_query("rust programming").await.unwrap();
433 assert_eq!(a, b);
434 }
435
436 #[tokio::test]
437 async fn test_bow_different_text_different_vector() {
438 let e = BagOfWordsEmbeddings::new(64);
439 let a = e.embed_query("rust programming").await.unwrap();
440 let b = e.embed_query("cooking recipe pasta").await.unwrap();
441 assert_ne!(a, b);
442 }
443
444 #[tokio::test]
445 async fn test_bow_shared_words_more_similar() {
446 let e = BagOfWordsEmbeddings::new(256);
447 let base = e.embed_query("rust programming language").await.unwrap();
448 let similar = e.embed_query("rust programming tutorial").await.unwrap();
449 let different = e.embed_query("cooking pasta recipe").await.unwrap();
450
451 let sim_similar = cosine_similarity(&base, &similar).unwrap_or(0.0);
452 let sim_different = cosine_similarity(&base, &different).unwrap_or(0.0);
453 assert!(
454 sim_similar > sim_different,
455 "Shared words should be more similar: {} vs {}",
456 sim_similar,
457 sim_different
458 );
459 }
460
461 #[tokio::test]
462 async fn test_bow_normalized() {
463 let e = BagOfWordsEmbeddings::new(64);
464 let v = e.embed_query("some text here").await.unwrap();
465 let norm: f32 = v.iter().map(|x| x * x).sum::<f32>().sqrt();
466 assert!((norm - 1.0).abs() < 1e-5, "norm = {}", norm);
467 }
468
469 #[tokio::test]
470 async fn test_bow_empty_text_returns_error() {
471 let e = BagOfWordsEmbeddings::new(64);
472 let result = e.embed_query("").await;
473 assert!(result.is_err());
474 assert!(matches!(result.unwrap_err(), EmbeddingError::EmptyInput));
475 }
476
477 #[tokio::test]
478 async fn test_bow_chinese_tokenize() {
479 let e = BagOfWordsEmbeddings::new(128);
480 let a = e.embed_query("机器学习").await.unwrap();
481 let b = e.embed_query("机器学习").await.unwrap();
482 assert_eq!(a, b);
483 let c = e.embed_query("深度学习").await.unwrap();
484 let sim = cosine_similarity(&a, &c).unwrap_or(0.0);
485 assert!(
486 sim > 0.0,
487 "Shared '学习' should have positive similarity: {}",
488 sim
489 );
490 }
491
492 #[test]
493 fn test_bow_tokenize_english() {
494 let t = BagOfWordsEmbeddings::tokenize("Hello, World! 123");
495 assert!(t.contains(&"hello".to_string()));
496 assert!(t.contains(&"world".to_string()));
497 assert!(t.contains(&"123".to_string()));
498 }
499
500 #[test]
501 fn test_bow_tokenize_chinese() {
502 let t = BagOfWordsEmbeddings::tokenize("机器学习");
503 assert!(t.contains(&"机".to_string()));
504 assert!(t.contains(&"学".to_string()));
505 assert_eq!(t.len(), 4);
506 }
507
508 #[test]
509 fn test_bow_model_name() {
510 let e = BagOfWordsEmbeddings::default_dim();
511 assert_eq!(e.model_name(), "local-bow");
512 }
513
514 #[tokio::test]
517 async fn test_local_embeddings_backward_compat() {
518 let e = LocalEmbeddings::new(64);
520 let v = e.embed_query("test backward compat").await.unwrap();
521 assert_eq!(v.len(), 64);
522 assert_eq!(e.model_name(), "local-bow");
523 }
524
525 #[cfg(feature = "local-embeddings")]
528 mod nn_tests {
529 use super::*;
530
531 #[test]
532 fn test_l2_normalize() {
533 let mut v = vec![3.0, 4.0];
534 LocalEmbeddings::l2_normalize(&mut v);
535 let norm: f32 = v.iter().map(|x| x * x).sum::<f32>().sqrt();
536 assert!((norm - 1.0).abs() < 1e-5);
537 assert!((v[0] - 0.6).abs() < 1e-5);
538 assert!((v[1] - 0.8).abs() < 1e-5);
539 }
540
541 #[test]
542 fn test_l2_normalize_zero() {
543 let mut v = vec![0.0, 0.0, 0.0];
544 LocalEmbeddings::l2_normalize(&mut v);
545 assert!(v.iter().all(|x| *x == 0.0));
546 }
547
548 #[test]
549 fn test_mean_pool_3d() {
550 let shape = vec![1usize, 2, 3];
552 let data = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0];
553 let result = LocalEmbeddings::mean_pool(&shape, &data).unwrap();
554 assert_eq!(result.len(), 3);
555 assert!((result[0] - 2.5).abs() < 1e-5);
556 assert!((result[1] - 3.5).abs() < 1e-5);
557 assert!((result[2] - 4.5).abs() < 1e-5);
558 }
559
560 #[test]
561 fn test_mean_pool_2d() {
562 let shape = vec![1usize, 3];
564 let data = vec![1.0, 2.0, 3.0];
565 let result = LocalEmbeddings::mean_pool(&shape, &data).unwrap();
566 assert_eq!(result.len(), 3);
567 assert!((result[0] - 1.0).abs() < 1e-5);
568 assert!((result[1] - 2.0).abs() < 1e-5);
569 assert!((result[2] - 3.0).abs() < 1e-5);
570 }
571
572 #[test]
573 fn test_simple_tokenize() {
574 let tokens = LocalEmbeddings::simple_tokenize("hello world test");
575 assert_eq!(tokens.len(), 3);
576 let tokens2 = LocalEmbeddings::simple_tokenize("hello");
578 assert_eq!(tokens[0], tokens2[0]);
579 }
580
581 #[test]
582 fn test_simple_tokenize_empty() {
583 let tokens = LocalEmbeddings::simple_tokenize("");
584 assert!(tokens.is_empty());
585 }
586 }
587}