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 "mmarco-minilm-l12-v2-int8",
321 ];
322
323 if CURATED.contains(&v.as_str()) {
324 return match fastembed_backend::FastembedReranker::try_open(model_id) {
325 Ok(r) => Some(Box::new(r) as Box<dyn Reranker>),
326 Err(err) => {
327 eprintln!(
328 "kimetsu-brain: reranker {model_id:?} unavailable ({err}); \
329 continuing without cross-encoder reranking"
330 );
331 None
332 }
333 };
334 }
335 if USER_DEFINED_ALIASES.contains(&v.as_str()) || v.contains('/') {
336 return match fastembed_backend::FastembedReranker::try_open_user_defined(model_id) {
337 Ok(r) => Some(Box::new(r) as Box<dyn Reranker>),
338 Err(err) => {
339 eprintln!(
340 "kimetsu-brain: reranker {model_id:?} unavailable ({err}); \
341 continuing without cross-encoder reranking"
342 );
343 None
344 }
345 };
346 }
347 eprintln!("kimetsu-brain: unknown reranker {model_id:?}");
348 None
349 }
350 #[cfg(not(feature = "embeddings"))]
351 {
352 let _ = v;
353 None
354 }
355}
356
357pub fn reranker_is_off(model_id: &str) -> bool {
358 matches!(
359 model_id.trim().to_ascii_lowercase().as_str(),
360 "" | "off" | "none" | "noop"
361 )
362}
363
364pub fn open_reranker_checked(model_id: &str) -> Result<Option<Box<dyn Reranker>>, String> {
366 if reranker_is_off(model_id) {
367 return Ok(None);
368 }
369 open_reranker_for_model(model_id).map(Some).ok_or_else(|| {
370 format!("requested reranker {model_id:?} unavailable; no cross-encoder measurement")
371 })
372}
373
374type CachedReranker = Result<Option<std::sync::Arc<dyn Reranker>>, String>;
375#[derive(Default)]
376struct RerankerCache(std::sync::Mutex<std::collections::HashMap<String, CachedReranker>>);
377impl RerankerCache {
378 fn get(&self, id: &str, load: impl FnOnce(&str) -> CachedReranker) -> CachedReranker {
379 if reranker_is_off(id) {
380 return Ok(None);
381 }
382 let mut entries = self.0.lock().unwrap_or_else(|e| e.into_inner());
383 entries
384 .entry(id.trim().to_string())
385 .or_insert_with(|| load(id))
386 .clone()
387 }
388}
389pub fn open_cached_reranker(model_id: &str) -> CachedReranker {
393 static CACHE: std::sync::OnceLock<RerankerCache> = std::sync::OnceLock::new();
394 CACHE
395 .get_or_init(RerankerCache::default)
396 .get(model_id, |id| {
397 #[cfg(feature = "embeddings")]
398 {
399 open_reranker_checked(id).map(|r| r.map(std::sync::Arc::from))
400 }
401 #[cfg(not(feature = "embeddings"))]
402 {
403 let _ = id;
404 Ok(None)
405 }
406 })
407}
408
409#[cfg(test)]
410mod configured_reranker_tests {
411 use super::*;
412 #[test]
413 fn configured_cache_reuses_model_and_off_never_loads() {
414 let cache = RerankerCache::default();
415 let calls = std::sync::atomic::AtomicUsize::new(0);
416 let load = |_: &str| {
417 calls.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
418 Ok(Some(
419 std::sync::Arc::new(StubReranker) as std::sync::Arc<dyn Reranker>
420 ))
421 };
422 let first = cache.get("configured", load).unwrap().unwrap();
423 let second = cache.get("configured", load).unwrap().unwrap();
424 assert!(std::sync::Arc::ptr_eq(&first, &second));
425 assert_eq!(calls.load(std::sync::atomic::Ordering::SeqCst), 1);
426 assert!(
427 cache
428 .get("off", |_| panic!("off must not load"))
429 .unwrap()
430 .is_none()
431 );
432 assert!(cache.get("failed", |_| Err("unavailable".into())).is_err());
433 assert!(
434 cache
435 .get("failed", |_| panic!("failure must remain explicit"))
436 .is_err()
437 );
438 }
439}
440
441pub fn open_default_embedder() -> &'static (dyn Embedder + Send + Sync) {
466 static CACHE: std::sync::OnceLock<Box<dyn Embedder + Send + Sync>> = std::sync::OnceLock::new();
467 let embedder = CACHE.get_or_init(build_default_embedder);
468 embedder.as_ref()
469}
470
471fn build_default_embedder() -> Box<dyn Embedder + Send + Sync> {
472 if env_disables_embedder() {
473 return Box::new(NoopEmbedder);
474 }
475 #[cfg(feature = "embeddings")]
476 {
477 match fastembed_backend::open_cached() {
478 Ok(handle) => return Box::new(handle),
479 Err(err) => {
480 eprintln!(
481 "kimetsu-brain: fastembed init failed ({err}); falling back to NoopEmbedder. \
482 Retrieval will stay FTS-only this session. Re-run with \
483 KIMETSU_BRAIN_EMBEDDER=noop to silence this warning."
484 );
485 }
486 }
487 }
488 Box::new(NoopEmbedder)
489}
490
491pub fn open_embedder_for(config_enabled: bool) -> &'static dyn Embedder {
504 if embedder_enabled_for_config(config_enabled) {
505 open_default_embedder()
506 } else {
507 &NoopEmbedder
508 }
509}
510
511pub fn open_embedder_for_checked(config_enabled: bool) -> Result<&'static dyn Embedder, String> {
514 let embedder = open_embedder_for(config_enabled);
515 validate_requested_embedder(
516 embedder,
517 embedder_enabled_for_config(config_enabled),
518 cfg!(feature = "embeddings"),
519 )?;
520 Ok(embedder)
521}
522
523fn validate_requested_embedder(
524 embedder: &dyn Embedder,
525 enabled: bool,
526 available: bool,
527) -> Result<(), String> {
528 if available && enabled && embedder.is_noop() {
529 return Err("requested embedder unavailable after initialization; no semantic measurement (explicitly disable embeddings for lexical-only serving)".into());
530 }
531 Ok(())
532}
533
534#[cfg(test)]
535mod checked_serving_loader_tests {
536 use super::*;
537 #[test]
538 fn failed_requested_model_is_not_an_intentional_lexical_measurement() {
539 assert!(validate_requested_embedder(&NoopEmbedder, true, true).is_err());
540 assert!(validate_requested_embedder(&NoopEmbedder, false, true).is_ok());
541 assert!(validate_requested_embedder(&NoopEmbedder, true, false).is_ok());
542 assert!(validate_requested_embedder(&StubEmbedder::default(), true, true).is_ok());
543 }
544}
545
546pub fn open_embedder_for_model(model_id: &str) -> Box<dyn Embedder + Send + Sync> {
555 let model_id = match canonical_embedder_id(model_id) {
556 Ok("noop") => return Box::new(NoopEmbedder),
557 Ok(id) => id,
558 Err(error) => {
559 eprintln!("{error}");
560 return Box::new(NoopEmbedder);
561 }
562 };
563 #[cfg(feature = "embeddings")]
564 {
565 match fastembed_backend::FastembedEmbedder::try_open(model_id) {
566 Ok(engine) => return Box::new(engine),
567 Err(err) => {
568 eprintln!(
569 "kimetsu-brain: failed to open embedder `{model_id}` ({err}); \
570 using NoopEmbedder (no vectors produced)."
571 );
572 }
573 }
574 }
575 #[cfg(not(feature = "embeddings"))]
576 {
577 let _ = model_id;
578 }
579 Box::new(NoopEmbedder)
580}
581
582pub fn canonical_embedder_id(id: &str) -> Result<&'static str, EmbedderError> {
585 match id.trim().to_ascii_lowercase().as_str() {
586 "noop" | "off" | "none" | "0" | "false" | "no" => Ok("noop"),
587 "" | "default" | "bge-small" | "bge-small-en-v1.5" => Ok("bge-small-en-v1.5"),
588 "bge-m3" | "m3" => Ok("bge-m3"),
589 "jina-code" | "jina-v2-base-code" | "jina-embeddings-v2-base-code" => {
590 Ok("jina-v2-base-code")
591 }
592 _ => Err(EmbedderError::LoadFailed(format!(
593 "unknown requested embedder {id:?}"
594 ))),
595 }
596}
597
598#[cfg(test)]
599mod explicit_embedder_tests {
600 use super::*;
601 #[test]
602 fn aliases_and_disable_have_one_effective_model_identity() {
603 assert_eq!(
604 canonical_embedder_id("jina-code").unwrap(),
605 "jina-v2-base-code"
606 );
607 assert_eq!(canonical_embedder_id("m3").unwrap(), "bge-m3");
608 assert_eq!(
609 canonical_embedder_id("bge-small").unwrap(),
610 "bge-small-en-v1.5"
611 );
612 for off in ["off", "noop", "false", "none", "0"] {
613 assert_eq!(canonical_embedder_id(off).unwrap(), "noop");
614 assert!(open_embedder_for_model(off).is_noop());
615 }
616 assert!(canonical_embedder_id("typo-not-a-model").is_err());
617 }
618}
619
620fn env_disables_embedder() -> bool {
625 match std::env::var("KIMETSU_BRAIN_EMBEDDER") {
626 Ok(value) => is_disable_value(&value.trim().to_ascii_lowercase()),
627 Err(_) => false,
628 }
629}
630
631pub fn embedder_enabled_for_config(config_enabled: bool) -> bool {
640 match std::env::var("KIMETSU_BRAIN_EMBEDDER") {
642 Ok(raw) => {
643 let v = raw.trim().to_ascii_lowercase();
644 if v.is_empty() {
645 config_enabled
647 } else if is_disable_value(&v) {
648 false
650 } else {
651 true
653 }
654 }
655 Err(_) => config_enabled,
657 }
658}
659
660fn is_disable_value(v: &str) -> bool {
661 matches!(v, "noop" | "off" | "none" | "0" | "false" | "no")
662}
663
664pub const BUILTIN_MODELS: &[(&str, usize, &str)] = &[
670 ("bge-small-en-v1.5", 384, "English, default, ~67 MB int8"),
671 ("bge-m3", 1024, "Multilingual, ~600 MB int8"),
672 (
673 "jina-v2-base-code",
674 768,
675 "English + code-tuned, ~165 MB int8",
676 ),
677];
678
679static EMBEDDER_OVERRIDE: std::sync::OnceLock<String> = std::sync::OnceLock::new();
685
686pub fn apply_embedder_selection(config_embedder: Option<&str>) {
692 if let Some(id) = config_embedder {
693 let id = id.trim();
694 if !id.is_empty() {
695 let _ = EMBEDDER_OVERRIDE.set(id.to_string());
696 }
697 }
698}
699
700fn map_builtin_id(v: &str) -> &'static str {
706 match v {
707 "" | "default" | "bge-small" | "bge-small-en-v1.5" => "bge-small-en-v1.5",
708 "bge-m3" | "m3" => "bge-m3",
709 "jina-code" | "jina-v2-base-code" | "jina-embeddings-v2-base-code" => "jina-v2-base-code",
710 "noop" | "off" | "none" | "0" | "false" | "no" => "bge-small-en-v1.5",
711 other => {
712 eprintln!(
713 "kimetsu-brain: unknown embedder {other:?}, \
714 falling back to bge-small-en-v1.5"
715 );
716 "bge-small-en-v1.5"
717 }
718 }
719}
720
721pub fn resolve_embedder_id(config_embedder: Option<&str>) -> &'static str {
727 if let Ok(raw) = std::env::var("KIMETSU_BRAIN_EMBEDDER") {
728 let v = raw.trim().to_ascii_lowercase();
729 if !v.is_empty() && !is_disable_value(&v) {
730 return map_builtin_id(&v);
731 }
732 }
735 let cfg = config_embedder
736 .map(str::to_string)
737 .or_else(|| EMBEDDER_OVERRIDE.get().cloned());
738 if let Some(c) = cfg {
739 let v = c.trim().to_ascii_lowercase();
740 if !v.is_empty() {
741 return map_builtin_id(&v);
742 }
743 }
744 "bge-small-en-v1.5"
745}
746
747pub fn pick_builtin_model_from_env() -> &'static str {
763 resolve_embedder_id(None)
767}
768
769#[cfg(feature = "embeddings")]
773mod fastembed_backend {
774 use super::{Embedder, EmbedderError, Reranker, pick_builtin_model_from_env};
775 use fastembed::{
776 EmbeddingModel, InitOptions, RerankInitOptions, RerankerModel, TextEmbedding, TextRerank,
777 };
778 use std::sync::{Arc, Mutex, OnceLock};
779
780 fn configure_runtime_threads() -> Result<(), EmbedderError> {
785 static CONFIGURED: OnceLock<Result<(), String>> = OnceLock::new();
786 CONFIGURED.get_or_init(|| {
787 let raw = std::env::var("KIMETSU_INTRA_THREADS").ok();
788 let Some(threads) = super::parse_runtime_threads(raw.as_deref())? else { return Ok(()) };
789 let pool = ort::environment::GlobalThreadPoolOptions::default()
790 .with_intra_threads(threads).map_err(|e| e.to_string())?
791 .with_inter_threads(1).map_err(|e| e.to_string())?
792 .with_spin_control(false).map_err(|e| e.to_string())?;
793 if !ort::init().with_global_thread_pool(pool).commit() {
794 return Err("KIMETSU_INTRA_THREADS cannot take effect: ONNX environment already configured; set it before the first model load".into());
795 }
796 eprintln!("kimetsu-brain: ONNX shared intra-op threads={threads}, inter-op=1, spinning=off");
797 Ok(())
798 }).clone().map_err(EmbedderError::LoadFailed)
799 }
800
801 fn hf_repo_for_alias(lowercased: &str) -> Option<&'static str> {
805 match lowercased {
806 "jina-reranker-v1-tiny-en" => Some("jinaai/jina-reranker-v1-tiny-en"),
807 "ms-marco-tinybert-l-2-v2" => Some("Xenova/ms-marco-TinyBERT-L-2-v2"),
808 "ms-marco-minilm-l-4-v2" => Some("Xenova/ms-marco-MiniLM-L-4-v2"),
809 "mmarco-minilm-l12-v2-int8" => Some("cross-encoder/mmarco-mMiniLMv2-L12-H384-v1"),
810 _ => None,
811 }
812 }
813
814 fn download_user_defined_reranker(
818 model_id: &str,
819 ) -> Result<(fastembed::OnnxSource, fastembed::TokenizerFiles), EmbedderError> {
820 use hf_hub::api::sync::ApiBuilder;
821
822 let lowercased = model_id.trim().to_ascii_lowercase();
823 let repo_id: String = if let Some(alias) = hf_repo_for_alias(&lowercased) {
824 alias.to_string()
825 } else if lowercased.contains('/') {
826 model_id.to_string()
828 } else {
829 return Err(EmbedderError::LoadFailed(format!(
830 "user-defined reranker: no HF repo mapping for {model_id:?}"
831 )));
832 };
833
834 let api = ApiBuilder::from_env().build().map_err(|e| {
838 EmbedderError::LoadFailed(format!("hf-hub ApiBuilder::from_env failed: {e}"))
839 })?;
840 let multilingual_int8 = lowercased == "mmarco-minilm-l12-v2-int8";
841 let repo = if multilingual_int8 {
842 api.repo(hf_hub::Repo::with_revision(
843 repo_id.clone(),
844 hf_hub::RepoType::Model,
845 "1427fd652930e4ba29e8149678df786c240d8825".into(),
846 ))
847 } else {
848 api.model(repo_id.clone())
849 };
850
851 let get_required = |filename: &str| -> Result<Vec<u8>, EmbedderError> {
853 let path = repo.get(filename).map_err(|e| {
854 EmbedderError::LoadFailed(format!("{repo_id}/{filename}: download failed: {e}"))
855 })?;
856 std::fs::read(&path).map_err(|e| {
857 EmbedderError::LoadFailed(format!("{repo_id}/{filename}: read failed: {e}"))
858 })
859 };
860
861 let tokenizer_file = get_required("tokenizer.json")?;
862 let config_file = get_required("config.json")?;
863 let tokenizer_config_file = get_required("tokenizer_config.json")?;
864 let special_tokens_map_file = get_required("special_tokens_map.json")?;
865
866 let onnx_path = if multilingual_int8 {
868 repo.get("onnx/model_quint8_avx2.onnx")
871 } else {
872 repo.get("onnx/model.onnx")
873 .or_else(|_| repo.get("model.onnx"))
874 }
875 .map_err(|e| {
876 EmbedderError::LoadFailed(format!(
877 "{repo_id}: could not find onnx/model.onnx or model.onnx: {e}"
878 ))
879 })?;
880
881 let tokenizer_files = fastembed::TokenizerFiles {
882 tokenizer_file,
883 config_file,
884 special_tokens_map_file,
885 tokenizer_config_file,
886 };
887
888 Ok((fastembed::OnnxSource::File(onnx_path), tokenizer_files))
889 }
890
891 pub struct FastembedEmbedder {
896 model_id: &'static str,
897 dim: usize,
898 engine: Mutex<TextEmbedding>,
899 }
900
901 impl FastembedEmbedder {
902 pub fn try_open(builtin_id: &str) -> Result<Self, EmbedderError> {
903 configure_runtime_threads()?;
904 let (kind, model_id, dim) = match builtin_id {
905 "bge-m3" => (EmbeddingModel::BGEM3, "bge-m3", 1024),
906 "jina-v2-base-code" => (
907 EmbeddingModel::JinaEmbeddingsV2BaseCode,
908 "jina-v2-base-code",
909 768,
910 ),
911 _ => (EmbeddingModel::BGESmallENV15, "bge-small-en-v1.5", 384),
913 };
914 let opts = InitOptions::new(kind).with_show_download_progress(false);
915 let engine = TextEmbedding::try_new(opts)
916 .map_err(|e| EmbedderError::LoadFailed(format!("fastembed init: {e}")))?;
917 Ok(Self {
918 model_id,
919 dim,
920 engine: Mutex::new(engine),
921 })
922 }
923 }
924
925 impl Embedder for FastembedEmbedder {
926 fn embed(&self, text: &str) -> Result<Vec<f32>, EmbedderError> {
927 let mut guard = self
928 .engine
929 .lock()
930 .unwrap_or_else(|poisoned| poisoned.into_inner());
931 let mut out = guard
932 .embed(vec![text], None)
933 .map_err(|e| EmbedderError::EmbedFailed(format!("fastembed embed: {e}")))?;
934 let vec = out
935 .pop()
936 .ok_or_else(|| EmbedderError::EmbedFailed("empty result".into()))?;
937 if vec.len() != self.dim {
938 return Err(EmbedderError::DimMismatch {
939 expected: self.dim,
940 got: vec.len(),
941 });
942 }
943 Ok(vec)
944 }
945
946 fn model_id(&self) -> &str {
947 self.model_id
948 }
949
950 fn dim(&self) -> usize {
951 self.dim
952 }
953
954 fn embed_batch(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>, EmbedderError> {
955 if texts.is_empty() {
956 return Ok(Vec::new());
957 }
958 let mut guard = self
959 .engine
960 .lock()
961 .unwrap_or_else(|poisoned| poisoned.into_inner());
962 let out = guard
963 .embed(texts, None)
964 .map_err(|e| EmbedderError::EmbedFailed(format!("fastembed embed_batch: {e}")))?;
965 if out.len() != texts.len() {
966 return Err(EmbedderError::EmbedFailed(format!(
967 "fastembed returned {} vectors for {} texts",
968 out.len(),
969 texts.len()
970 )));
971 }
972 for v in &out {
973 if v.len() != self.dim {
974 return Err(EmbedderError::DimMismatch {
975 expected: self.dim,
976 got: v.len(),
977 });
978 }
979 }
980 Ok(out)
981 }
982 }
983
984 #[derive(Clone)]
988 pub struct EmbedderHandle(Arc<FastembedEmbedder>);
989
990 impl Embedder for EmbedderHandle {
991 fn embed(&self, text: &str) -> Result<Vec<f32>, EmbedderError> {
992 self.0.embed(text)
993 }
994 fn model_id(&self) -> &str {
995 self.0.model_id()
996 }
997 fn dim(&self) -> usize {
998 self.0.dim()
999 }
1000 fn embed_batch(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>, EmbedderError> {
1001 self.0.embed_batch(texts)
1002 }
1003 }
1004
1005 pub struct FastembedReranker {
1015 model_id: String,
1016 engine: Mutex<TextRerank>,
1017 }
1018
1019 impl FastembedReranker {
1020 pub fn try_open(builtin_id: &str) -> Result<Self, EmbedderError> {
1024 configure_runtime_threads()?;
1025 let (kind, stable_id) = match builtin_id {
1026 "bge-reranker-base" => (RerankerModel::BGERerankerBase, "bge-reranker-base"),
1027 "bge-reranker-v2-m3" => (RerankerModel::BGERerankerV2M3, "bge-reranker-v2-m3"),
1028 "jina-reranker-v2-base-multilingual" => (
1029 RerankerModel::JINARerankerV2BaseMultiligual,
1030 "jina-reranker-v2-base-multilingual",
1031 ),
1032 _ => (
1034 RerankerModel::JINARerankerV1TurboEn,
1035 "jina-reranker-v1-turbo-en",
1036 ),
1037 };
1038 let opts = RerankInitOptions::new(kind).with_show_download_progress(false);
1039 let engine = TextRerank::try_new(opts)
1040 .map_err(|e| EmbedderError::LoadFailed(format!("fastembed reranker init: {e}")))?;
1041 Ok(Self {
1042 model_id: stable_id.to_string(),
1043 engine: Mutex::new(engine),
1044 })
1045 }
1046
1047 pub fn try_open_user_defined(alias_or_repo: &str) -> Result<Self, EmbedderError> {
1055 configure_runtime_threads()?;
1056 use fastembed::{RerankInitOptionsUserDefined, UserDefinedRerankingModel};
1057
1058 let (onnx_source, tokenizer_files) = download_user_defined_reranker(alias_or_repo)?;
1059
1060 let model = UserDefinedRerankingModel::new(onnx_source, tokenizer_files);
1061 let opts = RerankInitOptionsUserDefined::default();
1062 let engine = TextRerank::try_new_from_user_defined(model, opts).map_err(|e| {
1063 EmbedderError::LoadFailed(format!(
1064 "user-defined reranker {alias_or_repo:?} init: {e}"
1065 ))
1066 })?;
1067
1068 let model_id = alias_or_repo.trim().to_ascii_lowercase();
1070 Ok(Self {
1071 model_id,
1072 engine: Mutex::new(engine),
1073 })
1074 }
1075 }
1076
1077 impl Reranker for FastembedReranker {
1078 fn rerank(&self, query: &str, documents: &[&str]) -> Result<Vec<f32>, EmbedderError> {
1079 if documents.is_empty() {
1080 return Ok(Vec::new());
1081 }
1082 let mut guard = self
1083 .engine
1084 .lock()
1085 .unwrap_or_else(|poisoned| poisoned.into_inner());
1086 let raw_results = guard
1089 .rerank(query, documents, false, None)
1090 .map_err(|e| EmbedderError::EmbedFailed(format!("fastembed rerank: {e}")))?;
1091 let n = documents.len();
1092 let mut scores = vec![0.0f32; n];
1093 for result in raw_results {
1094 if result.index < n {
1095 scores[result.index] = 1.0 / (1.0 + (-result.score).exp());
1097 }
1098 }
1099 Ok(scores)
1100 }
1101
1102 fn model_id(&self) -> &str {
1103 &self.model_id
1104 }
1105 }
1106
1107 pub fn open_cached() -> Result<EmbedderHandle, EmbedderError> {
1112 static CELL: OnceLock<Result<Arc<FastembedEmbedder>, EmbedderError>> = OnceLock::new();
1113 let init = CELL.get_or_init(|| {
1114 let builtin = pick_builtin_model_from_env();
1115 FastembedEmbedder::try_open(builtin).map(Arc::new)
1116 });
1117 match init {
1118 Ok(arc) => Ok(EmbedderHandle(arc.clone())),
1119 Err(err) => Err(err.clone()),
1120 }
1121 }
1122}
1123
1124pub fn cosine_similarity(a: &[f32], b: &[f32]) -> f32 {
1130 if a.is_empty() || b.is_empty() || a.len() != b.len() {
1131 return 0.0;
1132 }
1133 let mut dot = 0.0f32;
1134 let mut na = 0.0f32;
1135 let mut nb = 0.0f32;
1136 for (x, y) in a.iter().zip(b.iter()) {
1137 dot += x * y;
1138 na += x * x;
1139 nb += y * y;
1140 }
1141 if na == 0.0 || nb == 0.0 {
1142 return 0.0;
1143 }
1144 dot / (na.sqrt() * nb.sqrt())
1145}
1146
1147pub fn embed_and_persist(
1165 conn: &rusqlite::Connection,
1166 memory_id: &str,
1167 text: &str,
1168 embedder: &dyn Embedder,
1169) -> KimetsuResult<Option<Vec<f32>>> {
1170 if embedder.is_noop() {
1171 return Ok(None);
1172 }
1173 use rusqlite::OptionalExtension;
1174 let expected_revision: Option<String> = conn.query_row(
1177 "SELECT COALESCE((SELECT event_id FROM memory_revisions WHERE memory_id=?1 ORDER BY revision_id DESC LIMIT 1),'baseline:' || memory_id)
1178 FROM memories WHERE memory_id=?1 AND text=?2 AND invalidated_at IS NULL AND superseded_by IS NULL",
1179 rusqlite::params![memory_id,text], |r|r.get(0)).optional()?;
1180 let Some(expected_revision) = expected_revision else {
1181 return Ok(None);
1182 };
1183 let vec = match embedder.embed(text) {
1184 Ok(v) => v,
1185 Err(EmbedderError::NotImplemented) => return Ok(None),
1188 Err(e) => return Err(format!("embed failed for memory {memory_id}: {e}").into()),
1189 };
1190 if vec.len() != embedder.dim() {
1191 return Err(format!(
1192 "embedder {} produced {} dims, expected {}",
1193 embedder.model_id(),
1194 vec.len(),
1195 embedder.dim()
1196 )
1197 .into());
1198 }
1199 let blob = encode_embedding(&vec);
1200 let changed = conn.execute(
1201 "UPDATE memories SET embedding=?1,embedding_model=?2 WHERE memory_id=?3 AND text=?4
1202 AND invalidated_at IS NULL AND superseded_by IS NULL
1203 AND COALESCE((SELECT event_id FROM memory_revisions WHERE memory_id=?3 ORDER BY revision_id DESC LIMIT 1),'baseline:' || memory_id)=?5",
1204 rusqlite::params![blob,embedder.model_id(),memory_id,text,expected_revision],
1205 )?;
1206 if changed == 0 {
1207 return Ok(None);
1208 }
1209 Ok(Some(vec))
1214}
1215
1216pub fn encode_embedding(vec: &[f32]) -> Vec<u8> {
1225 let mut out = Vec::with_capacity(vec.len() * 4);
1226 for v in vec {
1227 out.extend_from_slice(&v.to_le_bytes());
1228 }
1229 out
1230}
1231
1232pub fn decode_embedding(bytes: &[u8], expected_dim: Option<usize>) -> KimetsuResult<Vec<f32>> {
1235 if bytes.len() % 4 != 0 {
1236 return Err(format!("embedding blob length {} not a multiple of 4", bytes.len()).into());
1237 }
1238 let dim = bytes.len() / 4;
1239 if let Some(expected) = expected_dim
1240 && dim != expected
1241 {
1242 return Err(format!("embedding blob dim {dim} does not match expected {expected}").into());
1243 }
1244 let mut out = Vec::with_capacity(dim);
1245 for chunk in bytes.chunks_exact(4) {
1246 let mut buf = [0u8; 4];
1247 buf.copy_from_slice(chunk);
1248 out.push(f32::from_le_bytes(buf));
1249 }
1250 Ok(out)
1251}
1252
1253#[cfg(any(test, feature = "embeddings"))]
1254fn parse_runtime_threads(raw: Option<&str>) -> Result<Option<usize>, String> {
1255 let Some(raw) = raw else { return Ok(None) };
1256 let threads = raw
1257 .trim()
1258 .parse::<usize>()
1259 .map_err(|_| "KIMETSU_INTRA_THREADS must be an integer from 1 to 1024".to_string())?;
1260 if !(1..=1024).contains(&threads) {
1261 return Err("KIMETSU_INTRA_THREADS must be an integer from 1 to 1024".into());
1262 }
1263 Ok(Some(threads))
1264}
1265
1266#[cfg(test)]
1267mod tests {
1268 use super::*;
1269
1270 #[test]
1271 fn runtime_threads_are_explicit_bounded_and_invalid_values_are_errors() {
1272 assert_eq!(parse_runtime_threads(None).unwrap(), None);
1273 assert_eq!(parse_runtime_threads(Some(" 4 ")).unwrap(), Some(4));
1274 assert_eq!(parse_runtime_threads(Some("1")).unwrap(), Some(1));
1275 for value in ["0", "-1", "abc", "1025", "999999999999999999999999"] {
1276 assert!(
1277 parse_runtime_threads(Some(value)).is_err(),
1278 "invalid setting: {value}"
1279 );
1280 }
1281 }
1282
1283 #[test]
1284 fn map_builtin_id_maps_aliases_and_defaults_unknown() {
1285 assert_eq!(map_builtin_id("bge-small-en-v1.5"), "bge-small-en-v1.5");
1286 assert_eq!(map_builtin_id("default"), "bge-small-en-v1.5");
1287 assert_eq!(map_builtin_id("m3"), "bge-m3");
1288 assert_eq!(map_builtin_id("bge-m3"), "bge-m3");
1289 assert_eq!(map_builtin_id("jina-code"), "jina-v2-base-code");
1290 assert_eq!(
1291 map_builtin_id("jina-embeddings-v2-base-code"),
1292 "jina-v2-base-code"
1293 );
1294 assert_eq!(map_builtin_id("noop"), "bge-small-en-v1.5");
1297 assert_eq!(map_builtin_id("totally-made-up"), "bge-small-en-v1.5");
1299 }
1300
1301 #[test]
1302 fn builtin_models_table_is_consistent() {
1303 for (id, _dim, _blurb) in BUILTIN_MODELS {
1305 assert_eq!(map_builtin_id(id), *id, "id {id} must be stable");
1306 }
1307 }
1308
1309 #[test]
1310 fn resolve_embedder_id_uses_config_when_env_unset() {
1311 if std::env::var_os("KIMETSU_BRAIN_EMBEDDER").is_some() {
1316 return;
1317 }
1318 assert_eq!(resolve_embedder_id(Some("bge-m3")), "bge-m3");
1319 assert_eq!(resolve_embedder_id(Some("jina-code")), "jina-v2-base-code");
1320 assert_eq!(resolve_embedder_id(Some("nope")), "bge-small-en-v1.5");
1322 assert_eq!(resolve_embedder_id(None), "bge-small-en-v1.5");
1324 }
1325
1326 #[test]
1327 fn noop_embedder_returns_not_implemented_and_is_noop() {
1328 let e = NoopEmbedder;
1329 assert!(e.is_noop());
1330 assert_eq!(e.dim(), 0);
1331 assert_eq!(e.model_id(), "noop");
1332 assert!(matches!(
1333 e.embed("hello").unwrap_err(),
1334 EmbedderError::NotImplemented
1335 ));
1336 }
1337
1338 #[test]
1339 fn stub_embedder_is_deterministic() {
1340 let e = StubEmbedder::new();
1341 let a = e.embed("hello rust").expect("embed a");
1342 let b = e.embed("hello rust").expect("embed b");
1343 let c = e.embed("hello RUST").expect("embed c");
1344 assert_eq!(a, b, "same input -> same output");
1345 assert_eq!(
1346 a, c,
1347 "lowercasing means case differences collapse to the same vector"
1348 );
1349 assert_eq!(a.len(), 8);
1350 let norm = a.iter().map(|v| v * v).sum::<f32>().sqrt();
1352 assert!((norm - 1.0).abs() < 1e-5, "expected unit norm, got {norm}");
1353 }
1354
1355 #[test]
1356 fn stub_embedder_distinguishes_disjoint_inputs() {
1357 let e = StubEmbedder::new();
1358 let a = e.embed("foo bar").expect("a");
1359 let b = e.embed("qux quux").expect("b");
1360 let sim = cosine_similarity(&a, &b);
1361 assert!(
1364 sim < 0.99,
1365 "disjoint inputs should not be near-identical: {sim}"
1366 );
1367 }
1368
1369 #[test]
1370 fn stub_embedder_handles_empty_input() {
1371 let e = StubEmbedder::new();
1372 let v = e.embed("").expect("empty embed");
1373 assert_eq!(v.len(), 8);
1374 assert!(v.iter().all(|&x| x == 0.0));
1378 }
1379
1380 #[test]
1381 fn cosine_similarity_handles_edge_cases() {
1382 let a = [1.0f32, 0.0, 0.0];
1384 assert!((cosine_similarity(&a, &a) - 1.0).abs() < 1e-6);
1385
1386 let b = [0.0f32, 1.0, 0.0];
1388 assert!((cosine_similarity(&a, &b)).abs() < 1e-6);
1389
1390 let c = [-1.0f32, 0.0, 0.0];
1392 assert!((cosine_similarity(&a, &c) + 1.0).abs() < 1e-6);
1393
1394 assert_eq!(cosine_similarity(&[], &a), 0.0);
1396 assert_eq!(cosine_similarity(&a, &[0.0]), 0.0);
1397
1398 let zeros = [0.0f32, 0.0, 0.0];
1400 assert_eq!(cosine_similarity(&zeros, &a), 0.0);
1401 }
1402
1403 #[test]
1404 fn cosine_similarity_is_symmetric() {
1405 let a = [0.6f32, 0.8, 0.0];
1406 let b = [0.0f32, 1.0, 0.0];
1407 let ab = cosine_similarity(&a, &b);
1408 let ba = cosine_similarity(&b, &a);
1409 assert!((ab - ba).abs() < 1e-6);
1410 assert!((ab - 0.8).abs() < 1e-5);
1412 }
1413
1414 #[test]
1415 fn encode_decode_embedding_round_trip() {
1416 let vec = vec![0.1f32, -0.2, 3.125, -0.000_001, 42.0];
1417 let blob = encode_embedding(&vec);
1418 assert_eq!(blob.len(), vec.len() * 4);
1419 let back = decode_embedding(&blob, Some(vec.len())).expect("decode");
1420 assert_eq!(back.len(), vec.len());
1421 for (orig, got) in vec.iter().zip(back.iter()) {
1422 assert!(
1423 (orig - got).abs() < 1e-7,
1424 "f32 round-trip should be bit-exact"
1425 );
1426 }
1427 }
1428
1429 #[test]
1430 fn decode_embedding_rejects_unaligned_blob() {
1431 let bad = [0u8, 1, 2]; let err = decode_embedding(&bad, None).unwrap_err();
1433 assert!(err.to_string().contains("not a multiple of 4"));
1434 }
1435
1436 #[test]
1437 fn decode_embedding_rejects_dim_mismatch() {
1438 let vec = vec![1.0f32, 2.0, 3.0];
1439 let blob = encode_embedding(&vec);
1440 let err = decode_embedding(&blob, Some(5)).unwrap_err();
1441 assert!(err.to_string().contains("does not match expected"));
1442 }
1443
1444 #[test]
1448 fn stub_reranker_returns_doc_order_scores() {
1449 let r = StubReranker;
1450 let query = "rust async tokio";
1451 let docs = &["rust async tokio", "python django", "rust only"];
1452 let scores = r.rerank(query, docs).expect("rerank should succeed");
1453 assert_eq!(scores.len(), docs.len(), "one score per document");
1454 for (i, &s) in scores.iter().enumerate() {
1456 assert!(s > 0.0 && s < 1.0, "score[{i}] must be in (0,1), got {s}");
1457 }
1458 }
1459
1460 #[test]
1462 fn stub_reranker_higher_overlap_scores_higher() {
1463 let r = StubReranker;
1464 let query = "rust async tokio";
1465 let docs = &["rust async tokio runtime", "rust only", "python django"];
1469 let scores = r.rerank(query, docs).expect("rerank");
1470 assert!(
1471 scores[0] > scores[1],
1472 "3-token overlap must beat 1-token overlap: {} vs {}",
1473 scores[0],
1474 scores[1]
1475 );
1476 assert!(
1477 scores[1] > scores[2],
1478 "1-token overlap must beat 0-token overlap: {} vs {}",
1479 scores[1],
1480 scores[2]
1481 );
1482 }
1483
1484 #[test]
1486 fn stub_reranker_model_id() {
1487 let r = StubReranker;
1488 assert_eq!(r.model_id(), "stub-reranker");
1489 }
1490
1491 #[test]
1493 fn stub_reranker_empty_query_returns_floor() {
1494 let r = StubReranker;
1495 let docs = &["anything here", "another doc"];
1496 let scores = r.rerank("", docs).expect("rerank");
1497 for &s in &scores {
1498 assert!(
1499 (s - 0.05).abs() < 1e-6,
1500 "empty query must yield 0.05, got {s}"
1501 );
1502 }
1503 }
1504
1505 #[test]
1510 fn embed_batch_matches_per_row() {
1511 let e = StubEmbedder::new();
1512 let texts = ["foo bar", "qux", "hello world"];
1513 let batch = e.embed_batch(&texts).expect("embed_batch should succeed");
1514 assert_eq!(batch.len(), texts.len());
1515 for (i, text) in texts.iter().enumerate() {
1516 let single = e.embed(text).expect("per-row embed should succeed");
1517 assert_eq!(
1518 batch[i], single,
1519 "embed_batch[{i}] must match per-row embed for {text:?}"
1520 );
1521 }
1522 }
1523
1524 #[test]
1526 fn embed_batch_empty_is_empty() {
1527 let e = StubEmbedder::new();
1528 let result = e
1529 .embed_batch(&[])
1530 .expect("empty embed_batch should succeed");
1531 assert!(result.is_empty(), "expected empty Vec, got {result:?}");
1532 }
1533
1534 #[test]
1536 fn embed_batch_length_matches_input() {
1537 let e = StubEmbedder::new();
1538 let texts: Vec<&str> = vec!["alpha", "beta", "gamma", "delta", "epsilon"];
1539 let batch = e.embed_batch(&texts).expect("embed_batch should succeed");
1540 assert_eq!(batch.len(), texts.len(), "output len must equal input len");
1541 for (i, v) in batch.iter().enumerate() {
1542 assert_eq!(
1543 v.len(),
1544 e.dim(),
1545 "vector[{i}] len {} != dim {}",
1546 v.len(),
1547 e.dim()
1548 );
1549 }
1550 }
1551
1552 #[cfg(not(feature = "embeddings"))]
1559 #[test]
1560 fn open_default_embedder_returns_noop_on_default_build() {
1561 let e = open_default_embedder();
1562 assert!(e.is_noop());
1563 assert_eq!(e.dim(), 0);
1564 assert!(matches!(
1565 e.embed("anything").unwrap_err(),
1566 EmbedderError::NotImplemented
1567 ));
1568 }
1569
1570 #[test]
1577 fn env_disables_embedder_recognizes_off_values() {
1578 let lock = crate::user_brain::test_env_lock()
1579 .lock()
1580 .unwrap_or_else(|p| p.into_inner());
1581 let prev = std::env::var("KIMETSU_BRAIN_EMBEDDER").ok();
1582 for value in ["noop", "off", "NONE", "0", "false", "no"] {
1583 unsafe {
1585 std::env::set_var("KIMETSU_BRAIN_EMBEDDER", value);
1586 }
1587 assert!(env_disables_embedder(), "value {value:?} must disable");
1588 }
1589 for value in ["", "default", "bge-small", "bge-m3", "jina-code"] {
1590 unsafe {
1591 std::env::set_var("KIMETSU_BRAIN_EMBEDDER", value);
1592 }
1593 assert!(!env_disables_embedder(), "value {value:?} must NOT disable");
1594 }
1595 unsafe {
1597 match prev {
1598 Some(v) => std::env::set_var("KIMETSU_BRAIN_EMBEDDER", v),
1599 None => std::env::remove_var("KIMETSU_BRAIN_EMBEDDER"),
1600 }
1601 }
1602 drop(lock);
1603 }
1604
1605 #[test]
1609 fn w3_embedder_enabled_for_config_false_when_env_unset() {
1610 let lock = crate::user_brain::test_env_lock()
1611 .lock()
1612 .unwrap_or_else(|p| p.into_inner());
1613 let prev = std::env::var("KIMETSU_BRAIN_EMBEDDER").ok();
1614 unsafe {
1615 std::env::remove_var("KIMETSU_BRAIN_EMBEDDER");
1616 }
1617 assert!(
1619 !embedder_enabled_for_config(false),
1620 "config=false + env unset must be disabled"
1621 );
1622 assert!(
1624 embedder_enabled_for_config(true),
1625 "config=true + env unset must be enabled"
1626 );
1627 unsafe {
1628 match prev {
1629 Some(v) => std::env::set_var("KIMETSU_BRAIN_EMBEDDER", v),
1630 None => std::env::remove_var("KIMETSU_BRAIN_EMBEDDER"),
1631 }
1632 }
1633 drop(lock);
1634 }
1635
1636 #[test]
1638 fn w3_embedder_env_disable_overrides_config_true() {
1639 let lock = crate::user_brain::test_env_lock()
1640 .lock()
1641 .unwrap_or_else(|p| p.into_inner());
1642 let prev = std::env::var("KIMETSU_BRAIN_EMBEDDER").ok();
1643 unsafe {
1644 std::env::set_var("KIMETSU_BRAIN_EMBEDDER", "noop");
1645 }
1646 assert!(
1647 !embedder_enabled_for_config(true),
1648 "KIMETSU_BRAIN_EMBEDDER=noop must override config=true"
1649 );
1650 unsafe {
1651 match prev {
1652 Some(v) => std::env::set_var("KIMETSU_BRAIN_EMBEDDER", v),
1653 None => std::env::remove_var("KIMETSU_BRAIN_EMBEDDER"),
1654 }
1655 }
1656 drop(lock);
1657 }
1658
1659 #[test]
1662 fn w3_embedder_env_model_id_overrides_config_false() {
1663 let lock = crate::user_brain::test_env_lock()
1664 .lock()
1665 .unwrap_or_else(|p| p.into_inner());
1666 let prev = std::env::var("KIMETSU_BRAIN_EMBEDDER").ok();
1667 unsafe {
1668 std::env::set_var("KIMETSU_BRAIN_EMBEDDER", "bge-m3");
1669 }
1670 assert!(
1671 embedder_enabled_for_config(false),
1672 "real model id in env must override config=false → enabled"
1673 );
1674 unsafe {
1675 match prev {
1676 Some(v) => std::env::set_var("KIMETSU_BRAIN_EMBEDDER", v),
1677 None => std::env::remove_var("KIMETSU_BRAIN_EMBEDDER"),
1678 }
1679 }
1680 drop(lock);
1681 }
1682
1683 #[test]
1687 fn pick_builtin_model_from_env_handles_aliases() {
1688 let lock = crate::user_brain::test_env_lock()
1689 .lock()
1690 .unwrap_or_else(|p| p.into_inner());
1691 let prev = std::env::var("KIMETSU_BRAIN_EMBEDDER").ok();
1692 let cases = [
1693 ("", "bge-small-en-v1.5"),
1694 ("default", "bge-small-en-v1.5"),
1695 ("bge-small", "bge-small-en-v1.5"),
1696 ("BGE-SMALL-EN-V1.5", "bge-small-en-v1.5"),
1697 ("bge-m3", "bge-m3"),
1698 ("M3", "bge-m3"),
1699 ("jina-code", "jina-v2-base-code"),
1700 ("jina-v2-base-code", "jina-v2-base-code"),
1701 ("jina-embeddings-v2-base-code", "jina-v2-base-code"),
1702 ("totally-made-up", "bge-small-en-v1.5"),
1704 ];
1705 for (input, expected) in cases {
1706 unsafe {
1708 std::env::set_var("KIMETSU_BRAIN_EMBEDDER", input);
1709 }
1710 assert_eq!(
1711 pick_builtin_model_from_env(),
1712 expected,
1713 "input {input:?} -> expected {expected}"
1714 );
1715 }
1716 unsafe {
1717 match prev {
1718 Some(v) => std::env::set_var("KIMETSU_BRAIN_EMBEDDER", v),
1719 None => std::env::remove_var("KIMETSU_BRAIN_EMBEDDER"),
1720 }
1721 }
1722 drop(lock);
1723 }
1724}
1725
1726#[cfg(test)]
1727mod correction_race_tests {
1728 use super::*;
1729 #[test]
1730 fn slow_embedding_cannot_overwrite_a_newer_correction() {
1731 let dir = tempfile::tempdir().unwrap();
1732 let db = dir.path().join("brain.db");
1733 let writer = rusqlite::Connection::open(&db).unwrap();
1734 crate::schema::initialize(&writer).unwrap();
1735 let accepted = kimetsu_core::event::Event::new(
1736 kimetsu_core::ids::RunId::new(),
1737 "memory.accepted",
1738 serde_json::json!({"memory_id":"m","scope":"project","kind":"fact","text":"claim A"}),
1739 );
1740 crate::projector::apply_events(&writer, &[accepted]).unwrap();
1741 let (started_tx, started_rx) = std::sync::mpsc::channel();
1742 let (resume_tx, resume_rx) = std::sync::mpsc::channel();
1743 struct Blocking {
1744 started: std::sync::mpsc::Sender<()>,
1745 resume: std::sync::Mutex<std::sync::mpsc::Receiver<()>>,
1746 }
1747 impl Embedder for Blocking {
1748 fn embed(&self, _: &str) -> Result<Vec<f32>, EmbedderError> {
1749 self.started.send(()).unwrap();
1750 self.resume.lock().unwrap().recv().unwrap();
1751 Ok(vec![1.0, 0.0])
1752 }
1753 fn model_id(&self) -> &str {
1754 "stub"
1755 }
1756 fn dim(&self) -> usize {
1757 2
1758 }
1759 }
1760 let pending = std::thread::spawn(move || {
1761 let conn = rusqlite::Connection::open(db).unwrap();
1762 embed_and_persist(
1763 &conn,
1764 "m",
1765 "claim A",
1766 &Blocking {
1767 started: started_tx,
1768 resume: std::sync::Mutex::new(resume_rx),
1769 },
1770 )
1771 .unwrap()
1772 });
1773 started_rx.recv().unwrap();
1774 let correction = kimetsu_core::event::Event::new(
1775 kimetsu_core::ids::RunId::new(),
1776 "memory.corrected",
1777 serde_json::json!({"memory_id":"m","text":"claim B"}),
1778 );
1779 crate::projector::apply_events(&writer, &[correction]).unwrap();
1780 writer
1781 .execute(
1782 "UPDATE memories SET embedding=?1,embedding_model='stub' WHERE memory_id='m'",
1783 rusqlite::params![encode_embedding(&[0.0, 1.0])],
1784 )
1785 .unwrap();
1786 resume_tx.send(()).unwrap();
1787 assert!(
1788 pending.join().unwrap().is_none(),
1789 "stale computation must not be published"
1790 );
1791 let blob: Vec<u8> = writer
1792 .query_row(
1793 "SELECT embedding FROM memories WHERE memory_id='m'",
1794 [],
1795 |r| r.get(0),
1796 )
1797 .unwrap();
1798 assert_eq!(decode_embedding(&blob, Some(2)).unwrap(), vec![0.0, 1.0]);
1799 }
1800}