use anyhow::Result;
use std::path::PathBuf;
use std::sync::{Arc, OnceLock};
pub const EMBED_DIM: usize = 256;
pub const FULL_EMBED_DIM: usize = 1024;
pub const DEFAULT_MODEL_REPO: &str = "Qwen/Qwen3-Embedding-0.6B-GGUF";
pub const DEFAULT_MODEL_FILE: &str = "Qwen3-Embedding-0.6B-Q8_0.gguf";
pub const DEFAULT_TOKENIZER_REPO: &str = "Qwen/Qwen3-Embedding-0.6B";
pub const DEFAULT_TOKENIZER_FILE: &str = "tokenizer.json";
static EMBEDDER_CACHE: OnceLock<std::sync::Mutex<Option<Arc<dyn Embedder>>>> = OnceLock::new();
pub fn get_cached_embedder(
model_path: &PathBuf,
tokenizer_path: &PathBuf,
) -> Option<Arc<dyn Embedder>> {
let cache = EMBEDDER_CACHE.get_or_init(|| std::sync::Mutex::new(None));
let mut guard = cache.lock().ok()?;
if let Some(ref embedder) = *guard {
return Some(embedder.clone());
}
let config = EmbeddingConfig {
model_path: model_path.clone(),
tokenizer_path: tokenizer_path.clone(),
..Default::default()
};
let embedder = create_embedder(config);
let embedder_arc: Arc<dyn Embedder> = Arc::from(embedder);
match embedder_arc.embed("test") {
Ok(_) => {
*guard = Some(embedder_arc.clone());
Some(embedder_arc)
}
Err(_) => None,
}
}
pub trait Embedder: Send + Sync {
fn embed(&self, text: &str) -> Result<Vec<f32>>;
}
pub struct NoEmbedder;
impl Embedder for NoEmbedder {
fn embed(&self, _text: &str) -> Result<Vec<f32>> {
anyhow::bail!("embeddings feature is not enabled")
}
}
#[derive(Debug, Clone)]
pub struct EmbeddingConfig {
pub model_path: PathBuf,
pub tokenizer_path: PathBuf,
pub normalize: bool,
}
impl Default for EmbeddingConfig {
fn default() -> Self {
Self {
model_path: PathBuf::new(),
tokenizer_path: PathBuf::new(),
normalize: true,
}
}
}
#[cfg(feature = "embeddings")]
mod candle_embedder {
use super::*;
use anyhow::{Context, Result as AnyResult};
use candle_core::quantized::gguf_file;
use candle_core::{DType, Device, Tensor};
use candle_transformers::models::quantized_qwen2::ModelWeights;
use std::fs::File;
use std::io::BufReader;
use tokenizers::Tokenizer;
pub struct CandleEmbedder {
model: std::sync::Mutex<ModelWeights>,
tokenizer: Tokenizer,
device: Device,
config: EmbeddingConfig,
}
impl CandleEmbedder {
pub fn load(config: EmbeddingConfig) -> AnyResult<Self> {
let device = Device::Cpu;
let file = File::open(&config.model_path)
.with_context(|| format!("Failed to open GGUF file: {:?}", config.model_path))?;
let mut reader = BufReader::new(file);
let ct =
gguf_file::Content::read(&mut reader).context("Failed to read GGUF content")?;
let model = ModelWeights::from_gguf(ct, &mut reader, &device)
.context("Failed to build quantized Qwen2 model from GGUF")?;
let tokenizer = Tokenizer::from_file(&config.tokenizer_path).map_err(|e| {
anyhow::anyhow!(
"Failed to load tokenizer from {:?}: {}",
config.tokenizer_path,
e
)
})?;
Ok(Self {
model: std::sync::Mutex::new(model),
tokenizer,
device,
config,
})
}
fn mean_pool(
&self,
token_embeddings: &Tensor,
attention_mask: &Tensor,
) -> AnyResult<Tensor> {
let mask = attention_mask.to_dtype(DType::F32)?.unsqueeze(2)?;
let masked = token_embeddings.broadcast_mul(&mask)?;
let sum = masked.sum(1)?;
let mask_sum = mask.sum(1)?; let pooled = sum.broadcast_div(&mask_sum)?;
if self.config.normalize {
let norm = pooled.sqr()?.sum(1)?.sqrt()?;
let pooled = pooled.broadcast_div(&norm.unsqueeze(1)?)?;
Ok(pooled)
} else {
Ok(pooled)
}
}
}
impl Embedder for CandleEmbedder {
fn embed(&self, text: &str) -> AnyResult<Vec<f32>> {
let encoding = self
.tokenizer
.encode(text, true)
.map_err(|e| anyhow::anyhow!("Tokenization failed: {}", e))?;
let input_ids = encoding.get_ids();
let attention_mask = encoding.get_attention_mask();
let input_ids_tensor = Tensor::from_slice(
input_ids
.iter()
.map(|&v| v as u32)
.collect::<Vec<_>>()
.as_slice(),
(1, input_ids.len()),
&self.device,
)?;
let attention_mask_tensor = Tensor::from_slice(
attention_mask
.iter()
.map(|&v| v as u32)
.collect::<Vec<_>>()
.as_slice(),
(1, attention_mask.len()),
&self.device,
)?;
let mut model = self
.model
.lock()
.map_err(|e| anyhow::anyhow!("model lock poisoned: {}", e))?;
let embedded = model.forward(&input_ids_tensor, 0)?;
let pooled = self.mean_pool(&embedded, &attention_mask_tensor)?;
let full_embedding = pooled
.to_vec2::<f32>()?
.into_iter()
.next()
.unwrap_or_default();
let truncated: Vec<f32> = full_embedding.into_iter().take(EMBED_DIM).collect();
if self.config.normalize && !truncated.is_empty() {
let norm: f32 = truncated.iter().map(|v| v * v).sum::<f32>().sqrt();
if norm > 0.0 {
return Ok(truncated.iter().map(|v| v / norm).collect());
}
}
Ok(truncated)
}
}
}
#[cfg(feature = "embeddings")]
pub use candle_embedder::CandleEmbedder;
#[cfg(feature = "embeddings")]
pub fn create_embedder(config: EmbeddingConfig) -> Box<dyn Embedder> {
if config.model_path.exists() && config.tokenizer_path.exists() {
match CandleEmbedder::load(config) {
Ok(embedder) => {
tracing::info!("Local embedding model loaded successfully");
return Box::new(embedder);
}
Err(e) => {
tracing::warn!(
"Failed to load embedding model: {}, falling back to text search",
e
);
}
}
} else {
if !config.model_path.exists() {
tracing::info!(
"Embedding model not found at {:?}. Download with: huggingface-cli download {} {}",
config.model_path,
DEFAULT_MODEL_REPO,
DEFAULT_MODEL_FILE
);
}
if !config.tokenizer_path.exists() {
tracing::info!(
"Tokenizer not found at {:?}. Download with: huggingface-cli download {} {}",
config.tokenizer_path,
DEFAULT_TOKENIZER_REPO,
DEFAULT_TOKENIZER_FILE
);
}
}
Box::new(NoEmbedder)
}
#[cfg(not(feature = "embeddings"))]
pub fn create_embedder(_config: EmbeddingConfig) -> Box<dyn Embedder> {
Box::new(NoEmbedder)
}
#[cfg(feature = "embeddings")]
pub fn embeddings_available() -> bool {
true
}
#[cfg(not(feature = "embeddings"))]
pub fn embeddings_available() -> bool {
false
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_no_embedder_returns_error() {
let embedder = NoEmbedder;
assert!(embedder.embed("test").is_err());
}
#[test]
fn test_embed_dim() {
assert_eq!(EMBED_DIM, 256);
assert_eq!(FULL_EMBED_DIM, 1024);
}
#[test]
fn test_create_embedder_without_model_file() {
let config = EmbeddingConfig::default();
let embedder = create_embedder(config);
assert!(embedder.embed("test").is_err());
}
}