1use std::sync::Arc;
2
3use ::log::{debug, info, trace, warn};
4
5use crate::ports::{Embedder, Ranker};
6use crate::shared::SharedMemory;
7use crate::types::Fact;
8
9/// Cosine similarity of two vectors; 0 if either is empty, zero or they differ in length.
10pub fn cosine(a: &[f32], b: &[f32]) -> f32 {
11    if a.is_empty() || a.len() != b.len() {
12        return 0.0;
13    }
14    let (mut dot, mut na, mut nb) = (0.0f32, 0.0f32, 0.0f32);
15    for (x, y) in a.iter().zip(b) {
16        dot += x * y;
17        na += x * x;
18        nb += y * y;
19    }
20    if na == 0.0 || nb == 0.0 { 0.0 } else { dot / (na.sqrt() * nb.sqrt()) }
21}
22
23/// Calling a memory to mind: search by meaning, then let a second opinion order the best few.
24///
25/// Whiskers has little to remember at first, and a small memory is all relevant, so up to
26/// `keep` facts are simply returned. Past that, the message is embedded, the closest
27/// `candidates` by cosine are taken, and the ranker (Jev, which sees the message and the
28/// candidates together) orders them; the best `keep` are used. Either helper being down
29/// degrades the search, never the conversation: no embedder means the most recent facts, no
30/// ranker means the cosine order.
31pub struct Recall {
32    embedder: Arc<dyn Embedder>,
33    ranker: Arc<dyn Ranker>,
34    pub candidates: usize,
35    pub keep: usize,
36    /// Below this cosine a fact is not even a candidate. Measured 2026-10-04 with embeddinggemma-300M and no
37    /// prefixes: a fact that fits a question scores 0.58 to 0.67, unrelated ones 0.27 to 0.47.
38    pub min_cosine: f32,
39}
40
41impl Recall {
42    pub fn new(embedder: Arc<dyn Embedder>, ranker: Arc<dyn Ranker>) -> Self {
43        Self { embedder, ranker, candidates: 8, keep: 4, min_cosine: 0.45 }
44    }
45
46    pub fn search(&self, memory: &SharedMemory, query: &str) -> Vec<Fact> {
47        let mut facts = memory.usable();
48        debug!("recall: {} facts, query of {} chars, keep {}", facts.len(), query.len(), self.keep);
49        if facts.len() <= self.keep || query.trim().is_empty() {
50            debug!("recall: returning all {} facts without searching", facts.len());
51            return facts;
52        }
53        let q = match self.embedder.embed(&[query.to_owned()]) {
54            Ok(mut v) if v.len() == 1 => v.remove(0),
55            Ok(v) => {
56                warn!("recall: embedder returned {} vectors for one text; falling back to the most recent facts", v.len());
57                return most_recent(facts, self.keep);
58            }
59            Err(e) => {
60                warn!("recall: embedder failed ({}); falling back to the most recent facts", e.0);
61                return most_recent(facts, self.keep);
62            }
63        };
64        let mut scored: Vec<(f32, Fact)> = facts
65            .drain(..)
66            .map(|f| (cosine(&q, &f.embedding), f))
67            .filter(|(c, _)| *c >= self.min_cosine)
68            .collect();
69        scored.sort_by(|a, b| b.0.total_cmp(&a.0));
70        scored.truncate(self.candidates);
71        debug!("recall: {} candidates above cosine {}", scored.len(), self.min_cosine);
72        if scored.len() > self.keep {
73            let texts: Vec<String> = scored.iter().map(|(_, f)| f.text.clone()).collect();
74            match self.ranker.rank(query, &texts) {
75                Ok(p) if p.len() == scored.len() => {
76                    let mut paired: Vec<(f32, (f32, Fact))> = p.into_iter().zip(scored).collect();
77                    paired.sort_by(|a, b| b.0.total_cmp(&a.0));
78                    scored = paired.into_iter().map(|(_, sf)| sf).collect();
79                    trace!("recall: ranker ordered the candidates");
80                }
81                Ok(p) => warn!("recall: ranker returned {} scores for {} candidates; keeping the cosine order", p.len(), scored.len()),
82                Err(e) => warn!("recall: ranker failed ({}); keeping the cosine order", e.0),
83            }
84        }
85        let kept: Vec<Fact> = scored.into_iter().take(self.keep).map(|(_, f)| f).collect();
86        info!("recall: {} facts recalled", kept.len());
87        kept
88    }
89}
90
91fn most_recent(mut facts: Vec<Fact>, n: usize) -> Vec<Fact> {
92    trace!("recall: taking the {n} most recent of {} facts", facts.len());
93    facts.sort_by_key(|f| std::cmp::Reverse(f.learned_at_ms));
94    facts.truncate(n);
95    facts
96}