1use crate::record::{ChangePoint, ExperimentRecord, ResearchLine, Trend};
18use crate::retrieval::{RetrievalQuery, ScoredRecord, rank};
19use somatize_core::error::Result;
20use std::collections::HashMap;
21
22#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
24pub struct LineageNode {
25 pub record: ExperimentRecord,
27 pub depth: usize,
29}
30
31#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
33pub struct Lineage {
34 pub focus: ExperimentRecord,
36 pub ancestors: Vec<ExperimentRecord>,
38 pub descendants: Vec<LineageNode>,
40}
41
42impl Lineage {
43 pub fn root(&self) -> &ExperimentRecord {
45 self.ancestors.first().unwrap_or(&self.focus)
46 }
47}
48
49pub trait KnowledgeBase: Send + Sync {
51 fn record(&mut self, experiment: ExperimentRecord) -> Result<()>;
53
54 fn all(&self) -> Result<Vec<ExperimentRecord>>;
56
57 fn len(&self) -> usize;
59
60 fn refresh(&mut self) -> Result<usize> {
67 Ok(0)
68 }
69
70 fn is_empty(&self) -> bool {
72 self.len() == 0
73 }
74
75 fn get(&self, id: &str) -> Result<Option<ExperimentRecord>> {
77 Ok(self.all()?.into_iter().find(|e| e.id == id))
78 }
79
80 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 fn retrieve(&self, query: &RetrievalQuery) -> Result<Vec<ScoredRecord>> {
97 Ok(rank(&self.all()?, query))
98 }
99
100 fn lineage(&self, id: &str) -> Result<Option<Lineage>> {
102 Ok(build_lineage(&self.all()?, id))
103 }
104
105 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 fn research_lines(&self) -> Result<Vec<ResearchLine>> {
118 Ok(research_lines_of(&self.all()?))
119 }
120
121 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 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 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 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
179fn 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
190pub 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
225fn 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
253pub 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
287fn collect_descendants(
289 records: &[ExperimentRecord],
290 parent_id: &str,
291 depth: usize,
292 out: &mut Vec<LineageNode>,
293) {
294 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#[derive(Default)]
318pub struct MemoryKnowledgeBase {
319 experiments: Vec<ExperimentRecord>,
320 index: HashMap<String, usize>,
321}
322
323impl MemoryKnowledgeBase {
324 pub fn new() -> Self {
326 Self::default()
327 }
328
329 pub fn records(&self) -> &[ExperimentRecord] {
332 &self.experiments
333 }
334}
335
336impl KnowledgeBase for MemoryKnowledgeBase {
337 fn record(&mut self, experiment: ExperimentRecord) -> Result<()> {
338 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 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 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 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 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}