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",
"mmarco-minilm-l12-v2-int8",
];
if CURATED.contains(&v.as_str()) {
return match fastembed_backend::FastembedReranker::try_open(model_id) {
Ok(r) => Some(Box::new(r) as Box<dyn Reranker>),
Err(err) => {
eprintln!(
"kimetsu-brain: reranker {model_id:?} unavailable ({err}); \
continuing without cross-encoder reranking"
);
None
}
};
}
if USER_DEFINED_ALIASES.contains(&v.as_str()) || v.contains('/') {
return match fastembed_backend::FastembedReranker::try_open_user_defined(model_id) {
Ok(r) => Some(Box::new(r) as Box<dyn Reranker>),
Err(err) => {
eprintln!(
"kimetsu-brain: reranker {model_id:?} unavailable ({err}); \
continuing without cross-encoder reranking"
);
None
}
};
}
eprintln!("kimetsu-brain: unknown reranker {model_id:?}");
None
}
#[cfg(not(feature = "embeddings"))]
{
let _ = v;
None
}
}
pub fn reranker_is_off(model_id: &str) -> bool {
matches!(
model_id.trim().to_ascii_lowercase().as_str(),
"" | "off" | "none" | "noop"
)
}
pub fn open_reranker_checked(model_id: &str) -> Result<Option<Box<dyn Reranker>>, String> {
if reranker_is_off(model_id) {
return Ok(None);
}
open_reranker_for_model(model_id).map(Some).ok_or_else(|| {
format!("requested reranker {model_id:?} unavailable; no cross-encoder measurement")
})
}
type CachedReranker = Result<Option<std::sync::Arc<dyn Reranker>>, String>;
#[derive(Default)]
struct RerankerCache(std::sync::Mutex<std::collections::HashMap<String, CachedReranker>>);
impl RerankerCache {
fn get(&self, id: &str, load: impl FnOnce(&str) -> CachedReranker) -> CachedReranker {
if reranker_is_off(id) {
return Ok(None);
}
let mut entries = self.0.lock().unwrap_or_else(|e| e.into_inner());
entries
.entry(id.trim().to_string())
.or_insert_with(|| load(id))
.clone()
}
}
pub fn open_cached_reranker(model_id: &str) -> CachedReranker {
static CACHE: std::sync::OnceLock<RerankerCache> = std::sync::OnceLock::new();
CACHE
.get_or_init(RerankerCache::default)
.get(model_id, |id| {
#[cfg(feature = "embeddings")]
{
open_reranker_checked(id).map(|r| r.map(std::sync::Arc::from))
}
#[cfg(not(feature = "embeddings"))]
{
let _ = id;
Ok(None)
}
})
}
#[cfg(test)]
mod configured_reranker_tests {
use super::*;
#[test]
fn configured_cache_reuses_model_and_off_never_loads() {
let cache = RerankerCache::default();
let calls = std::sync::atomic::AtomicUsize::new(0);
let load = |_: &str| {
calls.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
Ok(Some(
std::sync::Arc::new(StubReranker) as std::sync::Arc<dyn Reranker>
))
};
let first = cache.get("configured", load).unwrap().unwrap();
let second = cache.get("configured", load).unwrap().unwrap();
assert!(std::sync::Arc::ptr_eq(&first, &second));
assert_eq!(calls.load(std::sync::atomic::Ordering::SeqCst), 1);
assert!(
cache
.get("off", |_| panic!("off must not load"))
.unwrap()
.is_none()
);
assert!(cache.get("failed", |_| Err("unavailable".into())).is_err());
assert!(
cache
.get("failed", |_| panic!("failure must remain explicit"))
.is_err()
);
}
}
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_checked(config_enabled: bool) -> Result<&'static dyn Embedder, String> {
let embedder = open_embedder_for(config_enabled);
validate_requested_embedder(
embedder,
embedder_enabled_for_config(config_enabled),
cfg!(feature = "embeddings"),
)?;
Ok(embedder)
}
fn validate_requested_embedder(
embedder: &dyn Embedder,
enabled: bool,
available: bool,
) -> Result<(), String> {
if available && enabled && embedder.is_noop() {
return Err("requested embedder unavailable after initialization; no semantic measurement (explicitly disable embeddings for lexical-only serving)".into());
}
Ok(())
}
#[cfg(test)]
mod checked_serving_loader_tests {
use super::*;
#[test]
fn failed_requested_model_is_not_an_intentional_lexical_measurement() {
assert!(validate_requested_embedder(&NoopEmbedder, true, true).is_err());
assert!(validate_requested_embedder(&NoopEmbedder, false, true).is_ok());
assert!(validate_requested_embedder(&NoopEmbedder, true, false).is_ok());
assert!(validate_requested_embedder(&StubEmbedder::default(), true, true).is_ok());
}
}
pub fn open_embedder_for_model(model_id: &str) -> Box<dyn Embedder + Send + Sync> {
let model_id = match canonical_embedder_id(model_id) {
Ok("noop") => return Box::new(NoopEmbedder),
Ok(id) => id,
Err(error) => {
eprintln!("{error}");
return Box::new(NoopEmbedder);
}
};
#[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)
}
pub fn canonical_embedder_id(id: &str) -> Result<&'static str, EmbedderError> {
match id.trim().to_ascii_lowercase().as_str() {
"noop" | "off" | "none" | "0" | "false" | "no" => Ok("noop"),
"" | "default" | "bge-small" | "bge-small-en-v1.5" => Ok("bge-small-en-v1.5"),
"bge-m3" | "m3" => Ok("bge-m3"),
"jina-code" | "jina-v2-base-code" | "jina-embeddings-v2-base-code" => {
Ok("jina-v2-base-code")
}
_ => Err(EmbedderError::LoadFailed(format!(
"unknown requested embedder {id:?}"
))),
}
}
#[cfg(test)]
mod explicit_embedder_tests {
use super::*;
#[test]
fn aliases_and_disable_have_one_effective_model_identity() {
assert_eq!(
canonical_embedder_id("jina-code").unwrap(),
"jina-v2-base-code"
);
assert_eq!(canonical_embedder_id("m3").unwrap(), "bge-m3");
assert_eq!(
canonical_embedder_id("bge-small").unwrap(),
"bge-small-en-v1.5"
);
for off in ["off", "noop", "false", "none", "0"] {
assert_eq!(canonical_embedder_id(off).unwrap(), "noop");
assert!(open_embedder_for_model(off).is_noop());
}
assert!(canonical_embedder_id("typo-not-a-model").is_err());
}
}
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 configure_runtime_threads() -> Result<(), EmbedderError> {
static CONFIGURED: OnceLock<Result<(), String>> = OnceLock::new();
CONFIGURED.get_or_init(|| {
let raw = std::env::var("KIMETSU_INTRA_THREADS").ok();
let Some(threads) = super::parse_runtime_threads(raw.as_deref())? else { return Ok(()) };
let pool = ort::environment::GlobalThreadPoolOptions::default()
.with_intra_threads(threads).map_err(|e| e.to_string())?
.with_inter_threads(1).map_err(|e| e.to_string())?
.with_spin_control(false).map_err(|e| e.to_string())?;
if !ort::init().with_global_thread_pool(pool).commit() {
return Err("KIMETSU_INTRA_THREADS cannot take effect: ONNX environment already configured; set it before the first model load".into());
}
eprintln!("kimetsu-brain: ONNX shared intra-op threads={threads}, inter-op=1, spinning=off");
Ok(())
}).clone().map_err(EmbedderError::LoadFailed)
}
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"),
"mmarco-minilm-l12-v2-int8" => Some("cross-encoder/mmarco-mMiniLMv2-L12-H384-v1"),
_ => None,
}
}
fn download_user_defined_reranker(
model_id: &str,
) -> Result<(fastembed::OnnxSource, fastembed::TokenizerFiles), EmbedderError> {
use hf_hub::api::sync::ApiBuilder;
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 = ApiBuilder::from_env().build().map_err(|e| {
EmbedderError::LoadFailed(format!("hf-hub ApiBuilder::from_env failed: {e}"))
})?;
let multilingual_int8 = lowercased == "mmarco-minilm-l12-v2-int8";
let repo = if multilingual_int8 {
api.repo(hf_hub::Repo::with_revision(
repo_id.clone(),
hf_hub::RepoType::Model,
"1427fd652930e4ba29e8149678df786c240d8825".into(),
))
} else {
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 = if multilingual_int8 {
repo.get("onnx/model_quint8_avx2.onnx")
} else {
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> {
configure_runtime_threads()?;
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> {
configure_runtime_threads()?;
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> {
configure_runtime_threads()?;
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);
}
use rusqlite::OptionalExtension;
let expected_revision: Option<String> = conn.query_row(
"SELECT COALESCE((SELECT event_id FROM memory_revisions WHERE memory_id=?1 ORDER BY revision_id DESC LIMIT 1),'baseline:' || memory_id)
FROM memories WHERE memory_id=?1 AND text=?2 AND invalidated_at IS NULL AND superseded_by IS NULL",
rusqlite::params![memory_id,text], |r|r.get(0)).optional()?;
let Some(expected_revision) = expected_revision else {
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);
let changed = conn.execute(
"UPDATE memories SET embedding=?1,embedding_model=?2 WHERE memory_id=?3 AND text=?4
AND invalidated_at IS NULL AND superseded_by IS NULL
AND COALESCE((SELECT event_id FROM memory_revisions WHERE memory_id=?3 ORDER BY revision_id DESC LIMIT 1),'baseline:' || memory_id)=?5",
rusqlite::params![blob,embedder.model_id(),memory_id,text,expected_revision],
)?;
if changed == 0 {
return Ok(None);
}
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(any(test, feature = "embeddings"))]
fn parse_runtime_threads(raw: Option<&str>) -> Result<Option<usize>, String> {
let Some(raw) = raw else { return Ok(None) };
let threads = raw
.trim()
.parse::<usize>()
.map_err(|_| "KIMETSU_INTRA_THREADS must be an integer from 1 to 1024".to_string())?;
if !(1..=1024).contains(&threads) {
return Err("KIMETSU_INTRA_THREADS must be an integer from 1 to 1024".into());
}
Ok(Some(threads))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn runtime_threads_are_explicit_bounded_and_invalid_values_are_errors() {
assert_eq!(parse_runtime_threads(None).unwrap(), None);
assert_eq!(parse_runtime_threads(Some(" 4 ")).unwrap(), Some(4));
assert_eq!(parse_runtime_threads(Some("1")).unwrap(), Some(1));
for value in ["0", "-1", "abc", "1025", "999999999999999999999999"] {
assert!(
parse_runtime_threads(Some(value)).is_err(),
"invalid setting: {value}"
);
}
}
#[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);
}
}
#[cfg(test)]
mod correction_race_tests {
use super::*;
#[test]
fn slow_embedding_cannot_overwrite_a_newer_correction() {
let dir = tempfile::tempdir().unwrap();
let db = dir.path().join("brain.db");
let writer = rusqlite::Connection::open(&db).unwrap();
crate::schema::initialize(&writer).unwrap();
let accepted = kimetsu_core::event::Event::new(
kimetsu_core::ids::RunId::new(),
"memory.accepted",
serde_json::json!({"memory_id":"m","scope":"project","kind":"fact","text":"claim A"}),
);
crate::projector::apply_events(&writer, &[accepted]).unwrap();
let (started_tx, started_rx) = std::sync::mpsc::channel();
let (resume_tx, resume_rx) = std::sync::mpsc::channel();
struct Blocking {
started: std::sync::mpsc::Sender<()>,
resume: std::sync::Mutex<std::sync::mpsc::Receiver<()>>,
}
impl Embedder for Blocking {
fn embed(&self, _: &str) -> Result<Vec<f32>, EmbedderError> {
self.started.send(()).unwrap();
self.resume.lock().unwrap().recv().unwrap();
Ok(vec![1.0, 0.0])
}
fn model_id(&self) -> &str {
"stub"
}
fn dim(&self) -> usize {
2
}
}
let pending = std::thread::spawn(move || {
let conn = rusqlite::Connection::open(db).unwrap();
embed_and_persist(
&conn,
"m",
"claim A",
&Blocking {
started: started_tx,
resume: std::sync::Mutex::new(resume_rx),
},
)
.unwrap()
});
started_rx.recv().unwrap();
let correction = kimetsu_core::event::Event::new(
kimetsu_core::ids::RunId::new(),
"memory.corrected",
serde_json::json!({"memory_id":"m","text":"claim B"}),
);
crate::projector::apply_events(&writer, &[correction]).unwrap();
writer
.execute(
"UPDATE memories SET embedding=?1,embedding_model='stub' WHERE memory_id='m'",
rusqlite::params![encode_embedding(&[0.0, 1.0])],
)
.unwrap();
resume_tx.send(()).unwrap();
assert!(
pending.join().unwrap().is_none(),
"stale computation must not be published"
);
let blob: Vec<u8> = writer
.query_row(
"SELECT embedding FROM memories WHERE memory_id='m'",
[],
|r| r.get(0),
)
.unwrap();
assert_eq!(decode_embedding(&blob, Some(2)).unwrap(), vec![0.0, 1.0]);
}
}