Skip to main content

kimetsu_brain/
embeddings.rs

1//! Embeddings + hybrid retrieval scaffolding (v0.4.2).
2//!
3//! The broker is FTS-only through v0.4.1: right answers with wrong
4//! words get missed. v0.4.2 lays the infrastructure for hybrid
5//! retrieval (lexical + semantic) without binding to a specific
6//! embedder. v0.4.3 wires fastembed-rs as the production default.
7//!
8//! Layers introduced here:
9//!   1. The [`Embedder`] trait — anything that can map a text to a
10//!      fixed-dimension float vector plus identify the model that
11//!      produced it.
12//!   2. [`NoopEmbedder`] — production default when no real embedder
13//!      is wired. `embed()` errors with `NotImplemented`; the write
14//!      path treats that as "store NULL" and the retrieval path
15//!      treats it as "skip the cosine blend, FTS only".
16//!   3. [`StubEmbedder`] — deterministic, dependency-free, test-only
17//!      pseudo-embedder. Lets us exercise the hybrid scoring path
18//!      end-to-end without depending on fastembed-rs or downloading
19//!      a model in CI.
20//!   4. [`cosine_similarity`] + BLOB codec helpers so the brain.db
21//!      schema can store embeddings as little-endian `f32` blobs.
22//!
23//! Wire compatibility:
24//!   * Embeddings are nullable. Pre-v0.4.2 rows have NULL embedding
25//!     + NULL embedding_model. The retrieval blender treats them as
26//!       "lexical-only" — they still score via FTS, they just don't
27//!       contribute to the cosine term.
28//!   * The `embedding_model` column carries an opaque string id
29//!     ("bge-small-en-v1.5", "stub-d8", etc.). Queries blend only
30//!     when the query's embedder id matches the row's stored id,
31//!     so mixing models inside one brain.db is safe (rows with a
32//!     different model fall back to lexical-only).
33//!
34//! Scoring (added in v0.4.2):
35//!   `final_relevance = (1 - alpha) * lexical + alpha * cosine`
36//!   where `alpha = brain.broker.hybrid_alpha` (defaulted to 0.5,
37//!   tuned later in v0.4.3 after live measurements). When no
38//!   cosine signal exists, `alpha = 0` effectively — i.e. pure
39//!   lexical. See [`context::memory_candidates`] for the wiring.
40
41use kimetsu_core::KimetsuResult;
42
43/// Default blend factor for hybrid scoring. 0.0 = pure lexical
44/// (v0.4.1 behavior), 1.0 = pure cosine. v0.4.2 ships 0.5 as a
45/// starting point; live data in v0.4.3 will tune it.
46pub const DEFAULT_HYBRID_ALPHA: f32 = 0.5;
47
48/// Embedder trait. Every implementation maps a text to a fixed
49/// dimension `Vec<f32>` and identifies itself via a stable model id.
50///
51/// `Send + Sync` so a single embedder instance can be shared across
52/// the chat REPL's threads (drainer, REPL, hook runner) without
53/// requiring per-call locking.
54pub trait Embedder: Send + Sync {
55    /// Compute an embedding for `text`. The returned vector MUST have
56    /// length == `self.dim()`. Implementations may normalize.
57    fn embed(&self, text: &str) -> Result<Vec<f32>, EmbedderError>;
58
59    /// Stable identifier for the model used. Stored alongside each
60    /// embedding so retrieval can detect cross-model mismatches and
61    /// skip the cosine blend rather than comparing apples-to-oranges.
62    fn model_id(&self) -> &str;
63
64    /// Embedding dimension. Used by the BLOB codec to validate
65    /// stored vectors against the active model on retrieval.
66    fn dim(&self) -> usize;
67
68    /// Convenience: true when this embedder is the production no-op.
69    /// Callers can short-circuit the write/retrieval cosine path
70    /// instead of allocating a vec only to discard it.
71    fn is_noop(&self) -> bool {
72        false
73    }
74
75    /// Embed many texts in as few backend calls as possible. The
76    /// production fastembed backend runs them through ONNX in batched
77    /// tensors (~10-40x faster than calling `embed` per text). The
78    /// returned Vec is 1:1 with `texts` (same order, same length).
79    /// The default impl loops `embed`, so non-batching embedders work
80    /// unchanged. On any failure the whole batch errors — callers that
81    /// want per-row resilience should fall back to `embed` per text.
82    fn embed_batch(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>, EmbedderError> {
83        texts.iter().map(|t| self.embed(t)).collect()
84    }
85}
86
87/// Implement `Embedder` for `Box<dyn Embedder>` so callers can hold
88/// an owned trait object without ceremony.
89impl 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/// Failure modes for an embedder.
108#[derive(Debug, Clone)]
109pub enum EmbedderError {
110    /// The embedder is intentionally a no-op — no embeddings will be
111    /// produced. Callers should fall back to a NULL embedding /
112    /// lexical-only retrieval.
113    NotImplemented,
114    /// Model failed to load. v0.4.3+ — e.g. the fastembed backend
115    /// can't download the model.
116    LoadFailed(String),
117    /// Inference failed (rare).
118    EmbedFailed(String),
119    /// Dimension mismatch between the embedder and a stored row.
120    /// The retrieval path skips this row's cosine contribution.
121    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/// Production default when no real embedder is configured.
140/// `embed()` returns `Err(NotImplemented)`; callers interpret that
141/// as "store NULL" / "skip the cosine blend".
142#[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/// Deterministic, dependency-free pseudo-embedder used in tests.
168///
169/// Hashes each word into a fixed-dim bucket (count-of-hash-buckets),
170/// then L2-normalizes the resulting vector. NOT semantic — texts
171/// that share words will be close; texts that don't share words
172/// will be far. Good enough to exercise the hybrid-scoring code
173/// path without depending on a real ML model.
174///
175/// Default dimension is 8 (small enough to keep tests fast).
176#[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            // Cheap, stable hash. Don't use DefaultHasher — its seed
204            // randomizes across processes and tests would be flaky.
205            // FNV-1a over the lowercased UTF-8 bytes is plenty.
206            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        // L2-normalize so cosine similarity reduces to a dot product.
216        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
234// ── Reranker trait + implementations ───────────────────────────────────────
235
236/// v1.0.0: cross-encoder reranker — scores (query, document) pairs jointly.
237/// Returns one sigmoid-normalized score in (0,1) per document, in DOCUMENT
238/// ORDER (not sorted). Implementations must be Send + Sync (the daemon
239/// shares one across worker threads).
240pub trait Reranker: Send + Sync {
241    fn rerank(&self, query: &str, documents: &[&str]) -> Result<Vec<f32>, EmbedderError>;
242    fn model_id(&self) -> &str;
243}
244
245/// Deterministic, dependency-free stub reranker for tests.
246///
247/// Scores a (query, document) pair by lowercase-tokenizing both on
248/// non-alphanumeric characters and computing token-overlap:
249///
250///   score = 0.05 + 0.9 * (|intersection| / |query_tokens|)
251///           clamped to (0,1)
252///
253/// This means no doc ever scores exactly 0 or 1, but docs with more query
254/// words in common always score higher. Safe to use in any test that doesn't
255/// need a real model.
256pub 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
290/// v1.0.0: open a reranker by curated or user-defined id.
291///
292/// Resolution:
293///   1. `"off"`/`"none"`/`"noop"`/empty → `None` (disable).
294///   2. One of the 4 curated ids → `FastembedReranker::try_open` (builtin ONNX).
295///   3. Known alias (`jina-reranker-v1-tiny-en`, `ms-marco-tinybert-l-2-v2`,
296///      `ms-marco-minilm-l-4-v2`) or any id containing `/` → user-defined ONNX
297///      via HuggingFace Hub download.
298///   4. Anything else → fallback to the default curated jina-turbo.
299///
300/// Lean builds (no `embeddings` feature) always return `None`.
301pub 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        // Curated ids → builtin path.
309        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        // User-defined alias ids → HF download path.
316        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        // Unknown → fallback to default curated turbo.
333        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
344/// Open the production-default embedder.
345///
346/// Resolution (v0.4.3):
347///   1. `KIMETSU_BRAIN_EMBEDDER=noop|off|none` → always `NoopEmbedder`,
348///      regardless of Cargo features. Useful for CI, hooks, and
349///      transient subprocesses that shouldn't pay the model-load
350///      cost.
351///   2. Cargo feature `embeddings` enabled →
352///      [`fastembed_backend::open_cached`] returns a process-wide
353///      cached [`FastembedEmbedder`] for the model picked by
354///      [`pick_builtin_model_from_env`] (default `bge-small-en-v1.5`,
355///      `bge-m3` or `jina-v2-base-code` opt-in via env). On model
356///      load failure (network, disk, ort runtime missing) we log
357///      and fall through to Noop so the brain stays usable on FTS
358///      alone.
359///   3. Cargo feature `embeddings` disabled → `NoopEmbedder`,
360///      identical to v0.4.2 build.
361///
362/// The returned trait object is borrowed from a process-static
363/// `OnceLock`; production callers get model-load cost paid exactly
364/// once over the process lifetime. Tests that need a different
365/// embedder must use [`crate::context::retrieve_context_with_embedder`]
366/// with an explicit [`StubEmbedder`] (or any other [`Embedder`])
367/// instead of going through this function.
368pub 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
394/// W3.1: config-aware embedder resolver.
395///
396/// Returns the shared real embedder when embeddings are enabled, or a
397/// static [`NoopEmbedder`] when they are disabled — so a project with
398/// `[embedder] enabled = false` gets FTS-only retrieval and writes no
399/// vectors, durably, without relying on the `KIMETSU_BRAIN_EMBEDDER`
400/// env var.
401///
402/// Precedence mirrors [`embedder_enabled_for_config`]:
403///   1. `KIMETSU_BRAIN_EMBEDDER` env disable value → Noop.
404///   2. `KIMETSU_BRAIN_EMBEDDER` real model id → real embedder.
405///   3. Env unset → `config_enabled` governs.
406pub 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
414/// v0.8: open a FRESH (uncached) embedder for an explicit built-in
415/// model id. Unlike [`open_default_embedder`], this bypasses the
416/// process-static cache AND the env/override resolution — the caller
417/// asked for a *specific* model (e.g. `kimetsu brain model set` and the
418/// MCP `model_set` reindex, which must re-embed with the newly-chosen
419/// model even though the running process may have a different default
420/// embedder cached). Returns [`NoopEmbedder`] on the lean build or if
421/// the model fails to load.
422pub 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
442/// v0.4.3: env-driven kill switch. Truthy values (1/true/yes/on)
443/// force-disable the embedder for this process; "noop", "off",
444/// "none" do the same. Anything else (or unset) leaves the
445/// `embeddings` feature in control.
446fn 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
453/// W3.1: config-aware enabled check. Resolution precedence:
454///   1. `KIMETSU_BRAIN_EMBEDDER` env is set to a disable value → false.
455///   2. `KIMETSU_BRAIN_EMBEDDER` env is set to a real model id → true
456///      (explicit model override = caller wants embeddings on).
457///   3. Env is unset → `config_enabled` governs.
458///
459/// Keep the env-only `env_disables_embedder()` working for back-compat
460/// callers (the `OnceLock` path).
461pub fn embedder_enabled_for_config(config_enabled: bool) -> bool {
462    // Precedence: env override > config > default.
463    match std::env::var("KIMETSU_BRAIN_EMBEDDER") {
464        Ok(raw) => {
465            let v = raw.trim().to_ascii_lowercase();
466            if v.is_empty() {
467                // Empty string — treat as unset, fall through to config.
468                config_enabled
469            } else if is_disable_value(&v) {
470                // Explicit disable in env wins.
471                false
472            } else {
473                // A real model id in env = caller wants embeddings on.
474                true
475            }
476        }
477        // Env unset → config governs.
478        Err(_) => config_enabled,
479    }
480}
481
482fn is_disable_value(v: &str) -> bool {
483    matches!(v, "noop" | "off" | "none" | "0" | "false" | "no")
484}
485
486/// v0.8: curated built-in embedding models, surfaced by
487/// `kimetsu brain model list` and `kimetsu_brain_model_list`. Tuple
488/// = (stable id, vector dimension, human blurb). This is the single
489/// source of truth for the selectable set; the fastembed backend
490/// maps these ids → `EmbeddingModel` in `try_open`.
491pub 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
501/// v0.8: process-global embedder override recorded by
502/// [`apply_embedder_selection`]. Brain-internal callers
503/// ([`pick_builtin_model_from_env`], the fastembed backend, reindex)
504/// have no `ProjectConfig` in hand, so the CLI/MCP layer stashes the
505/// config-selected id here once, early, before any embed happens.
506static EMBEDDER_OVERRIDE: std::sync::OnceLock<String> = std::sync::OnceLock::new();
507
508/// v0.8: record the config-provided embedder id so brain-internal
509/// callers resolve it when `KIMETSU_BRAIN_EMBEDDER` is unset (the env
510/// var always wins). Call once, early, before the first retrieval or
511/// embed — after the embedder `OnceLock` initializes this has no
512/// effect. No-op for `None`/empty, and only the first call sticks.
513pub 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
522/// v0.8: map any accepted alias (env value, config value, built-in
523/// id) to a stable built-in id. Unknown values warn and fall back to
524/// the lean English default. Disable values map to the default too;
525/// the *actual* disable is handled separately by
526/// [`env_disables_embedder`].
527fn 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
543/// v0.8: resolve the active built-in model id. Precedence:
544///   1. `KIMETSU_BRAIN_EMBEDDER` env (unless it's a disable value)
545///   2. the explicit `config_embedder` arg, else the override set by
546///      [`apply_embedder_selection`]
547///   3. `bge-small-en-v1.5` default
548pub 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        // empty / disable values fall through: the model *id* still
555        // resolves from config/default even when retrieval is off.
556    }
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
569/// v0.4.3: pick which builtin model to load from the env, returning
570/// a stable identifier. Used both by the fastembed backend (to map
571/// id → `EmbeddingModel`) and by `kimetsu brain reindex` (to label
572/// new rows with the right `embedding_model`).
573///
574/// Resolution:
575///   * unset / "" / "default" / "bge-small" / "bge-small-en-v1.5"
576///     → `"bge-small-en-v1.5"` (384 dim, ~67 MB int8, English)
577///   * "bge-m3"
578///     → `"bge-m3"` (1024 dim, ~600 MB int8, multilingual)
579///   * "jina-code" / "jina-v2-base-code" /
580///     "jina-embeddings-v2-base-code"
581///     → `"jina-v2-base-code"` (768 dim, ~165 MB int8, English +
582///     code-tuned)
583///   * anything else falls back to bge-small with a warning.
584pub fn pick_builtin_model_from_env() -> &'static str {
585    // v0.8: env > config-override (set via `apply_embedder_selection`)
586    // > default. Kept as a named entry point for the fastembed backend
587    // and `reindex`, which have no `ProjectConfig` to pass.
588    resolve_embedder_id(None)
589}
590
591// v0.4.3: real fastembed-backed embedder. Lives behind the
592// `embeddings` Cargo feature so the default build skips the
593// ~50-transitive-crate dep tree (ONNX runtime, tokenizers, etc).
594#[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    // ── HF Hub download helper (user-defined ONNX rerankers) ─────────────────
603
604    /// Alias table: lowercased stable id → HuggingFace repo id.
605    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    /// Download files for a user-defined reranker from HuggingFace Hub.
615    ///
616    /// Returns `(onnx_bytes, tokenizer_files)` or an `EmbedderError::LoadFailed`.
617    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            // Raw HF repo id passed directly.
627            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        // Helper: download a required file or return LoadFailed.
639        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        // Try `onnx/model.onnx` first, then `model.onnx` at root.
654        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    /// fastembed-backed embedder. Wraps the ONNX runtime in a
674    /// `Mutex` because `TextEmbedding::embed` takes `&mut self`.
675    /// The lock window is short (one inference per call); the
676    /// chat REPL's threads serialize cleanly through it.
677    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                // bge-small-en-v1.5 is the default + fallback.
693                _ => (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    /// Shared handle. `open_default_embedder` boxes this into a
766    /// `dyn Embedder` and stashes it in a process-static `OnceLock`,
767    /// so we only call into ONNX once per process.
768    #[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    /// fastembed-backed cross-encoder reranker.
787    ///
788    /// Wraps `TextRerank` in a `Mutex` because `rerank` takes `&mut self`.
789    /// The lock window is short (one rerank call per request); the daemon's
790    /// worker threads serialize cleanly through it.
791    ///
792    /// `model_id` is an owned `String` so both curated (`&'static str`
793    /// originates from a match arm) and user-defined (alias / HF repo id)
794    /// rerankers can share the same struct.
795    pub struct FastembedReranker {
796        model_id: String,
797        engine: Mutex<TextRerank>,
798    }
799
800    impl FastembedReranker {
801        /// Map a curated id to the corresponding `RerankerModel` variant and
802        /// initialize a `TextRerank` engine. Unknown ids fall back to the
803        /// jina-reranker-v1-turbo-en default.
804        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                // jina-reranker-v1-turbo-en is the default + fallback.
813                _ => (
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        /// Load a user-defined ONNX reranker by alias or raw HF repo id.
828        ///
829        /// Downloads the ONNX and tokenizer files from HuggingFace Hub (cached
830        /// locally) and constructs a `TextRerank` via
831        /// `try_new_from_user_defined`. The `model_id` stored in the struct is
832        /// the normalized alias (e.g. `"jina-reranker-v1-tiny-en"`) or the raw
833        /// repo id, lower-cased, so it is stable across calls.
834        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            // Normalise the stored id to lower-case alias or repo id.
848            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            // Pass documents as Vec<&str>; fastembed returns results sorted by
866            // score descending. Use .index to map back to document order.
867            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                    // Apply sigmoid to normalize logit → (0,1).
875                    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    /// Open (or return the cached) fastembed embedder for the model
887    /// picked by `KIMETSU_BRAIN_EMBEDDER`. Errors here propagate up
888    /// to `open_default_embedder`, which falls back to Noop +
889    /// prints a one-line warning.
890    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
903// --------- math helpers ---------
904
905/// Cosine similarity between two vectors. Returns 0.0 when either
906/// vector is empty or all-zeros. Does NOT assume the vectors are
907/// pre-normalized — divides by both norms.
908pub 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
926// --------- write-path helper ---------
927
928/// Compute the embedding for `text` and persist it onto an existing
929/// `memories.memory_id` row.
930///
931/// No-op when the embedder is intentionally a [`NoopEmbedder`] (it
932/// returns [`EmbedderError::NotImplemented`] which we silently
933/// swallow — the column stays NULL, retrieval falls back to
934/// FTS-only for the row, exact v0.4.1 behavior).
935///
936/// For other embedder errors we surface them up. The caller
937/// (`add_memory`, `add_user_memory`) can decide whether to fail the
938/// whole insert or log+continue — today they propagate.
939///
940/// Returns the computed embedding vector (so callers can reuse it
941/// for conflict detection without re-embedding). Returns `None` on
942/// Noop or NotImplemented — the same cases where the column stays NULL.
943pub 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        // NotImplemented is the contract for "skip silently". Treat
955        // any embedder that signals it the same way as NoopEmbedder.
956        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    // Tier-3: keep the warm usearch index current at add time. For in-memory
975    // DBs there is no cached handle — the rebuild-on-query path picks the row
976    // up, so we safely skip. Best-effort: an index failure must not abort a
977    // successful memory write.
978    #[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
1000// --------- BLOB codec ---------
1001//
1002// Embeddings are stored as little-endian f32 BLOBs. The encoder
1003// fixes byte order so brain.db files move between architectures.
1004// The decoder is strict: it returns Err if the byte length isn't a
1005// multiple of 4, or if the resulting dim doesn't match expectations.
1006
1007/// Serialize a float vector to little-endian bytes for storage.
1008pub 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
1016/// Decode a BLOB back into a float vector. Optionally validates the
1017/// expected dimension; pass `None` to accept any length.
1018pub 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        // disable values resolve to the lean default (the kill-switch
1053        // is handled separately by env_disables_embedder).
1054        assert_eq!(map_builtin_id("noop"), "bge-small-en-v1.5");
1055        // unknown -> warn + default.
1056        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        // Every advertised id must map back to itself.
1062        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        // Env mutation in parallel tests is racy + unsafe in edition
1070        // 2024, so we only assert the config/default paths when the env
1071        // var is genuinely absent. The env-wins path is exercised
1072        // manually (see the plan's verification section).
1073        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        // unknown config value -> default.
1079        assert_eq!(resolve_embedder_id(Some("nope")), "bge-small-en-v1.5");
1080        // None + no override stored in this test binary -> default.
1081        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        // L2-normalized: norm == 1 within float tolerance.
1109        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        // Disjoint word sets *can* still collide in the 8-bucket
1120        // hash, but on average should be low. Sanity bound: not 1.0.
1121        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        // All zeros: cosine similarity with self is 0 (we guard against
1133        // division by zero), which is exactly the behavior the
1134        // retrieval blender wants for content-free queries.
1135        assert!(v.iter().all(|&x| x == 0.0));
1136    }
1137
1138    #[test]
1139    fn cosine_similarity_handles_edge_cases() {
1140        // Identical normalized vectors -> 1.0.
1141        let a = [1.0f32, 0.0, 0.0];
1142        assert!((cosine_similarity(&a, &a) - 1.0).abs() < 1e-6);
1143
1144        // Orthogonal -> 0.0.
1145        let b = [0.0f32, 1.0, 0.0];
1146        assert!((cosine_similarity(&a, &b)).abs() < 1e-6);
1147
1148        // Anti-parallel -> -1.0.
1149        let c = [-1.0f32, 0.0, 0.0];
1150        assert!((cosine_similarity(&a, &c) + 1.0).abs() < 1e-6);
1151
1152        // Empty / mismatched dim -> 0.0 by contract.
1153        assert_eq!(cosine_similarity(&[], &a), 0.0);
1154        assert_eq!(cosine_similarity(&a, &[0.0]), 0.0);
1155
1156        // Zero norm -> 0.0 by contract (don't divide by zero).
1157        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        // Dot is 0.8, |a|=1, |b|=1 -> sim = 0.8.
1169        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]; // 3 bytes - not a multiple of 4
1190        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    // ── v1.0.0 StubReranker tests ─────────────────────────────────────────────
1203
1204    /// StubReranker returns one score per document in document order.
1205    #[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        // All scores in (0,1).
1213        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    /// Higher token overlap ⇒ higher score.
1219    #[test]
1220    fn stub_reranker_higher_overlap_scores_higher() {
1221        let r = StubReranker;
1222        let query = "rust async tokio";
1223        // doc0: 3/3 tokens shared → highest
1224        // doc1: 1/3 tokens shared → middle
1225        // doc2: 0/3 tokens shared → lowest (0.05)
1226        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    /// model_id is the stub constant.
1243    #[test]
1244    fn stub_reranker_model_id() {
1245        let r = StubReranker;
1246        assert_eq!(r.model_id(), "stub-reranker");
1247    }
1248
1249    /// Empty query → all docs get the floor score 0.05.
1250    #[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    // ── embed_batch contract tests (Stub-backed, no model download) ───────────
1264
1265    /// `embed_batch` returns the same vectors, in the same order, as
1266    /// calling `embed` on each text individually.
1267    #[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    /// `embed_batch(&[])` returns `Ok(vec![])` — empty slice, empty result.
1283    #[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    /// N texts → N vectors, each of length `dim()`.
1293    #[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    /// v0.4.3: under the default Cargo build (no `embeddings` feature)
1311    /// `open_default_embedder` MUST return Noop so a `cargo install
1312    /// kimetsu-cli` user doesn't accidentally start downloading a
1313    /// model from $HOME. Skip when `--features embeddings` is on —
1314    /// that build path has its own integration tests (run with
1315    /// `cargo test --features embeddings -- --ignored`).
1316    #[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    /// v0.4.3: env kill-switch works even when the `embeddings`
1329    /// feature is on — `KIMETSU_BRAIN_EMBEDDER=noop` returns Noop
1330    /// regardless. Tests the env parser directly rather than going
1331    /// through the cached `open_default_embedder`, which would
1332    /// otherwise be poisoned by whatever the previous test in the
1333    /// process initialized.
1334    #[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            // SAFETY: serialized via the shared brain test env lock.
1342            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        // Restore.
1354        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    // ── W3.1: embedder_enabled_for_config tests ──────────────────────
1364
1365    /// W3.1: config=false disables embedder when env is unset.
1366    #[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        // config=false + env unset → disabled.
1376        assert!(
1377            !embedder_enabled_for_config(false),
1378            "config=false + env unset must be disabled"
1379        );
1380        // config=true + env unset → enabled (default).
1381        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    /// W3.1: KIMETSU_BRAIN_EMBEDDER=noop overrides config=true.
1395    #[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    /// W3.1: a real model-id in env overrides config=false (explicit
1418    /// model = caller wants embeddings on).
1419    #[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    /// v0.4.3: model picker maps the user-facing env string onto a
1442    /// stable model id used by both fastembed init AND the
1443    /// `embedding_model` column on each memory row.
1444    #[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            // Unknown values fall back to bge-small with a warning.
1461            ("totally-made-up", "bge-small-en-v1.5"),
1462        ];
1463        for (input, expected) in cases {
1464            // SAFETY: serialized via the shared brain test env lock.
1465            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}