Skip to main content

somatize_memory/
knowledge_base.rs

1//! KnowledgeBase trait and in-memory implementation.
2//!
3//! The trait returns **owned** records, not references. That costs a
4//! clone per hit and buys two things a reference-returning trait cannot
5//! have: a backend that pages, streams or queries a remote store (a
6//! reference has to point at something the base already holds in
7//! memory), and an implementation that does not have to keep a
8//! duplicate `Vec` alive purely to have something to lend out.
9//!
10//! Only [`record`](KnowledgeBase::record), [`all`](KnowledgeBase::all)
11//! and [`len`](KnowledgeBase::len) are required. Every query — search,
12//! research lines, trends, change points, lineage, retrieval — has a
13//! default implementation over `all()`, so a new backend gets the whole
14//! analytics surface by implementing three methods, and the analytics
15//! themselves live in one place instead of once per backend.
16
17use crate::record::{ChangePoint, ExperimentRecord, ResearchLine, Trend};
18use crate::retrieval::{RetrievalQuery, ScoredRecord, rank};
19use somatize_core::error::Result;
20use std::collections::HashMap;
21
22/// One node of a lineage tree, with its distance from the focus.
23#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
24pub struct LineageNode {
25    /// The descendant record itself.
26    pub record: ExperimentRecord,
27    /// Generations below the focus (1 = direct child).
28    pub depth: usize,
29}
30
31/// An experiment in context: where it came from and what came of it.
32#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
33pub struct Lineage {
34    /// The experiment the lineage was asked about.
35    pub focus: ExperimentRecord,
36    /// Ancestors from the root down to the direct parent.
37    pub ancestors: Vec<ExperimentRecord>,
38    /// Descendants in pre-order, so the list reads as an indented tree.
39    pub descendants: Vec<LineageNode>,
40}
41
42impl Lineage {
43    /// The root of the focus's line.
44    pub fn root(&self) -> &ExperimentRecord {
45        self.ancestors.first().unwrap_or(&self.focus)
46    }
47}
48
49/// Trait for knowledge base backends.
50pub trait KnowledgeBase: Send + Sync {
51    /// Store a new experiment record.
52    fn record(&mut self, experiment: ExperimentRecord) -> Result<()>;
53
54    /// Every record this base holds, in insertion order.
55    fn all(&self) -> Result<Vec<ExperimentRecord>>;
56
57    /// Total number of experiments.
58    fn len(&self) -> usize;
59
60    /// Pick up records another process appended since this handle was
61    /// opened. Returns how many were newly loaded.
62    ///
63    /// Defaults to a no-op for bases with nothing behind them. A
64    /// long-lived reader (the MCP server) must call this before a read,
65    /// or it answers from a snapshot that gets staler by the hour.
66    fn refresh(&mut self) -> Result<usize> {
67        Ok(0)
68    }
69
70    /// Whether the base holds no records at all.
71    fn is_empty(&self) -> bool {
72        self.len() == 0
73    }
74
75    /// Get an experiment by ID.
76    fn get(&self, id: &str) -> Result<Option<ExperimentRecord>> {
77        Ok(self.all()?.into_iter().find(|e| e.id == id))
78    }
79
80    /// Substring search over name, hypothesis, notes, tags and pipeline.
81    ///
82    /// Kept as a cheap literal filter; [`retrieve`](Self::retrieve) is
83    /// the ranked one.
84    fn search(&self, query: &str, max_results: usize) -> Result<Vec<ExperimentRecord>> {
85        let q = query.to_lowercase();
86        Ok(self
87            .all()?
88            .into_iter()
89            .filter(|e| matches_literal(e, &q))
90            .take(max_results)
91            .collect())
92    }
93
94    /// Rank records against a query by relevance, structure, recency
95    /// and importance. See [`crate::retrieval`] for the formula.
96    fn retrieve(&self, query: &RetrievalQuery) -> Result<Vec<ScoredRecord>> {
97        Ok(rank(&self.all()?, query))
98    }
99
100    /// An experiment with its ancestors and descendants.
101    fn lineage(&self, id: &str) -> Result<Option<Lineage>> {
102        Ok(build_lineage(&self.all()?, id))
103    }
104
105    /// List experiments in a research line, ordered chronologically.
106    fn experiments_in_line(&self, line: &str) -> Result<Vec<ExperimentRecord>> {
107        let mut exps: Vec<ExperimentRecord> = self
108            .all()?
109            .into_iter()
110            .filter(|e| e.research_line.as_deref() == Some(line))
111            .collect();
112        exps.sort_by_key(|e| e.timestamp);
113        Ok(exps)
114    }
115
116    /// Get all research lines with trend analysis.
117    fn research_lines(&self) -> Result<Vec<ResearchLine>> {
118        Ok(research_lines_of(&self.all()?))
119    }
120
121    /// Find promising research lines (improving trend, high metric values).
122    fn promising_lines(&self, metric: &str) -> Result<Vec<ResearchLine>> {
123        Ok(self
124            .research_lines()?
125            .into_iter()
126            .filter(|l| {
127                l.trend == Trend::Improving || l.best_metric_name.as_deref() == Some(metric)
128            })
129            .collect())
130    }
131
132    /// Get the metric trajectory for a research line.
133    fn trajectory(&self, line: &str, metric: &str) -> Result<Vec<(String, f64)>> {
134        Ok(self
135            .experiments_in_line(line)?
136            .into_iter()
137            .filter_map(|e| e.metrics.get(metric).map(|&v| (e.id.clone(), v)))
138            .collect())
139    }
140
141    /// Detect change points where metrics shifted significantly.
142    fn change_points(&self, line: &str, metric: &str, threshold: f64) -> Result<Vec<ChangePoint>> {
143        let experiments = self.experiments_in_line(line)?;
144        let traj: Vec<(&ExperimentRecord, f64)> = experiments
145            .iter()
146            .filter_map(|e| e.metrics.get(metric).map(|&v| (e, v)))
147            .collect();
148        let mut points = Vec::new();
149        for window in traj.windows(2) {
150            let (_, val_before) = window[0];
151            let (exp_after, val_after) = window[1];
152            let delta = (val_after - val_before).abs();
153            if delta >= threshold {
154                points.push(ChangePoint {
155                    experiment_id: exp_after.id.clone(),
156                    timestamp: exp_after.timestamp,
157                    metric_name: metric.to_string(),
158                    value_before: val_before,
159                    value_after: val_after,
160                    description: format!(
161                        "{metric} changed from {val_before:.4} to {val_after:.4} (delta={delta:.4})"
162                    ),
163                });
164            }
165        }
166        Ok(points)
167    }
168
169    /// Direct children of an experiment.
170    fn children(&self, experiment_id: &str) -> Result<Vec<ExperimentRecord>> {
171        Ok(self
172            .all()?
173            .into_iter()
174            .filter(|e| e.parent.as_deref() == Some(experiment_id))
175            .collect())
176    }
177}
178
179/// Whether a lowercased query appears anywhere a human would look.
180fn matches_literal(e: &ExperimentRecord, q: &str) -> bool {
181    let contains = |s: &str| s.to_lowercase().contains(q);
182    contains(&e.name)
183        || contains(&e.pipeline_summary)
184        || e.hypothesis.as_deref().is_some_and(contains)
185        || e.notes.as_deref().is_some_and(contains)
186        || e.tags.iter().any(|t| contains(t))
187        || e.conclusion.as_ref().is_some_and(|c| contains(&c.headline))
188}
189
190/// Group records into research lines with a trend per line.
191pub fn research_lines_of(records: &[ExperimentRecord]) -> Vec<ResearchLine> {
192    let mut lines: HashMap<&str, Vec<&ExperimentRecord>> = HashMap::new();
193    for record in records {
194        if let Some(line) = record.research_line.as_deref() {
195            lines.entry(line).or_default().push(record);
196        }
197    }
198
199    let mut result: Vec<ResearchLine> = lines
200        .into_iter()
201        .map(|(name, members)| {
202            let mut best_value = None;
203            let mut best_name = None;
204            for record in &members {
205                for (metric, &value) in &record.metrics {
206                    if best_value.is_none_or(|best| value > best) {
207                        best_value = Some(value);
208                        best_name = Some(metric.clone());
209                    }
210                }
211            }
212            ResearchLine {
213                name: name.to_string(),
214                experiments: members.iter().map(|e| e.id.clone()).collect(),
215                trend: compute_trend(&members),
216                best_metric_value: best_value,
217                best_metric_name: best_name,
218            }
219        })
220        .collect();
221    result.sort_by(|a, b| a.name.cmp(&b.name));
222    result
223}
224
225/// Direction of the last few values of the line's first metric.
226fn compute_trend(members: &[&ExperimentRecord]) -> Trend {
227    if members.len() < 2 {
228        return Trend::Unknown;
229    }
230    let Some(key) = members[0].metrics.keys().next() else {
231        return Trend::Unknown;
232    };
233    let values: Vec<f64> = members
234        .iter()
235        .filter_map(|e| e.metrics.get(key).copied())
236        .collect();
237    if values.len() < 2 {
238        return Trend::Unknown;
239    }
240    let recent = &values[values.len().saturating_sub(3)..];
241    let diffs: Vec<f64> = recent.windows(2).map(|w| w[1] - w[0]).collect();
242    if diffs.is_empty() {
243        Trend::Unknown
244    } else if diffs.iter().all(|&d| d > 0.001) {
245        Trend::Improving
246    } else if diffs.iter().all(|&d| d < -0.001) {
247        Trend::Declining
248    } else {
249        Trend::Plateaued
250    }
251}
252
253/// Walk parents up and children down from `id`.
254///
255/// A parent cycle (only reachable from a hand-edited journal) stops the
256/// upward walk instead of hanging.
257pub fn build_lineage(records: &[ExperimentRecord], id: &str) -> Option<Lineage> {
258    let by_id: HashMap<&str, &ExperimentRecord> =
259        records.iter().map(|r| (r.id.as_str(), r)).collect();
260    let focus = *by_id.get(id)?;
261
262    let mut ancestors = Vec::new();
263    let mut seen: std::collections::HashSet<&str> = std::collections::HashSet::from([id]);
264    let mut cursor = focus.parent.as_deref();
265    while let Some(parent_id) = cursor {
266        if !seen.insert(parent_id) {
267            break;
268        }
269        let Some(parent) = by_id.get(parent_id) else {
270            break;
271        };
272        ancestors.push((*parent).clone());
273        cursor = parent.parent.as_deref();
274    }
275    ancestors.reverse();
276
277    let mut descendants = Vec::new();
278    collect_descendants(records, id, 1, &mut descendants);
279
280    Some(Lineage {
281        focus: focus.clone(),
282        ancestors,
283        descendants,
284    })
285}
286
287/// Pre-order walk so the flat list reads as an indented tree.
288fn collect_descendants(
289    records: &[ExperimentRecord],
290    parent_id: &str,
291    depth: usize,
292    out: &mut Vec<LineageNode>,
293) {
294    // A malformed journal must not recurse forever.
295    if depth > 64 {
296        return;
297    }
298    let mut children: Vec<&ExperimentRecord> = records
299        .iter()
300        .filter(|r| r.parent.as_deref() == Some(parent_id))
301        .collect();
302    children.sort_by_key(|r| r.timestamp);
303    for child in children {
304        out.push(LineageNode {
305            record: child.clone(),
306            depth,
307        });
308        collect_descendants(records, &child.id, depth + 1, out);
309    }
310}
311
312/// In-memory knowledge base implementation.
313///
314/// The reference backend and the storage half of [`FileKnowledgeBase`].
315///
316/// [`FileKnowledgeBase`]: crate::FileKnowledgeBase
317#[derive(Default)]
318pub struct MemoryKnowledgeBase {
319    experiments: Vec<ExperimentRecord>,
320    index: HashMap<String, usize>,
321}
322
323impl MemoryKnowledgeBase {
324    /// An empty base.
325    pub fn new() -> Self {
326        Self::default()
327    }
328
329    /// Borrowed view, for callers inside this crate that would rather
330    /// not clone (the trait itself hands out owned records on purpose).
331    pub fn records(&self) -> &[ExperimentRecord] {
332        &self.experiments
333    }
334}
335
336impl KnowledgeBase for MemoryKnowledgeBase {
337    fn record(&mut self, experiment: ExperimentRecord) -> Result<()> {
338        // Last write wins on the id index; the log keeps both.
339        self.index
340            .insert(experiment.id.clone(), self.experiments.len());
341        self.experiments.push(experiment);
342        Ok(())
343    }
344
345    fn all(&self) -> Result<Vec<ExperimentRecord>> {
346        Ok(self.experiments.clone())
347    }
348
349    fn len(&self) -> usize {
350        self.experiments.len()
351    }
352
353    /// Indexed, unlike the linear default.
354    fn get(&self, id: &str) -> Result<Option<ExperimentRecord>> {
355        Ok(self.index.get(id).map(|&idx| self.experiments[idx].clone()))
356    }
357}
358
359#[cfg(test)]
360mod tests {
361    use super::*;
362    use std::collections::BTreeMap;
363
364    fn make_experiment(id: &str, line: &str, metric_val: f64) -> ExperimentRecord {
365        let mut metrics = BTreeMap::new();
366        metrics.insert("f1".to_string(), metric_val);
367        ExperimentRecord::new(id, format!("Experiment {id}"))
368            .with_research_line(line)
369            .with_metrics(metrics)
370            .with_pipeline("Pipeline([Scaler, Classifier])")
371    }
372
373    #[test]
374    fn record_and_get() {
375        let mut kb = MemoryKnowledgeBase::new();
376        kb.record(make_experiment("exp_001", "line_a", 0.85))
377            .unwrap();
378
379        assert_eq!(kb.len(), 1);
380        assert_eq!(kb.get("exp_001").unwrap().unwrap().id, "exp_001");
381        assert_eq!(kb.all().unwrap().len(), 1);
382        assert_eq!(kb.refresh().unwrap(), 0, "an in-memory base has no backing");
383    }
384
385    #[test]
386    fn get_missing_returns_none() {
387        let kb = MemoryKnowledgeBase::new();
388        assert!(kb.get("nonexistent").unwrap().is_none());
389    }
390
391    #[test]
392    fn search_by_name() {
393        let mut kb = MemoryKnowledgeBase::new();
394        kb.record(make_experiment("exp_001", "line_a", 0.8))
395            .unwrap();
396        kb.record(ExperimentRecord::new("exp_002", "SVM experiment").with_research_line("line_b"))
397            .unwrap();
398
399        let results = kb.search("SVM", 10).unwrap();
400        assert_eq!(results.len(), 1);
401        assert_eq!(results[0].id, "exp_002");
402    }
403
404    #[test]
405    fn search_by_tag() {
406        let mut kb = MemoryKnowledgeBase::new();
407        kb.record(
408            make_experiment("exp_001", "line_a", 0.8)
409                .with_tags(vec!["normalization".into(), "time-series".into()]),
410        )
411        .unwrap();
412
413        assert_eq!(kb.search("time-series", 10).unwrap().len(), 1);
414    }
415
416    #[test]
417    fn search_reaches_the_conclusion_headline() {
418        let mut kb = MemoryKnowledgeBase::new();
419        let mut rec = make_experiment("e1", "line_a", 0.8);
420        rec.conclusion = Some(somatize_core::summary::RunConclusion {
421            headline: "completed in 2m · flags: DEAD_CHANNELS".into(),
422            ..Default::default()
423        });
424        kb.record(rec).unwrap();
425        assert_eq!(kb.search("dead_channels", 10).unwrap().len(), 1);
426    }
427
428    #[test]
429    fn experiments_in_line_ordered() {
430        let mut kb = MemoryKnowledgeBase::new();
431        kb.record(make_experiment("exp_003", "line_a", 0.9))
432            .unwrap();
433        kb.record(make_experiment("exp_001", "line_a", 0.7))
434            .unwrap();
435        kb.record(make_experiment("exp_002", "line_b", 0.8))
436            .unwrap();
437
438        assert_eq!(kb.experiments_in_line("line_a").unwrap().len(), 2);
439    }
440
441    #[test]
442    fn research_lines_detected() {
443        let mut kb = MemoryKnowledgeBase::new();
444        for (i, v) in [0.7, 0.8, 0.85].iter().enumerate() {
445            kb.record(make_experiment(&format!("e{i}"), "rocket_znorm", *v))
446                .unwrap();
447        }
448        kb.record(make_experiment("e4", "inception_minmax", 0.6))
449            .unwrap();
450
451        let lines = kb.research_lines().unwrap();
452        assert_eq!(lines.len(), 2);
453        let rocket = lines.iter().find(|l| l.name == "rocket_znorm").unwrap();
454        assert_eq!(rocket.experiments.len(), 3);
455        assert_eq!(rocket.trend, Trend::Improving);
456    }
457
458    #[test]
459    fn trajectory_returns_metric_values() {
460        let mut kb = MemoryKnowledgeBase::new();
461        kb.record(make_experiment("e1", "line_a", 0.7)).unwrap();
462        kb.record(make_experiment("e2", "line_a", 0.8)).unwrap();
463        kb.record(make_experiment("e3", "line_a", 0.85)).unwrap();
464
465        let traj = kb.trajectory("line_a", "f1").unwrap();
466        assert_eq!(traj.len(), 3);
467        assert!((traj[0].1 - 0.7).abs() < 0.001);
468        assert!((traj[2].1 - 0.85).abs() < 0.001);
469    }
470
471    #[test]
472    fn change_points_detected() {
473        let mut kb = MemoryKnowledgeBase::new();
474        for (id, v) in [("e1", 0.50), ("e2", 0.51), ("e3", 0.80), ("e4", 0.82)] {
475            kb.record(make_experiment(id, "line_a", v)).unwrap();
476        }
477
478        let points = kb.change_points("line_a", "f1", 0.1).unwrap();
479        assert_eq!(points.len(), 1);
480        assert_eq!(points[0].experiment_id, "e3");
481        assert!((points[0].value_before - 0.51).abs() < 0.001);
482        assert!((points[0].value_after - 0.80).abs() < 0.001);
483    }
484
485    #[test]
486    fn change_points_skip_experiments_missing_the_metric() {
487        let mut kb = MemoryKnowledgeBase::new();
488        kb.record(make_experiment("e1", "line_a", 0.5)).unwrap();
489        // No metrics at all: it must not be treated as a zero.
490        kb.record(ExperimentRecord::new("e2", "gap").with_research_line("line_a"))
491            .unwrap();
492        kb.record(make_experiment("e3", "line_a", 0.9)).unwrap();
493
494        let points = kb.change_points("line_a", "f1", 0.1).unwrap();
495        assert_eq!(points.len(), 1);
496        assert_eq!(points[0].experiment_id, "e3");
497    }
498
499    #[test]
500    fn children_found() {
501        let mut kb = MemoryKnowledgeBase::new();
502        kb.record(make_experiment("parent", "line_a", 0.7)).unwrap();
503        kb.record(make_experiment("child_1", "line_a", 0.8).with_parent("parent"))
504            .unwrap();
505        kb.record(make_experiment("child_2", "line_a", 0.75).with_parent("parent"))
506            .unwrap();
507        kb.record(make_experiment("unrelated", "line_b", 0.6))
508            .unwrap();
509
510        assert_eq!(kb.children("parent").unwrap().len(), 2);
511    }
512
513    #[test]
514    fn lineage_walks_up_and_down() {
515        let mut kb = MemoryKnowledgeBase::new();
516        kb.record(make_experiment("root", "l", 0.5)).unwrap();
517        kb.record(make_experiment("mid", "l", 0.6).with_parent("root"))
518            .unwrap();
519        kb.record(make_experiment("leaf_a", "l", 0.7).with_parent("mid"))
520            .unwrap();
521        kb.record(make_experiment("leaf_b", "l", 0.8).with_parent("mid"))
522            .unwrap();
523        kb.record(make_experiment("elsewhere", "l", 0.9)).unwrap();
524
525        let lineage = kb.lineage("mid").unwrap().unwrap();
526        assert_eq!(lineage.focus.id, "mid");
527        assert_eq!(
528            lineage.ancestors.iter().map(|a| &a.id).collect::<Vec<_>>(),
529            vec!["root"]
530        );
531        assert_eq!(lineage.root().id, "root");
532        assert_eq!(
533            lineage
534                .descendants
535                .iter()
536                .map(|d| (d.record.id.as_str(), d.depth))
537                .collect::<Vec<_>>(),
538            vec![("leaf_a", 1), ("leaf_b", 1)]
539        );
540
541        // From the root, the whole tree in pre-order.
542        let from_root = kb.lineage("root").unwrap().unwrap();
543        assert!(from_root.ancestors.is_empty());
544        assert_eq!(from_root.root().id, "root");
545        assert_eq!(
546            from_root
547                .descendants
548                .iter()
549                .map(|d| (d.record.id.as_str(), d.depth))
550                .collect::<Vec<_>>(),
551            vec![("mid", 1), ("leaf_a", 2), ("leaf_b", 2)]
552        );
553
554        assert!(kb.lineage("does-not-exist").unwrap().is_none());
555    }
556
557    #[test]
558    fn lineage_survives_a_hand_edited_parent_cycle() {
559        let records = vec![
560            make_experiment("a", "l", 0.1).with_parent("b"),
561            make_experiment("b", "l", 0.2).with_parent("a"),
562        ];
563        let lineage = build_lineage(&records, "a").unwrap();
564        // Walks up once, then refuses to loop.
565        assert_eq!(lineage.ancestors.len(), 1);
566        assert_eq!(lineage.ancestors[0].id, "b");
567    }
568
569    #[test]
570    fn promising_lines() {
571        let mut kb = MemoryKnowledgeBase::new();
572        for (i, v) in [0.7, 0.8, 0.9].iter().enumerate() {
573            kb.record(make_experiment(&format!("g{i}"), "good_line", *v))
574                .unwrap();
575        }
576        for (i, v) in [0.9, 0.8, 0.7].iter().enumerate() {
577            kb.record(make_experiment(&format!("b{i}"), "bad_line", *v))
578                .unwrap();
579        }
580
581        let promising = kb.promising_lines("f1").unwrap();
582        assert!(promising.iter().any(|l| l.name == "good_line"));
583    }
584
585    #[test]
586    fn trend_detection() {
587        let mut kb = MemoryKnowledgeBase::new();
588        for (line, values) in [
589            ("improving", [0.5, 0.7, 0.9]),
590            ("declining", [0.9, 0.7, 0.5]),
591            ("plateau", [0.8, 0.8, 0.8]),
592        ] {
593            for (i, v) in values.iter().enumerate() {
594                kb.record(make_experiment(&format!("{line}{i}"), line, *v))
595                    .unwrap();
596            }
597        }
598
599        let lines = kb.research_lines().unwrap();
600        let of = |name: &str| lines.iter().find(|l| l.name == name).unwrap().trend;
601        assert_eq!(of("improving"), Trend::Improving);
602        assert_eq!(of("declining"), Trend::Declining);
603        assert_eq!(of("plateau"), Trend::Plateaued);
604    }
605
606    #[test]
607    fn empty_kb() {
608        let kb = MemoryKnowledgeBase::new();
609        assert!(kb.is_empty());
610        assert!(kb.research_lines().unwrap().is_empty());
611        assert!(kb.search("anything", 10).unwrap().is_empty());
612        assert!(kb.all().unwrap().is_empty());
613    }
614
615    #[test]
616    fn serde_roundtrip() {
617        let exp = make_experiment("e1", "line_a", 0.85)
618            .with_hypothesis("Z-norm improves accuracy")
619            .with_notes("First attempt with z-normalization")
620            .with_tags(vec!["normalization".into()]);
621
622        let json = serde_json::to_string(&exp).unwrap();
623        let deserialized: ExperimentRecord = serde_json::from_str(&json).unwrap();
624        assert_eq!(deserialized.id, "e1");
625        assert_eq!(deserialized.hypothesis.unwrap(), "Z-norm improves accuracy");
626    }
627}