khive_pack_memory/
recall_feedback.rs1use khive_brain_core::{BalancedRecallState, BetaPosterior, FeedbackEventKind, FeedbackSignal};
9use uuid::Uuid;
10
11const FAST_US: i64 = 50_000;
16
17pub fn on_recall_hit(state: &mut BalancedRecallState, target_id: Uuid, latency_us: i64) {
23 state.total_events += 1;
24 state.relevance.update_success();
25 if latency_us <= FAST_US {
26 state.temporal.update_success();
27 } else {
28 state.temporal.update_failure();
29 }
30 let posterior = state
31 .entity_posteriors
32 .get_or_insert(target_id, || BetaPosterior::new(1.0, 1.0));
33 posterior.update_success();
34}
35
36pub fn on_recall_miss(state: &mut BalancedRecallState) {
41 state.total_events += 1;
42 state.relevance.update_failure();
43 state.temporal.update_failure();
44}
45
46pub fn on_explicit_feedback(state: &mut BalancedRecallState, target_id: Uuid, signal: &str) {
55 if let Some(event_kind) = FeedbackEventKind::from_signal_str(signal) {
57 let w = event_kind.update_weight();
58 if event_kind.is_positive() {
59 state.salience.update_success_weighted(w);
60 } else {
61 state.salience.update_failure_weighted(w);
62 }
63 if event_kind == FeedbackEventKind::Correction {
65 state.relevance.update_failure_weighted(w);
66 }
67 let posterior = state
69 .entity_posteriors
70 .get_or_insert(target_id, || BetaPosterior::new(1.0, 1.0));
71 if event_kind.is_positive() {
72 posterior.update_success_weighted(w);
73 } else {
74 posterior.update_failure_weighted(w);
75 }
76 state.total_events += 1;
77 } else if let Ok(fb) =
78 serde_json::from_value::<FeedbackSignal>(serde_json::Value::String(signal.to_owned()))
79 {
80 let positive = matches!(fb, FeedbackSignal::Useful);
82 if positive {
83 state.salience.update_success();
84 } else {
85 state.salience.update_failure();
86 }
87 let posterior = state
88 .entity_posteriors
89 .get_or_insert(target_id, || BetaPosterior::new(1.0, 1.0));
90 if positive {
91 posterior.update_success();
92 } else {
93 posterior.update_failure();
94 }
95 state.total_events += 1;
96 }
97 }
99
100#[cfg(test)]
101mod tests {
102 use super::*;
103
104 fn fresh() -> BalancedRecallState {
105 BalancedRecallState::new(100)
106 }
107
108 #[test]
109 fn recall_hit_fast_increments_relevance_temporal_and_entity_alpha() {
110 let mut s = fresh();
111 let id = Uuid::new_v4();
112 on_recall_hit(&mut s, id, 10_000); assert_eq!(s.total_events, 1);
115 assert!((s.relevance.alpha() - 8.0).abs() < 1e-12);
117 assert!((s.relevance.beta() - 3.0).abs() < 1e-12);
118 assert!((s.temporal.alpha() - 2.0).abs() < 1e-12);
120 assert!((s.temporal.beta() - 9.0).abs() < 1e-12);
121 let ep = s.entity_posteriors.get(&id).unwrap();
123 assert!((ep.alpha() - 2.0).abs() < 1e-12);
124 }
125
126 #[test]
127 fn recall_hit_slow_increments_relevance_alpha_but_temporal_beta() {
128 let mut s = fresh();
129 let id = Uuid::new_v4();
130 on_recall_hit(&mut s, id, 100_000); assert!((s.relevance.alpha() - 8.0).abs() < 1e-12);
133 assert!((s.temporal.beta() - 10.0).abs() < 1e-12);
135 assert!((s.temporal.alpha() - 1.0).abs() < 1e-12);
136 }
137
138 #[test]
139 fn recall_miss_increments_relevance_beta_and_temporal_beta() {
140 let mut s = fresh();
141 on_recall_miss(&mut s);
142
143 assert_eq!(s.total_events, 1);
144 assert!((s.relevance.beta() - 4.0).abs() < 1e-12); assert!((s.temporal.beta() - 10.0).abs() < 1e-12); assert!(s.entity_posteriors.is_empty());
147 }
148
149 #[test]
150 fn explicit_feedback_useful_increments_salience_alpha() {
151 let mut s = fresh();
152 let id = Uuid::new_v4();
153 on_explicit_feedback(&mut s, id, "useful");
154
155 assert_eq!(s.total_events, 1);
156 assert!((s.salience.alpha() - 3.0).abs() < 1e-12);
158 assert!((s.salience.beta() - 8.0).abs() < 1e-12);
159 let ep = s.entity_posteriors.get(&id).unwrap();
160 assert!((ep.alpha() - 2.0).abs() < 1e-12);
161 }
162
163 #[test]
164 fn explicit_feedback_not_useful_increments_salience_beta() {
165 let mut s = fresh();
166 let id = Uuid::new_v4();
167 on_explicit_feedback(&mut s, id, "not_useful");
168
169 assert!((s.salience.beta() - 9.0).abs() < 1e-12); let ep = s.entity_posteriors.get(&id).unwrap();
171 assert!((ep.beta() - 2.0).abs() < 1e-12);
172 }
173
174 #[test]
175 fn explicit_feedback_wrong_increments_salience_beta() {
176 let mut s = fresh();
177 let id = Uuid::new_v4();
178 on_explicit_feedback(&mut s, id, "wrong");
179
180 assert!((s.salience.beta() - 9.0).abs() < 1e-12); }
182
183 #[test]
184 fn explicit_feedback_explicit_positive_applies_weight_1_5() {
185 let mut s = fresh();
186 let id = Uuid::new_v4();
187 on_explicit_feedback(&mut s, id, "explicit_positive");
188
189 assert!((s.salience.alpha() - 3.5).abs() < 1e-12); assert!((s.salience.beta() - 8.0).abs() < 1e-12);
192 let ep = s.entity_posteriors.get(&id).unwrap();
193 assert!((ep.alpha() - 2.5).abs() < 1e-12); }
195
196 #[test]
197 fn explicit_feedback_correction_updates_salience_beta_and_relevance_beta() {
198 let mut s = fresh();
199 let id = Uuid::new_v4();
200 on_explicit_feedback(&mut s, id, "correction");
201
202 assert!((s.salience.beta() - 10.0).abs() < 1e-12); assert!((s.relevance.beta() - 5.0).abs() < 1e-12); let ep = s.entity_posteriors.get(&id).unwrap();
207 assert!((ep.beta() - 3.0).abs() < 1e-12); }
209
210 #[test]
211 fn explicit_feedback_unknown_signal_is_noop() {
212 let mut s = fresh();
213 let id = Uuid::new_v4();
214 let sal_before = (s.salience.alpha(), s.salience.beta());
215 on_explicit_feedback(&mut s, id, "bad_value");
216
217 assert_eq!(s.total_events, 0);
218 assert!((s.salience.alpha() - sal_before.0).abs() < 1e-12);
219 assert!((s.salience.beta() - sal_before.1).abs() < 1e-12);
220 assert!(s.entity_posteriors.is_empty());
221 }
222}