1use std::collections::HashMap;
8
9use uuid::Uuid;
10
11use khive_fold::objective::{Objective, ObjectiveContext};
12use khive_fold::ordering::HasId;
13
14#[derive(Debug, Clone)]
20pub struct RetrievalCandidate {
21 pub id: Uuid,
23 pub vector_score: Option<f64>,
25 pub text_score: Option<f64>,
27 pub graph_distance: Option<u32>,
29 pub rrf_score: Option<f64>,
31}
32
33impl HasId for RetrievalCandidate {
34 #[inline]
35 fn id(&self) -> Uuid {
36 self.id
37 }
38}
39
40pub struct VectorSimilarityObjective;
46
47impl Objective<RetrievalCandidate> for VectorSimilarityObjective {
48 #[inline]
49 fn score(&self, candidate: &RetrievalCandidate, _context: &ObjectiveContext) -> f64 {
50 candidate.vector_score.unwrap_or(0.0)
51 }
52
53 fn name(&self) -> &str {
54 "VectorSimilarityObjective"
55 }
56}
57
58pub struct TextRelevanceObjective;
64
65impl Objective<RetrievalCandidate> for TextRelevanceObjective {
66 #[inline]
67 fn score(&self, candidate: &RetrievalCandidate, _context: &ObjectiveContext) -> f64 {
68 candidate.text_score.unwrap_or(0.0)
69 }
70
71 fn name(&self) -> &str {
72 "TextRelevanceObjective"
73 }
74}
75
76pub struct GraphProximityObjective {
91 pub max_distance: u32,
93}
94
95impl Objective<RetrievalCandidate> for GraphProximityObjective {
96 fn score(&self, candidate: &RetrievalCandidate, _context: &ObjectiveContext) -> f64 {
97 let d = match candidate.graph_distance {
98 Some(d) => d,
99 None => return 0.0,
100 };
101 if self.max_distance == 0 || d >= self.max_distance {
102 return 0.0;
103 }
104 1.0 - (d as f64 / self.max_distance as f64)
105 }
106
107 fn name(&self) -> &str {
108 "GraphProximityObjective"
109 }
110}
111
112pub struct RrfFusionObjective;
121
122impl Objective<RetrievalCandidate> for RrfFusionObjective {
123 #[inline]
124 fn score(&self, candidate: &RetrievalCandidate, _context: &ObjectiveContext) -> f64 {
125 candidate.rrf_score.unwrap_or(0.0)
126 }
127
128 fn name(&self) -> &str {
129 "RrfFusionObjective"
130 }
131}
132
133impl Objective<NoteCandidate> for RrfFusionObjective {
134 #[inline]
135 fn score(&self, candidate: &NoteCandidate, _context: &ObjectiveContext) -> f64 {
136 candidate.rrf_score.unwrap_or(0.0)
137 }
138
139 fn name(&self) -> &str {
140 "RrfFusionObjective"
141 }
142}
143
144#[derive(Debug, Clone)]
153pub struct NoteCandidate {
154 pub id: Uuid,
156 pub rrf_score: Option<f64>,
158 pub salience: f64,
160 pub decay_factor: f64,
162 pub age_days: f64,
164 pub effective_salience: f64,
170 pub rerank_scores: HashMap<String, f64>,
173}
174
175impl HasId for NoteCandidate {
176 #[inline]
177 fn id(&self) -> Uuid {
178 self.id
179 }
180}
181
182pub struct DecayAwareSalienceObjective {
195 decay_rate: f64,
196}
197
198impl DecayAwareSalienceObjective {
199 pub fn new(decay_rate: f64) -> Self {
204 assert!(
205 decay_rate.is_finite() && decay_rate >= 0.0,
206 "decay_rate must be finite and non-negative, got {decay_rate}"
207 );
208 Self { decay_rate }
209 }
210
211 pub fn default_memory() -> Self {
213 Self::new(0.01)
214 }
215}
216
217impl Objective<NoteCandidate> for DecayAwareSalienceObjective {
218 #[inline]
219 fn score(&self, candidate: &NoteCandidate, _context: &ObjectiveContext) -> f64 {
220 candidate.salience * (-self.decay_rate * candidate.age_days).exp()
221 }
222
223 fn name(&self) -> &str {
224 "DecayAwareSalienceObjective"
225 }
226}
227
228pub struct AmplifiedDecayAwareSalienceObjective {
243 pub alpha: f64,
245}
246
247impl AmplifiedDecayAwareSalienceObjective {
248 pub fn new(alpha: f64) -> Self {
250 Self { alpha }
251 }
252
253 pub fn default_memory() -> Self {
255 Self::new(1.5)
256 }
257}
258
259impl Objective<NoteCandidate> for AmplifiedDecayAwareSalienceObjective {
260 #[inline]
261 fn score(&self, candidate: &NoteCandidate, _context: &ObjectiveContext) -> f64 {
262 candidate.effective_salience.powf(self.alpha)
265 }
266
267 fn name(&self) -> &str {
268 "AmplifiedDecayAwareSalienceObjective"
269 }
270}
271
272pub struct TemporalRecencyObjective {
284 pub half_life_days: f64,
286}
287
288impl TemporalRecencyObjective {
289 pub fn default_memory() -> Self {
291 Self {
292 half_life_days: 30.0,
293 }
294 }
295}
296
297impl Objective<NoteCandidate> for TemporalRecencyObjective {
298 #[inline]
299 fn score(&self, candidate: &NoteCandidate, _context: &ObjectiveContext) -> f64 {
300 let k = std::f64::consts::LN_2 / self.half_life_days.max(f64::EPSILON);
301 (-k * candidate.age_days).exp()
302 }
303
304 fn name(&self) -> &str {
305 "TemporalRecencyObjective"
306 }
307}
308
309pub struct RerankerObjective {
318 pub reranker_name: String,
320}
321
322impl RerankerObjective {
323 pub fn new(name: impl Into<String>) -> Self {
325 Self {
326 reranker_name: name.into(),
327 }
328 }
329}
330
331impl Objective<NoteCandidate> for RerankerObjective {
332 #[inline]
333 fn score(&self, candidate: &NoteCandidate, _context: &ObjectiveContext) -> f64 {
334 candidate
335 .rerank_scores
336 .get(&self.reranker_name)
337 .copied()
338 .unwrap_or(0.0)
339 }
340
341 fn name(&self) -> &str {
342 "RerankerObjective"
343 }
344}
345
346pub struct MemoryRecallPipeline {
355 pipeline: khive_fold::WeightedObjective<NoteCandidate>,
356}
357
358impl MemoryRecallPipeline {
359 pub fn new(
366 relevance_weight: f64,
367 salience_weight: f64,
368 temporal_weight: f64,
369 half_life_days: f64,
370 salience_alpha: f64,
371 ) -> Self {
372 use khive_fold::WeightedObjective;
373 let pipeline = WeightedObjective::<NoteCandidate>::new()
374 .add(Box::new(RrfFusionObjective), relevance_weight)
375 .add(
376 Box::new(AmplifiedDecayAwareSalienceObjective::new(salience_alpha)),
377 salience_weight,
378 )
379 .add(
380 Box::new(TemporalRecencyObjective { half_life_days }),
381 temporal_weight,
382 );
383 Self { pipeline }
384 }
385
386 pub fn default_memory() -> Self {
390 Self::new(0.70, 0.20, 0.10, 30.0, 1.5)
391 }
392
393 pub fn score(&self, candidate: &NoteCandidate) -> f64 {
398 let ctx = ObjectiveContext::new();
399 use khive_fold::objective::Objective;
400 self.pipeline.score(candidate, &ctx).clamp(0.0, 1.0)
401 }
402}
403
404#[cfg(test)]
409mod tests {
410 use super::*;
411 use khive_fold::objective::{Objective, ObjectiveContext};
412 use khive_fold::WeightedObjective;
413 use uuid::Uuid;
414
415 fn ctx() -> ObjectiveContext {
416 ObjectiveContext::new()
417 }
418
419 fn candidate(
420 vector: Option<f64>,
421 text: Option<f64>,
422 dist: Option<u32>,
423 rrf: Option<f64>,
424 ) -> RetrievalCandidate {
425 RetrievalCandidate {
426 id: Uuid::new_v4(),
427 vector_score: vector,
428 text_score: text,
429 graph_distance: dist,
430 rrf_score: rrf,
431 }
432 }
433
434 fn note_candidate(
435 rrf: Option<f64>,
436 salience: f64,
437 decay_factor: f64,
438 age_days: f64,
439 ) -> NoteCandidate {
440 let effective_salience = salience * (-decay_factor * age_days).exp();
442 NoteCandidate {
443 id: Uuid::new_v4(),
444 rrf_score: rrf,
445 salience,
446 decay_factor,
447 age_days,
448 effective_salience,
449 rerank_scores: HashMap::new(),
450 }
451 }
452
453 #[test]
456 fn vector_present_returns_signal() {
457 let c = candidate(Some(0.85), None, None, None);
458 let score = VectorSimilarityObjective.score(&c, &ctx());
459 assert!((score - 0.85).abs() < 1e-12);
460 }
461
462 #[test]
463 fn vector_absent_returns_zero() {
464 let c = candidate(None, None, None, None);
465 assert_eq!(VectorSimilarityObjective.score(&c, &ctx()), 0.0);
466 }
467
468 #[test]
469 fn vector_zero_score_returns_zero() {
470 let c = candidate(Some(0.0), None, None, None);
471 assert_eq!(VectorSimilarityObjective.score(&c, &ctx()), 0.0);
472 }
473
474 #[test]
477 fn text_present_returns_signal() {
478 let c = candidate(None, Some(0.6), None, None);
479 let score = TextRelevanceObjective.score(&c, &ctx());
480 assert!((score - 0.6).abs() < 1e-12);
481 }
482
483 #[test]
484 fn text_absent_returns_zero() {
485 let c = candidate(None, None, None, None);
486 assert_eq!(TextRelevanceObjective.score(&c, &ctx()), 0.0);
487 }
488
489 #[test]
492 fn graph_anchor_hit_scores_one() {
493 let c = candidate(None, None, Some(0), None);
494 let obj = GraphProximityObjective { max_distance: 3 };
495 assert!((obj.score(&c, &ctx()) - 1.0).abs() < 1e-12);
496 }
497
498 #[test]
499 fn graph_midpoint_scores_half() {
500 let c = candidate(None, None, Some(1), None);
501 let obj = GraphProximityObjective { max_distance: 2 };
502 assert!((obj.score(&c, &ctx()) - 0.5).abs() < 1e-12);
503 }
504
505 #[test]
506 fn graph_at_boundary_scores_zero() {
507 let c = candidate(None, None, Some(3), None);
508 let obj = GraphProximityObjective { max_distance: 3 };
509 assert_eq!(obj.score(&c, &ctx()), 0.0);
510 }
511
512 #[test]
513 fn graph_beyond_boundary_scores_zero() {
514 let c = candidate(None, None, Some(10), None);
515 let obj = GraphProximityObjective { max_distance: 3 };
516 assert_eq!(obj.score(&c, &ctx()), 0.0);
517 }
518
519 #[test]
520 fn graph_absent_scores_zero() {
521 let c = candidate(None, None, None, None);
522 let obj = GraphProximityObjective { max_distance: 3 };
523 assert_eq!(obj.score(&c, &ctx()), 0.0);
524 }
525
526 #[test]
527 fn graph_max_distance_zero_always_scores_zero() {
528 let c = candidate(None, None, Some(0), None);
530 let obj = GraphProximityObjective { max_distance: 0 };
531 assert_eq!(obj.score(&c, &ctx()), 0.0);
532 }
533
534 #[test]
537 fn rrf_present_returns_signal() {
538 let c = candidate(None, None, None, Some(0.0327));
539 let score = RrfFusionObjective.score(&c, &ctx());
540 assert!((score - 0.0327).abs() < 1e-12);
541 }
542
543 #[test]
544 fn rrf_absent_returns_zero() {
545 let c = candidate(None, None, None, None);
546 assert_eq!(RrfFusionObjective.score(&c, &ctx()), 0.0);
547 }
548
549 #[test]
552 fn weighted_composition_vector_and_text() {
553 let c = candidate(Some(0.8), Some(0.6), None, None);
554
555 let obj = WeightedObjective::<RetrievalCandidate>::new()
556 .add(Box::new(VectorSimilarityObjective), 0.5)
557 .add(Box::new(TextRelevanceObjective), 0.5);
558
559 let score = obj.score(&c, &ctx());
560 assert!((score - 0.7).abs() < 1e-12);
562 }
563
564 #[test]
565 fn weighted_composition_with_graph() {
566 let c = candidate(Some(1.0), Some(0.0), Some(1), None);
567
568 let obj = WeightedObjective::<RetrievalCandidate>::new()
569 .add(Box::new(VectorSimilarityObjective), 0.4)
570 .add(Box::new(TextRelevanceObjective), 0.3)
571 .add(Box::new(GraphProximityObjective { max_distance: 4 }), 0.3);
572
573 let score = obj.score(&c, &ctx());
574 assert!((score - 0.625).abs() < 1e-12);
575 }
576
577 #[test]
578 fn weighted_all_absent_returns_zero() {
579 let c = candidate(None, None, None, None);
580
581 let obj = WeightedObjective::<RetrievalCandidate>::new()
582 .add(Box::new(VectorSimilarityObjective), 0.5)
583 .add(Box::new(TextRelevanceObjective), 0.5);
584
585 assert_eq!(obj.score(&c, &ctx()), 0.0);
587 }
588
589 #[test]
592 fn has_id_returns_candidate_uuid() {
593 let id = Uuid::new_v4();
594 let c = RetrievalCandidate {
595 id,
596 vector_score: None,
597 text_score: None,
598 graph_distance: None,
599 rrf_score: None,
600 };
601 assert_eq!(c.id(), id);
602 }
603
604 #[test]
607 fn select_top_orders_by_vector_score() {
608 use khive_fold::DeterministicObjective;
609
610 let candidates = vec![
611 candidate(Some(0.3), None, None, None),
612 candidate(Some(0.9), None, None, None),
613 candidate(Some(0.6), None, None, None),
614 ];
615
616 let top = VectorSimilarityObjective.select_top_deterministic(&candidates, 2, &ctx());
617
618 assert_eq!(top.len(), 2);
619 assert!((top[0].score - 0.9).abs() < 1e-12);
620 assert!((top[1].score - 0.6).abs() < 1e-12);
621 }
622
623 #[test]
626 fn note_candidate_has_id_returns_uuid() {
627 let id = Uuid::new_v4();
628 let c = NoteCandidate {
629 id,
630 rrf_score: None,
631 salience: 0.5,
632 decay_factor: 0.01,
633 age_days: 0.0,
634 effective_salience: 0.5,
635 rerank_scores: HashMap::new(),
636 };
637 assert_eq!(c.id(), id);
638 }
639
640 #[test]
643 fn decay_aware_zero_age_returns_full_salience() {
644 let obj = DecayAwareSalienceObjective::new(0.01);
645 let c = note_candidate(None, 0.8, 0.01, 0.0);
646 let score = obj.score(&c, &ctx());
647 assert!((score - 0.8).abs() < 1e-12, "got {score}");
648 }
649
650 #[test]
651 fn decay_aware_configured_rates_produce_expected_distinct_scores() {
652 let slow = DecayAwareSalienceObjective::new(0.01);
653 let fast = DecayAwareSalienceObjective::new(0.1);
654 let c = note_candidate(None, 0.8, 0.99, 10.0);
655
656 let slow_score = slow.score(&c, &ctx());
657 let fast_score = fast.score(&c, &ctx());
658 let expected_slow = 0.8 * (-0.01_f64 * 10.0).exp();
659 let expected_fast = 0.8 * (-0.1_f64 * 10.0).exp();
660
661 assert!(
662 (slow_score - expected_slow).abs() < 1e-12,
663 "got {slow_score}, expected {expected_slow}"
664 );
665 assert!(
666 (fast_score - expected_fast).abs() < 1e-12,
667 "got {fast_score}, expected {expected_fast}"
668 );
669 assert!(
670 slow_score > fast_score,
671 "slower configured decay should score higher: {slow_score} vs {fast_score}"
672 );
673 }
674
675 #[test]
676 fn decay_aware_does_not_read_candidate_decay_factor() {
677 let obj = DecayAwareSalienceObjective::new(0.05);
678 let slow = note_candidate(None, 1.0, 0.001, 100.0);
679 let fast = note_candidate(None, 1.0, 0.1, 100.0);
680 let score_slow = obj.score(&slow, &ctx());
681 let score_fast = obj.score(&fast, &ctx());
682 assert!(
683 (score_slow - score_fast).abs() < 1e-12,
684 "candidate decay_factor changed fixed-rate score: {score_slow} vs {score_fast}"
685 );
686 }
687
688 #[test]
689 fn decay_aware_rejects_invalid_configured_rates() {
690 for bad in [-0.1, f64::NAN, f64::INFINITY, f64::NEG_INFINITY] {
691 let result = std::panic::catch_unwind(|| DecayAwareSalienceObjective::new(bad));
692 assert!(result.is_err(), "accepted invalid decay_rate {bad}");
693 }
694 }
695
696 #[test]
699 fn temporal_score_one_at_zero_age() {
700 let obj = TemporalRecencyObjective {
701 half_life_days: 30.0,
702 };
703 let c = note_candidate(None, 0.5, 0.01, 0.0);
704 let score = obj.score(&c, &ctx());
705 assert!((score - 1.0).abs() < 1e-12, "got {score}");
706 }
707
708 #[test]
709 fn temporal_score_half_at_half_life() {
710 let half_life = 30.0;
711 let obj = TemporalRecencyObjective {
712 half_life_days: half_life,
713 };
714 let c = note_candidate(None, 0.5, 0.01, half_life);
715 let score = obj.score(&c, &ctx());
716 assert!(
717 (score - 0.5).abs() < 1e-10,
718 "expected 0.5 at half_life, got {score}"
719 );
720 }
721
722 #[test]
723 fn temporal_score_decreases_with_age() {
724 let obj = TemporalRecencyObjective {
725 half_life_days: 30.0,
726 };
727 let young = note_candidate(None, 1.0, 0.01, 10.0);
728 let old = note_candidate(None, 1.0, 0.01, 100.0);
729 let score_young = obj.score(&young, &ctx());
730 let score_old = obj.score(&old, &ctx());
731 assert!(
732 score_young > score_old,
733 "younger note should score higher: {score_young} vs {score_old}"
734 );
735 }
736
737 #[test]
740 fn reranker_returns_named_score() {
741 let mut c = note_candidate(None, 0.5, 0.01, 0.0);
742 c.rerank_scores.insert("cross_encoder".to_string(), 0.9);
743 let obj = RerankerObjective::new("cross_encoder");
744 let score = obj.score(&c, &ctx());
745 assert!((score - 0.9).abs() < 1e-12, "got {score}");
746 }
747
748 #[test]
749 fn reranker_absent_key_returns_zero() {
750 let c = note_candidate(None, 0.5, 0.01, 0.0);
751 let obj = RerankerObjective::new("cross_encoder");
752 let score = obj.score(&c, &ctx());
753 assert_eq!(score, 0.0);
754 }
755
756 #[test]
757 fn reranker_different_keys_independent() {
758 let mut c = note_candidate(None, 0.5, 0.01, 0.0);
759 c.rerank_scores.insert("salience".to_string(), 0.7);
760 let obj_ce = RerankerObjective::new("cross_encoder");
761 let obj_sal = RerankerObjective::new("salience");
762 assert_eq!(obj_ce.score(&c, &ctx()), 0.0);
763 assert!((obj_sal.score(&c, &ctx()) - 0.7).abs() < 1e-12);
764 }
765
766 #[test]
769 fn memory_pipeline_weighted_composition() {
770 let c = NoteCandidate {
772 id: Uuid::new_v4(),
773 rrf_score: Some(0.5),
774 salience: 0.8,
775 decay_factor: 0.01,
776 age_days: 0.0,
777 effective_salience: 0.8,
778 rerank_scores: HashMap::new(),
779 };
780 let pipeline = WeightedObjective::<NoteCandidate>::new()
781 .add(Box::new(RrfFusionObjective), 0.70)
782 .add(Box::new(DecayAwareSalienceObjective::new(0.0)), 0.20)
783 .add(
784 Box::new(TemporalRecencyObjective {
785 half_life_days: 30.0,
786 }),
787 0.10,
788 );
789 let score = pipeline.score(&c, &ctx());
790 assert!((score - 0.61).abs() < 1e-10, "got {score}");
791 }
792}