use async_trait::async_trait;
#[cfg(feature = "local-embeddings")]
use std::path::Path;
use crate::{EmbeddingError, Embeddings};
pub struct BagOfWordsEmbeddings {
dim: usize,
}
impl BagOfWordsEmbeddings {
pub fn new(dim: usize) -> Self {
Self { dim: dim.max(1) }
}
pub fn default_dim() -> Self {
Self::new(256)
}
fn tokenize(text: &str) -> Vec<String> {
let mut tokens = Vec::new();
let mut current = String::new();
for c in text.chars() {
if c.is_alphanumeric() {
if c.is_ascii() {
current.push(c.to_ascii_lowercase());
} else {
if !current.is_empty() {
tokens.push(std::mem::take(&mut current));
}
tokens.push(c.to_string());
}
} else if !current.is_empty() {
tokens.push(std::mem::take(&mut current));
}
}
if !current.is_empty() {
tokens.push(current);
}
tokens
}
fn hash(s: &str) -> u64 {
let mut h: u64 = 0xcbf29ce484222325;
for b in s.bytes() {
h ^= b as u64;
h = h.wrapping_mul(0x100000001b3);
}
h
}
fn embed(&self, text: &str) -> Vec<f32> {
let mut v = vec![0.0f32; self.dim];
for token in Self::tokenize(text) {
let idx = (Self::hash(&token) as usize) % self.dim;
v[idx] += 1.0;
}
let norm: f32 = v.iter().map(|x| x * x).sum::<f32>().sqrt();
if norm > 0.0 {
for x in &mut v {
*x /= norm;
}
}
v
}
}
impl Default for BagOfWordsEmbeddings {
fn default() -> Self {
Self::default_dim()
}
}
#[async_trait]
impl Embeddings for BagOfWordsEmbeddings {
async fn embed_query(&self, text: &str) -> Result<Vec<f32>, EmbeddingError> {
if text.trim().is_empty() {
return Err(EmbeddingError::EmptyInput);
}
Ok(self.embed(text))
}
fn dimension(&self) -> usize {
self.dim
}
fn model_name(&self) -> &str {
"local-bow"
}
}
#[cfg(feature = "local-embeddings")]
mod nn {
use super::*;
use ort::value::Tensor;
use std::path::PathBuf;
use std::sync::{Arc, Condvar, Mutex};
const DEFAULT_MAX_BATCH: usize = 32;
const DEFAULT_MAX_SEQ_LEN: usize = 512;
fn default_pool_size() -> usize {
std::thread::available_parallelism()
.map(|n| n.get().clamp(1, 8))
.unwrap_or(2)
}
fn lock_error<T>(_: std::sync::PoisonError<T>) -> EmbeddingError {
EmbeddingError::ApiError("session pool lock poisoned".to_string())
}
struct SessionPool {
model_bytes: Arc<Vec<u8>>,
idle: Mutex<Vec<ort::session::Session>>,
live: Mutex<usize>,
capacity: usize,
notify: Condvar,
}
impl SessionPool {
fn new(model_bytes: Vec<u8>, capacity: usize) -> Result<Arc<Self>, EmbeddingError> {
let capacity = capacity.max(1);
let pool = Arc::new(Self {
model_bytes: Arc::new(model_bytes),
idle: Mutex::new(Vec::new()),
live: Mutex::new(0),
capacity,
notify: Condvar::new(),
});
let session = pool.build_session()?;
*pool.live.lock().map_err(lock_error)? = 1;
pool.idle.lock().map_err(lock_error)?.push(session);
Ok(pool)
}
fn build_session(&self) -> Result<ort::session::Session, EmbeddingError> {
ort::session::Session::builder()
.map_err(|e| {
EmbeddingError::ApiError(format!("Failed to create ONNX SessionBuilder: {e}"))
})?
.commit_from_memory(&self.model_bytes)
.map_err(|e| {
EmbeddingError::ApiError(format!("Failed to load ONNX model from memory: {e}"))
})
}
fn acquire(self: &Arc<Self>) -> Result<SessionGuard, EmbeddingError> {
{
let mut idle = self.idle.lock().map_err(lock_error)?;
if let Some(session) = idle.pop() {
return Ok(SessionGuard {
pool: self.clone(),
session: Some(session),
});
}
}
let mut live = self.live.lock().map_err(lock_error)?;
loop {
{
let mut idle = self.idle.lock().map_err(lock_error)?;
if let Some(session) = idle.pop() {
return Ok(SessionGuard {
pool: self.clone(),
session: Some(session),
});
}
}
if *live < self.capacity {
*live += 1;
match self.build_session() {
Ok(session) => {
return Ok(SessionGuard {
pool: self.clone(),
session: Some(session),
});
}
Err(e) => {
*live -= 1;
self.notify.notify_one();
return Err(e);
}
}
}
live = self.notify.wait(live).map_err(lock_error)?;
}
}
fn release(&self, session: ort::session::Session) {
let Ok(mut idle) = self.idle.lock() else {
return;
};
idle.push(session);
drop(idle);
self.notify.notify_one();
}
}
struct SessionGuard {
pool: Arc<SessionPool>,
session: Option<ort::session::Session>,
}
impl SessionGuard {
fn session(&mut self) -> Result<&mut ort::session::Session, EmbeddingError> {
self.session.as_mut().ok_or_else(|| {
EmbeddingError::ApiError("session guard already released".to_string())
})
}
}
impl Drop for SessionGuard {
fn drop(&mut self) {
if let Some(session) = self.session.take() {
self.pool.release(session);
}
}
}
struct LocalInner {
pool: Arc<SessionPool>,
tokenizer: tokenizers::Tokenizer,
dim: usize,
model_name: String,
seq_limit: usize,
max_batch: usize,
}
impl LocalInner {
fn tokenize(&self, texts: &[String]) -> Result<Vec<tokenizers::Encoding>, EmbeddingError> {
self.tokenizer
.encode_batch(texts.to_vec(), true)
.map_err(|e| {
EmbeddingError::Config(format!("Tokenizer failed to encode input: {e}"))
})
}
fn build_batch_tensors(
encodings: &[tokenizers::Encoding],
seq_limit: usize,
pad_id: u32,
) -> (Vec<i64>, Vec<i64>, Vec<i64>, usize) {
let batch = encodings.len();
let lens: Vec<usize> = encodings
.iter()
.map(|e| e.get_ids().len().min(seq_limit))
.collect();
let max_len = lens.iter().copied().max().unwrap_or(0);
let pad = pad_id as i64;
let mut input_ids = vec![pad; batch * max_len];
let mut attention_mask = vec![0i64; batch * max_len];
let mut token_type_ids = vec![0i64; batch * max_len];
for (b, enc) in encodings.iter().enumerate() {
let ids = enc.get_ids();
let mask = enc.get_attention_mask();
let types = enc.get_type_ids();
let len = lens[b];
let row = b * max_len;
for i in 0..len {
input_ids[row + i] = ids[i] as i64;
attention_mask[row + i] = mask.get(i).copied().unwrap_or(1) as i64;
token_type_ids[row + i] = types.get(i).copied().unwrap_or(0) as i64;
}
}
(input_ids, attention_mask, token_type_ids, max_len)
}
fn infer_rows(
&self,
encodings: &[tokenizers::Encoding],
) -> Result<Vec<Vec<f32>>, EmbeddingError> {
if encodings.is_empty() {
return Err(EmbeddingError::EmptyInput);
}
let pad_id = Self::resolve_pad_id(&self.tokenizer);
let (input_ids, attention_mask, token_type_ids, fed_seq_len) =
Self::build_batch_tensors(encodings, self.seq_limit, pad_id);
if fed_seq_len == 0 {
return Err(EmbeddingError::EmptyInput);
}
let batch = encodings.len();
let mask_rows: Vec<Vec<i64>> = (0..batch)
.map(|b| attention_mask[b * fed_seq_len..(b + 1) * fed_seq_len].to_vec())
.collect();
let input_shape = vec![batch as i64, fed_seq_len as i64];
let input_tensor = Tensor::from_array((input_shape, input_ids)).map_err(|e| {
EmbeddingError::ApiError(format!("Failed to construct input_ids tensor: {e}"))
})?;
let attention_tensor = Tensor::from_array((input_shape.clone(), attention_mask))
.map_err(|e| {
EmbeddingError::ApiError(format!(
"Failed to construct attention_mask tensor: {e}"
))
})?;
let type_tensor =
Tensor::from_array((input_shape.clone(), token_type_ids)).map_err(|e| {
EmbeddingError::ApiError(format!(
"Failed to construct token_type_ids tensor: {e}"
))
})?;
let mut guard = self.pool.acquire()?;
let input_names: Vec<String> = {
let session = guard.session()?;
session
.inputs()
.iter()
.map(|o| o.name().to_string())
.collect()
};
let mut input_ids_slot = Some(input_tensor);
let mut attention_slot = Some(attention_tensor);
let mut type_slot = Some(type_tensor);
let mut named: Vec<(String, Tensor<i64>)> = Vec::with_capacity(input_names.len());
for name in input_names {
let tensor = match name.as_str() {
"input_ids" => input_ids_slot.take().ok_or_else(|| {
EmbeddingError::ParseError(
"ONNX model declares duplicate 'input_ids' input".to_string(),
)
})?,
"attention_mask" => attention_slot.take().ok_or_else(|| {
EmbeddingError::ParseError(
"ONNX model declares duplicate 'attention_mask' input".to_string(),
)
})?,
"token_type_ids" => type_slot.take().ok_or_else(|| {
EmbeddingError::ParseError(
"ONNX model declares duplicate 'token_type_ids' input".to_string(),
)
})?,
other => {
return Err(EmbeddingError::ParseError(format!(
"Unsupported ONNX model input '{other}': the `local-embeddings` \
feature only supports input_ids / attention_mask / token_type_ids"
)));
}
};
named.push((name, tensor));
}
if named.is_empty() {
return Err(EmbeddingError::ParseError(
"ONNX model declares no supported inputs".to_string(),
));
}
let outputs = guard
.session()?
.run(named)
.map_err(|e| EmbeddingError::ApiError(format!("ONNX inference failed: {e}")))?;
let output_value = outputs.get(0).ok_or_else(|| {
EmbeddingError::ParseError("ONNX model has no output".to_string())
})?;
let (shape, data) = output_value.try_extract_tensor::<f32>().map_err(|e| {
EmbeddingError::ParseError(format!("Failed to extract output tensor: {e}"))
})?;
let shape_vec: Vec<usize> = shape.iter().map(|&d| d as usize).collect();
let mut rows = Self::pool_rows(&shape_vec, data, &mask_rows, batch, fed_seq_len)?;
for row in &mut rows {
crate::l2_normalize(row);
}
Ok(rows)
}
fn pool_rows(
shape: &[usize],
data: &[f32],
masks: &[Vec<i64>],
batch: usize,
fed_seq_len: usize,
) -> Result<Vec<Vec<f32>>, EmbeddingError> {
match shape.len() {
3 => {
let out_batch = shape[0];
if out_batch != batch {
return Err(EmbeddingError::BatchMismatch {
expected: batch,
actual: out_batch,
});
}
let out_seq = shape[1];
let dim = shape[2];
let seq = out_seq.min(fed_seq_len);
let mut result = vec![vec![0.0f32; dim]; batch];
for b in 0..batch {
let mask = &masks[b];
let mut count = 0usize;
for s in 0..seq {
if s >= mask.len() || mask[s] == 0 {
continue; }
count += 1;
let base = (b * out_seq + s) * dim;
for d in 0..dim {
result[b][d] += data[base + d];
}
}
if count > 0 {
for d in 0..dim {
result[b][d] /= count as f32;
}
}
}
Ok(result)
}
2 => {
let out_batch = shape[0];
if out_batch != batch {
return Err(EmbeddingError::BatchMismatch {
expected: batch,
actual: out_batch,
});
}
let dim = shape[1];
let mut result = Vec::with_capacity(batch);
for b in 0..batch {
let base = b * dim;
result.push(data[base..base + dim].to_vec());
}
Ok(result)
}
_ => Err(EmbeddingError::ParseError(format!(
"Unsupported output dimension count: {}",
shape.len()
))),
}
}
fn embed_batch(&self, texts: &[String]) -> Result<Vec<Vec<f32>>, EmbeddingError> {
if texts.is_empty() {
return Ok(Vec::new());
}
let mut results = Vec::with_capacity(texts.len());
for chunk in texts.chunks(self.max_batch.max(1)) {
let encodings = self.tokenize(chunk)?;
results.extend(self.infer_rows(&encodings)?);
}
Ok(results)
}
fn embed_single(&self, text: &str) -> Result<Vec<f32>, EmbeddingError> {
let mut rows = self.embed_batch(&[text.to_string()])?;
rows.pop().ok_or_else(|| {
EmbeddingError::ParseError(
"model returned no embedding for single text".to_string(),
)
})
}
fn infer_dimension(session: &ort::session::Session) -> Result<usize, EmbeddingError> {
let outputs = session.outputs();
if outputs.is_empty() {
return Err(EmbeddingError::ParseError(
"ONNX model has no output nodes".to_string(),
));
}
let dtype = outputs[0].dtype();
let shape = dtype.tensor_shape().ok_or_else(|| {
EmbeddingError::ParseError("Output is not a Tensor type".to_string())
})?;
let dim = shape
.iter()
.rev()
.find_map(|&d| if d > 0 { Some(d as usize) } else { None })
.ok_or_else(|| {
EmbeddingError::ParseError(format!(
"Cannot infer embedding dimension from model output shape: {:?}",
*shape
))
})?;
Ok(dim)
}
fn infer_input_capability(
session: &ort::session::Session,
max_batch: usize,
max_seq_len: usize,
) -> Result<(usize, usize), EmbeddingError> {
let inputs = session.inputs();
let input = inputs.first().ok_or_else(|| {
EmbeddingError::ParseError("ONNX model has no input nodes".to_string())
})?;
let shape = input.dtype().tensor_shape().ok_or_else(|| {
EmbeddingError::ParseError("Model input is not a Tensor type".to_string())
})?;
let batch_dim = shape.first().copied().unwrap_or(-1);
let seq_dim = shape.get(1).copied().unwrap_or(-1);
let batch_cap = if batch_dim > 0 {
(batch_dim as usize).min(max_batch)
} else {
max_batch
};
let seq_limit = if seq_dim > 0 {
seq_dim as usize
} else {
max_seq_len
};
Ok((batch_cap, seq_limit))
}
fn resolve_pad_id(tokenizer: &tokenizers::Tokenizer) -> u32 {
tokenizer
.get_padding()
.map(|p| p.pad_id)
.or_else(|| tokenizer.token_to_id("[PAD]"))
.unwrap_or(0)
}
}
fn discover_tokenizer(model_path: &Path) -> Result<PathBuf, EmbeddingError> {
let dir = model_path.parent().unwrap_or_else(|| Path::new("."));
let stem = model_path
.file_stem()
.map(|s| s.to_string_lossy().to_string())
.unwrap_or_default();
for candidate in [dir.join(format!("{stem}.json")), dir.join("tokenizer.json")] {
if candidate.is_file() {
return Ok(candidate);
}
}
Err(EmbeddingError::Config(format!(
"No tokenizer.json found next to ONNX model '{}'. The `local-embeddings` feature \
uses the HuggingFace `tokenizers` crate and requires a real tokenizer.json \
(WordPiece/BPE vocab). Place one as '{stem}.json' or 'tokenizer.json' next to the \
model, or pass it explicitly via `LocalEmbeddings::from_file_with_tokenizer`. The old \
byte-hash fake tokenizer was removed in P2-2: feeding fake token IDs to a neural \
embedding model yields garbage vectors.",
model_path.display()
)))
}
pub struct LocalEmbeddings {
inner: Arc<LocalInner>,
}
impl LocalEmbeddings {
pub fn from_file(model_path: impl AsRef<Path>) -> Result<Self, EmbeddingError> {
Self::builder().model_path(model_path).build()
}
pub fn from_file_with_tokenizer(
model_path: impl AsRef<Path>,
tokenizer_path: impl AsRef<Path>,
) -> Result<Self, EmbeddingError> {
Self::builder()
.model_path(model_path)
.tokenizer_path(tokenizer_path)
.build()
}
pub fn builder() -> LocalEmbeddingsBuilder {
LocalEmbeddingsBuilder {
model_path: None,
tokenizer_path: None,
pool_size: default_pool_size(),
max_batch: DEFAULT_MAX_BATCH,
max_seq_len: DEFAULT_MAX_SEQ_LEN,
}
}
}
pub struct LocalEmbeddingsBuilder {
model_path: Option<PathBuf>,
tokenizer_path: Option<PathBuf>,
pool_size: usize,
max_batch: usize,
max_seq_len: usize,
}
impl LocalEmbeddingsBuilder {
pub fn model_path(mut self, path: impl AsRef<Path>) -> Self {
self.model_path = Some(path.as_ref().to_path_buf());
self
}
pub fn tokenizer_path(mut self, path: impl AsRef<Path>) -> Self {
self.tokenizer_path = Some(path.as_ref().to_path_buf());
self
}
pub fn pool_size(mut self, size: usize) -> Self {
self.pool_size = size;
self
}
pub fn max_batch(mut self, n: usize) -> Self {
self.max_batch = n;
self
}
pub fn max_seq_len(mut self, n: usize) -> Self {
self.max_seq_len = n;
self
}
pub fn build(self) -> Result<LocalEmbeddings, EmbeddingError> {
let model_path = self.model_path.ok_or_else(|| {
EmbeddingError::Config(
"model path is required: call LocalEmbeddings::builder().model_path(path)"
.to_string(),
)
})?;
let model_bytes = std::fs::read(&model_path).map_err(|e| {
EmbeddingError::ApiError(format!(
"Failed to read ONNX model '{}': {e}",
model_path.display()
))
})?;
let tokenizer_path = match self.tokenizer_path {
Some(p) => p,
None => discover_tokenizer(&model_path)?,
};
let tokenizer = tokenizers::Tokenizer::from_file(&tokenizer_path).map_err(|e| {
EmbeddingError::Config(format!(
"Failed to load tokenizer '{}': {e}. Expected a HuggingFace tokenizer.json \
(WordPiece/BPE).",
tokenizer_path.display()
))
})?;
let model_name = model_path
.file_stem()
.and_then(|s| s.to_str())
.unwrap_or("unknown")
.to_string();
let pool = SessionPool::new(model_bytes, self.pool_size)?;
let (dim, max_batch, seq_limit) = {
let mut guard = pool.acquire()?;
let session = guard.session()?;
let dim = LocalInner::infer_dimension(session)?;
let (max_batch, seq_limit) =
LocalInner::infer_input_capability(session, self.max_batch, self.max_seq_len)?;
(dim, max_batch, seq_limit)
};
Ok(LocalEmbeddings {
inner: Arc::new(LocalInner {
pool,
tokenizer,
dim,
model_name,
seq_limit,
max_batch,
}),
})
}
}
#[async_trait]
impl Embeddings for LocalEmbeddings {
async fn embed_query(&self, text: &str) -> Result<Vec<f32>, EmbeddingError> {
let inner = self.inner.clone();
let text = text.to_string();
tokio::task::spawn_blocking(move || inner.embed_single(&text))
.await
.map_err(|e| EmbeddingError::ApiError(format!("Task execution failed: {e}")))?
}
async fn embed_documents(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>, EmbeddingError> {
if texts.is_empty() {
return Ok(Vec::new());
}
if texts.iter().any(|t| t.trim().is_empty()) {
return Err(EmbeddingError::EmptyInput);
}
let inner = self.inner.clone();
let texts: Vec<String> = texts.iter().map(|s| s.to_string()).collect();
tokio::task::spawn_blocking(move || inner.embed_batch(&texts))
.await
.map_err(|e| EmbeddingError::ApiError(format!("Task execution failed: {e}")))?
}
fn dimension(&self) -> usize {
self.inner.dim
}
fn model_name(&self) -> &str {
&self.inner.model_name
}
}
#[cfg(test)]
mod nn_tests {
use super::*;
fn tiny_tokenizer(with_pad_in_vocab: bool) -> tokenizers::Tokenizer {
let mut vocab = serde_json::json!({
"[UNK]": 0,
"hello": 1,
"world": 2,
});
if with_pad_in_vocab {
vocab["[PAD]"] = serde_json::json!(3);
}
let json = serde_json::json!({
"version": "1.0",
"truncation": null,
"padding": null,
"added_tokens": [],
"normalizer": null,
"pre_tokenizer": { "type": "Whitespace" },
"post_processor": null,
"decoder": null,
"model": {
"type": "WordPiece",
"vocab": vocab,
"unk_token": "[UNK]",
"continuing_subword_prefix": "##",
"max_input_chars_per_word": 100
}
});
tokenizers::Tokenizer::from_bytes(json.to_string().as_bytes())
.expect("tiny tokenizer should deserialize")
}
#[test]
fn test_l2_normalize() {
let mut v = vec![3.0, 4.0];
crate::l2_normalize(&mut v);
let norm: f32 = v.iter().map(|x| x * x).sum::<f32>().sqrt();
assert!((norm - 1.0).abs() < 1e-5);
assert!((v[0] - 0.6).abs() < 1e-5);
assert!((v[1] - 0.8).abs() < 1e-5);
}
#[test]
fn test_l2_normalize_zero() {
let mut v = vec![0.0, 0.0, 0.0];
crate::l2_normalize(&mut v);
assert!(v.iter().all(|x| *x == 0.0));
}
#[test]
fn test_tokenize_real_wordpiece() {
let tok = tiny_tokenizer(false);
let enc = tok.encode("hello world", true).unwrap();
assert_eq!(enc.get_ids(), &[1u32, 2u32]);
assert_eq!(enc.get_attention_mask(), &[1u32, 1u32]);
}
#[test]
fn test_tokenize_unknown_word_uses_unk() {
let tok = tiny_tokenizer(false);
let enc = tok.encode("zzzznotinvocab", true).unwrap();
assert_eq!(enc.get_ids(), &[0u32]);
}
#[test]
fn test_build_batch_tensors_pads_to_longest() {
let tok = tiny_tokenizer(false);
let encodings = tok
.encode_batch(vec!["hello".to_string(), "hello world".to_string()], true)
.unwrap();
let (input_ids, attention_mask, token_type_ids, max_len) =
LocalInner::build_batch_tensors(&encodings, 8, 0);
assert_eq!(max_len, 2);
assert_eq!(input_ids, vec![1, 0, 1, 2]);
assert_eq!(attention_mask, vec![1, 0, 1, 1]);
assert_eq!(token_type_ids, vec![0, 0, 0, 0]);
}
#[test]
fn test_resolve_pad_id_with_pad_token() {
let tok = tiny_tokenizer(true);
assert_eq!(LocalInner::resolve_pad_id(&tok), 3);
}
#[test]
fn test_resolve_pad_id_defaults_zero() {
let tok = tiny_tokenizer(false);
assert_eq!(LocalInner::resolve_pad_id(&tok), 0);
}
#[test]
fn test_pool_rows_3d_masked() {
let shape = vec![2usize, 3, 2];
let data = vec![
1.0, 10.0, 2.0, 20.0, 3.0, 30.0, 4.0, 40.0, 5.0, 50.0, 6.0, 60.0,
];
let masks = vec![vec![1, 1, 1], vec![1, 0, 0]];
let rows = LocalInner::pool_rows(&shape, &data, &masks, 2, 3).unwrap();
assert_eq!(rows.len(), 2);
assert!((rows[0][0] - 2.0).abs() < 1e-5);
assert!((rows[0][1] - 20.0).abs() < 1e-5);
assert!((rows[1][0] - 4.0).abs() < 1e-5);
assert!((rows[1][1] - 40.0).abs() < 1e-5);
}
#[test]
fn test_pool_rows_2d() {
let shape = vec![2usize, 2];
let data = vec![1.0, 2.0, 3.0, 4.0];
let masks = vec![vec![1], vec![1]];
let rows = LocalInner::pool_rows(&shape, &data, &masks, 2, 1).unwrap();
assert_eq!(rows, vec![vec![1.0, 2.0], vec![3.0, 4.0]]);
}
#[test]
fn test_pool_rows_batch_mismatch() {
let shape = vec![3usize, 2, 2];
let data = vec![0.0; 12];
let masks = vec![vec![1], vec![1]];
let err = LocalInner::pool_rows(&shape, &data, &masks, 2, 1).unwrap_err();
assert!(matches!(
err,
EmbeddingError::BatchMismatch {
expected: 2,
actual: 3
}
));
}
}
}
#[cfg(feature = "local-embeddings")]
pub use nn::{LocalEmbeddings, LocalEmbeddingsBuilder};
#[cfg(not(feature = "local-embeddings"))]
#[deprecated(
note = "LocalEmbeddings without the `local-embeddings` feature degrades to \
BagOfWordsEmbeddings (bag-of-words hash), not semantic neural embedding. \
Enable the `local-embeddings` feature, or use BagOfWordsEmbeddings explicitly."
)]
pub type LocalEmbeddings = BagOfWordsEmbeddings;
#[cfg(test)]
mod tests {
use super::*;
use crate::cosine_similarity;
#[tokio::test]
async fn test_bow_dimension() {
let e = BagOfWordsEmbeddings::new(128);
let v = e.embed_query("hello world").await.unwrap();
assert_eq!(v.len(), 128);
assert_eq!(e.dimension(), 128);
}
#[tokio::test]
async fn test_bow_same_text_same_vector() {
let e = BagOfWordsEmbeddings::new(64);
let a = e.embed_query("rust programming").await.unwrap();
let b = e.embed_query("rust programming").await.unwrap();
assert_eq!(a, b);
}
#[tokio::test]
async fn test_bow_different_text_different_vector() {
let e = BagOfWordsEmbeddings::new(64);
let a = e.embed_query("rust programming").await.unwrap();
let b = e.embed_query("cooking recipe pasta").await.unwrap();
assert_ne!(a, b);
}
#[tokio::test]
async fn test_bow_shared_words_more_similar() {
let e = BagOfWordsEmbeddings::new(256);
let base = e.embed_query("rust programming language").await.unwrap();
let similar = e.embed_query("rust programming tutorial").await.unwrap();
let different = e.embed_query("cooking pasta recipe").await.unwrap();
let sim_similar = cosine_similarity(&base, &similar).unwrap_or(0.0);
let sim_different = cosine_similarity(&base, &different).unwrap_or(0.0);
assert!(
sim_similar > sim_different,
"Shared words should be more similar: {} vs {}",
sim_similar,
sim_different
);
}
#[tokio::test]
async fn test_bow_normalized() {
let e = BagOfWordsEmbeddings::new(64);
let v = e.embed_query("some text here").await.unwrap();
let norm: f32 = v.iter().map(|x| x * x).sum::<f32>().sqrt();
assert!((norm - 1.0).abs() < 1e-5, "norm = {}", norm);
}
#[tokio::test]
async fn test_bow_empty_text_returns_error() {
let e = BagOfWordsEmbeddings::new(64);
let result = e.embed_query("").await;
assert!(result.is_err());
assert!(matches!(result.unwrap_err(), EmbeddingError::EmptyInput));
}
#[tokio::test]
async fn test_bow_chinese_tokenize() {
let e = BagOfWordsEmbeddings::new(128);
let a = e.embed_query("机器学习").await.unwrap();
let b = e.embed_query("机器学习").await.unwrap();
assert_eq!(a, b);
let c = e.embed_query("深度学习").await.unwrap();
let sim = cosine_similarity(&a, &c).unwrap_or(0.0);
assert!(
sim > 0.0,
"Shared '学习' should have positive similarity: {}",
sim
);
}
#[test]
fn test_bow_tokenize_english() {
let t = BagOfWordsEmbeddings::tokenize("Hello, World! 123");
assert!(t.contains(&"hello".to_string()));
assert!(t.contains(&"world".to_string()));
assert!(t.contains(&"123".to_string()));
}
#[test]
fn test_bow_tokenize_chinese() {
let t = BagOfWordsEmbeddings::tokenize("机器学习");
assert!(t.contains(&"机".to_string()));
assert!(t.contains(&"学".to_string()));
assert_eq!(t.len(), 4);
}
#[test]
fn test_bow_model_name() {
let e = BagOfWordsEmbeddings::default_dim();
assert_eq!(e.model_name(), "local-bow");
}
#[allow(deprecated)]
#[tokio::test]
async fn test_local_embeddings_backward_compat() {
let e = LocalEmbeddings::new(64);
let v = e.embed_query("test backward compat").await.unwrap();
assert_eq!(v.len(), 64);
assert_eq!(e.model_name(), "local-bow");
}
}