use std::sync::Mutex; use std::sync::OnceLock; use candle_core::{DType, Device, Tensor}; use candle_nn::VarBuilder; use candle_transformers::models::bert::{BertModel, Config}; use tokenizers::Tokenizer; struct CandleEmbeddingModel { model: BertModel, tokenizer: Tokenizer, device: Device, } impl CandleEmbeddingModel { fn new() -> Result { if cfg!(test) || std::env::var("MCP_OFFLINE_EMBEDDINGS").is_ok() { return Err("Offline embedding mode active (tests/offline flag)".to_string()); } let client = hf_hub::HFClientSync::new().map_err(|e| e.to_string())?; let repo = client.model("sentence-transformers", "all-MiniLM-L6-v2"); let config_file = repo .download_file() .filename("config.json") .send() .map_err(|e| format!("Failed to download config.json: {}", e))?; let tokenizer_file = repo .download_file() .filename("tokenizer.json") .send() .map_err(|e| format!("Failed to download tokenizer.json: {}", e))?; let weights_file = repo .download_file() .filename("model.safetensors") .send() .map_err(|e| format!("Failed to download model.safetensors: {}", e))?; let config_str = std::fs::read_to_string(&config_file) .map_err(|e| format!("Failed to read config.json: {}", e))?; let config: Config = serde_json::from_str(&config_str) .map_err(|e| format!("Failed to parse config.json: {}", e))?; let mut tokenizer = Tokenizer::from_file(&tokenizer_file) .map_err(|e| format!("Failed to load tokenizer: {}", e))?; tokenizer.with_padding(Some(tokenizers::PaddingParams::default())); let device = Device::Cpu; let vb = unsafe { VarBuilder::from_mmaped_safetensors(&[weights_file], DType::F32, &device) .map_err(|e| format!("Failed to load safetensors: {}", e))? }; let model = BertModel::load(vb, &config).map_err(|e| format!("Failed to load BertModel: {}", e))?; Ok(Self { model, tokenizer, device, }) } fn embed(&self, texts: &[String]) -> Result>, String> { if texts.is_empty() { return Ok(Vec::new()); } let encodings = self .tokenizer .encode_batch(texts.to_vec(), true) .map_err(|e| format!("Failed to encode texts: {}", e))?; let batch_size = encodings.len(); if batch_size == 0 { return Ok(Vec::new()); } let seq_len = encodings[0].get_ids().len(); if seq_len == 0 { return Ok(vec![vec![0.0; 384]; batch_size]); } let mut all_ids: Vec = Vec::with_capacity(batch_size * seq_len); let mut all_type_ids: Vec = Vec::with_capacity(batch_size * seq_len); let mut all_attention_mask: Vec = Vec::with_capacity(batch_size * seq_len); for enc in &encodings { all_ids.extend(enc.get_ids()); all_type_ids.extend(enc.get_type_ids()); all_attention_mask.extend(enc.get_attention_mask()); } let input_ids = Tensor::from_vec(all_ids, (batch_size, seq_len), &self.device) .map_err(|e| format!("Failed to build input_ids tensor: {}", e))?; let token_type_ids = Tensor::from_vec(all_type_ids, (batch_size, seq_len), &self.device) .map_err(|e| format!("Failed to build token_type_ids tensor: {}", e))?; let attention_mask = Tensor::from_vec(all_attention_mask, (batch_size, seq_len), &self.device) .map_err(|e| format!("Failed to build attention_mask tensor: {}", e))?; let sequence_output = self .model .forward(&input_ids, &token_type_ids, Some(&attention_mask)) .map_err(|e| format!("Bert forward failed: {}", e))?; // Mean pooling: sum(sequence_output * mask) / clamp(sum(mask), min=1e-9) let mask_f32 = attention_mask .to_dtype(DType::F32) .map_err(|e| e.to_string())? .unsqueeze(2) .map_err(|e| e.to_string())?; let sum_embeddings = sequence_output .broadcast_mul(&mask_f32) .map_err(|e| e.to_string())? .sum(1) .map_err(|e| e.to_string())?; let sum_mask = mask_f32 .sum(1) .map_err(|e| e.to_string())? .clamp(1e-9, f32::MAX) .map_err(|e| e.to_string())?; let mean_pooled = sum_embeddings .broadcast_div(&sum_mask) .map_err(|e| e.to_string())?; // L2 Normalization let norm = mean_pooled .sqr() .map_err(|e| e.to_string())? .sum_keepdim(1) .map_err(|e| e.to_string())? .sqrt() .map_err(|e| e.to_string())?; let normalized = mean_pooled .broadcast_div(&norm) .map_err(|e| e.to_string())?; normalized.to_vec2::().map_err(|e| e.to_string()) } } fn fallback_embed(text: &str) -> Vec { const DIM: usize = 384; let mut vec = vec![0.0f32; DIM]; let words: Vec<&str> = text.split_whitespace().collect(); if words.is_empty() { vec[0] = 1.0; return vec; } use std::hash::{Hash, Hasher}; for word in words { let clean: String = word .chars() .filter(|c| c.is_alphanumeric()) .flat_map(|c| c.to_lowercase()) .collect(); if clean.is_empty() { continue; } let mut hasher = std::collections::hash_map::DefaultHasher::new(); clean.hash(&mut hasher); let h = hasher.finish(); let idx = (h as usize) % DIM; let sign = if (h >> 32) & 1 == 0 { 1.0f32 } else { -1.0f32 }; vec[idx] += sign; let chars: Vec = clean.chars().collect(); for window in chars.windows(3) { let mut h2 = std::collections::hash_map::DefaultHasher::new(); window.hash(&mut h2); let hv = h2.finish(); let idx2 = (hv as usize) % DIM; let s2 = if (hv >> 32) & 1 == 0 { 0.5f32 } else { -0.5f32 }; vec[idx2] += s2; } } let norm_sq: f32 = vec.iter().map(|x| x * x).sum(); if norm_sq > 0.0 { let norm = norm_sq.sqrt(); for x in vec.iter_mut() { *x /= norm; } } else { vec[0] = 1.0; } vec } enum EmbeddingModel { Candle(CandleEmbeddingModel), Fallback, } impl EmbeddingModel { fn new() -> Self { match CandleEmbeddingModel::new() { Ok(model) => EmbeddingModel::Candle(model), Err(e) => { tracing::warn!( "Failed to initialize Candle BERT model ({e}); falling back to deterministic offline embeddings." ); EmbeddingModel::Fallback } } } fn embed(&self, texts: &[String]) -> Result>, String> { match self { EmbeddingModel::Candle(model) => model.embed(texts), EmbeddingModel::Fallback => Ok(texts.iter().map(|t| fallback_embed(t)).collect()), } } } static EMBEDDING_MODEL: OnceLock> = OnceLock::new(); static INIT_MUTEX: Mutex<()> = Mutex::new(()); fn get_embedding_model() -> Result<&'static Mutex, String> { if let Some(model) = EMBEDDING_MODEL.get() { return Ok(model); } let _guard = INIT_MUTEX.lock().map_err(|e| e.to_string())?; if let Some(model) = EMBEDDING_MODEL.get() { return Ok(model); } let model = EmbeddingModel::new(); let _ = EMBEDDING_MODEL.set(Mutex::new(model)); Ok(EMBEDDING_MODEL.get().unwrap()) } pub async fn generate_embedding_async(text: String) -> Result, String> { tokio::task::spawn_blocking(move || { let model_mutex = get_embedding_model()?; let model = model_mutex.lock().map_err(|e| e.to_string())?; let embeddings = model.embed(&[text])?; let emb = embeddings .into_iter() .next() .ok_or_else(|| "Embedding model returned no embeddings".to_string())?; if emb.is_empty() { return Err("Embedding model generated a 0-length vector".to_string()); } Ok(emb) }) .await .map_err(|e| e.to_string())? } pub fn cosine_similarity(a: &[f32], b: &[f32]) -> f32 { if a.is_empty() || b.is_empty() || a.len() != b.len() { return 0.0; } let mut dot_product = 0.0f32; let mut norm_a_sq = 0.0f32; let mut norm_b_sq = 0.0f32; let chunks_a = a.chunks_exact(8); let chunks_b = b.chunks_exact(8); let remainder_a = chunks_a.remainder(); let remainder_b = chunks_b.remainder(); for (ca, cb) in chunks_a.zip(chunks_b) { for i in 0..8 { let x = ca[i]; let y = cb[i]; dot_product += x * y; norm_a_sq += x * x; norm_b_sq += y * y; } } for (&x, &y) in remainder_a.iter().zip(remainder_b.iter()) { dot_product += x * y; norm_a_sq += x * x; norm_b_sq += y * y; } // Fast path: If vectors are already normalized (Candle & fallback embeddings), skip square roots if (norm_a_sq - 1.0).abs() < 1e-4 && (norm_b_sq - 1.0).abs() < 1e-4 { return dot_product.clamp(-1.0, 1.0); } let norm_product = norm_a_sq * norm_b_sq; if norm_product <= 0.0 { 0.0 } else { (dot_product / norm_product.sqrt()).clamp(-1.0, 1.0) } } pub async fn generate_embeddings_async(texts: Vec) -> Result>, String> { if texts.is_empty() { return Ok(Vec::new()); } tokio::task::spawn_blocking(move || { let model_mutex = get_embedding_model()?; let model = model_mutex.lock().map_err(|e| e.to_string())?; let mut all_embeddings = Vec::with_capacity(texts.len()); let mut current_chunk = Vec::new(); let mut current_chars = 0; const MAX_CHARS_PER_BATCH: usize = 16384; for text in texts { let text_len = text.len(); if !current_chunk.is_empty() && (current_chunk.len() >= 64 || current_chars + text_len > MAX_CHARS_PER_BATCH) { let chunk_vec = std::mem::take(&mut current_chunk); let chunk_embeddings = model.embed(&chunk_vec)?; all_embeddings.extend(chunk_embeddings); current_chars = 0; } current_chars += text_len; current_chunk.push(text); } if !current_chunk.is_empty() { let chunk_embeddings = model.embed(¤t_chunk)?; all_embeddings.extend(chunk_embeddings); } Ok(all_embeddings) }) .await .map_err(|e| e.to_string())? } #[cfg(test)] mod tests { use super::*; #[test] fn test_cosine_similarity_edge_cases() { assert_eq!(cosine_similarity(&[], &[]), 0.0); assert_eq!(cosine_similarity(&[1.0, 2.0], &[1.0]), 0.0); assert_eq!(cosine_similarity(&[0.0, 0.0], &[0.0, 0.0]), 0.0); let v1 = vec![1.0, 0.0, 0.0]; let v2 = vec![1.0, 0.0, 0.0]; assert!((cosine_similarity(&v1, &v2) - 1.0).abs() < 1e-5); let v3 = vec![0.0, 1.0, 0.0]; assert!((cosine_similarity(&v1, &v3) - 0.0).abs() < 1e-5); } #[tokio::test] async fn test_generate_embeddings_async_empty() { let res = generate_embeddings_async(vec![]).await.unwrap(); assert!(res.is_empty()); } #[tokio::test] async fn test_generate_embeddings_async_single_text() { let text = "test text".to_string(); let res = generate_embeddings_async(vec![text.clone()]).await.unwrap(); assert_eq!(res.len(), 1); assert_eq!(res[0].len(), 384); } #[tokio::test] async fn test_generate_embeddings_async_multiple_texts() { let texts = vec![ "test text 1".to_string(), "test text 2".to_string(), "test text 3".to_string(), ]; let res = generate_embeddings_async(texts.clone()).await.unwrap(); assert_eq!(res.len(), 3); for embedding in &res { assert_eq!(embedding.len(), 384); } } }