use fastembed::{EmbeddingModel, Error as FastEmbedError, InitOptions, TextEmbedding};
use serde::{Deserialize, Serialize};
pub struct EmbeddingGenerator {
model: TextEmbedding,
}
impl EmbeddingGenerator {
pub fn new(model_name: EmbeddingModel, cache_dir: Option<std::path::PathBuf>) -> Result<Self, FastEmbedError> {
let mut opts = InitOptions::new(model_name);
if let Some(dir) = cache_dir {
opts = opts.with_cache_dir(dir);
}
let model = TextEmbedding::try_new(opts)?;
Ok(EmbeddingGenerator { model })
}
pub fn generate_embeddings(&self, documents: &[&str]) -> Result<Vec<Vec<f32>>, FastEmbedError> {
self.model.embed(documents.to_vec(), None)
}
}
#[derive(Debug, Clone, Serialize, Deserialize)] pub struct DocumentToUpsert {
pub file_path: String,
pub vector: Vec<f32>,
pub source: String, }
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_embedding_generator_init_and_embed() -> Result<(), FastEmbedError> {
let model_name = EmbeddingModel::AllMiniLML6V2; let generator = EmbeddingGenerator::new(model_name.clone(), None)?;
let documents = vec!["This is a test document.", "Another document."];
let embeddings = generator.generate_embeddings(&documents)?;
assert_eq!(embeddings.len(), 2);
let expected_dim = TextEmbedding::list_supported_models()
.iter()
.find(|m| m.model == model_name)
.map(|m| m.dim)
.unwrap_or(0);
if expected_dim > 0 {
assert_eq!(embeddings[0].len(), expected_dim);
assert_eq!(embeddings[1].len(), expected_dim);
}
Ok(())
}
}