use kimetsu_core::KimetsuResult;
pub const DEFAULT_HYBRID_ALPHA: f32 = 0.5;
pub trait Embedder: Send + Sync {
fn embed(&self, text: &str) -> Result<Vec<f32>, EmbedderError>;
fn model_id(&self) -> &str;
fn dim(&self) -> usize;
fn is_noop(&self) -> bool {
false
}
fn embed_batch(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>, EmbedderError> {
texts.iter().map(|t| self.embed(t)).collect()
}
}
impl Embedder for Box<dyn Embedder> {
fn embed(&self, text: &str) -> Result<Vec<f32>, EmbedderError> {
(**self).embed(text)
}
fn model_id(&self) -> &str {
(**self).model_id()
}
fn dim(&self) -> usize {
(**self).dim()
}
fn is_noop(&self) -> bool {
(**self).is_noop()
}
fn embed_batch(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>, EmbedderError> {
(**self).embed_batch(texts)
}
}
#[derive(Debug, Clone)]
pub enum EmbedderError {
NotImplemented,
LoadFailed(String),
EmbedFailed(String),
DimMismatch { expected: usize, got: usize },
}
impl std::fmt::Display for EmbedderError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::NotImplemented => write!(f, "embedder not implemented"),
Self::LoadFailed(msg) => write!(f, "embedder load failed: {msg}"),
Self::EmbedFailed(msg) => write!(f, "embed call failed: {msg}"),
Self::DimMismatch { expected, got } => {
write!(f, "embedding dim mismatch: expected {expected}, got {got}")
}
}
}
}
impl std::error::Error for EmbedderError {}
#[derive(Debug, Default, Clone, Copy)]
pub struct NoopEmbedder;
impl NoopEmbedder {
pub const MODEL_ID: &'static str = "noop";
}
impl Embedder for NoopEmbedder {
fn embed(&self, _text: &str) -> Result<Vec<f32>, EmbedderError> {
Err(EmbedderError::NotImplemented)
}
fn model_id(&self) -> &str {
Self::MODEL_ID
}
fn dim(&self) -> usize {
0
}
fn is_noop(&self) -> bool {
true
}
}
#[derive(Debug, Clone, Copy)]
pub struct StubEmbedder {
dim: usize,
}
impl StubEmbedder {
pub const MODEL_ID: &'static str = "stub-d8";
pub const fn new() -> Self {
Self { dim: 8 }
}
pub const fn with_dim(dim: usize) -> Self {
Self { dim }
}
}
impl Default for StubEmbedder {
fn default() -> Self {
Self::new()
}
}
impl Embedder for StubEmbedder {
fn embed(&self, text: &str) -> Result<Vec<f32>, EmbedderError> {
let mut bucket = vec![0.0f32; self.dim];
for word in text.split_whitespace() {
let normalized = word.to_lowercase();
let mut h: u64 = 0xcbf2_9ce4_8422_2325;
for byte in normalized.bytes() {
h ^= byte as u64;
h = h.wrapping_mul(0x0000_0100_0000_01B3);
}
let idx = (h as usize) % self.dim.max(1);
bucket[idx] += 1.0;
}
let norm = bucket.iter().map(|v| v * v).sum::<f32>().sqrt();
if norm > 0.0 {
for v in &mut bucket {
*v /= norm;
}
}
Ok(bucket)
}
fn model_id(&self) -> &str {
Self::MODEL_ID
}
fn dim(&self) -> usize {
self.dim
}
}
pub trait Reranker: Send + Sync {
fn rerank(&self, query: &str, documents: &[&str]) -> Result<Vec<f32>, EmbedderError>;
fn model_id(&self) -> &str;
}
pub struct StubReranker;
impl Reranker for StubReranker {
fn rerank(&self, query: &str, documents: &[&str]) -> Result<Vec<f32>, EmbedderError> {
let query_tokens: std::collections::HashSet<String> = query
.split(|c: char| !c.is_alphanumeric())
.filter(|t| !t.is_empty())
.map(|t| t.to_lowercase())
.collect();
let q_len = query_tokens.len();
let scores = documents
.iter()
.map(|doc| {
if q_len == 0 {
return 0.05_f32;
}
let doc_tokens: std::collections::HashSet<String> = doc
.split(|c: char| !c.is_alphanumeric())
.filter(|t| !t.is_empty())
.map(|t| t.to_lowercase())
.collect();
let intersection = query_tokens.intersection(&doc_tokens).count();
let overlap = intersection as f32 / q_len as f32;
(0.05 + 0.9 * overlap).clamp(0.0, 1.0)
})
.collect();
Ok(scores)
}
fn model_id(&self) -> &str {
"stub-reranker"
}
}
pub fn open_reranker_for_model(model_id: &str) -> Option<Box<dyn Reranker>> {
let v = model_id.trim().to_ascii_lowercase();
if v.is_empty() || matches!(v.as_str(), "off" | "none" | "noop") {
return None;
}
#[cfg(feature = "embeddings")]
{
const CURATED: &[&str] = &[
"jina-reranker-v1-turbo-en",
"bge-reranker-base",
"bge-reranker-v2-m3",
"jina-reranker-v2-base-multilingual",
];
const USER_DEFINED_ALIASES: &[&str] = &[
"jina-reranker-v1-tiny-en",
"ms-marco-tinybert-l-2-v2",
"ms-marco-minilm-l-4-v2",
];
if CURATED.contains(&v.as_str()) {
return fastembed_backend::FastembedReranker::try_open(model_id)
.ok()
.map(|r| Box::new(r) as Box<dyn Reranker>);
}
if USER_DEFINED_ALIASES.contains(&v.as_str()) || v.contains('/') {
return fastembed_backend::FastembedReranker::try_open_user_defined(model_id)
.ok()
.map(|r| Box::new(r) as Box<dyn Reranker>);
}
fastembed_backend::FastembedReranker::try_open("jina-reranker-v1-turbo-en")
.ok()
.map(|r| Box::new(r) as Box<dyn Reranker>)
}
#[cfg(not(feature = "embeddings"))]
{
let _ = v;
None
}
}
pub fn open_default_embedder() -> &'static (dyn Embedder + Send + Sync) {
static CACHE: std::sync::OnceLock<Box<dyn Embedder + Send + Sync>> = std::sync::OnceLock::new();
let embedder = CACHE.get_or_init(build_default_embedder);
embedder.as_ref()
}
fn build_default_embedder() -> Box<dyn Embedder + Send + Sync> {
if env_disables_embedder() {
return Box::new(NoopEmbedder);
}
#[cfg(feature = "embeddings")]
{
match fastembed_backend::open_cached() {
Ok(handle) => return Box::new(handle),
Err(err) => {
eprintln!(
"kimetsu-brain: fastembed init failed ({err}); falling back to NoopEmbedder. \
Retrieval will stay FTS-only this session. Re-run with \
KIMETSU_BRAIN_EMBEDDER=noop to silence this warning."
);
}
}
}
Box::new(NoopEmbedder)
}
pub fn open_embedder_for(config_enabled: bool) -> &'static dyn Embedder {
if embedder_enabled_for_config(config_enabled) {
open_default_embedder()
} else {
&NoopEmbedder
}
}
pub fn open_embedder_for_model(model_id: &str) -> Box<dyn Embedder + Send + Sync> {
#[cfg(feature = "embeddings")]
{
match fastembed_backend::FastembedEmbedder::try_open(model_id) {
Ok(engine) => return Box::new(engine),
Err(err) => {
eprintln!(
"kimetsu-brain: failed to open embedder `{model_id}` ({err}); \
using NoopEmbedder (no vectors produced)."
);
}
}
}
#[cfg(not(feature = "embeddings"))]
{
let _ = model_id;
}
Box::new(NoopEmbedder)
}
fn env_disables_embedder() -> bool {
match std::env::var("KIMETSU_BRAIN_EMBEDDER") {
Ok(value) => is_disable_value(&value.trim().to_ascii_lowercase()),
Err(_) => false,
}
}
pub fn embedder_enabled_for_config(config_enabled: bool) -> bool {
match std::env::var("KIMETSU_BRAIN_EMBEDDER") {
Ok(raw) => {
let v = raw.trim().to_ascii_lowercase();
if v.is_empty() {
config_enabled
} else if is_disable_value(&v) {
false
} else {
true
}
}
Err(_) => config_enabled,
}
}
fn is_disable_value(v: &str) -> bool {
matches!(v, "noop" | "off" | "none" | "0" | "false" | "no")
}
pub const BUILTIN_MODELS: &[(&str, usize, &str)] = &[
("bge-small-en-v1.5", 384, "English, default, ~67 MB int8"),
("bge-m3", 1024, "Multilingual, ~600 MB int8"),
(
"jina-v2-base-code",
768,
"English + code-tuned, ~165 MB int8",
),
];
static EMBEDDER_OVERRIDE: std::sync::OnceLock<String> = std::sync::OnceLock::new();
pub fn apply_embedder_selection(config_embedder: Option<&str>) {
if let Some(id) = config_embedder {
let id = id.trim();
if !id.is_empty() {
let _ = EMBEDDER_OVERRIDE.set(id.to_string());
}
}
}
fn map_builtin_id(v: &str) -> &'static str {
match v {
"" | "default" | "bge-small" | "bge-small-en-v1.5" => "bge-small-en-v1.5",
"bge-m3" | "m3" => "bge-m3",
"jina-code" | "jina-v2-base-code" | "jina-embeddings-v2-base-code" => "jina-v2-base-code",
"noop" | "off" | "none" | "0" | "false" | "no" => "bge-small-en-v1.5",
other => {
eprintln!(
"kimetsu-brain: unknown embedder {other:?}, \
falling back to bge-small-en-v1.5"
);
"bge-small-en-v1.5"
}
}
}
pub fn resolve_embedder_id(config_embedder: Option<&str>) -> &'static str {
if let Ok(raw) = std::env::var("KIMETSU_BRAIN_EMBEDDER") {
let v = raw.trim().to_ascii_lowercase();
if !v.is_empty() && !is_disable_value(&v) {
return map_builtin_id(&v);
}
}
let cfg = config_embedder
.map(str::to_string)
.or_else(|| EMBEDDER_OVERRIDE.get().cloned());
if let Some(c) = cfg {
let v = c.trim().to_ascii_lowercase();
if !v.is_empty() {
return map_builtin_id(&v);
}
}
"bge-small-en-v1.5"
}
pub fn pick_builtin_model_from_env() -> &'static str {
resolve_embedder_id(None)
}
#[cfg(feature = "embeddings")]
mod fastembed_backend {
use super::{Embedder, EmbedderError, Reranker, pick_builtin_model_from_env};
use fastembed::{
EmbeddingModel, InitOptions, RerankInitOptions, RerankerModel, TextEmbedding, TextRerank,
};
use std::sync::{Arc, Mutex, OnceLock};
fn hf_repo_for_alias(lowercased: &str) -> Option<&'static str> {
match lowercased {
"jina-reranker-v1-tiny-en" => Some("jinaai/jina-reranker-v1-tiny-en"),
"ms-marco-tinybert-l-2-v2" => Some("Xenova/ms-marco-TinyBERT-L-2-v2"),
"ms-marco-minilm-l-4-v2" => Some("Xenova/ms-marco-MiniLM-L-4-v2"),
_ => None,
}
}
fn download_user_defined_reranker(
model_id: &str,
) -> Result<(fastembed::OnnxSource, fastembed::TokenizerFiles), EmbedderError> {
use hf_hub::api::sync::Api;
let lowercased = model_id.trim().to_ascii_lowercase();
let repo_id: String = if let Some(alias) = hf_repo_for_alias(&lowercased) {
alias.to_string()
} else if lowercased.contains('/') {
model_id.to_string()
} else {
return Err(EmbedderError::LoadFailed(format!(
"user-defined reranker: no HF repo mapping for {model_id:?}"
)));
};
let api = Api::new()
.map_err(|e| EmbedderError::LoadFailed(format!("hf-hub Api::new failed: {e}")))?;
let repo = api.model(repo_id.clone());
let get_required = |filename: &str| -> Result<Vec<u8>, EmbedderError> {
let path = repo.get(filename).map_err(|e| {
EmbedderError::LoadFailed(format!("{repo_id}/{filename}: download failed: {e}"))
})?;
std::fs::read(&path).map_err(|e| {
EmbedderError::LoadFailed(format!("{repo_id}/{filename}: read failed: {e}"))
})
};
let tokenizer_file = get_required("tokenizer.json")?;
let config_file = get_required("config.json")?;
let tokenizer_config_file = get_required("tokenizer_config.json")?;
let special_tokens_map_file = get_required("special_tokens_map.json")?;
let onnx_path = repo
.get("onnx/model.onnx")
.or_else(|_| repo.get("model.onnx"))
.map_err(|e| {
EmbedderError::LoadFailed(format!(
"{repo_id}: could not find onnx/model.onnx or model.onnx: {e}"
))
})?;
let tokenizer_files = fastembed::TokenizerFiles {
tokenizer_file,
config_file,
special_tokens_map_file,
tokenizer_config_file,
};
Ok((fastembed::OnnxSource::File(onnx_path), tokenizer_files))
}
pub struct FastembedEmbedder {
model_id: &'static str,
dim: usize,
engine: Mutex<TextEmbedding>,
}
impl FastembedEmbedder {
pub fn try_open(builtin_id: &str) -> Result<Self, EmbedderError> {
let (kind, model_id, dim) = match builtin_id {
"bge-m3" => (EmbeddingModel::BGEM3, "bge-m3", 1024),
"jina-v2-base-code" => (
EmbeddingModel::JinaEmbeddingsV2BaseCode,
"jina-v2-base-code",
768,
),
_ => (EmbeddingModel::BGESmallENV15, "bge-small-en-v1.5", 384),
};
let opts = InitOptions::new(kind).with_show_download_progress(false);
let engine = TextEmbedding::try_new(opts)
.map_err(|e| EmbedderError::LoadFailed(format!("fastembed init: {e}")))?;
Ok(Self {
model_id,
dim,
engine: Mutex::new(engine),
})
}
}
impl Embedder for FastembedEmbedder {
fn embed(&self, text: &str) -> Result<Vec<f32>, EmbedderError> {
let mut guard = self
.engine
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
let mut out = guard
.embed(vec![text], None)
.map_err(|e| EmbedderError::EmbedFailed(format!("fastembed embed: {e}")))?;
let vec = out
.pop()
.ok_or_else(|| EmbedderError::EmbedFailed("empty result".into()))?;
if vec.len() != self.dim {
return Err(EmbedderError::DimMismatch {
expected: self.dim,
got: vec.len(),
});
}
Ok(vec)
}
fn model_id(&self) -> &str {
self.model_id
}
fn dim(&self) -> usize {
self.dim
}
fn embed_batch(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>, EmbedderError> {
if texts.is_empty() {
return Ok(Vec::new());
}
let mut guard = self
.engine
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
let out = guard
.embed(texts, None)
.map_err(|e| EmbedderError::EmbedFailed(format!("fastembed embed_batch: {e}")))?;
if out.len() != texts.len() {
return Err(EmbedderError::EmbedFailed(format!(
"fastembed returned {} vectors for {} texts",
out.len(),
texts.len()
)));
}
for v in &out {
if v.len() != self.dim {
return Err(EmbedderError::DimMismatch {
expected: self.dim,
got: v.len(),
});
}
}
Ok(out)
}
}
#[derive(Clone)]
pub struct EmbedderHandle(Arc<FastembedEmbedder>);
impl Embedder for EmbedderHandle {
fn embed(&self, text: &str) -> Result<Vec<f32>, EmbedderError> {
self.0.embed(text)
}
fn model_id(&self) -> &str {
self.0.model_id()
}
fn dim(&self) -> usize {
self.0.dim()
}
fn embed_batch(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>, EmbedderError> {
self.0.embed_batch(texts)
}
}
pub struct FastembedReranker {
model_id: String,
engine: Mutex<TextRerank>,
}
impl FastembedReranker {
pub fn try_open(builtin_id: &str) -> Result<Self, EmbedderError> {
let (kind, stable_id) = match builtin_id {
"bge-reranker-base" => (RerankerModel::BGERerankerBase, "bge-reranker-base"),
"bge-reranker-v2-m3" => (RerankerModel::BGERerankerV2M3, "bge-reranker-v2-m3"),
"jina-reranker-v2-base-multilingual" => (
RerankerModel::JINARerankerV2BaseMultiligual,
"jina-reranker-v2-base-multilingual",
),
_ => (
RerankerModel::JINARerankerV1TurboEn,
"jina-reranker-v1-turbo-en",
),
};
let opts = RerankInitOptions::new(kind).with_show_download_progress(false);
let engine = TextRerank::try_new(opts)
.map_err(|e| EmbedderError::LoadFailed(format!("fastembed reranker init: {e}")))?;
Ok(Self {
model_id: stable_id.to_string(),
engine: Mutex::new(engine),
})
}
pub fn try_open_user_defined(alias_or_repo: &str) -> Result<Self, EmbedderError> {
use fastembed::{RerankInitOptionsUserDefined, UserDefinedRerankingModel};
let (onnx_source, tokenizer_files) = download_user_defined_reranker(alias_or_repo)?;
let model = UserDefinedRerankingModel::new(onnx_source, tokenizer_files);
let opts = RerankInitOptionsUserDefined::default();
let engine = TextRerank::try_new_from_user_defined(model, opts).map_err(|e| {
EmbedderError::LoadFailed(format!(
"user-defined reranker {alias_or_repo:?} init: {e}"
))
})?;
let model_id = alias_or_repo.trim().to_ascii_lowercase();
Ok(Self {
model_id,
engine: Mutex::new(engine),
})
}
}
impl Reranker for FastembedReranker {
fn rerank(&self, query: &str, documents: &[&str]) -> Result<Vec<f32>, EmbedderError> {
if documents.is_empty() {
return Ok(Vec::new());
}
let mut guard = self
.engine
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
let raw_results = guard
.rerank(query, documents, false, None)
.map_err(|e| EmbedderError::EmbedFailed(format!("fastembed rerank: {e}")))?;
let n = documents.len();
let mut scores = vec![0.0f32; n];
for result in raw_results {
if result.index < n {
scores[result.index] = 1.0 / (1.0 + (-result.score).exp());
}
}
Ok(scores)
}
fn model_id(&self) -> &str {
&self.model_id
}
}
pub fn open_cached() -> Result<EmbedderHandle, EmbedderError> {
static CELL: OnceLock<Result<Arc<FastembedEmbedder>, EmbedderError>> = OnceLock::new();
let init = CELL.get_or_init(|| {
let builtin = pick_builtin_model_from_env();
FastembedEmbedder::try_open(builtin).map(Arc::new)
});
match init {
Ok(arc) => Ok(EmbedderHandle(arc.clone())),
Err(err) => Err(err.clone()),
}
}
}
pub fn cosine_similarity(a: &[f32], b: &[f32]) -> f32 {
if a.is_empty() || b.is_empty() || a.len() != b.len() {
return 0.0;
}
let mut dot = 0.0f32;
let mut na = 0.0f32;
let mut nb = 0.0f32;
for (x, y) in a.iter().zip(b.iter()) {
dot += x * y;
na += x * x;
nb += y * y;
}
if na == 0.0 || nb == 0.0 {
return 0.0;
}
dot / (na.sqrt() * nb.sqrt())
}
pub fn embed_and_persist(
conn: &rusqlite::Connection,
memory_id: &str,
text: &str,
embedder: &dyn Embedder,
) -> KimetsuResult<Option<Vec<f32>>> {
if embedder.is_noop() {
return Ok(None);
}
let vec = match embedder.embed(text) {
Ok(v) => v,
Err(EmbedderError::NotImplemented) => return Ok(None),
Err(e) => return Err(format!("embed failed for memory {memory_id}: {e}").into()),
};
if vec.len() != embedder.dim() {
return Err(format!(
"embedder {} produced {} dims, expected {}",
embedder.model_id(),
vec.len(),
embedder.dim()
)
.into());
}
let blob = encode_embedding(&vec);
conn.execute(
"UPDATE memories SET embedding = ?1, embedding_model = ?2 WHERE memory_id = ?3",
rusqlite::params![blob, embedder.model_id(), memory_id],
)?;
#[cfg(feature = "embeddings")]
if let Some(handle) = crate::ann::cached_handle(conn) {
let rowid: Option<i64> = conn
.query_row(
"SELECT rowid FROM memories WHERE memory_id = ?1",
rusqlite::params![memory_id],
|r| r.get(0),
)
.ok();
if let Some(rowid) = rowid {
let mut guard = handle.write().unwrap_or_else(|p| p.into_inner());
if let Err(e) = guard.add(rowid, &vec) {
eprintln!(
"kimetsu-brain: ann add failed for memory {memory_id}: {e} (index will reconcile on next open)"
);
}
}
}
Ok(Some(vec))
}
pub fn encode_embedding(vec: &[f32]) -> Vec<u8> {
let mut out = Vec::with_capacity(vec.len() * 4);
for v in vec {
out.extend_from_slice(&v.to_le_bytes());
}
out
}
pub fn decode_embedding(bytes: &[u8], expected_dim: Option<usize>) -> KimetsuResult<Vec<f32>> {
if bytes.len() % 4 != 0 {
return Err(format!("embedding blob length {} not a multiple of 4", bytes.len()).into());
}
let dim = bytes.len() / 4;
if let Some(expected) = expected_dim
&& dim != expected
{
return Err(format!("embedding blob dim {dim} does not match expected {expected}").into());
}
let mut out = Vec::with_capacity(dim);
for chunk in bytes.chunks_exact(4) {
let mut buf = [0u8; 4];
buf.copy_from_slice(chunk);
out.push(f32::from_le_bytes(buf));
}
Ok(out)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn map_builtin_id_maps_aliases_and_defaults_unknown() {
assert_eq!(map_builtin_id("bge-small-en-v1.5"), "bge-small-en-v1.5");
assert_eq!(map_builtin_id("default"), "bge-small-en-v1.5");
assert_eq!(map_builtin_id("m3"), "bge-m3");
assert_eq!(map_builtin_id("bge-m3"), "bge-m3");
assert_eq!(map_builtin_id("jina-code"), "jina-v2-base-code");
assert_eq!(
map_builtin_id("jina-embeddings-v2-base-code"),
"jina-v2-base-code"
);
assert_eq!(map_builtin_id("noop"), "bge-small-en-v1.5");
assert_eq!(map_builtin_id("totally-made-up"), "bge-small-en-v1.5");
}
#[test]
fn builtin_models_table_is_consistent() {
for (id, _dim, _blurb) in BUILTIN_MODELS {
assert_eq!(map_builtin_id(id), *id, "id {id} must be stable");
}
}
#[test]
fn resolve_embedder_id_uses_config_when_env_unset() {
if std::env::var_os("KIMETSU_BRAIN_EMBEDDER").is_some() {
return;
}
assert_eq!(resolve_embedder_id(Some("bge-m3")), "bge-m3");
assert_eq!(resolve_embedder_id(Some("jina-code")), "jina-v2-base-code");
assert_eq!(resolve_embedder_id(Some("nope")), "bge-small-en-v1.5");
assert_eq!(resolve_embedder_id(None), "bge-small-en-v1.5");
}
#[test]
fn noop_embedder_returns_not_implemented_and_is_noop() {
let e = NoopEmbedder;
assert!(e.is_noop());
assert_eq!(e.dim(), 0);
assert_eq!(e.model_id(), "noop");
assert!(matches!(
e.embed("hello").unwrap_err(),
EmbedderError::NotImplemented
));
}
#[test]
fn stub_embedder_is_deterministic() {
let e = StubEmbedder::new();
let a = e.embed("hello rust").expect("embed a");
let b = e.embed("hello rust").expect("embed b");
let c = e.embed("hello RUST").expect("embed c");
assert_eq!(a, b, "same input -> same output");
assert_eq!(
a, c,
"lowercasing means case differences collapse to the same vector"
);
assert_eq!(a.len(), 8);
let norm = a.iter().map(|v| v * v).sum::<f32>().sqrt();
assert!((norm - 1.0).abs() < 1e-5, "expected unit norm, got {norm}");
}
#[test]
fn stub_embedder_distinguishes_disjoint_inputs() {
let e = StubEmbedder::new();
let a = e.embed("foo bar").expect("a");
let b = e.embed("qux quux").expect("b");
let sim = cosine_similarity(&a, &b);
assert!(
sim < 0.99,
"disjoint inputs should not be near-identical: {sim}"
);
}
#[test]
fn stub_embedder_handles_empty_input() {
let e = StubEmbedder::new();
let v = e.embed("").expect("empty embed");
assert_eq!(v.len(), 8);
assert!(v.iter().all(|&x| x == 0.0));
}
#[test]
fn cosine_similarity_handles_edge_cases() {
let a = [1.0f32, 0.0, 0.0];
assert!((cosine_similarity(&a, &a) - 1.0).abs() < 1e-6);
let b = [0.0f32, 1.0, 0.0];
assert!((cosine_similarity(&a, &b)).abs() < 1e-6);
let c = [-1.0f32, 0.0, 0.0];
assert!((cosine_similarity(&a, &c) + 1.0).abs() < 1e-6);
assert_eq!(cosine_similarity(&[], &a), 0.0);
assert_eq!(cosine_similarity(&a, &[0.0]), 0.0);
let zeros = [0.0f32, 0.0, 0.0];
assert_eq!(cosine_similarity(&zeros, &a), 0.0);
}
#[test]
fn cosine_similarity_is_symmetric() {
let a = [0.6f32, 0.8, 0.0];
let b = [0.0f32, 1.0, 0.0];
let ab = cosine_similarity(&a, &b);
let ba = cosine_similarity(&b, &a);
assert!((ab - ba).abs() < 1e-6);
assert!((ab - 0.8).abs() < 1e-5);
}
#[test]
fn encode_decode_embedding_round_trip() {
let vec = vec![0.1f32, -0.2, 3.125, -0.000_001, 42.0];
let blob = encode_embedding(&vec);
assert_eq!(blob.len(), vec.len() * 4);
let back = decode_embedding(&blob, Some(vec.len())).expect("decode");
assert_eq!(back.len(), vec.len());
for (orig, got) in vec.iter().zip(back.iter()) {
assert!(
(orig - got).abs() < 1e-7,
"f32 round-trip should be bit-exact"
);
}
}
#[test]
fn decode_embedding_rejects_unaligned_blob() {
let bad = [0u8, 1, 2]; let err = decode_embedding(&bad, None).unwrap_err();
assert!(err.to_string().contains("not a multiple of 4"));
}
#[test]
fn decode_embedding_rejects_dim_mismatch() {
let vec = vec![1.0f32, 2.0, 3.0];
let blob = encode_embedding(&vec);
let err = decode_embedding(&blob, Some(5)).unwrap_err();
assert!(err.to_string().contains("does not match expected"));
}
#[test]
fn stub_reranker_returns_doc_order_scores() {
let r = StubReranker;
let query = "rust async tokio";
let docs = &["rust async tokio", "python django", "rust only"];
let scores = r.rerank(query, docs).expect("rerank should succeed");
assert_eq!(scores.len(), docs.len(), "one score per document");
for (i, &s) in scores.iter().enumerate() {
assert!(s > 0.0 && s < 1.0, "score[{i}] must be in (0,1), got {s}");
}
}
#[test]
fn stub_reranker_higher_overlap_scores_higher() {
let r = StubReranker;
let query = "rust async tokio";
let docs = &["rust async tokio runtime", "rust only", "python django"];
let scores = r.rerank(query, docs).expect("rerank");
assert!(
scores[0] > scores[1],
"3-token overlap must beat 1-token overlap: {} vs {}",
scores[0],
scores[1]
);
assert!(
scores[1] > scores[2],
"1-token overlap must beat 0-token overlap: {} vs {}",
scores[1],
scores[2]
);
}
#[test]
fn stub_reranker_model_id() {
let r = StubReranker;
assert_eq!(r.model_id(), "stub-reranker");
}
#[test]
fn stub_reranker_empty_query_returns_floor() {
let r = StubReranker;
let docs = &["anything here", "another doc"];
let scores = r.rerank("", docs).expect("rerank");
for &s in &scores {
assert!(
(s - 0.05).abs() < 1e-6,
"empty query must yield 0.05, got {s}"
);
}
}
#[test]
fn embed_batch_matches_per_row() {
let e = StubEmbedder::new();
let texts = ["foo bar", "qux", "hello world"];
let batch = e.embed_batch(&texts).expect("embed_batch should succeed");
assert_eq!(batch.len(), texts.len());
for (i, text) in texts.iter().enumerate() {
let single = e.embed(text).expect("per-row embed should succeed");
assert_eq!(
batch[i], single,
"embed_batch[{i}] must match per-row embed for {text:?}"
);
}
}
#[test]
fn embed_batch_empty_is_empty() {
let e = StubEmbedder::new();
let result = e
.embed_batch(&[])
.expect("empty embed_batch should succeed");
assert!(result.is_empty(), "expected empty Vec, got {result:?}");
}
#[test]
fn embed_batch_length_matches_input() {
let e = StubEmbedder::new();
let texts: Vec<&str> = vec!["alpha", "beta", "gamma", "delta", "epsilon"];
let batch = e.embed_batch(&texts).expect("embed_batch should succeed");
assert_eq!(batch.len(), texts.len(), "output len must equal input len");
for (i, v) in batch.iter().enumerate() {
assert_eq!(
v.len(),
e.dim(),
"vector[{i}] len {} != dim {}",
v.len(),
e.dim()
);
}
}
#[cfg(not(feature = "embeddings"))]
#[test]
fn open_default_embedder_returns_noop_on_default_build() {
let e = open_default_embedder();
assert!(e.is_noop());
assert_eq!(e.dim(), 0);
assert!(matches!(
e.embed("anything").unwrap_err(),
EmbedderError::NotImplemented
));
}
#[test]
fn env_disables_embedder_recognizes_off_values() {
let lock = crate::user_brain::test_env_lock()
.lock()
.unwrap_or_else(|p| p.into_inner());
let prev = std::env::var("KIMETSU_BRAIN_EMBEDDER").ok();
for value in ["noop", "off", "NONE", "0", "false", "no"] {
unsafe {
std::env::set_var("KIMETSU_BRAIN_EMBEDDER", value);
}
assert!(env_disables_embedder(), "value {value:?} must disable");
}
for value in ["", "default", "bge-small", "bge-m3", "jina-code"] {
unsafe {
std::env::set_var("KIMETSU_BRAIN_EMBEDDER", value);
}
assert!(!env_disables_embedder(), "value {value:?} must NOT disable");
}
unsafe {
match prev {
Some(v) => std::env::set_var("KIMETSU_BRAIN_EMBEDDER", v),
None => std::env::remove_var("KIMETSU_BRAIN_EMBEDDER"),
}
}
drop(lock);
}
#[test]
fn w3_embedder_enabled_for_config_false_when_env_unset() {
let lock = crate::user_brain::test_env_lock()
.lock()
.unwrap_or_else(|p| p.into_inner());
let prev = std::env::var("KIMETSU_BRAIN_EMBEDDER").ok();
unsafe {
std::env::remove_var("KIMETSU_BRAIN_EMBEDDER");
}
assert!(
!embedder_enabled_for_config(false),
"config=false + env unset must be disabled"
);
assert!(
embedder_enabled_for_config(true),
"config=true + env unset must be enabled"
);
unsafe {
match prev {
Some(v) => std::env::set_var("KIMETSU_BRAIN_EMBEDDER", v),
None => std::env::remove_var("KIMETSU_BRAIN_EMBEDDER"),
}
}
drop(lock);
}
#[test]
fn w3_embedder_env_disable_overrides_config_true() {
let lock = crate::user_brain::test_env_lock()
.lock()
.unwrap_or_else(|p| p.into_inner());
let prev = std::env::var("KIMETSU_BRAIN_EMBEDDER").ok();
unsafe {
std::env::set_var("KIMETSU_BRAIN_EMBEDDER", "noop");
}
assert!(
!embedder_enabled_for_config(true),
"KIMETSU_BRAIN_EMBEDDER=noop must override config=true"
);
unsafe {
match prev {
Some(v) => std::env::set_var("KIMETSU_BRAIN_EMBEDDER", v),
None => std::env::remove_var("KIMETSU_BRAIN_EMBEDDER"),
}
}
drop(lock);
}
#[test]
fn w3_embedder_env_model_id_overrides_config_false() {
let lock = crate::user_brain::test_env_lock()
.lock()
.unwrap_or_else(|p| p.into_inner());
let prev = std::env::var("KIMETSU_BRAIN_EMBEDDER").ok();
unsafe {
std::env::set_var("KIMETSU_BRAIN_EMBEDDER", "bge-m3");
}
assert!(
embedder_enabled_for_config(false),
"real model id in env must override config=false → enabled"
);
unsafe {
match prev {
Some(v) => std::env::set_var("KIMETSU_BRAIN_EMBEDDER", v),
None => std::env::remove_var("KIMETSU_BRAIN_EMBEDDER"),
}
}
drop(lock);
}
#[test]
fn pick_builtin_model_from_env_handles_aliases() {
let lock = crate::user_brain::test_env_lock()
.lock()
.unwrap_or_else(|p| p.into_inner());
let prev = std::env::var("KIMETSU_BRAIN_EMBEDDER").ok();
let cases = [
("", "bge-small-en-v1.5"),
("default", "bge-small-en-v1.5"),
("bge-small", "bge-small-en-v1.5"),
("BGE-SMALL-EN-V1.5", "bge-small-en-v1.5"),
("bge-m3", "bge-m3"),
("M3", "bge-m3"),
("jina-code", "jina-v2-base-code"),
("jina-v2-base-code", "jina-v2-base-code"),
("jina-embeddings-v2-base-code", "jina-v2-base-code"),
("totally-made-up", "bge-small-en-v1.5"),
];
for (input, expected) in cases {
unsafe {
std::env::set_var("KIMETSU_BRAIN_EMBEDDER", input);
}
assert_eq!(
pick_builtin_model_from_env(),
expected,
"input {input:?} -> expected {expected}"
);
}
unsafe {
match prev {
Some(v) => std::env::set_var("KIMETSU_BRAIN_EMBEDDER", v),
None => std::env::remove_var("KIMETSU_BRAIN_EMBEDDER"),
}
}
drop(lock);
}
}