use anyhow::Result;
use tracing::info;
use crate::models::{EmbeddedEntity, ParsedEntity};
use anyhow::Context;
use fastembed::{EmbeddingModel, InitOptions, TextEmbedding};
const DEFAULT_MODEL: EmbeddingModel = EmbeddingModel::AllMiniLML6V2;
pub struct Embedder {
model: TextEmbedding,
}
impl Embedder {
pub fn init(cache_dir: std::path::PathBuf) -> Result<Self> {
info!(
"Initialising fastembed model ({DEFAULT_MODEL:?}) in {}…",
cache_dir.display()
);
std::fs::create_dir_all(&cache_dir).context("Failed to create fastembed cache dir")?;
let model = TextEmbedding::try_new(
InitOptions::new(DEFAULT_MODEL)
.with_cache_dir(cache_dir)
.with_show_download_progress(true),
)
.context("Failed to initialise fastembed TextEmbedding model")?;
info!("Embedding model ready");
Ok(Self { model })
}
pub fn embed(
&mut self,
entities: Vec<ParsedEntity>,
batch_size: usize,
) -> Result<Vec<EmbeddedEntity>> {
if entities.is_empty() {
return Ok(vec![]);
}
let texts: Vec<&str> = entities.iter().map(|e| e.embed_text.as_str()).collect();
info!(
"Embedding {} entities (batch_size={})…",
texts.len(),
batch_size
);
let vectors = self
.model
.embed(texts, Some(batch_size))
.context("fastembed embedding failed")?;
debug_assert_eq!(
vectors.len(),
entities.len(),
"Mismatch between entity count and vector count"
);
let embedded: Vec<EmbeddedEntity> = entities
.into_iter()
.zip(vectors)
.map(|(entity, vector)| EmbeddedEntity { entity, vector })
.collect();
info!("Embedding complete — {} vectors produced", embedded.len());
Ok(embedded)
}
pub fn embed_query(&mut self, query: &str) -> Result<Vec<f32>> {
let vectors = self
.model
.embed(vec![query], Some(1))
.context("fastembed query embedding failed")?;
vectors
.into_iter()
.next()
.ok_or_else(|| anyhow::anyhow!("No vector returned for query"))
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::models::{EntityKind, ParsedEntity};
#[ignore = "Downloads ONNX model (~23MB) and requires significant memory/CPU"]
#[test]
fn test_embedder_init_and_embed_basic() {
let temp_dir = tempfile::tempdir().unwrap();
let mut embedder =
Embedder::init(temp_dir.path().to_path_buf()).expect("Failed to init embedder");
let entity = ParsedEntity::new(
"TestClass",
EntityKind::Class,
"TestClass",
None,
None,
"java",
"Test.java",
1,
10,
None,
"test-repo",
);
let mut entities = vec![entity];
entities[0].embed_text = "[class] TestClass\nFile: Test.java:1".to_string();
let results = embedder.embed(entities, 1).expect("Failed to embed");
assert_eq!(results.len(), 1);
assert_eq!(results[0].vector.len(), 384); }
#[ignore = "Downloads ONNX model (~23MB) and requires significant memory/CPU"]
#[test]
fn test_embedder_embed_query() {
let temp_dir = tempfile::tempdir().unwrap();
let mut embedder =
Embedder::init(temp_dir.path().to_path_buf()).expect("Failed to init embedder");
let vector = embedder
.embed_query("How to implement a singleton in Java?")
.expect("Failed to embed query");
assert_eq!(vector.len(), 384);
}
}