Skip to main content

semtree_embed/
fastembed.rs

1use async_trait::async_trait;
2use fastembed::{EmbeddingModel, InitOptions, TextEmbedding};
3
4use crate::{EmbedError, Embedder, Embedding};
5
6pub struct FastEmbedder {
7    model: TextEmbedding,
8    model_id: String,
9    dimension: usize,
10}
11
12impl FastEmbedder {
13    pub fn new() -> Result<Self, EmbedError> {
14        Self::with_model(EmbeddingModel::AllMiniLML6V2)
15    }
16
17    pub fn with_model(model: EmbeddingModel) -> Result<Self, EmbedError> {
18        // `{model:?}` yields the variant name (e.g. `AllMiniLML6V2`), which is
19        // stable across releases and unique per model.
20        let model_id = format!("fastembed:{model:?}");
21        let te = TextEmbedding::try_new(InitOptions::new(model))
22            .map_err(|e| EmbedError::ModelLoad(e.to_string()))?;
23
24        // Probe the real dimension instead of hard-coding a per-model table.
25        let dimension = te
26            .embed(vec!["dimension probe".to_string()], None)
27            .map_err(|e| EmbedError::EmbedFailed(e.to_string()))?
28            .first()
29            .map(|v| v.len())
30            .ok_or_else(|| EmbedError::EmbedFailed("empty probe embedding".to_string()))?;
31
32        Ok(Self {
33            model: te,
34            model_id,
35            dimension,
36        })
37    }
38}
39
40#[async_trait]
41impl Embedder for FastEmbedder {
42    async fn embed(&self, texts: &[&str]) -> Result<Vec<Embedding>, EmbedError> {
43        let texts: Vec<String> = texts.iter().map(|s| s.to_string()).collect();
44        self.model
45            .embed(texts, None)
46            .map_err(|e| EmbedError::EmbedFailed(e.to_string()))
47    }
48
49    fn dimension(&self) -> usize {
50        self.dimension
51    }
52
53    fn model_id(&self) -> &str {
54        &self.model_id
55    }
56}