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
}
}
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()
}
}
#[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 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_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,
}
}
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, pick_builtin_model_from_env};
use fastembed::{EmbeddingModel, InitOptions, TextEmbedding};
use std::sync::{Arc, Mutex, OnceLock};
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
}
}
#[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()
}
}
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<()> {
if embedder.is_noop() {
return Ok(());
}
let vec = match embedder.embed(text) {
Ok(v) => v,
Err(EmbedderError::NotImplemented) => return Ok(()),
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],
)?;
Ok(())
}
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().is_multiple_of(4) {
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"));
}
#[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 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);
}
}