1use khive_fold::{Fold, FoldContext};
2use khive_storage::event::Event;
3
4use crate::event::{entity_signal, interpret, is_recall_positive};
5use crate::state::{BalancedRecallState, BetaPosterior};
6
7pub struct BalancedRecallFold {
16 entity_capacity: usize,
17}
18
19impl BalancedRecallFold {
20 pub fn new(entity_capacity: usize) -> Self {
21 Self { entity_capacity }
22 }
23}
24
25impl Fold<Event, BalancedRecallState> for BalancedRecallFold {
26 fn init(&self, _context: &FoldContext) -> BalancedRecallState {
27 BalancedRecallState::new(self.entity_capacity)
28 }
29
30 fn reduce(
31 &self,
32 mut state: BalancedRecallState,
33 event: &Event,
34 _ctx: &FoldContext,
35 ) -> BalancedRecallState {
36 let signal = interpret(event);
37
38 state.total_events += 1;
39
40 if let Some(positive) = is_recall_positive(&signal) {
42 if positive {
43 state.relevance.update_success();
44 } else {
45 state.relevance.update_failure();
46 }
47 }
48
49 if let Some((entity_id, positive)) = entity_signal(&signal) {
51 let posterior = state
52 .entity_posteriors
53 .get_or_insert(entity_id, || BetaPosterior::new(1.0, 1.0));
54 if positive {
55 posterior.update_success();
56 } else {
57 posterior.update_failure();
58 }
59 }
60
61 state
62 }
63
64 fn finalize(&self, state: BalancedRecallState, _context: &FoldContext) -> BalancedRecallState {
65 state
66 }
67}
68
69#[cfg(test)]
70mod tests {
71 use super::*;
72 use khive_types::{EventKind, EventOutcome, SubstrateKind};
73 use uuid::Uuid;
74
75 fn make_event(verb: &str, outcome: EventOutcome, target: Option<Uuid>) -> Event {
76 let mut e = Event::new("test", verb, EventKind::Audit, SubstrateKind::Note, "brain");
77 e.outcome = outcome;
78 e.target_id = target;
79 e
80 }
81
82 #[test]
83 fn initial_state_has_informative_priors() {
84 let fold = BalancedRecallFold::new(100);
85 let ctx = FoldContext::new();
86 let state = fold.init(&ctx);
87 assert!((state.relevance.alpha - 7.0).abs() < 1e-12);
89 assert!((state.relevance.beta - 3.0).abs() < 1e-12);
90 assert!((state.importance.alpha - 2.0).abs() < 1e-12);
92 assert!((state.importance.beta - 8.0).abs() < 1e-12);
93 assert!((state.temporal.alpha - 1.0).abs() < 1e-12);
95 assert!((state.temporal.beta - 9.0).abs() < 1e-12);
96 }
97
98 #[test]
99 fn recall_hit_updates_relevance_and_entity() {
100 let fold = BalancedRecallFold::new(100);
101 let ctx = FoldContext::new();
102 let mut state = fold.init(&ctx);
103
104 let id = Uuid::new_v4();
105 let event = make_event("recall", EventOutcome::Success, Some(id));
106 state = fold.reduce(state, &event, &ctx);
107
108 assert_eq!(state.total_events, 1);
109 assert!((state.relevance.alpha - 8.0).abs() < 1e-12); let ep = state.entity_posteriors.get(&id).unwrap();
111 assert!((ep.alpha - 2.0).abs() < 1e-12); }
113
114 #[test]
115 fn recall_miss_updates_relevance_beta() {
116 let fold = BalancedRecallFold::new(100);
117 let ctx = FoldContext::new();
118 let mut state = fold.init(&ctx);
119
120 let event = make_event("recall", EventOutcome::Success, None);
121 state = fold.reduce(state, &event, &ctx);
122
123 assert!((state.relevance.beta - 4.0).abs() < 1e-12); assert!(state.entity_posteriors.is_empty());
126 }
127
128 #[test]
129 fn irrelevant_event_increments_counter_only() {
130 let fold = BalancedRecallFold::new(100);
131 let ctx = FoldContext::new();
132 let mut state = fold.init(&ctx);
133
134 let event = make_event("link", EventOutcome::Success, Some(Uuid::new_v4()));
135 state = fold.reduce(state, &event, &ctx);
136
137 assert_eq!(state.total_events, 1);
138 assert!((state.relevance.alpha - 7.0).abs() < 1e-12); }
140
141 #[test]
142 fn feedback_not_useful_increments_entity_beta() {
143 let fold = BalancedRecallFold::new(100);
144 let ctx = FoldContext::new();
145 let mut state = fold.init(&ctx);
146
147 let id = Uuid::new_v4();
148 let mut event = make_event("brain.feedback", EventOutcome::Success, Some(id));
149 event.payload = serde_json::json!({"signal": "not_useful"});
150 state = fold.reduce(state, &event, &ctx);
151
152 assert_eq!(state.total_events, 1);
153 let ep = state.entity_posteriors.get(&id).unwrap();
154 assert!((ep.alpha - 1.0).abs() < 1e-12);
155 assert!((ep.beta - 2.0).abs() < 1e-12);
156 }
157
158 #[test]
159 fn brain_emit_legacy_does_not_update_entity() {
160 let fold = BalancedRecallFold::new(100);
162 let ctx = FoldContext::new();
163 let mut state = fold.init(&ctx);
164
165 let id = Uuid::new_v4();
166 let mut event = make_event("brain.emit", EventOutcome::Success, Some(id));
167 event.payload = serde_json::json!({"signal": "useful"});
168 state = fold.reduce(state, &event, &ctx);
169
170 assert_eq!(state.total_events, 1);
171 assert!(state.entity_posteriors.is_empty()); }
173
174 #[test]
175 fn deterministic_replay() {
176 let fold = BalancedRecallFold::new(100);
177 let ctx = FoldContext::new();
178
179 let id = Uuid::new_v4();
180 let events = vec![
181 make_event("recall", EventOutcome::Success, Some(id)),
182 make_event("recall", EventOutcome::Success, None),
183 make_event("search", EventOutcome::Success, None),
184 make_event("recall", EventOutcome::Success, Some(id)),
185 ];
186
187 let mut s1 = fold.init(&ctx);
188 for e in &events {
189 s1 = fold.reduce(s1, e, &ctx);
190 }
191
192 let mut s2 = fold.init(&ctx);
193 for e in &events {
194 s2 = fold.reduce(s2, e, &ctx);
195 }
196
197 let snap1 = s1.to_snapshot();
198 let snap2 = s2.to_snapshot();
199 assert_eq!(snap1.total_events, snap2.total_events);
200 assert_eq!(snap1.relevance, snap2.relevance);
201 assert_eq!(snap1.entity_posteriors, snap2.entity_posteriors);
202 }
203}