use async_trait::async_trait;
use fastembed::{EmbeddingModel, InitOptions, TextEmbedding};
use klieo_core::error::MemoryError;
use std::sync::{Arc, Mutex};
use crate::Embedder;
pub const DEFAULT_MODEL: EmbeddingModel = EmbeddingModel::AllMiniLML6V2;
pub const DEFAULT_DIM: usize = 384;
pub struct FastEmbedEmbedder {
inner: Arc<Mutex<TextEmbedding>>,
dim: usize,
}
impl FastEmbedEmbedder {
pub fn new() -> Result<Self, MemoryError> {
Self::with_model(DEFAULT_MODEL, DEFAULT_DIM)
}
pub fn with_model(model: EmbeddingModel, dim: usize) -> Result<Self, MemoryError> {
let inner = TextEmbedding::try_new(InitOptions::new(model))
.map_err(|e| MemoryError::Embedding(format!("fastembed init: {e}")))?;
Ok(Self {
inner: Arc::new(Mutex::new(inner)),
dim,
})
}
}
#[async_trait]
impl Embedder for FastEmbedEmbedder {
fn dimension(&self) -> usize {
self.dim
}
async fn embed(&self, texts: &[String]) -> Result<Vec<Vec<f32>>, MemoryError> {
let texts_owned: Vec<String> = texts.to_vec();
let inner = self.inner.clone();
let dim = self.dim;
tokio::task::spawn_blocking(move || -> Result<Vec<Vec<f32>>, MemoryError> {
let mut guard = inner.lock().unwrap_or_else(|p| p.into_inner());
let docs = guard
.embed(texts_owned, None)
.map_err(|e| MemoryError::Embedding(format!("fastembed embed: {e}")))?;
for v in &docs {
if v.len() != dim {
return Err(MemoryError::Embedding(format!(
"fastembed produced {}-dim vector, expected {dim}",
v.len()
)));
}
}
Ok(docs)
})
.await
.map_err(|e| MemoryError::Embedding(format!("blocking task join: {e}")))?
}
}