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,
cache_dir: std::path::PathBuf,
}
impl Embedder {
pub fn cache_dir_path(p: &std::path::Path) -> std::path::PathBuf {
p.to_path_buf()
}
pub fn reinit(&mut self) -> Result<()> {
let fresh = TextEmbedding::try_new(
InitOptions::new(DEFAULT_MODEL)
.with_cache_dir(self.cache_dir.clone())
.with_show_download_progress(false),
)
.context("Failed to reinit fastembed TextEmbedding model")?;
self.model = fresh;
Ok(())
}
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.clone())
.with_show_download_progress(true),
)
.context("Failed to initialise fastembed TextEmbedding model")?;
info!("Embedding model ready");
Ok(Self { model, cache_dir })
}
#[expect(
clippy::cognitive_complexity,
reason = "function is verbose but correct — extraction deferred"
)]
pub fn embed(
&mut self,
entities: Vec<ParsedEntity>,
batch_size: usize,
) -> Result<Vec<EmbeddedEntity>> {
if entities.is_empty() {
return Ok(vec![]);
}
let repo_name = entities[0].repo_name.clone();
let texts: Vec<&str> = entities.iter().map(|e| e.embed_text.as_str()).collect();
info!(
"[{repo_name}] 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!(
"[{repo_name}] 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);
}
}
pub fn needs_reset(batch_count: usize, interval: usize) -> bool {
interval > 0 && batch_count > 0 && batch_count.is_multiple_of(interval)
}
#[cfg(test)]
mod reset_tests {
use super::*;
#[test]
fn test_needs_reset_disabled_when_interval_zero() {
assert!(!needs_reset(500, 0));
assert!(!needs_reset(1000, 0));
}
#[test]
fn test_needs_reset_true_exactly_at_interval() {
assert!(needs_reset(500, 500));
assert!(needs_reset(1000, 500));
assert!(needs_reset(250, 250));
}
#[test]
fn test_needs_reset_false_before_interval() {
assert!(!needs_reset(499, 500));
assert!(!needs_reset(1, 500));
}
#[test]
fn test_needs_reset_false_between_intervals() {
assert!(!needs_reset(501, 500));
assert!(!needs_reset(999, 500));
}
#[test]
fn test_needs_reset_multiples_of_interval() {
for multiplier in 1..=10 {
assert!(needs_reset(500 * multiplier, 500));
}
}
#[test]
fn test_needs_reset_batch_count_zero_never_resets() {
assert!(!needs_reset(0, 500));
assert!(!needs_reset(0, 1));
}
}