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::path::PathBuf;
130 use std::sync::{Arc, Condvar, Mutex};
131
132 const DEFAULT_MAX_BATCH: usize = 32;
152 const DEFAULT_MAX_SEQ_LEN: usize = 512;
154
155 fn default_pool_size() -> usize {
158 std::thread::available_parallelism()
159 .map(|n| n.get().clamp(1, 8))
160 .unwrap_or(2)
161 }
162
163 fn lock_error<T>(_: std::sync::PoisonError<T>) -> EmbeddingError {
165 EmbeddingError::ApiError("session pool lock poisoned".to_string())
166 }
167
168 struct SessionPool {
176 model_bytes: Arc<Vec<u8>>,
177 idle: Mutex<Vec<ort::session::Session>>,
178 live: Mutex<usize>,
179 capacity: usize,
180 notify: Condvar,
181 }
182
183 impl SessionPool {
184 fn new(model_bytes: Vec<u8>, capacity: usize) -> Result<Arc<Self>, EmbeddingError> {
186 let capacity = capacity.max(1);
187 let pool = Arc::new(Self {
188 model_bytes: Arc::new(model_bytes),
189 idle: Mutex::new(Vec::new()),
190 live: Mutex::new(0),
191 capacity,
192 notify: Condvar::new(),
193 });
194 let session = pool.build_session()?;
195 *pool.live.lock().map_err(lock_error)? = 1;
196 pool.idle.lock().map_err(lock_error)?.push(session);
197 Ok(pool)
198 }
199
200 fn build_session(&self) -> Result<ort::session::Session, EmbeddingError> {
201 ort::session::Session::builder()
202 .map_err(|e| {
203 EmbeddingError::ApiError(format!("Failed to create ONNX SessionBuilder: {e}"))
204 })?
205 .commit_from_memory(&self.model_bytes)
206 .map_err(|e| {
207 EmbeddingError::ApiError(format!("Failed to load ONNX model from memory: {e}"))
208 })
209 }
210
211 fn acquire(self: &Arc<Self>) -> Result<SessionGuard, EmbeddingError> {
214 {
216 let mut idle = self.idle.lock().map_err(lock_error)?;
217 if let Some(session) = idle.pop() {
218 return Ok(SessionGuard {
219 pool: self.clone(),
220 session: Some(session),
221 });
222 }
223 }
224 let mut live = self.live.lock().map_err(lock_error)?;
226 loop {
227 {
229 let mut idle = self.idle.lock().map_err(lock_error)?;
230 if let Some(session) = idle.pop() {
231 return Ok(SessionGuard {
232 pool: self.clone(),
233 session: Some(session),
234 });
235 }
236 }
237 if *live < self.capacity {
238 *live += 1;
239 match self.build_session() {
240 Ok(session) => {
241 return Ok(SessionGuard {
242 pool: self.clone(),
243 session: Some(session),
244 });
245 }
246 Err(e) => {
247 *live -= 1;
249 self.notify.notify_one();
250 return Err(e);
251 }
252 }
253 }
254 live = self.notify.wait(live).map_err(lock_error)?;
255 }
256 }
257
258 fn release(&self, session: ort::session::Session) {
260 let Ok(mut idle) = self.idle.lock() else {
261 return;
262 };
263 idle.push(session);
264 drop(idle);
265 self.notify.notify_one();
266 }
267 }
268
269 struct SessionGuard {
271 pool: Arc<SessionPool>,
272 session: Option<ort::session::Session>,
273 }
274
275 impl SessionGuard {
276 fn session(&mut self) -> Result<&mut ort::session::Session, EmbeddingError> {
277 self.session.as_mut().ok_or_else(|| {
278 EmbeddingError::ApiError("session guard already released".to_string())
279 })
280 }
281 }
282
283 impl Drop for SessionGuard {
284 fn drop(&mut self) {
285 if let Some(session) = self.session.take() {
286 self.pool.release(session);
287 }
288 }
289 }
290
291 struct LocalInner {
296 pool: Arc<SessionPool>,
297 tokenizer: tokenizers::Tokenizer,
298 dim: usize,
299 model_name: String,
300 seq_limit: usize,
301 max_batch: usize,
302 }
303
304 impl LocalInner {
305 fn tokenize(&self, texts: &[String]) -> Result<Vec<tokenizers::Encoding>, EmbeddingError> {
308 self.tokenizer
309 .encode_batch(texts.to_vec(), true)
310 .map_err(|e| {
311 EmbeddingError::Config(format!("Tokenizer failed to encode input: {e}"))
312 })
313 }
314
315 fn build_batch_tensors(
320 encodings: &[tokenizers::Encoding],
321 seq_limit: usize,
322 pad_id: u32,
323 ) -> (Vec<i64>, Vec<i64>, Vec<i64>, usize) {
324 let batch = encodings.len();
325 let lens: Vec<usize> = encodings
326 .iter()
327 .map(|e| e.get_ids().len().min(seq_limit))
328 .collect();
329 let max_len = lens.iter().copied().max().unwrap_or(0);
330 let pad = pad_id as i64;
331 let mut input_ids = vec![pad; batch * max_len];
332 let mut attention_mask = vec![0i64; batch * max_len];
333 let mut token_type_ids = vec![0i64; batch * max_len];
334 for (b, enc) in encodings.iter().enumerate() {
335 let ids = enc.get_ids();
336 let mask = enc.get_attention_mask();
337 let types = enc.get_type_ids();
338 let len = lens[b];
339 let row = b * max_len;
340 for i in 0..len {
341 input_ids[row + i] = ids[i] as i64;
342 attention_mask[row + i] = mask.get(i).copied().unwrap_or(1) as i64;
343 token_type_ids[row + i] = types.get(i).copied().unwrap_or(0) as i64;
344 }
345 }
346 (input_ids, attention_mask, token_type_ids, max_len)
347 }
348
349 fn infer_rows(
351 &self,
352 encodings: &[tokenizers::Encoding],
353 ) -> Result<Vec<Vec<f32>>, EmbeddingError> {
354 if encodings.is_empty() {
355 return Err(EmbeddingError::EmptyInput);
356 }
357 let pad_id = Self::resolve_pad_id(&self.tokenizer);
358 let (input_ids, attention_mask, token_type_ids, fed_seq_len) =
359 Self::build_batch_tensors(encodings, self.seq_limit, pad_id);
360 if fed_seq_len == 0 {
361 return Err(EmbeddingError::EmptyInput);
362 }
363 let batch = encodings.len();
364 let mask_rows: Vec<Vec<i64>> = (0..batch)
365 .map(|b| attention_mask[b * fed_seq_len..(b + 1) * fed_seq_len].to_vec())
366 .collect();
367
368 let input_shape = vec![batch as i64, fed_seq_len as i64];
369 let input_tensor = Tensor::from_array((input_shape, input_ids)).map_err(|e| {
370 EmbeddingError::ApiError(format!("Failed to construct input_ids tensor: {e}"))
371 })?;
372 let attention_tensor = Tensor::from_array((input_shape.clone(), attention_mask))
373 .map_err(|e| {
374 EmbeddingError::ApiError(format!(
375 "Failed to construct attention_mask tensor: {e}"
376 ))
377 })?;
378 let type_tensor =
379 Tensor::from_array((input_shape.clone(), token_type_ids)).map_err(|e| {
380 EmbeddingError::ApiError(format!(
381 "Failed to construct token_type_ids tensor: {e}"
382 ))
383 })?;
384
385 let mut guard = self.pool.acquire()?;
386 let input_names: Vec<String> = {
387 let session = guard.session()?;
388 session
389 .inputs()
390 .iter()
391 .map(|o| o.name().to_string())
392 .collect()
393 };
394 let mut input_ids_slot = Some(input_tensor);
396 let mut attention_slot = Some(attention_tensor);
397 let mut type_slot = Some(type_tensor);
398 let mut named: Vec<(String, Tensor<i64>)> = Vec::with_capacity(input_names.len());
399 for name in input_names {
400 let tensor = match name.as_str() {
401 "input_ids" => input_ids_slot.take().ok_or_else(|| {
402 EmbeddingError::ParseError(
403 "ONNX model declares duplicate 'input_ids' input".to_string(),
404 )
405 })?,
406 "attention_mask" => attention_slot.take().ok_or_else(|| {
407 EmbeddingError::ParseError(
408 "ONNX model declares duplicate 'attention_mask' input".to_string(),
409 )
410 })?,
411 "token_type_ids" => type_slot.take().ok_or_else(|| {
412 EmbeddingError::ParseError(
413 "ONNX model declares duplicate 'token_type_ids' input".to_string(),
414 )
415 })?,
416 other => {
417 return Err(EmbeddingError::ParseError(format!(
418 "Unsupported ONNX model input '{other}': the `local-embeddings` \
419 feature only supports input_ids / attention_mask / token_type_ids"
420 )));
421 }
422 };
423 named.push((name, tensor));
424 }
425 if named.is_empty() {
426 return Err(EmbeddingError::ParseError(
427 "ONNX model declares no supported inputs".to_string(),
428 ));
429 }
430
431 let outputs = guard
432 .session()?
433 .run(named)
434 .map_err(|e| EmbeddingError::ApiError(format!("ONNX inference failed: {e}")))?;
435
436 let output_value = outputs.get(0).ok_or_else(|| {
437 EmbeddingError::ParseError("ONNX model has no output".to_string())
438 })?;
439 let (shape, data) = output_value.try_extract_tensor::<f32>().map_err(|e| {
440 EmbeddingError::ParseError(format!("Failed to extract output tensor: {e}"))
441 })?;
442 let shape_vec: Vec<usize> = shape.iter().map(|&d| d as usize).collect();
443
444 let mut rows = Self::pool_rows(&shape_vec, data, &mask_rows, batch, fed_seq_len)?;
445 for row in &mut rows {
446 crate::l2_normalize(row);
447 }
448 Ok(rows)
449 }
450
451 fn pool_rows(
458 shape: &[usize],
459 data: &[f32],
460 masks: &[Vec<i64>],
461 batch: usize,
462 fed_seq_len: usize,
463 ) -> Result<Vec<Vec<f32>>, EmbeddingError> {
464 match shape.len() {
465 3 => {
466 let out_batch = shape[0];
467 if out_batch != batch {
468 return Err(EmbeddingError::BatchMismatch {
469 expected: batch,
470 actual: out_batch,
471 });
472 }
473 let out_seq = shape[1];
474 let dim = shape[2];
475 let seq = out_seq.min(fed_seq_len);
476 let mut result = vec![vec![0.0f32; dim]; batch];
477 for b in 0..batch {
478 let mask = &masks[b];
479 let mut count = 0usize;
480 for s in 0..seq {
481 if s >= mask.len() || mask[s] == 0 {
482 continue; }
484 count += 1;
485 let base = (b * out_seq + s) * dim;
486 for d in 0..dim {
487 result[b][d] += data[base + d];
488 }
489 }
490 if count > 0 {
491 for d in 0..dim {
492 result[b][d] /= count as f32;
493 }
494 }
495 }
496 Ok(result)
497 }
498 2 => {
499 let out_batch = shape[0];
500 if out_batch != batch {
501 return Err(EmbeddingError::BatchMismatch {
502 expected: batch,
503 actual: out_batch,
504 });
505 }
506 let dim = shape[1];
507 let mut result = Vec::with_capacity(batch);
508 for b in 0..batch {
509 let base = b * dim;
510 result.push(data[base..base + dim].to_vec());
511 }
512 Ok(result)
513 }
514 _ => Err(EmbeddingError::ParseError(format!(
515 "Unsupported output dimension count: {}",
516 shape.len()
517 ))),
518 }
519 }
520
521 fn embed_batch(&self, texts: &[String]) -> Result<Vec<Vec<f32>>, EmbeddingError> {
523 if texts.is_empty() {
524 return Ok(Vec::new());
525 }
526 let mut results = Vec::with_capacity(texts.len());
527 for chunk in texts.chunks(self.max_batch.max(1)) {
528 let encodings = self.tokenize(chunk)?;
529 results.extend(self.infer_rows(&encodings)?);
530 }
531 Ok(results)
532 }
533
534 fn embed_single(&self, text: &str) -> Result<Vec<f32>, EmbeddingError> {
536 let mut rows = self.embed_batch(&[text.to_string()])?;
537 rows.pop().ok_or_else(|| {
538 EmbeddingError::ParseError(
539 "model returned no embedding for single text".to_string(),
540 )
541 })
542 }
543
544 fn infer_dimension(session: &ort::session::Session) -> Result<usize, EmbeddingError> {
546 let outputs = session.outputs();
547 if outputs.is_empty() {
548 return Err(EmbeddingError::ParseError(
549 "ONNX model has no output nodes".to_string(),
550 ));
551 }
552 let dtype = outputs[0].dtype();
553 let shape = dtype.tensor_shape().ok_or_else(|| {
554 EmbeddingError::ParseError("Output is not a Tensor type".to_string())
555 })?;
556 let dim = shape
557 .iter()
558 .rev()
559 .find_map(|&d| if d > 0 { Some(d as usize) } else { None })
560 .ok_or_else(|| {
561 EmbeddingError::ParseError(format!(
562 "Cannot infer embedding dimension from model output shape: {:?}",
563 *shape
564 ))
565 })?;
566 Ok(dim)
567 }
568
569 fn infer_input_capability(
576 session: &ort::session::Session,
577 max_batch: usize,
578 max_seq_len: usize,
579 ) -> Result<(usize, usize), EmbeddingError> {
580 let inputs = session.inputs();
581 let input = inputs.first().ok_or_else(|| {
582 EmbeddingError::ParseError("ONNX model has no input nodes".to_string())
583 })?;
584 let shape = input.dtype().tensor_shape().ok_or_else(|| {
585 EmbeddingError::ParseError("Model input is not a Tensor type".to_string())
586 })?;
587 let batch_dim = shape.first().copied().unwrap_or(-1);
588 let seq_dim = shape.get(1).copied().unwrap_or(-1);
589 let batch_cap = if batch_dim > 0 {
590 (batch_dim as usize).min(max_batch)
591 } else {
592 max_batch
593 };
594 let seq_limit = if seq_dim > 0 {
595 seq_dim as usize
596 } else {
597 max_seq_len
598 };
599 Ok((batch_cap, seq_limit))
600 }
601
602 fn resolve_pad_id(tokenizer: &tokenizers::Tokenizer) -> u32 {
605 tokenizer
606 .get_padding()
607 .map(|p| p.pad_id)
608 .or_else(|| tokenizer.token_to_id("[PAD]"))
609 .unwrap_or(0)
610 }
611 }
612
613 fn discover_tokenizer(model_path: &Path) -> Result<PathBuf, EmbeddingError> {
615 let dir = model_path.parent().unwrap_or_else(|| Path::new("."));
616 let stem = model_path
617 .file_stem()
618 .map(|s| s.to_string_lossy().to_string())
619 .unwrap_or_default();
620 for candidate in [dir.join(format!("{stem}.json")), dir.join("tokenizer.json")] {
621 if candidate.is_file() {
622 return Ok(candidate);
623 }
624 }
625 Err(EmbeddingError::Config(format!(
626 "No tokenizer.json found next to ONNX model '{}'. The `local-embeddings` feature \
627 uses the HuggingFace `tokenizers` crate and requires a real tokenizer.json \
628 (WordPiece/BPE vocab). Place one as '{stem}.json' or 'tokenizer.json' next to the \
629 model, or pass it explicitly via `LocalEmbeddings::from_file_with_tokenizer`. The old \
630 byte-hash fake tokenizer was removed in P2-2: feeding fake token IDs to a neural \
631 embedding model yields garbage vectors.",
632 model_path.display()
633 )))
634 }
635
636 pub struct LocalEmbeddings {
655 inner: Arc<LocalInner>,
656 }
657
658 impl LocalEmbeddings {
659 pub fn from_file(model_path: impl AsRef<Path>) -> Result<Self, EmbeddingError> {
662 Self::builder().model_path(model_path).build()
663 }
664
665 pub fn from_file_with_tokenizer(
667 model_path: impl AsRef<Path>,
668 tokenizer_path: impl AsRef<Path>,
669 ) -> Result<Self, EmbeddingError> {
670 Self::builder()
671 .model_path(model_path)
672 .tokenizer_path(tokenizer_path)
673 .build()
674 }
675
676 pub fn builder() -> LocalEmbeddingsBuilder {
678 LocalEmbeddingsBuilder {
679 model_path: None,
680 tokenizer_path: None,
681 pool_size: default_pool_size(),
682 max_batch: DEFAULT_MAX_BATCH,
683 max_seq_len: DEFAULT_MAX_SEQ_LEN,
684 }
685 }
686 }
687
688 pub struct LocalEmbeddingsBuilder {
690 model_path: Option<PathBuf>,
691 tokenizer_path: Option<PathBuf>,
692 pool_size: usize,
693 max_batch: usize,
694 max_seq_len: usize,
695 }
696
697 impl LocalEmbeddingsBuilder {
698 pub fn model_path(mut self, path: impl AsRef<Path>) -> Self {
700 self.model_path = Some(path.as_ref().to_path_buf());
701 self
702 }
703
704 pub fn tokenizer_path(mut self, path: impl AsRef<Path>) -> Self {
706 self.tokenizer_path = Some(path.as_ref().to_path_buf());
707 self
708 }
709
710 pub fn pool_size(mut self, size: usize) -> Self {
712 self.pool_size = size;
713 self
714 }
715
716 pub fn max_batch(mut self, n: usize) -> Self {
718 self.max_batch = n;
719 self
720 }
721
722 pub fn max_seq_len(mut self, n: usize) -> Self {
724 self.max_seq_len = n;
725 self
726 }
727
728 pub fn build(self) -> Result<LocalEmbeddings, EmbeddingError> {
730 let model_path = self.model_path.ok_or_else(|| {
731 EmbeddingError::Config(
732 "model path is required: call LocalEmbeddings::builder().model_path(path)"
733 .to_string(),
734 )
735 })?;
736 let model_bytes = std::fs::read(&model_path).map_err(|e| {
737 EmbeddingError::ApiError(format!(
738 "Failed to read ONNX model '{}': {e}",
739 model_path.display()
740 ))
741 })?;
742 let tokenizer_path = match self.tokenizer_path {
743 Some(p) => p,
744 None => discover_tokenizer(&model_path)?,
745 };
746 let tokenizer = tokenizers::Tokenizer::from_file(&tokenizer_path).map_err(|e| {
747 EmbeddingError::Config(format!(
748 "Failed to load tokenizer '{}': {e}. Expected a HuggingFace tokenizer.json \
749 (WordPiece/BPE).",
750 tokenizer_path.display()
751 ))
752 })?;
753
754 let model_name = model_path
755 .file_stem()
756 .and_then(|s| s.to_str())
757 .unwrap_or("unknown")
758 .to_string();
759
760 let pool = SessionPool::new(model_bytes, self.pool_size)?;
761
762 let (dim, max_batch, seq_limit) = {
764 let mut guard = pool.acquire()?;
765 let session = guard.session()?;
766 let dim = LocalInner::infer_dimension(session)?;
767 let (max_batch, seq_limit) =
768 LocalInner::infer_input_capability(session, self.max_batch, self.max_seq_len)?;
769 (dim, max_batch, seq_limit)
770 };
771
772 Ok(LocalEmbeddings {
773 inner: Arc::new(LocalInner {
774 pool,
775 tokenizer,
776 dim,
777 model_name,
778 seq_limit,
779 max_batch,
780 }),
781 })
782 }
783 }
784
785 #[async_trait]
786 impl Embeddings for LocalEmbeddings {
787 async fn embed_query(&self, text: &str) -> Result<Vec<f32>, EmbeddingError> {
788 let inner = self.inner.clone();
791 let text = text.to_string();
792 tokio::task::spawn_blocking(move || inner.embed_single(&text))
793 .await
794 .map_err(|e| EmbeddingError::ApiError(format!("Task execution failed: {e}")))?
795 }
796
797 async fn embed_documents(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>, EmbeddingError> {
798 if texts.is_empty() {
799 return Ok(Vec::new());
800 }
801 if texts.iter().any(|t| t.trim().is_empty()) {
803 return Err(EmbeddingError::EmptyInput);
804 }
805
806 let inner = self.inner.clone();
807 let texts: Vec<String> = texts.iter().map(|s| s.to_string()).collect();
808 tokio::task::spawn_blocking(move || inner.embed_batch(&texts))
809 .await
810 .map_err(|e| EmbeddingError::ApiError(format!("Task execution failed: {e}")))?
811 }
812
813 fn dimension(&self) -> usize {
814 self.inner.dim
815 }
816
817 fn model_name(&self) -> &str {
818 &self.inner.model_name
819 }
820 }
821
822 #[cfg(test)]
823 mod nn_tests {
824 use super::*;
825
826 fn tiny_tokenizer(with_pad_in_vocab: bool) -> tokenizers::Tokenizer {
830 let mut vocab = serde_json::json!({
831 "[UNK]": 0,
832 "hello": 1,
833 "world": 2,
834 });
835 if with_pad_in_vocab {
836 vocab["[PAD]"] = serde_json::json!(3);
837 }
838 let json = serde_json::json!({
839 "version": "1.0",
840 "truncation": null,
841 "padding": null,
842 "added_tokens": [],
843 "normalizer": null,
844 "pre_tokenizer": { "type": "Whitespace" },
845 "post_processor": null,
846 "decoder": null,
847 "model": {
848 "type": "WordPiece",
849 "vocab": vocab,
850 "unk_token": "[UNK]",
851 "continuing_subword_prefix": "##",
852 "max_input_chars_per_word": 100
853 }
854 });
855 tokenizers::Tokenizer::from_bytes(json.to_string().as_bytes())
856 .expect("tiny tokenizer should deserialize")
857 }
858
859 #[test]
860 fn test_l2_normalize() {
861 let mut v = vec![3.0, 4.0];
862 crate::l2_normalize(&mut v);
863 let norm: f32 = v.iter().map(|x| x * x).sum::<f32>().sqrt();
864 assert!((norm - 1.0).abs() < 1e-5);
865 assert!((v[0] - 0.6).abs() < 1e-5);
866 assert!((v[1] - 0.8).abs() < 1e-5);
867 }
868
869 #[test]
870 fn test_l2_normalize_zero() {
871 let mut v = vec![0.0, 0.0, 0.0];
872 crate::l2_normalize(&mut v);
873 assert!(v.iter().all(|x| *x == 0.0));
874 }
875
876 #[test]
879 fn test_tokenize_real_wordpiece() {
880 let tok = tiny_tokenizer(false);
881 let enc = tok.encode("hello world", true).unwrap();
882 assert_eq!(enc.get_ids(), &[1u32, 2u32]);
883 assert_eq!(enc.get_attention_mask(), &[1u32, 1u32]);
884 }
885
886 #[test]
888 fn test_tokenize_unknown_word_uses_unk() {
889 let tok = tiny_tokenizer(false);
890 let enc = tok.encode("zzzznotinvocab", true).unwrap();
891 assert_eq!(enc.get_ids(), &[0u32]);
892 }
893
894 #[test]
896 fn test_build_batch_tensors_pads_to_longest() {
897 let tok = tiny_tokenizer(false);
898 let encodings = tok
899 .encode_batch(vec!["hello".to_string(), "hello world".to_string()], true)
900 .unwrap();
901 let (input_ids, attention_mask, token_type_ids, max_len) =
902 LocalInner::build_batch_tensors(&encodings, 8, 0);
903 assert_eq!(max_len, 2);
904 assert_eq!(input_ids, vec![1, 0, 1, 2]);
906 assert_eq!(attention_mask, vec![1, 0, 1, 1]);
907 assert_eq!(token_type_ids, vec![0, 0, 0, 0]);
908 }
909
910 #[test]
911 fn test_resolve_pad_id_with_pad_token() {
912 let tok = tiny_tokenizer(true);
913 assert_eq!(LocalInner::resolve_pad_id(&tok), 3);
914 }
915
916 #[test]
917 fn test_resolve_pad_id_defaults_zero() {
918 let tok = tiny_tokenizer(false);
919 assert_eq!(LocalInner::resolve_pad_id(&tok), 0);
920 }
921
922 #[test]
924 fn test_pool_rows_3d_masked() {
925 let shape = vec![2usize, 3, 2];
927 let data = vec![
928 1.0, 10.0, 2.0, 20.0, 3.0, 30.0, 4.0, 40.0, 5.0, 50.0, 6.0, 60.0,
931 ];
932 let masks = vec![vec![1, 1, 1], vec![1, 0, 0]];
933 let rows = LocalInner::pool_rows(&shape, &data, &masks, 2, 3).unwrap();
934 assert_eq!(rows.len(), 2);
935 assert!((rows[0][0] - 2.0).abs() < 1e-5);
937 assert!((rows[0][1] - 20.0).abs() < 1e-5);
938 assert!((rows[1][0] - 4.0).abs() < 1e-5);
940 assert!((rows[1][1] - 40.0).abs() < 1e-5);
941 }
942
943 #[test]
945 fn test_pool_rows_2d() {
946 let shape = vec![2usize, 2];
947 let data = vec![1.0, 2.0, 3.0, 4.0];
948 let masks = vec![vec![1], vec![1]];
949 let rows = LocalInner::pool_rows(&shape, &data, &masks, 2, 1).unwrap();
950 assert_eq!(rows, vec![vec![1.0, 2.0], vec![3.0, 4.0]]);
951 }
952
953 #[test]
955 fn test_pool_rows_batch_mismatch() {
956 let shape = vec![3usize, 2, 2];
957 let data = vec![0.0; 12];
958 let masks = vec![vec![1], vec![1]];
959 let err = LocalInner::pool_rows(&shape, &data, &masks, 2, 1).unwrap_err();
960 assert!(matches!(
961 err,
962 EmbeddingError::BatchMismatch {
963 expected: 2,
964 actual: 3
965 }
966 ));
967 }
968 }
969}
970
971#[cfg(feature = "local-embeddings")]
973pub use nn::{LocalEmbeddings, LocalEmbeddingsBuilder};
974
975#[cfg(not(feature = "local-embeddings"))]
989#[deprecated(
990 note = "LocalEmbeddings without the `local-embeddings` feature degrades to \
991 BagOfWordsEmbeddings (bag-of-words hash), not semantic neural embedding. \
992 Enable the `local-embeddings` feature, or use BagOfWordsEmbeddings explicitly."
993)]
994pub type LocalEmbeddings = BagOfWordsEmbeddings;
995
996#[cfg(test)]
1001mod tests {
1002 use super::*;
1003 use crate::cosine_similarity;
1004
1005 #[tokio::test]
1008 async fn test_bow_dimension() {
1009 let e = BagOfWordsEmbeddings::new(128);
1010 let v = e.embed_query("hello world").await.unwrap();
1011 assert_eq!(v.len(), 128);
1012 assert_eq!(e.dimension(), 128);
1013 }
1014
1015 #[tokio::test]
1016 async fn test_bow_same_text_same_vector() {
1017 let e = BagOfWordsEmbeddings::new(64);
1018 let a = e.embed_query("rust programming").await.unwrap();
1019 let b = e.embed_query("rust programming").await.unwrap();
1020 assert_eq!(a, b);
1021 }
1022
1023 #[tokio::test]
1024 async fn test_bow_different_text_different_vector() {
1025 let e = BagOfWordsEmbeddings::new(64);
1026 let a = e.embed_query("rust programming").await.unwrap();
1027 let b = e.embed_query("cooking recipe pasta").await.unwrap();
1028 assert_ne!(a, b);
1029 }
1030
1031 #[tokio::test]
1032 async fn test_bow_shared_words_more_similar() {
1033 let e = BagOfWordsEmbeddings::new(256);
1034 let base = e.embed_query("rust programming language").await.unwrap();
1035 let similar = e.embed_query("rust programming tutorial").await.unwrap();
1036 let different = e.embed_query("cooking pasta recipe").await.unwrap();
1037
1038 let sim_similar = cosine_similarity(&base, &similar).unwrap_or(0.0);
1039 let sim_different = cosine_similarity(&base, &different).unwrap_or(0.0);
1040 assert!(
1041 sim_similar > sim_different,
1042 "Shared words should be more similar: {} vs {}",
1043 sim_similar,
1044 sim_different
1045 );
1046 }
1047
1048 #[tokio::test]
1049 async fn test_bow_normalized() {
1050 let e = BagOfWordsEmbeddings::new(64);
1051 let v = e.embed_query("some text here").await.unwrap();
1052 let norm: f32 = v.iter().map(|x| x * x).sum::<f32>().sqrt();
1053 assert!((norm - 1.0).abs() < 1e-5, "norm = {}", norm);
1054 }
1055
1056 #[tokio::test]
1057 async fn test_bow_empty_text_returns_error() {
1058 let e = BagOfWordsEmbeddings::new(64);
1059 let result = e.embed_query("").await;
1060 assert!(result.is_err());
1061 assert!(matches!(result.unwrap_err(), EmbeddingError::EmptyInput));
1062 }
1063
1064 #[tokio::test]
1065 async fn test_bow_chinese_tokenize() {
1066 let e = BagOfWordsEmbeddings::new(128);
1067 let a = e.embed_query("机器学习").await.unwrap();
1068 let b = e.embed_query("机器学习").await.unwrap();
1069 assert_eq!(a, b);
1070 let c = e.embed_query("深度学习").await.unwrap();
1071 let sim = cosine_similarity(&a, &c).unwrap_or(0.0);
1072 assert!(
1073 sim > 0.0,
1074 "Shared '学习' should have positive similarity: {}",
1075 sim
1076 );
1077 }
1078
1079 #[test]
1080 fn test_bow_tokenize_english() {
1081 let t = BagOfWordsEmbeddings::tokenize("Hello, World! 123");
1082 assert!(t.contains(&"hello".to_string()));
1083 assert!(t.contains(&"world".to_string()));
1084 assert!(t.contains(&"123".to_string()));
1085 }
1086
1087 #[test]
1088 fn test_bow_tokenize_chinese() {
1089 let t = BagOfWordsEmbeddings::tokenize("机器学习");
1090 assert!(t.contains(&"机".to_string()));
1091 assert!(t.contains(&"学".to_string()));
1092 assert_eq!(t.len(), 4);
1093 }
1094
1095 #[test]
1096 fn test_bow_model_name() {
1097 let e = BagOfWordsEmbeddings::default_dim();
1098 assert_eq!(e.model_name(), "local-bow");
1099 }
1100
1101 #[allow(deprecated)]
1106 #[tokio::test]
1107 async fn test_local_embeddings_backward_compat() {
1108 let e = LocalEmbeddings::new(64);
1110 let v = e.embed_query("test backward compat").await.unwrap();
1111 assert_eq!(v.len(), 64);
1112 assert_eq!(e.model_name(), "local-bow");
1113 }
1114}