Files
mcp-memory/server/src/embedding.rs
T

390 lines
13 KiB
Rust

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<Self, String> {
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<Vec<Vec<f32>>, 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<u32> = Vec::with_capacity(batch_size * seq_len);
let mut all_type_ids: Vec<u32> = Vec::with_capacity(batch_size * seq_len);
let mut all_attention_mask: Vec<u32> = 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::<f32>().map_err(|e| e.to_string())
}
}
fn fallback_embed(text: &str) -> Vec<f32> {
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;
}
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 hash_bytes = blake3::hash(clean.as_bytes());
let h = u64::from_le_bytes(hash_bytes.as_bytes()[0..8].try_into().unwrap());
let idx = (h as usize) % DIM;
let sign = if (h >> 32) & 1 == 0 { 1.0f32 } else { -1.0f32 };
vec[idx] += sign;
let chars: Vec<char> = clean.chars().collect();
for window in chars.windows(3) {
let window_str: String = window.iter().collect();
let h2_bytes = blake3::hash(window_str.as_bytes());
let hv = u64::from_le_bytes(h2_bytes.as_bytes()[0..8].try_into().unwrap());
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<Vec<Vec<f32>>, String> {
match self {
EmbeddingModel::Candle(model) => model.embed(texts),
EmbeddingModel::Fallback => Ok(texts.iter().map(|t| fallback_embed(t)).collect()),
}
}
}
static EMBEDDING_MODEL: OnceLock<Mutex<EmbeddingModel>> = OnceLock::new();
static INIT_MUTEX: Mutex<()> = Mutex::new(());
fn get_embedding_model() -> Result<&'static Mutex<EmbeddingModel>, 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<Vec<f32>, 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<String>) -> Result<Vec<Vec<f32>>, 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(&current_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);
}
}
#[test]
fn test_fallback_embed_deterministic_stability() {
let text = "The quick brown fox jumps over the lazy dog";
let emb1 = fallback_embed(text);
let emb2 = fallback_embed(text);
assert_eq!(emb1.len(), 384);
assert_eq!(emb1, emb2);
assert!((cosine_similarity(&emb1, &emb2) - 1.0).abs() < 1e-5);
}
}