1use kimetsu_core::KimetsuResult;
42
43pub const DEFAULT_HYBRID_ALPHA: f32 = 0.5;
47
48pub trait Embedder: Send + Sync {
55 fn embed(&self, text: &str) -> Result<Vec<f32>, EmbedderError>;
58
59 fn model_id(&self) -> &str;
63
64 fn dim(&self) -> usize;
67
68 fn is_noop(&self) -> bool {
72 false
73 }
74
75 fn embed_batch(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>, EmbedderError> {
83 texts.iter().map(|t| self.embed(t)).collect()
84 }
85}
86
87impl Embedder for Box<dyn Embedder> {
90 fn embed(&self, text: &str) -> Result<Vec<f32>, EmbedderError> {
91 (**self).embed(text)
92 }
93 fn model_id(&self) -> &str {
94 (**self).model_id()
95 }
96 fn dim(&self) -> usize {
97 (**self).dim()
98 }
99 fn is_noop(&self) -> bool {
100 (**self).is_noop()
101 }
102 fn embed_batch(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>, EmbedderError> {
103 (**self).embed_batch(texts)
104 }
105}
106
107#[derive(Debug, Clone)]
109pub enum EmbedderError {
110 NotImplemented,
114 LoadFailed(String),
117 EmbedFailed(String),
119 DimMismatch { expected: usize, got: usize },
122}
123
124impl std::fmt::Display for EmbedderError {
125 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
126 match self {
127 Self::NotImplemented => write!(f, "embedder not implemented"),
128 Self::LoadFailed(msg) => write!(f, "embedder load failed: {msg}"),
129 Self::EmbedFailed(msg) => write!(f, "embed call failed: {msg}"),
130 Self::DimMismatch { expected, got } => {
131 write!(f, "embedding dim mismatch: expected {expected}, got {got}")
132 }
133 }
134 }
135}
136
137impl std::error::Error for EmbedderError {}
138
139#[derive(Debug, Default, Clone, Copy)]
143pub struct NoopEmbedder;
144
145impl NoopEmbedder {
146 pub const MODEL_ID: &'static str = "noop";
147}
148
149impl Embedder for NoopEmbedder {
150 fn embed(&self, _text: &str) -> Result<Vec<f32>, EmbedderError> {
151 Err(EmbedderError::NotImplemented)
152 }
153
154 fn model_id(&self) -> &str {
155 Self::MODEL_ID
156 }
157
158 fn dim(&self) -> usize {
159 0
160 }
161
162 fn is_noop(&self) -> bool {
163 true
164 }
165}
166
167#[derive(Debug, Clone, Copy)]
177pub struct StubEmbedder {
178 dim: usize,
179}
180
181impl StubEmbedder {
182 pub const MODEL_ID: &'static str = "stub-d8";
183
184 pub const fn new() -> Self {
185 Self { dim: 8 }
186 }
187
188 pub const fn with_dim(dim: usize) -> Self {
189 Self { dim }
190 }
191}
192
193impl Default for StubEmbedder {
194 fn default() -> Self {
195 Self::new()
196 }
197}
198
199impl Embedder for StubEmbedder {
200 fn embed(&self, text: &str) -> Result<Vec<f32>, EmbedderError> {
201 let mut bucket = vec![0.0f32; self.dim];
202 for word in text.split_whitespace() {
203 let normalized = word.to_lowercase();
207 let mut h: u64 = 0xcbf2_9ce4_8422_2325;
208 for byte in normalized.bytes() {
209 h ^= byte as u64;
210 h = h.wrapping_mul(0x0000_0100_0000_01B3);
211 }
212 let idx = (h as usize) % self.dim.max(1);
213 bucket[idx] += 1.0;
214 }
215 let norm = bucket.iter().map(|v| v * v).sum::<f32>().sqrt();
217 if norm > 0.0 {
218 for v in &mut bucket {
219 *v /= norm;
220 }
221 }
222 Ok(bucket)
223 }
224
225 fn model_id(&self) -> &str {
226 Self::MODEL_ID
227 }
228
229 fn dim(&self) -> usize {
230 self.dim
231 }
232}
233
234pub trait Reranker: Send + Sync {
241 fn rerank(&self, query: &str, documents: &[&str]) -> Result<Vec<f32>, EmbedderError>;
242 fn model_id(&self) -> &str;
243}
244
245pub struct StubReranker;
257
258impl Reranker for StubReranker {
259 fn rerank(&self, query: &str, documents: &[&str]) -> Result<Vec<f32>, EmbedderError> {
260 let query_tokens: std::collections::HashSet<String> = query
261 .split(|c: char| !c.is_alphanumeric())
262 .filter(|t| !t.is_empty())
263 .map(|t| t.to_lowercase())
264 .collect();
265 let q_len = query_tokens.len();
266 let scores = documents
267 .iter()
268 .map(|doc| {
269 if q_len == 0 {
270 return 0.05_f32;
271 }
272 let doc_tokens: std::collections::HashSet<String> = doc
273 .split(|c: char| !c.is_alphanumeric())
274 .filter(|t| !t.is_empty())
275 .map(|t| t.to_lowercase())
276 .collect();
277 let intersection = query_tokens.intersection(&doc_tokens).count();
278 let overlap = intersection as f32 / q_len as f32;
279 (0.05 + 0.9 * overlap).clamp(0.0, 1.0)
280 })
281 .collect();
282 Ok(scores)
283 }
284
285 fn model_id(&self) -> &str {
286 "stub-reranker"
287 }
288}
289
290pub fn open_reranker_for_model(model_id: &str) -> Option<Box<dyn Reranker>> {
302 let v = model_id.trim().to_ascii_lowercase();
303 if v.is_empty() || matches!(v.as_str(), "off" | "none" | "noop") {
304 return None;
305 }
306 #[cfg(feature = "embeddings")]
307 {
308 const CURATED: &[&str] = &[
310 "jina-reranker-v1-turbo-en",
311 "bge-reranker-base",
312 "bge-reranker-v2-m3",
313 "jina-reranker-v2-base-multilingual",
314 ];
315 const USER_DEFINED_ALIASES: &[&str] = &[
317 "jina-reranker-v1-tiny-en",
318 "ms-marco-tinybert-l-2-v2",
319 "ms-marco-minilm-l-4-v2",
320 ];
321
322 if CURATED.contains(&v.as_str()) {
323 return fastembed_backend::FastembedReranker::try_open(model_id)
324 .ok()
325 .map(|r| Box::new(r) as Box<dyn Reranker>);
326 }
327 if USER_DEFINED_ALIASES.contains(&v.as_str()) || v.contains('/') {
328 return fastembed_backend::FastembedReranker::try_open_user_defined(model_id)
329 .ok()
330 .map(|r| Box::new(r) as Box<dyn Reranker>);
331 }
332 fastembed_backend::FastembedReranker::try_open("jina-reranker-v1-turbo-en")
334 .ok()
335 .map(|r| Box::new(r) as Box<dyn Reranker>)
336 }
337 #[cfg(not(feature = "embeddings"))]
338 {
339 let _ = v;
340 None
341 }
342}
343
344pub fn open_default_embedder() -> &'static (dyn Embedder + Send + Sync) {
369 static CACHE: std::sync::OnceLock<Box<dyn Embedder + Send + Sync>> = std::sync::OnceLock::new();
370 let embedder = CACHE.get_or_init(build_default_embedder);
371 embedder.as_ref()
372}
373
374fn build_default_embedder() -> Box<dyn Embedder + Send + Sync> {
375 if env_disables_embedder() {
376 return Box::new(NoopEmbedder);
377 }
378 #[cfg(feature = "embeddings")]
379 {
380 match fastembed_backend::open_cached() {
381 Ok(handle) => return Box::new(handle),
382 Err(err) => {
383 eprintln!(
384 "kimetsu-brain: fastembed init failed ({err}); falling back to NoopEmbedder. \
385 Retrieval will stay FTS-only this session. Re-run with \
386 KIMETSU_BRAIN_EMBEDDER=noop to silence this warning."
387 );
388 }
389 }
390 }
391 Box::new(NoopEmbedder)
392}
393
394pub fn open_embedder_for(config_enabled: bool) -> &'static dyn Embedder {
407 if embedder_enabled_for_config(config_enabled) {
408 open_default_embedder()
409 } else {
410 &NoopEmbedder
411 }
412}
413
414pub fn open_embedder_for_model(model_id: &str) -> Box<dyn Embedder + Send + Sync> {
423 #[cfg(feature = "embeddings")]
424 {
425 match fastembed_backend::FastembedEmbedder::try_open(model_id) {
426 Ok(engine) => return Box::new(engine),
427 Err(err) => {
428 eprintln!(
429 "kimetsu-brain: failed to open embedder `{model_id}` ({err}); \
430 using NoopEmbedder (no vectors produced)."
431 );
432 }
433 }
434 }
435 #[cfg(not(feature = "embeddings"))]
436 {
437 let _ = model_id;
438 }
439 Box::new(NoopEmbedder)
440}
441
442fn env_disables_embedder() -> bool {
447 match std::env::var("KIMETSU_BRAIN_EMBEDDER") {
448 Ok(value) => is_disable_value(&value.trim().to_ascii_lowercase()),
449 Err(_) => false,
450 }
451}
452
453pub fn embedder_enabled_for_config(config_enabled: bool) -> bool {
462 match std::env::var("KIMETSU_BRAIN_EMBEDDER") {
464 Ok(raw) => {
465 let v = raw.trim().to_ascii_lowercase();
466 if v.is_empty() {
467 config_enabled
469 } else if is_disable_value(&v) {
470 false
472 } else {
473 true
475 }
476 }
477 Err(_) => config_enabled,
479 }
480}
481
482fn is_disable_value(v: &str) -> bool {
483 matches!(v, "noop" | "off" | "none" | "0" | "false" | "no")
484}
485
486pub const BUILTIN_MODELS: &[(&str, usize, &str)] = &[
492 ("bge-small-en-v1.5", 384, "English, default, ~67 MB int8"),
493 ("bge-m3", 1024, "Multilingual, ~600 MB int8"),
494 (
495 "jina-v2-base-code",
496 768,
497 "English + code-tuned, ~165 MB int8",
498 ),
499];
500
501static EMBEDDER_OVERRIDE: std::sync::OnceLock<String> = std::sync::OnceLock::new();
507
508pub fn apply_embedder_selection(config_embedder: Option<&str>) {
514 if let Some(id) = config_embedder {
515 let id = id.trim();
516 if !id.is_empty() {
517 let _ = EMBEDDER_OVERRIDE.set(id.to_string());
518 }
519 }
520}
521
522fn map_builtin_id(v: &str) -> &'static str {
528 match v {
529 "" | "default" | "bge-small" | "bge-small-en-v1.5" => "bge-small-en-v1.5",
530 "bge-m3" | "m3" => "bge-m3",
531 "jina-code" | "jina-v2-base-code" | "jina-embeddings-v2-base-code" => "jina-v2-base-code",
532 "noop" | "off" | "none" | "0" | "false" | "no" => "bge-small-en-v1.5",
533 other => {
534 eprintln!(
535 "kimetsu-brain: unknown embedder {other:?}, \
536 falling back to bge-small-en-v1.5"
537 );
538 "bge-small-en-v1.5"
539 }
540 }
541}
542
543pub fn resolve_embedder_id(config_embedder: Option<&str>) -> &'static str {
549 if let Ok(raw) = std::env::var("KIMETSU_BRAIN_EMBEDDER") {
550 let v = raw.trim().to_ascii_lowercase();
551 if !v.is_empty() && !is_disable_value(&v) {
552 return map_builtin_id(&v);
553 }
554 }
557 let cfg = config_embedder
558 .map(str::to_string)
559 .or_else(|| EMBEDDER_OVERRIDE.get().cloned());
560 if let Some(c) = cfg {
561 let v = c.trim().to_ascii_lowercase();
562 if !v.is_empty() {
563 return map_builtin_id(&v);
564 }
565 }
566 "bge-small-en-v1.5"
567}
568
569pub fn pick_builtin_model_from_env() -> &'static str {
585 resolve_embedder_id(None)
589}
590
591#[cfg(feature = "embeddings")]
595mod fastembed_backend {
596 use super::{Embedder, EmbedderError, Reranker, pick_builtin_model_from_env};
597 use fastembed::{
598 EmbeddingModel, InitOptions, RerankInitOptions, RerankerModel, TextEmbedding, TextRerank,
599 };
600 use std::sync::{Arc, Mutex, OnceLock};
601
602 fn hf_repo_for_alias(lowercased: &str) -> Option<&'static str> {
606 match lowercased {
607 "jina-reranker-v1-tiny-en" => Some("jinaai/jina-reranker-v1-tiny-en"),
608 "ms-marco-tinybert-l-2-v2" => Some("Xenova/ms-marco-TinyBERT-L-2-v2"),
609 "ms-marco-minilm-l-4-v2" => Some("Xenova/ms-marco-MiniLM-L-4-v2"),
610 _ => None,
611 }
612 }
613
614 fn download_user_defined_reranker(
618 model_id: &str,
619 ) -> Result<(fastembed::OnnxSource, fastembed::TokenizerFiles), EmbedderError> {
620 use hf_hub::api::sync::Api;
621
622 let lowercased = model_id.trim().to_ascii_lowercase();
623 let repo_id: String = if let Some(alias) = hf_repo_for_alias(&lowercased) {
624 alias.to_string()
625 } else if lowercased.contains('/') {
626 model_id.to_string()
628 } else {
629 return Err(EmbedderError::LoadFailed(format!(
630 "user-defined reranker: no HF repo mapping for {model_id:?}"
631 )));
632 };
633
634 let api = Api::new()
635 .map_err(|e| EmbedderError::LoadFailed(format!("hf-hub Api::new failed: {e}")))?;
636 let repo = api.model(repo_id.clone());
637
638 let get_required = |filename: &str| -> Result<Vec<u8>, EmbedderError> {
640 let path = repo.get(filename).map_err(|e| {
641 EmbedderError::LoadFailed(format!("{repo_id}/{filename}: download failed: {e}"))
642 })?;
643 std::fs::read(&path).map_err(|e| {
644 EmbedderError::LoadFailed(format!("{repo_id}/{filename}: read failed: {e}"))
645 })
646 };
647
648 let tokenizer_file = get_required("tokenizer.json")?;
649 let config_file = get_required("config.json")?;
650 let tokenizer_config_file = get_required("tokenizer_config.json")?;
651 let special_tokens_map_file = get_required("special_tokens_map.json")?;
652
653 let onnx_path = repo
655 .get("onnx/model.onnx")
656 .or_else(|_| repo.get("model.onnx"))
657 .map_err(|e| {
658 EmbedderError::LoadFailed(format!(
659 "{repo_id}: could not find onnx/model.onnx or model.onnx: {e}"
660 ))
661 })?;
662
663 let tokenizer_files = fastembed::TokenizerFiles {
664 tokenizer_file,
665 config_file,
666 special_tokens_map_file,
667 tokenizer_config_file,
668 };
669
670 Ok((fastembed::OnnxSource::File(onnx_path), tokenizer_files))
671 }
672
673 pub struct FastembedEmbedder {
678 model_id: &'static str,
679 dim: usize,
680 engine: Mutex<TextEmbedding>,
681 }
682
683 impl FastembedEmbedder {
684 pub fn try_open(builtin_id: &str) -> Result<Self, EmbedderError> {
685 let (kind, model_id, dim) = match builtin_id {
686 "bge-m3" => (EmbeddingModel::BGEM3, "bge-m3", 1024),
687 "jina-v2-base-code" => (
688 EmbeddingModel::JinaEmbeddingsV2BaseCode,
689 "jina-v2-base-code",
690 768,
691 ),
692 _ => (EmbeddingModel::BGESmallENV15, "bge-small-en-v1.5", 384),
694 };
695 let opts = InitOptions::new(kind).with_show_download_progress(false);
696 let engine = TextEmbedding::try_new(opts)
697 .map_err(|e| EmbedderError::LoadFailed(format!("fastembed init: {e}")))?;
698 Ok(Self {
699 model_id,
700 dim,
701 engine: Mutex::new(engine),
702 })
703 }
704 }
705
706 impl Embedder for FastembedEmbedder {
707 fn embed(&self, text: &str) -> Result<Vec<f32>, EmbedderError> {
708 let mut guard = self
709 .engine
710 .lock()
711 .unwrap_or_else(|poisoned| poisoned.into_inner());
712 let mut out = guard
713 .embed(vec![text], None)
714 .map_err(|e| EmbedderError::EmbedFailed(format!("fastembed embed: {e}")))?;
715 let vec = out
716 .pop()
717 .ok_or_else(|| EmbedderError::EmbedFailed("empty result".into()))?;
718 if vec.len() != self.dim {
719 return Err(EmbedderError::DimMismatch {
720 expected: self.dim,
721 got: vec.len(),
722 });
723 }
724 Ok(vec)
725 }
726
727 fn model_id(&self) -> &str {
728 self.model_id
729 }
730
731 fn dim(&self) -> usize {
732 self.dim
733 }
734
735 fn embed_batch(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>, EmbedderError> {
736 if texts.is_empty() {
737 return Ok(Vec::new());
738 }
739 let mut guard = self
740 .engine
741 .lock()
742 .unwrap_or_else(|poisoned| poisoned.into_inner());
743 let out = guard
744 .embed(texts, None)
745 .map_err(|e| EmbedderError::EmbedFailed(format!("fastembed embed_batch: {e}")))?;
746 if out.len() != texts.len() {
747 return Err(EmbedderError::EmbedFailed(format!(
748 "fastembed returned {} vectors for {} texts",
749 out.len(),
750 texts.len()
751 )));
752 }
753 for v in &out {
754 if v.len() != self.dim {
755 return Err(EmbedderError::DimMismatch {
756 expected: self.dim,
757 got: v.len(),
758 });
759 }
760 }
761 Ok(out)
762 }
763 }
764
765 #[derive(Clone)]
769 pub struct EmbedderHandle(Arc<FastembedEmbedder>);
770
771 impl Embedder for EmbedderHandle {
772 fn embed(&self, text: &str) -> Result<Vec<f32>, EmbedderError> {
773 self.0.embed(text)
774 }
775 fn model_id(&self) -> &str {
776 self.0.model_id()
777 }
778 fn dim(&self) -> usize {
779 self.0.dim()
780 }
781 fn embed_batch(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>, EmbedderError> {
782 self.0.embed_batch(texts)
783 }
784 }
785
786 pub struct FastembedReranker {
796 model_id: String,
797 engine: Mutex<TextRerank>,
798 }
799
800 impl FastembedReranker {
801 pub fn try_open(builtin_id: &str) -> Result<Self, EmbedderError> {
805 let (kind, stable_id) = match builtin_id {
806 "bge-reranker-base" => (RerankerModel::BGERerankerBase, "bge-reranker-base"),
807 "bge-reranker-v2-m3" => (RerankerModel::BGERerankerV2M3, "bge-reranker-v2-m3"),
808 "jina-reranker-v2-base-multilingual" => (
809 RerankerModel::JINARerankerV2BaseMultiligual,
810 "jina-reranker-v2-base-multilingual",
811 ),
812 _ => (
814 RerankerModel::JINARerankerV1TurboEn,
815 "jina-reranker-v1-turbo-en",
816 ),
817 };
818 let opts = RerankInitOptions::new(kind).with_show_download_progress(false);
819 let engine = TextRerank::try_new(opts)
820 .map_err(|e| EmbedderError::LoadFailed(format!("fastembed reranker init: {e}")))?;
821 Ok(Self {
822 model_id: stable_id.to_string(),
823 engine: Mutex::new(engine),
824 })
825 }
826
827 pub fn try_open_user_defined(alias_or_repo: &str) -> Result<Self, EmbedderError> {
835 use fastembed::{RerankInitOptionsUserDefined, UserDefinedRerankingModel};
836
837 let (onnx_source, tokenizer_files) = download_user_defined_reranker(alias_or_repo)?;
838
839 let model = UserDefinedRerankingModel::new(onnx_source, tokenizer_files);
840 let opts = RerankInitOptionsUserDefined::default();
841 let engine = TextRerank::try_new_from_user_defined(model, opts).map_err(|e| {
842 EmbedderError::LoadFailed(format!(
843 "user-defined reranker {alias_or_repo:?} init: {e}"
844 ))
845 })?;
846
847 let model_id = alias_or_repo.trim().to_ascii_lowercase();
849 Ok(Self {
850 model_id,
851 engine: Mutex::new(engine),
852 })
853 }
854 }
855
856 impl Reranker for FastembedReranker {
857 fn rerank(&self, query: &str, documents: &[&str]) -> Result<Vec<f32>, EmbedderError> {
858 if documents.is_empty() {
859 return Ok(Vec::new());
860 }
861 let mut guard = self
862 .engine
863 .lock()
864 .unwrap_or_else(|poisoned| poisoned.into_inner());
865 let raw_results = guard
868 .rerank(query, documents, false, None)
869 .map_err(|e| EmbedderError::EmbedFailed(format!("fastembed rerank: {e}")))?;
870 let n = documents.len();
871 let mut scores = vec![0.0f32; n];
872 for result in raw_results {
873 if result.index < n {
874 scores[result.index] = 1.0 / (1.0 + (-result.score).exp());
876 }
877 }
878 Ok(scores)
879 }
880
881 fn model_id(&self) -> &str {
882 &self.model_id
883 }
884 }
885
886 pub fn open_cached() -> Result<EmbedderHandle, EmbedderError> {
891 static CELL: OnceLock<Result<Arc<FastembedEmbedder>, EmbedderError>> = OnceLock::new();
892 let init = CELL.get_or_init(|| {
893 let builtin = pick_builtin_model_from_env();
894 FastembedEmbedder::try_open(builtin).map(Arc::new)
895 });
896 match init {
897 Ok(arc) => Ok(EmbedderHandle(arc.clone())),
898 Err(err) => Err(err.clone()),
899 }
900 }
901}
902
903pub fn cosine_similarity(a: &[f32], b: &[f32]) -> f32 {
909 if a.is_empty() || b.is_empty() || a.len() != b.len() {
910 return 0.0;
911 }
912 let mut dot = 0.0f32;
913 let mut na = 0.0f32;
914 let mut nb = 0.0f32;
915 for (x, y) in a.iter().zip(b.iter()) {
916 dot += x * y;
917 na += x * x;
918 nb += y * y;
919 }
920 if na == 0.0 || nb == 0.0 {
921 return 0.0;
922 }
923 dot / (na.sqrt() * nb.sqrt())
924}
925
926pub fn embed_and_persist(
944 conn: &rusqlite::Connection,
945 memory_id: &str,
946 text: &str,
947 embedder: &dyn Embedder,
948) -> KimetsuResult<Option<Vec<f32>>> {
949 if embedder.is_noop() {
950 return Ok(None);
951 }
952 let vec = match embedder.embed(text) {
953 Ok(v) => v,
954 Err(EmbedderError::NotImplemented) => return Ok(None),
957 Err(e) => return Err(format!("embed failed for memory {memory_id}: {e}").into()),
958 };
959 if vec.len() != embedder.dim() {
960 return Err(format!(
961 "embedder {} produced {} dims, expected {}",
962 embedder.model_id(),
963 vec.len(),
964 embedder.dim()
965 )
966 .into());
967 }
968 let blob = encode_embedding(&vec);
969 conn.execute(
970 "UPDATE memories SET embedding = ?1, embedding_model = ?2 WHERE memory_id = ?3",
971 rusqlite::params![blob, embedder.model_id(), memory_id],
972 )?;
973
974 #[cfg(feature = "embeddings")]
979 if let Some(handle) = crate::ann::cached_handle(conn) {
980 let rowid: Option<i64> = conn
981 .query_row(
982 "SELECT rowid FROM memories WHERE memory_id = ?1",
983 rusqlite::params![memory_id],
984 |r| r.get(0),
985 )
986 .ok();
987 if let Some(rowid) = rowid {
988 let mut guard = handle.write().unwrap_or_else(|p| p.into_inner());
989 if let Err(e) = guard.add(rowid, &vec) {
990 eprintln!(
991 "kimetsu-brain: ann add failed for memory {memory_id}: {e} (index will reconcile on next open)"
992 );
993 }
994 }
995 }
996
997 Ok(Some(vec))
998}
999
1000pub fn encode_embedding(vec: &[f32]) -> Vec<u8> {
1009 let mut out = Vec::with_capacity(vec.len() * 4);
1010 for v in vec {
1011 out.extend_from_slice(&v.to_le_bytes());
1012 }
1013 out
1014}
1015
1016pub fn decode_embedding(bytes: &[u8], expected_dim: Option<usize>) -> KimetsuResult<Vec<f32>> {
1019 if bytes.len() % 4 != 0 {
1020 return Err(format!("embedding blob length {} not a multiple of 4", bytes.len()).into());
1021 }
1022 let dim = bytes.len() / 4;
1023 if let Some(expected) = expected_dim
1024 && dim != expected
1025 {
1026 return Err(format!("embedding blob dim {dim} does not match expected {expected}").into());
1027 }
1028 let mut out = Vec::with_capacity(dim);
1029 for chunk in bytes.chunks_exact(4) {
1030 let mut buf = [0u8; 4];
1031 buf.copy_from_slice(chunk);
1032 out.push(f32::from_le_bytes(buf));
1033 }
1034 Ok(out)
1035}
1036
1037#[cfg(test)]
1038mod tests {
1039 use super::*;
1040
1041 #[test]
1042 fn map_builtin_id_maps_aliases_and_defaults_unknown() {
1043 assert_eq!(map_builtin_id("bge-small-en-v1.5"), "bge-small-en-v1.5");
1044 assert_eq!(map_builtin_id("default"), "bge-small-en-v1.5");
1045 assert_eq!(map_builtin_id("m3"), "bge-m3");
1046 assert_eq!(map_builtin_id("bge-m3"), "bge-m3");
1047 assert_eq!(map_builtin_id("jina-code"), "jina-v2-base-code");
1048 assert_eq!(
1049 map_builtin_id("jina-embeddings-v2-base-code"),
1050 "jina-v2-base-code"
1051 );
1052 assert_eq!(map_builtin_id("noop"), "bge-small-en-v1.5");
1055 assert_eq!(map_builtin_id("totally-made-up"), "bge-small-en-v1.5");
1057 }
1058
1059 #[test]
1060 fn builtin_models_table_is_consistent() {
1061 for (id, _dim, _blurb) in BUILTIN_MODELS {
1063 assert_eq!(map_builtin_id(id), *id, "id {id} must be stable");
1064 }
1065 }
1066
1067 #[test]
1068 fn resolve_embedder_id_uses_config_when_env_unset() {
1069 if std::env::var_os("KIMETSU_BRAIN_EMBEDDER").is_some() {
1074 return;
1075 }
1076 assert_eq!(resolve_embedder_id(Some("bge-m3")), "bge-m3");
1077 assert_eq!(resolve_embedder_id(Some("jina-code")), "jina-v2-base-code");
1078 assert_eq!(resolve_embedder_id(Some("nope")), "bge-small-en-v1.5");
1080 assert_eq!(resolve_embedder_id(None), "bge-small-en-v1.5");
1082 }
1083
1084 #[test]
1085 fn noop_embedder_returns_not_implemented_and_is_noop() {
1086 let e = NoopEmbedder;
1087 assert!(e.is_noop());
1088 assert_eq!(e.dim(), 0);
1089 assert_eq!(e.model_id(), "noop");
1090 assert!(matches!(
1091 e.embed("hello").unwrap_err(),
1092 EmbedderError::NotImplemented
1093 ));
1094 }
1095
1096 #[test]
1097 fn stub_embedder_is_deterministic() {
1098 let e = StubEmbedder::new();
1099 let a = e.embed("hello rust").expect("embed a");
1100 let b = e.embed("hello rust").expect("embed b");
1101 let c = e.embed("hello RUST").expect("embed c");
1102 assert_eq!(a, b, "same input -> same output");
1103 assert_eq!(
1104 a, c,
1105 "lowercasing means case differences collapse to the same vector"
1106 );
1107 assert_eq!(a.len(), 8);
1108 let norm = a.iter().map(|v| v * v).sum::<f32>().sqrt();
1110 assert!((norm - 1.0).abs() < 1e-5, "expected unit norm, got {norm}");
1111 }
1112
1113 #[test]
1114 fn stub_embedder_distinguishes_disjoint_inputs() {
1115 let e = StubEmbedder::new();
1116 let a = e.embed("foo bar").expect("a");
1117 let b = e.embed("qux quux").expect("b");
1118 let sim = cosine_similarity(&a, &b);
1119 assert!(
1122 sim < 0.99,
1123 "disjoint inputs should not be near-identical: {sim}"
1124 );
1125 }
1126
1127 #[test]
1128 fn stub_embedder_handles_empty_input() {
1129 let e = StubEmbedder::new();
1130 let v = e.embed("").expect("empty embed");
1131 assert_eq!(v.len(), 8);
1132 assert!(v.iter().all(|&x| x == 0.0));
1136 }
1137
1138 #[test]
1139 fn cosine_similarity_handles_edge_cases() {
1140 let a = [1.0f32, 0.0, 0.0];
1142 assert!((cosine_similarity(&a, &a) - 1.0).abs() < 1e-6);
1143
1144 let b = [0.0f32, 1.0, 0.0];
1146 assert!((cosine_similarity(&a, &b)).abs() < 1e-6);
1147
1148 let c = [-1.0f32, 0.0, 0.0];
1150 assert!((cosine_similarity(&a, &c) + 1.0).abs() < 1e-6);
1151
1152 assert_eq!(cosine_similarity(&[], &a), 0.0);
1154 assert_eq!(cosine_similarity(&a, &[0.0]), 0.0);
1155
1156 let zeros = [0.0f32, 0.0, 0.0];
1158 assert_eq!(cosine_similarity(&zeros, &a), 0.0);
1159 }
1160
1161 #[test]
1162 fn cosine_similarity_is_symmetric() {
1163 let a = [0.6f32, 0.8, 0.0];
1164 let b = [0.0f32, 1.0, 0.0];
1165 let ab = cosine_similarity(&a, &b);
1166 let ba = cosine_similarity(&b, &a);
1167 assert!((ab - ba).abs() < 1e-6);
1168 assert!((ab - 0.8).abs() < 1e-5);
1170 }
1171
1172 #[test]
1173 fn encode_decode_embedding_round_trip() {
1174 let vec = vec![0.1f32, -0.2, 3.125, -0.000_001, 42.0];
1175 let blob = encode_embedding(&vec);
1176 assert_eq!(blob.len(), vec.len() * 4);
1177 let back = decode_embedding(&blob, Some(vec.len())).expect("decode");
1178 assert_eq!(back.len(), vec.len());
1179 for (orig, got) in vec.iter().zip(back.iter()) {
1180 assert!(
1181 (orig - got).abs() < 1e-7,
1182 "f32 round-trip should be bit-exact"
1183 );
1184 }
1185 }
1186
1187 #[test]
1188 fn decode_embedding_rejects_unaligned_blob() {
1189 let bad = [0u8, 1, 2]; let err = decode_embedding(&bad, None).unwrap_err();
1191 assert!(err.to_string().contains("not a multiple of 4"));
1192 }
1193
1194 #[test]
1195 fn decode_embedding_rejects_dim_mismatch() {
1196 let vec = vec![1.0f32, 2.0, 3.0];
1197 let blob = encode_embedding(&vec);
1198 let err = decode_embedding(&blob, Some(5)).unwrap_err();
1199 assert!(err.to_string().contains("does not match expected"));
1200 }
1201
1202 #[test]
1206 fn stub_reranker_returns_doc_order_scores() {
1207 let r = StubReranker;
1208 let query = "rust async tokio";
1209 let docs = &["rust async tokio", "python django", "rust only"];
1210 let scores = r.rerank(query, docs).expect("rerank should succeed");
1211 assert_eq!(scores.len(), docs.len(), "one score per document");
1212 for (i, &s) in scores.iter().enumerate() {
1214 assert!(s > 0.0 && s < 1.0, "score[{i}] must be in (0,1), got {s}");
1215 }
1216 }
1217
1218 #[test]
1220 fn stub_reranker_higher_overlap_scores_higher() {
1221 let r = StubReranker;
1222 let query = "rust async tokio";
1223 let docs = &["rust async tokio runtime", "rust only", "python django"];
1227 let scores = r.rerank(query, docs).expect("rerank");
1228 assert!(
1229 scores[0] > scores[1],
1230 "3-token overlap must beat 1-token overlap: {} vs {}",
1231 scores[0],
1232 scores[1]
1233 );
1234 assert!(
1235 scores[1] > scores[2],
1236 "1-token overlap must beat 0-token overlap: {} vs {}",
1237 scores[1],
1238 scores[2]
1239 );
1240 }
1241
1242 #[test]
1244 fn stub_reranker_model_id() {
1245 let r = StubReranker;
1246 assert_eq!(r.model_id(), "stub-reranker");
1247 }
1248
1249 #[test]
1251 fn stub_reranker_empty_query_returns_floor() {
1252 let r = StubReranker;
1253 let docs = &["anything here", "another doc"];
1254 let scores = r.rerank("", docs).expect("rerank");
1255 for &s in &scores {
1256 assert!(
1257 (s - 0.05).abs() < 1e-6,
1258 "empty query must yield 0.05, got {s}"
1259 );
1260 }
1261 }
1262
1263 #[test]
1268 fn embed_batch_matches_per_row() {
1269 let e = StubEmbedder::new();
1270 let texts = ["foo bar", "qux", "hello world"];
1271 let batch = e.embed_batch(&texts).expect("embed_batch should succeed");
1272 assert_eq!(batch.len(), texts.len());
1273 for (i, text) in texts.iter().enumerate() {
1274 let single = e.embed(text).expect("per-row embed should succeed");
1275 assert_eq!(
1276 batch[i], single,
1277 "embed_batch[{i}] must match per-row embed for {text:?}"
1278 );
1279 }
1280 }
1281
1282 #[test]
1284 fn embed_batch_empty_is_empty() {
1285 let e = StubEmbedder::new();
1286 let result = e
1287 .embed_batch(&[])
1288 .expect("empty embed_batch should succeed");
1289 assert!(result.is_empty(), "expected empty Vec, got {result:?}");
1290 }
1291
1292 #[test]
1294 fn embed_batch_length_matches_input() {
1295 let e = StubEmbedder::new();
1296 let texts: Vec<&str> = vec!["alpha", "beta", "gamma", "delta", "epsilon"];
1297 let batch = e.embed_batch(&texts).expect("embed_batch should succeed");
1298 assert_eq!(batch.len(), texts.len(), "output len must equal input len");
1299 for (i, v) in batch.iter().enumerate() {
1300 assert_eq!(
1301 v.len(),
1302 e.dim(),
1303 "vector[{i}] len {} != dim {}",
1304 v.len(),
1305 e.dim()
1306 );
1307 }
1308 }
1309
1310 #[cfg(not(feature = "embeddings"))]
1317 #[test]
1318 fn open_default_embedder_returns_noop_on_default_build() {
1319 let e = open_default_embedder();
1320 assert!(e.is_noop());
1321 assert_eq!(e.dim(), 0);
1322 assert!(matches!(
1323 e.embed("anything").unwrap_err(),
1324 EmbedderError::NotImplemented
1325 ));
1326 }
1327
1328 #[test]
1335 fn env_disables_embedder_recognizes_off_values() {
1336 let lock = crate::user_brain::test_env_lock()
1337 .lock()
1338 .unwrap_or_else(|p| p.into_inner());
1339 let prev = std::env::var("KIMETSU_BRAIN_EMBEDDER").ok();
1340 for value in ["noop", "off", "NONE", "0", "false", "no"] {
1341 unsafe {
1343 std::env::set_var("KIMETSU_BRAIN_EMBEDDER", value);
1344 }
1345 assert!(env_disables_embedder(), "value {value:?} must disable");
1346 }
1347 for value in ["", "default", "bge-small", "bge-m3", "jina-code"] {
1348 unsafe {
1349 std::env::set_var("KIMETSU_BRAIN_EMBEDDER", value);
1350 }
1351 assert!(!env_disables_embedder(), "value {value:?} must NOT disable");
1352 }
1353 unsafe {
1355 match prev {
1356 Some(v) => std::env::set_var("KIMETSU_BRAIN_EMBEDDER", v),
1357 None => std::env::remove_var("KIMETSU_BRAIN_EMBEDDER"),
1358 }
1359 }
1360 drop(lock);
1361 }
1362
1363 #[test]
1367 fn w3_embedder_enabled_for_config_false_when_env_unset() {
1368 let lock = crate::user_brain::test_env_lock()
1369 .lock()
1370 .unwrap_or_else(|p| p.into_inner());
1371 let prev = std::env::var("KIMETSU_BRAIN_EMBEDDER").ok();
1372 unsafe {
1373 std::env::remove_var("KIMETSU_BRAIN_EMBEDDER");
1374 }
1375 assert!(
1377 !embedder_enabled_for_config(false),
1378 "config=false + env unset must be disabled"
1379 );
1380 assert!(
1382 embedder_enabled_for_config(true),
1383 "config=true + env unset must be enabled"
1384 );
1385 unsafe {
1386 match prev {
1387 Some(v) => std::env::set_var("KIMETSU_BRAIN_EMBEDDER", v),
1388 None => std::env::remove_var("KIMETSU_BRAIN_EMBEDDER"),
1389 }
1390 }
1391 drop(lock);
1392 }
1393
1394 #[test]
1396 fn w3_embedder_env_disable_overrides_config_true() {
1397 let lock = crate::user_brain::test_env_lock()
1398 .lock()
1399 .unwrap_or_else(|p| p.into_inner());
1400 let prev = std::env::var("KIMETSU_BRAIN_EMBEDDER").ok();
1401 unsafe {
1402 std::env::set_var("KIMETSU_BRAIN_EMBEDDER", "noop");
1403 }
1404 assert!(
1405 !embedder_enabled_for_config(true),
1406 "KIMETSU_BRAIN_EMBEDDER=noop must override config=true"
1407 );
1408 unsafe {
1409 match prev {
1410 Some(v) => std::env::set_var("KIMETSU_BRAIN_EMBEDDER", v),
1411 None => std::env::remove_var("KIMETSU_BRAIN_EMBEDDER"),
1412 }
1413 }
1414 drop(lock);
1415 }
1416
1417 #[test]
1420 fn w3_embedder_env_model_id_overrides_config_false() {
1421 let lock = crate::user_brain::test_env_lock()
1422 .lock()
1423 .unwrap_or_else(|p| p.into_inner());
1424 let prev = std::env::var("KIMETSU_BRAIN_EMBEDDER").ok();
1425 unsafe {
1426 std::env::set_var("KIMETSU_BRAIN_EMBEDDER", "bge-m3");
1427 }
1428 assert!(
1429 embedder_enabled_for_config(false),
1430 "real model id in env must override config=false → enabled"
1431 );
1432 unsafe {
1433 match prev {
1434 Some(v) => std::env::set_var("KIMETSU_BRAIN_EMBEDDER", v),
1435 None => std::env::remove_var("KIMETSU_BRAIN_EMBEDDER"),
1436 }
1437 }
1438 drop(lock);
1439 }
1440
1441 #[test]
1445 fn pick_builtin_model_from_env_handles_aliases() {
1446 let lock = crate::user_brain::test_env_lock()
1447 .lock()
1448 .unwrap_or_else(|p| p.into_inner());
1449 let prev = std::env::var("KIMETSU_BRAIN_EMBEDDER").ok();
1450 let cases = [
1451 ("", "bge-small-en-v1.5"),
1452 ("default", "bge-small-en-v1.5"),
1453 ("bge-small", "bge-small-en-v1.5"),
1454 ("BGE-SMALL-EN-V1.5", "bge-small-en-v1.5"),
1455 ("bge-m3", "bge-m3"),
1456 ("M3", "bge-m3"),
1457 ("jina-code", "jina-v2-base-code"),
1458 ("jina-v2-base-code", "jina-v2-base-code"),
1459 ("jina-embeddings-v2-base-code", "jina-v2-base-code"),
1460 ("totally-made-up", "bge-small-en-v1.5"),
1462 ];
1463 for (input, expected) in cases {
1464 unsafe {
1466 std::env::set_var("KIMETSU_BRAIN_EMBEDDER", input);
1467 }
1468 assert_eq!(
1469 pick_builtin_model_from_env(),
1470 expected,
1471 "input {input:?} -> expected {expected}"
1472 );
1473 }
1474 unsafe {
1475 match prev {
1476 Some(v) => std::env::set_var("KIMETSU_BRAIN_EMBEDDER", v),
1477 None => std::env::remove_var("KIMETSU_BRAIN_EMBEDDER"),
1478 }
1479 }
1480 drop(lock);
1481 }
1482}