use crate::util::UnwrapPoison;
use crate::util::model_state::{AtomicModelState, ModelLoadGuard, ModelState};
use anyhow::{Context, Result, anyhow};
use candle_core::quantized::{QMatMul, gguf_file};
use candle_core::{DType, Device, Tensor};
use candle_nn::rotary_emb::rope_slow;
use candle_nn::{Embedding, Module};
use std::collections::HashMap;
use std::fs::File;
use std::path::Path;
use std::sync::RwLock;
use std::sync::atomic::Ordering;
use std::time::Duration;
use tokenizers::Tokenizer;
use tracing::{debug, info, warn};
const MAX_SEQ_LEN: usize = 8192;
const ROPE_FREQ_BASE: f32 = 1_000_000.0;
const MODEL_DOWNLOAD_TIMEOUT: Duration = Duration::from_mins(10);
const DEFAULT_PAD_ID: u32 = 128_001;
const MODEL_URL: &str = "https://huggingface.co/jinaai/jina-embeddings-v5-text-nano-retrieval/resolve/main/v5-nano-retrieval-Q4_K_M.gguf";
const MODEL_SHA256: &str = "f50822244ba0c7a348c5455b99bb8a0afd182511e8a816888c5dc65d972e51d5";
const TOKENIZER_URL: &str = "https://huggingface.co/jinaai/jina-embeddings-v5-text-nano-retrieval/resolve/main/tokenizer.json";
const TOKENIZER_SHA256: &str = "98d4a1d32152d6cedf85b5e88f3b205106dca1fe72aaab34e0ac13c238421069";
const MODEL_FILENAME: &str = "v5-nano-retrieval-Q4_K_M.gguf";
const TOKENIZER_FILENAME: &str = "embed_tokenizer.json";
static GLOBAL_EMBEDDER: RwLock<Option<Embedder>> = RwLock::new(None);
static STATE: AtomicModelState = AtomicModelState::new(ModelState::Uninit);
#[inline]
fn set_embedder_ready(emb: Embedder) {
*GLOBAL_EMBEDDER.write().unwrap_poison() = Some(emb);
STATE.store(ModelState::Ready, Ordering::Release);
}
fn ensure_embedder() -> bool {
match STATE.load(Ordering::Acquire) {
ModelState::Ready => return true,
ModelState::Uninit => {} _ => return false, }
if !STATE.transition(ModelState::Uninit, ModelState::Loading) {
return false;
}
let models_dir = crate::util::models_dir()
.expect("CONFIG must be initialized before embedder initialization");
let model_path = models_dir.join(MODEL_FILENAME);
let tokenizer_path = models_dir.join(TOKENIZER_FILENAME);
std::fs::create_dir_all(&models_dir).ok();
if model_path.exists()
&& tokenizer_path.exists()
&& let Some(emb) = try_load_embedder(&model_path, &tokenizer_path, "from cached files")
{
set_embedder_ready(emb);
return true;
}
if tokio::runtime::Handle::try_current().is_err() {
warn!("Embedder: no tokio runtime available — transitioning to FAILED");
STATE.store(ModelState::Failed, Ordering::Release);
return false;
}
tokio::spawn(download_retry_loop());
false
}
#[must_use]
pub fn embed_query(text: &str) -> Option<Vec<f32>> {
with_embedder(|emb| emb.embed_queries(&[text]).ok()?.into_iter().next())
}
#[must_use]
pub fn embed_document(text: &str) -> Option<Vec<f32>> {
with_embedder(|emb| emb.embed_documents(&[text]).ok()?.into_iter().next())
}
fn with_embedder<F>(f: F) -> Option<Vec<f32>>
where
F: FnOnce(&Embedder) -> Option<Vec<f32>>,
{
if !ensure_embedder() {
return None;
}
let guard = GLOBAL_EMBEDDER.read().unwrap_poison();
let emb = guard.as_ref()?;
f(emb)
}
async fn download_retry_loop() {
let _guard = ModelLoadGuard::new(&STATE);
let models_dir = crate::util::models_dir()
.expect("CONFIG must be initialized before embedder initialization");
std::fs::create_dir_all(&models_dir).ok();
let model_dest = models_dir.join(MODEL_FILENAME);
let tokenizer_dest = models_dir.join(TOKENIZER_FILENAME);
let Ok(client) = crate::util::http::build_download_client(MODEL_DOWNLOAD_TIMEOUT) else {
warn!("Embedder: failed to build HTTP client — background download cancelled");
STATE.store(ModelState::Failed, Ordering::Release);
return;
};
let mut delay = Duration::from_mins(1);
let max_delay = Duration::from_mins(30);
loop {
if model_dest.exists() && tokenizer_dest.exists() {
if let Some(emb) = try_load_embedder(&model_dest, &tokenizer_dest, "from cached files")
{
set_embedder_ready(emb);
return;
}
warn!("Deleting cached model files after load failure, forcing re-download");
let _ = std::fs::remove_file(&model_dest);
let _ = std::fs::remove_file(&tokenizer_dest);
}
let (model_result, tokenizer_result) = tokio::join!(
maybe_download(&client, MODEL_URL, &model_dest, Some(MODEL_SHA256)),
maybe_download(
&client,
TOKENIZER_URL,
&tokenizer_dest,
Some(TOKENIZER_SHA256)
),
);
let model_ok = model_result.is_ok();
let tokenizer_ok = tokenizer_result.is_ok();
if let Err(e) = &model_result {
warn!(error = %e, retry_after_secs = delay.as_secs(), "Failed to download embedding model, retrying");
let _ = std::fs::remove_file(model_dest.with_extension("tmp"));
}
if let Err(e) = &tokenizer_result {
warn!(error = %e, retry_after_secs = delay.as_secs(), "Failed to download tokenizer, retrying");
let _ = std::fs::remove_file(tokenizer_dest.with_extension("tmp"));
}
if model_ok && tokenizer_ok {
if let Some(emb) = try_load_embedder(&model_dest, &tokenizer_dest, "after download") {
set_embedder_ready(emb);
return;
}
warn!(
"Giving up on embedding model: freshly downloaded, SHA256-verified files \
failed to load. This indicates a code-level issue with Embedder::load(). \
The model will remain unavailable for this session (FTS-only fallback)."
);
STATE.store(ModelState::Failed, Ordering::Release);
return;
}
tokio::time::sleep(delay).await;
delay = (delay * 2).min(max_delay);
}
}
fn try_load_embedder(
model_path: &Path,
tokenizer_path: &Path,
context: &'static str,
) -> Option<Embedder> {
match Embedder::load(model_path, tokenizer_path) {
Ok(emb) => {
info!("Embedding model loaded successfully ({context})");
Some(emb)
}
Err(e) => {
warn!(reason = %e, context, "Failed to load embedding model");
None
}
}
}
async fn maybe_download(
client: &reqwest::Client,
url: &str,
dest: &Path,
expected_sha256: Option<&str>,
) -> Result<()> {
if dest.exists() {
return Ok(());
}
download_file(client, url, dest, expected_sha256).await
}
async fn download_file(
client: &reqwest::Client,
url: &str,
dest: &Path,
expected_sha256: Option<&str>,
) -> Result<()> {
let mut size = 0u64;
crate::util::http::download_verified(
client,
url,
dest,
expected_sha256.unwrap_or(""),
None,
crate::util::http::DownloadSizeCheck::Exact,
|downloaded, _| size = downloaded,
)
.await?;
info!(path = %dest.display(), size, "Downloaded model file");
Ok(())
}
fn get_meta<T>(
metadata: &HashMap<String, gguf_file::Value>,
key: &str,
extract: impl FnOnce(&gguf_file::Value) -> std::result::Result<T, candle_core::Error>,
) -> Result<T> {
let value = metadata
.get(key)
.ok_or_else(|| anyhow!("Missing metadata key '{key}'"))?;
extract(value).map_err(|e| anyhow!("Failed to read metadata '{key}': {e}"))
}
#[derive(Debug)]
struct Layer {
attn_q: QMatMul,
attn_k: QMatMul,
attn_v: QMatMul,
attn_o: QMatMul,
attn_norm: Tensor,
ffn_gate: QMatMul,
ffn_up: QMatMul,
ffn_down: QMatMul,
ffn_norm: Tensor,
}
impl Layer {
#[expect(clippy::many_single_char_names)]
fn forward(
&self,
x: &Tensor,
mask: &Tensor,
cos: &Tensor,
sin: &Tensor,
n_head: usize,
head_dim: usize,
) -> Result<Tensor> {
let residual = x;
let h = candle_nn::ops::rms_norm(x, &self.attn_norm, 1e-5)?;
let q = self.attn_q.forward(&h)?;
let k = self.attn_k.forward(&h)?;
let v = self.attn_v.forward(&h)?;
let h = Layer::apply_attention(&q, &k, &v, mask, cos, sin, n_head, head_dim)?;
let h = self.attn_o.forward(&h)?;
let h = (h + residual)?;
let residual = &h;
let h = candle_nn::ops::rms_norm(&h, &self.ffn_norm, 1e-5)?;
let gate = self.ffn_gate.forward(&h)?;
let up = self.ffn_up.forward(&h)?;
let h = self
.ffn_down
.forward(&(candle_nn::ops::silu(&gate)? * up)?)?;
let h = (h + residual)?;
Ok(h)
}
#[expect(clippy::too_many_arguments)]
fn apply_attention(
q: &Tensor,
k: &Tensor,
v: &Tensor,
mask: &Tensor,
cos: &Tensor,
sin: &Tensor,
n_head: usize,
head_dim: usize,
) -> Result<Tensor> {
let (b_sz, seq_len, n_embd) = q.shape().dims3()?;
let q = q
.reshape((b_sz, seq_len, n_head, head_dim))?
.transpose(1, 2)?;
let k = k
.reshape((b_sz, seq_len, n_head, head_dim))?
.transpose(1, 2)?;
let v = v
.reshape((b_sz, seq_len, n_head, head_dim))?
.transpose(1, 2)?
.contiguous()?;
let q = rope_slow(&q, cos, sin)?;
let k = rope_slow(&k, cos, sin)?;
#[expect(clippy::cast_precision_loss)]
let scale = 1.0_f64 / (head_dim as f64).sqrt();
let att = q.matmul(&k.t()?)?;
let att = (att * scale)?;
let mask = mask.broadcast_as(att.shape())?;
let att = (att + mask)?;
let att = candle_nn::ops::softmax_last_dim(&att)?;
let y = att.matmul(&v)?;
let y = y.transpose(1, 2)?.reshape((b_sz, seq_len, n_embd))?;
Ok(y)
}
}
fn load_qmatmul(
content: &gguf_file::Content,
file: &mut File,
prefix: &str,
name: &str,
device: &Device,
) -> Result<QMatMul> {
QMatMul::from_qtensor(
content
.tensor(file, &format!("{prefix}.{name}.weight"), device)
.with_context(|| format!("Failed to load {prefix}.{name}.weight"))?,
)
.with_context(|| format!("Failed to create QMatMul for {name}"))
}
fn load_norm(
content: &gguf_file::Content,
file: &mut File,
prefix: &str,
name: &str,
device: &Device,
) -> Result<Tensor> {
content
.tensor(file, &format!("{prefix}.{name}.weight"), device)
.with_context(|| format!("Failed to load {prefix}.{name}.weight"))?
.dequantize(device)
.with_context(|| format!("Failed to dequantize {name}"))
}
pub struct Embedder {
tokenizer: Tokenizer,
tok_embeddings: Embedding,
layers: Vec<Layer>,
output_norm: Tensor,
cos: Tensor,
sin: Tensor,
head_dim: usize,
n_head: usize,
pad_id: u32,
}
impl Embedder {
#[expect(clippy::too_many_lines)]
pub fn load(model_path: &Path, tokenizer_path: &Path) -> Result<Self> {
let tokenizer = Tokenizer::from_file(tokenizer_path).map_err(|e| {
anyhow!(
"Failed to load tokenizer from {}: {e}",
tokenizer_path.display()
)
})?;
let pad_id = tokenizer
.token_to_id("<|end_of_text|>")
.or_else(|| tokenizer.token_to_id("<|pad|>"))
.or_else(|| tokenizer.token_to_id("[PAD]"))
.unwrap_or(DEFAULT_PAD_ID);
debug!(pad_id, "Discovered pad token ID from tokenizer");
let device = Device::Cpu;
let mut file = std::fs::File::open(model_path)
.map_err(|e| anyhow!("Failed to open model file {}: {e}", model_path.display()))?;
let content = gguf_file::Content::read(&mut file)
.map_err(|e| anyhow!("Failed to read GGUF file: {e}"))?;
let hidden_size = get_meta(&content.metadata, "eurobert.embedding_length", |v| {
v.to_u32()
})? as usize;
let n_head = get_meta(&content.metadata, "eurobert.attention.head_count", |v| {
v.to_u32()
})? as usize;
let head_dim = get_meta(&content.metadata, "eurobert.attention.value_length", |v| {
v.to_u32()
})? as usize;
let rope_freq_base = get_meta(
&content.metadata,
"eurobert.rope.freq_base",
candle_core::quantized::gguf_file::Value::to_f32,
)
.unwrap_or(ROPE_FREQ_BASE);
let n_layers = content
.tensor_infos
.keys()
.filter_map(|name| {
let name = name.as_str();
if name.starts_with("blk.") && name.ends_with(".attn_q.weight") {
name.trim_start_matches("blk.")
.split('.')
.next()?
.parse::<usize>()
.ok()
} else {
None
}
})
.max()
.map(|max| max + 1)
.context("No layer tensors found in GGUF file")?;
info!(
hidden_size,
n_head, head_dim, n_layers, rope_freq_base, "Loading EuroBERT embedding model"
);
let tok_embd_qt = content
.tensor(&mut file, "token_embd.weight", &device)
.context("Failed to load token_embd.weight")?;
let tok_embd_f32 = tok_embd_qt
.dequantize(&device)
.context("Failed to dequantize token_embd.weight")?;
let tok_embeddings = Embedding::new(tok_embd_f32, hidden_size);
let output_norm_qt = content
.tensor(&mut file, "output_norm.weight", &device)
.context("Failed to load output_norm.weight")?;
let output_norm = output_norm_qt
.dequantize(&device)
.context("Failed to dequantize output_norm.weight")?;
let mut layers = Vec::with_capacity(n_layers);
for i in 0..n_layers {
let prefix = format!("blk.{i}");
let attn_q = load_qmatmul(&content, &mut file, &prefix, "attn_q", &device)?;
let attn_k = load_qmatmul(&content, &mut file, &prefix, "attn_k", &device)?;
let attn_v = load_qmatmul(&content, &mut file, &prefix, "attn_v", &device)?;
let attn_o = load_qmatmul(&content, &mut file, &prefix, "attn_output", &device)?;
let attn_norm = load_norm(&content, &mut file, &prefix, "attn_norm", &device)?;
let ffn_gate = load_qmatmul(&content, &mut file, &prefix, "ffn_gate", &device)?;
let ffn_up = load_qmatmul(&content, &mut file, &prefix, "ffn_up", &device)?;
let ffn_down = load_qmatmul(&content, &mut file, &prefix, "ffn_down", &device)?;
let ffn_norm = load_norm(&content, &mut file, &prefix, "ffn_norm", &device)?;
layers.push(Layer {
attn_q,
attn_k,
attn_v,
attn_o,
attn_norm,
ffn_gate,
ffn_up,
ffn_down,
ffn_norm,
});
}
let (cos, sin) = precompute_freqs_cis(head_dim, rope_freq_base, &device)?;
let emb = Self {
tokenizer,
tok_embeddings,
layers,
output_norm,
cos,
sin,
head_dim,
n_head,
pad_id,
};
let v = emb.embed_documents(&["."])?;
anyhow::ensure!(
!v.is_empty() && !v[0].is_empty(),
"Embedder warm-up produced empty output"
);
anyhow::ensure!(
v[0].len() == hidden_size,
"Embedder warm-up produced wrong dimension: expected {hidden_size}, got {}",
v[0].len()
);
info!("Embedder initialized successfully (hidden_size={hidden_size}, layers={n_layers})");
Ok(emb)
}
pub fn embed_queries(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>> {
self.embed_prefixed("Query: ", texts)
}
pub fn embed_documents(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>> {
self.embed_prefixed("Document: ", texts)
}
fn embed_prefixed(&self, prefix: &str, texts: &[&str]) -> Result<Vec<Vec<f32>>> {
let encodings: Vec<_> = texts
.iter()
.map(|t| {
let input = format!("{prefix}{t}");
self.tokenizer.encode(input, true)
})
.collect::<std::result::Result<Vec<_>, _>>()
.map_err(|e| anyhow!("Tokenization error: {e}"))?;
if encodings.is_empty() {
anyhow::bail!("Empty input");
}
let max_len = encodings
.iter()
.map(|e| e.len().min(MAX_SEQ_LEN))
.max()
.context("Empty encoding")?;
let batch_size = encodings.len();
let mut input_ids_vec = vec![i64::from(self.pad_id); batch_size * max_len];
let mut attention_mask_vec = vec![0i64; batch_size * max_len];
for (row, enc) in encodings.iter().enumerate() {
let ids = enc.get_ids();
let mask = enc.get_attention_mask();
let len = ids.len().min(MAX_SEQ_LEN);
for col in 0..len {
input_ids_vec[row * max_len + col] = i64::from(ids[col]);
attention_mask_vec[row * max_len + col] = i64::from(mask[col]);
}
}
let input_ids = Tensor::from_vec(input_ids_vec, (batch_size, max_len), &Device::Cpu)?;
let attention_mask =
Tensor::from_vec(attention_mask_vec, (batch_size, max_len), &Device::Cpu)?;
let embeddings = self.forward(&input_ids, &attention_mask)?;
let result = last_token_pool_and_l2_normalize(&embeddings, &attention_mask)?;
Ok(result)
}
fn forward(&self, input_ids: &Tensor, attention_mask: &Tensor) -> Result<Tensor> {
let (_batch_size, _seq_len) = input_ids.shape().dims2()?;
let mask = build_attn_mask(attention_mask, &Device::Cpu)?;
let mut h = self.tok_embeddings.forward(input_ids)?;
for layer in &self.layers {
h = layer.forward(&h, &mask, &self.cos, &self.sin, self.n_head, self.head_dim)?;
}
h = candle_nn::ops::rms_norm(&h, &self.output_norm, 1e-5)?;
Ok(h)
}
}
fn precompute_freqs_cis(
head_dim: usize,
freq_base: f32,
device: &Device,
) -> Result<(Tensor, Tensor)> {
#[expect(clippy::cast_precision_loss)]
let theta: Vec<f32> = (0..head_dim)
.step_by(2)
.map(|i| 1.0_f32 / freq_base.powf(i as f32 / head_dim as f32))
.collect();
let theta = Tensor::from_vec(theta, (head_dim / 2,), device)?;
#[expect(clippy::cast_possible_truncation)]
let positions = Tensor::arange(0u32, MAX_SEQ_LEN as u32, device)?
.to_dtype(DType::F32)?
.reshape((MAX_SEQ_LEN, 1))?;
let idx_theta = positions.matmul(&theta.reshape((1, theta.elem_count()))?)?;
let cos = idx_theta.cos()?;
let sin = idx_theta.sin()?;
Ok((cos, sin))
}
fn build_attn_mask(attention_mask: &Tensor, device: &Device) -> Result<Tensor> {
let (batch_size, seq_len) = attention_mask.shape().dims2()?;
let mask_f32 = attention_mask.to_dtype(DType::F32)?;
let mask_a = mask_f32.reshape((batch_size, 1, 1, seq_len))?;
let mask_b = mask_f32.reshape((batch_size, 1, seq_len, 1))?;
let pairwise = mask_a.broadcast_mul(&mask_b)?;
let large_neg = Tensor::new(-1e10_f32, device)?.broadcast_as(pairwise.shape())?;
let zero = Tensor::new(0.0_f32, device)?.broadcast_as(pairwise.shape())?;
let mask_cond = pairwise.eq(&zero)?;
let mask = mask_cond.where_cond(&large_neg, &zero)?;
Ok(mask)
}
fn last_token_pool_and_l2_normalize(
embeddings: &Tensor,
attention_mask: &Tensor,
) -> Result<Vec<Vec<f32>>> {
let (batch_size, _seq_len, hidden_size) = embeddings.shape().dims3()?;
let mut results = Vec::with_capacity(batch_size);
let seq_lengths: Vec<i64> = attention_mask.sum(1)?.to_vec1()?;
for (i, &seq_len) in seq_lengths.iter().enumerate().take(batch_size) {
#[expect(clippy::cast_possible_truncation, clippy::cast_sign_loss)]
let last_pos = (seq_len - 1).max(0) as usize;
let token_emb = embeddings.narrow(0, i, 1)?.narrow(1, last_pos, 1)?;
let token_emb = token_emb.reshape(hidden_size)?;
let norm = token_emb
.sqr()?
.sum_all()?
.sqrt()?
.to_scalar::<f32>()?
.max(1e-12);
let normalized = token_emb.broadcast_div(&Tensor::new(norm, token_emb.device())?)?;
let vec: Vec<f32> = normalized.to_vec1()?;
results.push(vec);
}
Ok(results)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::vector::cosine_similarity;
use std::sync::Mutex;
static TEST_GLOBAL_STATE_MUTEX: Mutex<()> = Mutex::new(());
fn init_test_config() -> std::path::PathBuf {
let root = crate::util::test::test_root().clone();
let _ = crate::config::CONFIG.try_set_storage_root(root.clone());
root
}
fn test_embedder_opt() -> Option<&'static Embedder> {
use std::sync::OnceLock;
static TEST_EMBEDDER: OnceLock<Option<Embedder>> = OnceLock::new();
TEST_EMBEDDER
.get_or_init(|| {
let mut candidates = Vec::new();
if let Some(root) = crate::config::CONFIG.try_storage_root() {
candidates.push(root.join("models"));
}
if let Some(home) = std::env::var("HOME").ok().filter(|h| !h.is_empty()) {
let real = std::path::PathBuf::from(&home)
.join(".mahbot")
.join("models");
if !candidates.contains(&real) {
candidates.push(real);
}
}
for models_dir in &candidates {
let model_path = models_dir.join(MODEL_FILENAME);
let tokenizer_path = models_dir.join(TOKENIZER_FILENAME);
if model_path.exists() && tokenizer_path.exists() {
match Embedder::load(&model_path, &tokenizer_path) {
Ok(emb) => return Some(emb),
Err(e) => {
eprintln!("WARNING: Failed to load test embedder: {e}");
return None;
}
}
}
}
let last_candidate = candidates.last().map(|p| p.display().to_string());
eprintln!(
"Embedder model files not found. Looked in: {}. \
Run the application first to download embedding models (~150 MB).",
last_candidate.as_deref().unwrap_or("<none>")
);
None
})
.as_ref()
}
fn reset_global_state() {
*GLOBAL_EMBEDDER.write().unwrap_poison() = None;
STATE.store(ModelState::Uninit, Ordering::Release);
}
#[tokio::test]
async fn test_embedder_graceful_degradation() {
let _lock = TEST_GLOBAL_STATE_MUTEX.lock().unwrap();
let _root = init_test_config();
reset_global_state();
let result = embed_document("test");
assert!(
result.is_none(),
"embed_document() should return None when model not available"
);
let guard = GLOBAL_EMBEDDER.read().unwrap_poison();
assert!(guard.is_none(), "global embedder should remain None");
}
#[test]
fn test_ensure_embedder_terminally_failed() {
let _lock = TEST_GLOBAL_STATE_MUTEX.lock().unwrap();
let _root = init_test_config();
reset_global_state();
STATE.store(ModelState::Failed, Ordering::Release);
let r1 = ensure_embedder();
assert!(
!r1,
"ensure_embedder should return false in ModelState::Failed"
);
assert_eq!(
STATE.load(Ordering::Acquire),
ModelState::Failed,
"STATE should remain FAILED after ensure_embedder()"
);
let r2 = ensure_embedder();
assert!(!r2, "ensure_embedder should still return false");
assert_eq!(
STATE.load(Ordering::Acquire),
ModelState::Failed,
"STATE should remain FAILED after second call"
);
reset_global_state();
}
#[test]
fn test_ensure_embedder_no_runtime() {
let _lock = TEST_GLOBAL_STATE_MUTEX.lock().unwrap();
let _root = init_test_config();
reset_global_state();
let result = ensure_embedder();
assert!(
!result,
"ensure_embedder should return false when no tokio runtime is available"
);
assert_eq!(
STATE.load(Ordering::Acquire),
ModelState::Failed,
"STATE should be FAILED after detecting no tokio runtime"
);
let r2 = ensure_embedder();
assert!(!r2, "ensure_embedder should still return false");
assert_eq!(
STATE.load(Ordering::Acquire),
ModelState::Failed,
"STATE should remain FAILED"
);
reset_global_state();
}
#[test]
fn test_load_corrupted_files_fails() {
let tmp = tempfile::TempDir::new().expect("failed to create temp dir");
let model_path = tmp.path().join(MODEL_FILENAME);
let tokenizer_path = tmp.path().join(TOKENIZER_FILENAME);
std::fs::write(&model_path, b"not a valid gguf file").unwrap();
std::fs::write(&tokenizer_path, b"not valid json at all").unwrap();
let result = Embedder::load(&model_path, &tokenizer_path);
assert!(
result.is_err(),
"Embedder::load() should fail on corrupted files"
);
let err = result.err().expect("just checked is_err");
let err_msg = format!("{err:#}");
let model_path_str = model_path.to_string_lossy();
let tokenizer_path_str = tokenizer_path.to_string_lossy();
assert!(
err_msg.contains(model_path_str.as_ref())
|| err_msg.contains(tokenizer_path_str.as_ref()),
"Error should mention one of the file paths, got: {err_msg}"
);
}
#[ignore = "loads the real cached 157 MB GGUF model (~3-4 s); runs only when explicitly invoked"]
#[test]
fn test_embedding_produces_unit_vectors() {
let Some(emb) = test_embedder_opt() else {
eprintln!("SKIP: embedder model files not cached — skipping model-backed test");
return;
};
let v = emb.embed_documents(&["hello world"]).unwrap();
assert_eq!(v.len(), 1);
assert_eq!(v[0].len(), 768);
let norm: f32 = v[0].iter().map(|x| x * x).sum::<f32>().sqrt();
assert!(
(norm - 1.0).abs() < 1e-5,
"expected unit vector, got norm={norm}"
);
let docs = &["first document", "second document about something"];
let v = emb.embed_documents(docs).unwrap();
assert_eq!(v.len(), 2);
for vec in &v {
assert_eq!(vec.len(), 768);
let norm: f32 = vec.iter().map(|x| x * x).sum::<f32>().sqrt();
assert!(
(norm - 1.0).abs() < 1e-5,
"expected unit vector, got norm={norm}"
);
}
let v = emb.embed_queries(&["what is rust?"]).unwrap();
assert_eq!(v.len(), 1);
assert_eq!(v[0].len(), 768);
let norm: f32 = v[0].iter().map(|x| x * x).sum::<f32>().sqrt();
assert!(
(norm - 1.0).abs() < 1e-5,
"expected unit vector, got norm={norm}"
);
}
#[ignore = "loads the real cached 157 MB GGUF model (~3-4 s); runs only when explicitly invoked"]
#[test]
fn test_similar_embeddings_are_similar() {
let Some(emb) = test_embedder_opt() else {
eprintln!("SKIP: embedder model files not cached — skipping model-backed test");
return;
};
let v = emb
.embed_documents(&[
"rust programming language",
"the rust programming language",
"python programming language",
])
.unwrap();
let sim_01 = cosine_similarity(&v[0], &v[1]);
let sim_02 = cosine_similarity(&v[0], &v[2]);
assert!(
sim_01 > sim_02,
"rust/rust ({sim_01}) should be more similar than rust/python ({sim_02})"
);
}
#[ignore = "loads the real cached 157 MB GGUF model (~3-4 s); runs only when explicitly invoked"]
#[test]
fn test_empty_input_fails() {
let Some(emb) = test_embedder_opt() else {
eprintln!("SKIP: embedder model files not cached — skipping model-backed test");
return;
};
let result = emb.embed_documents(&[]);
assert!(result.is_err(), "empty input should produce an error");
}
}