use std::sync::Arc;
pub trait Embedder: Send + Sync {
fn embed(&self, text: &str) -> Result<Vec<f32>, String>;
fn model_id(&self) -> &'static str;
fn dim(&self) -> usize;
}
pub struct FakeEmbedder {
pub dim: usize,
}
impl Default for FakeEmbedder {
fn default() -> Self {
Self { dim: 64 }
}
}
impl Embedder for FakeEmbedder {
fn embed(&self, text: &str) -> Result<Vec<f32>, String> {
let mut v = vec![0.0f32; self.dim];
for word in text.split_whitespace() {
let mut h = blake3::Hasher::new();
h.update(word.as_bytes());
let hash = h.finalize();
let bytes = hash.as_bytes();
let bucket = u32::from_le_bytes([bytes[0], bytes[1], bytes[2], bytes[3]]) as usize;
let sign = if bytes[4] & 1 == 0 { 1.0 } else { -1.0 };
v[bucket % self.dim] += sign;
let bucket2 = u32::from_le_bytes([bytes[5], bytes[6], bytes[7], bytes[8]]) as usize;
let sign2 = if bytes[9] & 1 == 0 { 1.0 } else { -1.0 };
v[bucket2 % self.dim] += sign2;
}
let norm: f32 = v.iter().map(|x| x * x).sum::<f32>().sqrt();
if norm > 0.0 {
for x in &mut v {
*x /= norm;
}
}
Ok(v)
}
fn model_id(&self) -> &'static str {
"fake-deterministic"
}
fn dim(&self) -> usize {
self.dim
}
}
#[cfg(feature = "fastembed")]
pub struct FastEmbedder {
model: std::sync::Mutex<fastembed::TextEmbedding>,
}
#[cfg(feature = "fastembed")]
impl FastEmbedder {
pub fn bge_small() -> Result<Self, String> {
let model = fastembed::TextEmbedding::try_new(
fastembed::InitOptions::new(fastembed::EmbeddingModel::BGESmallENV15)
.with_show_download_progress(false),
)
.map_err(|e| e.to_string())?;
Ok(Self {
model: std::sync::Mutex::new(model),
})
}
}
#[cfg(feature = "fastembed")]
impl Embedder for FastEmbedder {
fn embed(&self, text: &str) -> Result<Vec<f32>, String> {
let mut model = self.model.lock().map_err(|e| e.to_string())?;
let out = model
.embed(vec![text], None)
.map_err(|e: fastembed::Error| e.to_string())?;
out.into_iter()
.next()
.ok_or_else(|| "empty embedding".to_string())
}
fn model_id(&self) -> &'static str {
"bge-small"
}
fn dim(&self) -> usize {
384
}
}
pub type EmbedderRef = Arc<dyn Embedder>;