use fastembed::{
EmbeddingModel, InitOptions, RerankInitOptions, RerankerModel, TextEmbedding, TextRerank,
};
use std::path::PathBuf;
pub fn cache_dir() -> PathBuf {
dirs::cache_dir()
.unwrap_or_else(|| std::path::PathBuf::from(".leankg-cache"))
.join("leankg")
.join("models")
}
pub const DEFAULT_EMBEDDING_MODEL: EmbeddingModel = EmbeddingModel::BGESmallENV15;
pub const DEFAULT_RERANKER_MODEL: RerankerModel = RerankerModel::BGERerankerV2M3;
pub const EMBEDDING_DIM: usize = 384;
pub struct Embedder {
inner: TextEmbedding,
}
impl Embedder {
pub fn new() -> Result<Self, Box<dyn std::error::Error>> {
Self::with_model(DEFAULT_EMBEDDING_MODEL)
}
pub fn with_model(model: EmbeddingModel) -> Result<Self, Box<dyn std::error::Error>> {
let opts = InitOptions::new(model)
.with_cache_dir(cache_dir())
.with_show_download_progress(true);
let inner = TextEmbedding::try_new(opts)?;
Ok(Self { inner })
}
pub fn embed(&self, texts: &[String]) -> Result<Vec<Vec<f32>>, Box<dyn std::error::Error>> {
let borrowed: Vec<&str> = texts.iter().map(|s| s.as_str()).collect();
let vectors = self.inner.embed(borrowed, None)?;
Ok(vectors)
}
pub fn dim(&self) -> usize {
EMBEDDING_DIM
}
}
pub struct Reranker {
inner: TextRerank,
}
impl Reranker {
pub fn new() -> Result<Self, Box<dyn std::error::Error>> {
Self::with_model(DEFAULT_RERANKER_MODEL)
}
pub fn with_model(model: RerankerModel) -> Result<Self, Box<dyn std::error::Error>> {
let opts = RerankInitOptions::new(model)
.with_cache_dir(cache_dir())
.with_show_download_progress(true);
let inner = TextRerank::try_new(opts)?;
Ok(Self { inner })
}
pub fn rerank(
&self,
query: &str,
documents: Vec<String>,
) -> Result<Vec<RerankScore>, Box<dyn std::error::Error>> {
let borrowed: Vec<&str> = documents.iter().map(|s| s.as_str()).collect();
let results = self.inner.rerank(query, borrowed, false, None)?;
Ok(results
.into_iter()
.map(|r| RerankScore {
document_idx: r.index,
score: r.score,
})
.collect())
}
}
#[derive(Debug, Clone)]
pub struct RerankScore {
pub document_idx: usize,
pub score: f32,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum RerankerStatus {
Active,
Fallback,
}
pub fn init_models() -> Result<InitReport, Box<dyn std::error::Error>> {
tracing::info!(
"initializing embedding + reranker models at {}",
cache_dir().display()
);
let _embedder = Embedder::new()?;
let _reranker = Reranker::new()?;
Ok(InitReport {
cache_dir: cache_dir(),
})
}
#[derive(Debug, Clone)]
pub struct InitReport {
pub cache_dir: PathBuf,
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn cache_dir_ends_with_leankg_models() {
let dir = cache_dir();
let components: Vec<_> = dir.components().collect();
let last_two: Vec<String> = components
.into_iter()
.rev()
.take(2)
.map(|c| c.as_os_str().to_string_lossy().to_string())
.collect();
assert_eq!(last_two, vec!["models".to_string(), "leankg".to_string()]);
}
#[test]
fn embedding_dim_matches_bge_small() {
assert_eq!(EMBEDDING_DIM, 384);
}
}