1use khive_fold::{Fold, FoldContext};
2use khive_storage::event::Event;
3
4use crate::event::{
5 entity_signal, interpret, is_recall_positive, BrainSignal, FeedbackEventKind, FeedbackSignal,
6};
7use crate::state::{BalancedRecallState, BetaPosterior, SectionPosteriorState, DEFAULT_ESS_CAP};
8
9pub struct BalancedRecallFold {
18 entity_capacity: usize,
19}
20
21impl BalancedRecallFold {
22 pub fn new(entity_capacity: usize) -> Self {
23 Self { entity_capacity }
24 }
25}
26
27impl Fold<Event, BalancedRecallState> for BalancedRecallFold {
28 fn init(&self, _context: &FoldContext) -> BalancedRecallState {
29 BalancedRecallState::new(self.entity_capacity)
30 }
31
32 fn reduce(
33 &self,
34 mut state: BalancedRecallState,
35 event: &Event,
36 _ctx: &FoldContext,
37 ) -> BalancedRecallState {
38 let signal = interpret(event);
39
40 state.total_events += 1;
41
42 if let Some(positive) = is_recall_positive(&signal) {
44 if positive {
45 state.relevance.update_success();
46 } else {
47 state.relevance.update_failure();
48 }
49 }
50
51 if let BrainSignal::Feedback { signal: ref fb, .. } = signal {
55 match fb {
56 FeedbackSignal::Useful => state.salience.update_success(),
57 FeedbackSignal::NotUseful | FeedbackSignal::Wrong => {
58 state.salience.update_failure()
59 }
60 }
61 }
62
63 if let BrainSignal::SemanticFeedback {
66 event_kind: ref ek, ..
67 } = signal
68 {
69 let w = ek.update_weight();
70 if ek.is_positive() {
71 state.salience.update_success_weighted(w);
72 } else {
73 state.salience.update_failure_weighted(w);
74 }
75 if *ek == FeedbackEventKind::Correction {
77 state.relevance.update_failure_weighted(w);
78 }
79 }
80
81 const FAST_US: i64 = 50_000;
96 match &signal {
97 BrainSignal::RecallHit { latency_us, .. } => {
98 if *latency_us <= FAST_US {
99 state.temporal.update_success();
100 } else {
101 state.temporal.update_failure();
102 }
103 }
104 BrainSignal::RecallMiss => state.temporal.update_failure(),
105 _ => {}
106 }
107
108 if let BrainSignal::SemanticFeedback {
112 target_id: eid,
113 event_kind: ref ek,
114 ..
115 } = signal
116 {
117 let posterior = state
118 .entity_posteriors
119 .get_or_insert(eid, || BetaPosterior::new(1.0, 1.0));
120 let w = ek.update_weight();
121 if ek.is_positive() {
122 posterior.update_success_weighted(w);
123 } else {
124 posterior.update_failure_weighted(w);
125 }
126 } else if let Some((entity_id, positive)) = entity_signal(&signal) {
127 let posterior = state
128 .entity_posteriors
129 .get_or_insert(entity_id, || BetaPosterior::new(1.0, 1.0));
130 if positive {
131 posterior.update_success();
132 } else {
133 posterior.update_failure();
134 }
135 }
136
137 state
138 }
139
140 fn finalize(&self, state: BalancedRecallState, _context: &FoldContext) -> BalancedRecallState {
141 state
142 }
143}
144
145pub struct SectionPosteriorFold;
152
153impl SectionPosteriorFold {
154 pub fn new() -> Self {
155 Self
156 }
157}
158
159impl Default for SectionPosteriorFold {
160 fn default() -> Self {
161 Self::new()
162 }
163}
164
165impl Fold<Event, SectionPosteriorState> for SectionPosteriorFold {
166 fn init(&self, _context: &FoldContext) -> SectionPosteriorState {
167 SectionPosteriorState::new()
168 }
169
170 fn reduce(
171 &self,
172 mut state: SectionPosteriorState,
173 event: &Event,
174 _ctx: &FoldContext,
175 ) -> SectionPosteriorState {
176 let signal = interpret(event);
177
178 if let BrainSignal::Feedback {
179 section_signals: Some(ref signals),
180 ..
181 } = signal
182 {
183 state.total_events += 1;
184
185 for (section_type, feedback_signal) in signals {
186 if let Some(posterior) = state.posteriors.get_mut(section_type) {
187 match feedback_signal {
188 FeedbackSignal::Useful => posterior.alpha += 1.0,
189 FeedbackSignal::NotUseful => posterior.beta += 1.0,
190 FeedbackSignal::Wrong => posterior.beta += 2.0,
191 }
192 if let Some(prior) = state.priors.get(section_type) {
193 posterior.apply_ess_cap(&prior.clone(), DEFAULT_ESS_CAP);
194 }
195 }
196 }
197
198 if state.exploration_epoch > 0 {
199 state.exploration_epoch -= 1;
200 }
201 }
202
203 state
204 }
205
206 fn finalize(
207 &self,
208 state: SectionPosteriorState,
209 _context: &FoldContext,
210 ) -> SectionPosteriorState {
211 state
212 }
213}
214
215#[cfg(test)]
216mod tests {
217 use super::*;
218 use khive_types::{EventKind, EventOutcome, SubstrateKind};
219 use uuid::Uuid;
220
221 fn make_event(verb: &str, outcome: EventOutcome, target: Option<Uuid>) -> Event {
222 let mut e = Event::new("test", verb, EventKind::Audit, SubstrateKind::Note, "brain");
223 e.outcome = outcome;
224 e.target_id = target;
225 e
226 }
227
228 #[test]
229 fn initial_state_has_informative_priors() {
230 let fold = BalancedRecallFold::new(100);
231 let ctx = FoldContext::new();
232 let state = fold.init(&ctx);
233 assert!((state.relevance.alpha - 7.0).abs() < 1e-12);
235 assert!((state.relevance.beta - 3.0).abs() < 1e-12);
236 assert!((state.salience.alpha - 2.0).abs() < 1e-12);
238 assert!((state.salience.beta - 8.0).abs() < 1e-12);
239 assert!((state.temporal.alpha - 1.0).abs() < 1e-12);
241 assert!((state.temporal.beta - 9.0).abs() < 1e-12);
242 }
243
244 #[test]
245 fn recall_hit_updates_relevance_and_entity() {
246 let fold = BalancedRecallFold::new(100);
247 let ctx = FoldContext::new();
248 let mut state = fold.init(&ctx);
249
250 let id = Uuid::new_v4();
251 let event = make_event("recall", EventOutcome::Success, Some(id));
252 state = fold.reduce(state, &event, &ctx);
253
254 assert_eq!(state.total_events, 1);
255 assert!((state.relevance.alpha - 8.0).abs() < 1e-12); let ep = state.entity_posteriors.get(&id).unwrap();
257 assert!((ep.alpha - 2.0).abs() < 1e-12); }
259
260 #[test]
261 fn recall_miss_updates_relevance_beta() {
262 let fold = BalancedRecallFold::new(100);
263 let ctx = FoldContext::new();
264 let mut state = fold.init(&ctx);
265
266 let event = make_event("recall", EventOutcome::Success, None);
267 state = fold.reduce(state, &event, &ctx);
268
269 assert!((state.relevance.beta - 4.0).abs() < 1e-12); assert!(state.entity_posteriors.is_empty());
272 }
273
274 #[test]
275 fn irrelevant_event_increments_counter_only() {
276 let fold = BalancedRecallFold::new(100);
277 let ctx = FoldContext::new();
278 let mut state = fold.init(&ctx);
279
280 let event = make_event("link", EventOutcome::Success, Some(Uuid::new_v4()));
281 state = fold.reduce(state, &event, &ctx);
282
283 assert_eq!(state.total_events, 1);
284 assert!((state.relevance.alpha - 7.0).abs() < 1e-12); }
286
287 #[test]
288 fn feedback_not_useful_increments_entity_beta() {
289 let fold = BalancedRecallFold::new(100);
290 let ctx = FoldContext::new();
291 let mut state = fold.init(&ctx);
292
293 let id = Uuid::new_v4();
294 let mut event = make_event("brain.feedback", EventOutcome::Success, Some(id));
295 event.payload = serde_json::json!({"signal": "not_useful"});
296 state = fold.reduce(state, &event, &ctx);
297
298 assert_eq!(state.total_events, 1);
299 let ep = state.entity_posteriors.get(&id).unwrap();
300 assert!((ep.alpha - 1.0).abs() < 1e-12);
301 assert!((ep.beta - 2.0).abs() < 1e-12);
302 }
303
304 #[test]
305 fn brain_emit_legacy_does_not_update_entity() {
306 let fold = BalancedRecallFold::new(100);
308 let ctx = FoldContext::new();
309 let mut state = fold.init(&ctx);
310
311 let id = Uuid::new_v4();
312 let mut event = make_event("brain.emit", EventOutcome::Success, Some(id));
313 event.payload = serde_json::json!({"signal": "useful"});
314 state = fold.reduce(state, &event, &ctx);
315
316 assert_eq!(state.total_events, 1);
317 assert!(state.entity_posteriors.is_empty()); }
319
320 #[test]
321 fn deterministic_replay() {
322 let fold = BalancedRecallFold::new(100);
323 let ctx = FoldContext::new();
324
325 let id = Uuid::new_v4();
326 let events = vec![
327 make_event("recall", EventOutcome::Success, Some(id)),
328 make_event("recall", EventOutcome::Success, None),
329 make_event("search", EventOutcome::Success, None),
330 make_event("recall", EventOutcome::Success, Some(id)),
331 ];
332
333 let mut s1 = fold.init(&ctx);
334 for e in &events {
335 s1 = fold.reduce(s1, e, &ctx);
336 }
337
338 let mut s2 = fold.init(&ctx);
339 for e in &events {
340 s2 = fold.reduce(s2, e, &ctx);
341 }
342
343 let snap1 = s1.to_snapshot();
344 let snap2 = s2.to_snapshot();
345 assert_eq!(snap1.total_events, snap2.total_events);
346 assert_eq!(snap1.relevance, snap2.relevance);
347 assert_eq!(snap1.entity_posteriors, snap2.entity_posteriors);
348 }
349
350 fn make_semantic_feedback_event(signal: &str, target: Uuid) -> Event {
353 let mut e = Event::new(
354 "test",
355 "brain.feedback",
356 khive_types::EventKind::Audit,
357 SubstrateKind::Note,
358 "brain",
359 );
360 e.outcome = EventOutcome::Success;
361 e.target_id = Some(target);
362 e.payload = serde_json::json!({"signal": signal});
363 e
364 }
365
366 #[test]
367 fn semantic_feedback_explicit_positive_updates_salience_alpha_and_entity_alpha() {
368 let fold = BalancedRecallFold::new(100);
369 let ctx = FoldContext::new();
370 let state = fold.init(&ctx);
371
372 let sal_alpha_prior = state.salience.alpha; let sal_beta_prior = state.salience.beta; let id = Uuid::new_v4();
376 let event = make_semantic_feedback_event("explicit_positive", id);
377 let state = fold.reduce(state, &event, &ctx);
378
379 assert!(
381 (state.salience.alpha - (sal_alpha_prior + 1.5)).abs() < 1e-12,
382 "explicit_positive must add 1.5 to salience.alpha: expected {}, got {}",
383 sal_alpha_prior + 1.5,
384 state.salience.alpha
385 );
386 assert!(
387 (state.salience.beta - sal_beta_prior).abs() < 1e-12,
388 "explicit_positive must not change salience.beta"
389 );
390 let rel_beta_prior = state.relevance.beta;
392 let _ = rel_beta_prior;
394
395 let ep = state.entity_posteriors.get(&id).unwrap();
397 assert!(
398 (ep.alpha - 2.5).abs() < 1e-12,
399 "entity posterior alpha must be 1.0 + 1.5 = 2.5, got {}",
400 ep.alpha
401 );
402 assert!(
403 (ep.beta - 1.0).abs() < 1e-12,
404 "entity posterior beta must remain at 1.0, got {}",
405 ep.beta
406 );
407 }
408
409 #[test]
410 fn semantic_feedback_implicit_negative_updates_salience_beta_and_entity_beta() {
411 let fold = BalancedRecallFold::new(100);
412 let ctx = FoldContext::new();
413 let state = fold.init(&ctx);
414
415 let sal_alpha_prior = state.salience.alpha; let sal_beta_prior = state.salience.beta; let id = Uuid::new_v4();
419 let event = make_semantic_feedback_event("implicit_negative", id);
420 let state = fold.reduce(state, &event, &ctx);
421
422 assert!(
424 (state.salience.alpha - sal_alpha_prior).abs() < 1e-12,
425 "implicit_negative must not change salience.alpha"
426 );
427 assert!(
428 (state.salience.beta - (sal_beta_prior + 0.5)).abs() < 1e-12,
429 "implicit_negative must add 0.5 to salience.beta: expected {}, got {}",
430 sal_beta_prior + 0.5,
431 state.salience.beta
432 );
433
434 let ep = state.entity_posteriors.get(&id).unwrap();
436 assert!(
437 (ep.alpha - 1.0).abs() < 1e-12,
438 "entity posterior alpha must remain at 1.0, got {}",
439 ep.alpha
440 );
441 assert!(
442 (ep.beta - 1.5).abs() < 1e-12,
443 "entity posterior beta must be 1.0 + 0.5 = 1.5, got {}",
444 ep.beta
445 );
446 }
447
448 #[test]
449 fn semantic_feedback_correction_updates_salience_beta_relevance_beta_and_entity_beta() {
450 let fold = BalancedRecallFold::new(100);
451 let ctx = FoldContext::new();
452 let state = fold.init(&ctx);
453
454 let sal_alpha_prior = state.salience.alpha; let sal_beta_prior = state.salience.beta; let rel_alpha_prior = state.relevance.alpha; let rel_beta_prior = state.relevance.beta; let id = Uuid::new_v4();
460 let event = make_semantic_feedback_event("correction", id);
461 let state = fold.reduce(state, &event, &ctx);
462
463 assert!(
465 (state.salience.alpha - sal_alpha_prior).abs() < 1e-12,
466 "correction must not change salience.alpha"
467 );
468 assert!(
469 (state.salience.beta - (sal_beta_prior + 2.0)).abs() < 1e-12,
470 "correction must add 2.0 to salience.beta: expected {}, got {}",
471 sal_beta_prior + 2.0,
472 state.salience.beta
473 );
474
475 assert!(
477 (state.relevance.alpha - rel_alpha_prior).abs() < 1e-12,
478 "correction must not change relevance.alpha"
479 );
480 assert!(
481 (state.relevance.beta - (rel_beta_prior + 2.0)).abs() < 1e-12,
482 "correction must add 2.0 to relevance.beta: expected {}, got {}",
483 rel_beta_prior + 2.0,
484 state.relevance.beta
485 );
486
487 let ep = state.entity_posteriors.get(&id).unwrap();
489 assert!(
490 (ep.alpha - 1.0).abs() < 1e-12,
491 "entity posterior alpha must remain at 1.0, got {}",
492 ep.alpha
493 );
494 assert!(
495 (ep.beta - 3.0).abs() < 1e-12,
496 "entity posterior beta must be 1.0 + 2.0 = 3.0, got {}",
497 ep.beta
498 );
499 }
500
501 #[test]
505 fn test_355_posteriors_update_after_dispatch() {
506 let fold = BalancedRecallFold::new(100);
507 let ctx = FoldContext::new();
508 let state = fold.init(&ctx);
509
510 let sal_alpha_prior = state.salience.alpha; let sal_beta_prior = state.salience.beta; let tmp_alpha_prior = state.temporal.alpha; let tmp_beta_prior = state.temporal.beta; let id = Uuid::new_v4();
518 let mut fb_useful = make_event("brain.feedback", EventOutcome::Success, Some(id));
519 fb_useful.payload = serde_json::json!({"signal": "useful"});
520 let state = fold.reduce(state, &fb_useful, &ctx);
521
522 assert!(
523 (state.salience.alpha - (sal_alpha_prior + 1.0)).abs() < 1e-12,
524 "useful feedback must increment salience.alpha: expected {}, got {}",
525 sal_alpha_prior + 1.0,
526 state.salience.alpha
527 );
528 assert!(
529 (state.salience.beta - sal_beta_prior).abs() < 1e-12,
530 "useful feedback must not change salience.beta"
531 );
532
533 let mut hit = make_event("recall", EventOutcome::Success, Some(id));
535 hit.duration_us = 0;
536 let state = fold.reduce(state, &hit, &ctx);
537
538 assert!(
539 (state.temporal.alpha - (tmp_alpha_prior + 1.0)).abs() < 1e-12,
540 "fast recall hit must increment temporal.alpha: expected {}, got {}",
541 tmp_alpha_prior + 1.0,
542 state.temporal.alpha
543 );
544 assert!(
545 (state.temporal.beta - tmp_beta_prior).abs() < 1e-12,
546 "fast recall hit must not change temporal.beta"
547 );
548
549 let mut slow_hit = make_event("recall", EventOutcome::Success, Some(id));
551 slow_hit.duration_us = 100_000;
552 let state = fold.reduce(state, &slow_hit, &ctx);
553
554 assert!(
555 (state.temporal.beta - (tmp_beta_prior + 1.0)).abs() < 1e-12,
556 "slow recall hit must increment temporal.beta"
557 );
558
559 let mut fb_bad = make_event("brain.feedback", EventOutcome::Success, Some(id));
561 fb_bad.payload = serde_json::json!({"signal": "not_useful"});
562 let state = fold.reduce(state, &fb_bad, &ctx);
563
564 assert!(
565 (state.salience.beta - (sal_beta_prior + 1.0)).abs() < 1e-12,
566 "not_useful feedback must increment salience.beta"
567 );
568 }
569
570 use crate::state::SectionType as ST;
573
574 fn make_section_feedback_event(section_signals: serde_json::Value) -> Event {
575 let id = Uuid::new_v4();
576 let mut e = make_event("brain.feedback", EventOutcome::Success, Some(id));
577 e.payload = serde_json::json!({
578 "signal": "useful",
579 "section_signals": section_signals
580 });
581 e
582 }
583
584 #[test]
585 fn section_fold_useful_increments_alpha() {
586 let fold = SectionPosteriorFold::new();
587 let ctx = FoldContext::new();
588 let state = fold.init(&ctx);
589
590 let alpha_before = state.posteriors[&ST::Overview].alpha;
591
592 let event = make_section_feedback_event(serde_json::json!({
593 "overview": "useful"
594 }));
595 let state = fold.reduce(state, &event, &ctx);
596
597 assert!(
598 (state.posteriors[&ST::Overview].alpha - (alpha_before + 1.0)).abs() < 1e-12,
599 "useful must increment alpha"
600 );
601 }
602
603 #[test]
604 fn section_fold_not_useful_increments_beta() {
605 let fold = SectionPosteriorFold::new();
606 let ctx = FoldContext::new();
607 let state = fold.init(&ctx);
608
609 let beta_before = state.posteriors[&ST::Formalism].beta;
610
611 let event = make_section_feedback_event(serde_json::json!({
612 "formalism": "not_useful"
613 }));
614 let state = fold.reduce(state, &event, &ctx);
615
616 assert!(
617 (state.posteriors[&ST::Formalism].beta - (beta_before + 1.0)).abs() < 1e-12,
618 "not_useful must increment beta by 1"
619 );
620 }
621
622 #[test]
623 fn section_fold_wrong_increments_beta_by_two() {
624 let fold = SectionPosteriorFold::new();
625 let ctx = FoldContext::new();
626 let state = fold.init(&ctx);
627
628 let beta_before = state.posteriors[&ST::Examples].beta;
629
630 let event = make_section_feedback_event(serde_json::json!({
631 "examples": "wrong"
632 }));
633 let state = fold.reduce(state, &event, &ctx);
634
635 assert!(
636 (state.posteriors[&ST::Examples].beta - (beta_before + 2.0)).abs() < 1e-12,
637 "wrong must increment beta by 2"
638 );
639 }
640
641 #[test]
642 fn section_fold_no_section_signals_is_noop() {
643 let fold = SectionPosteriorFold::new();
644 let ctx = FoldContext::new();
645 let state = fold.init(&ctx);
646 let total_before = state.total_events;
647
648 let id = Uuid::new_v4();
650 let mut e = make_event("brain.feedback", EventOutcome::Success, Some(id));
651 e.payload = serde_json::json!({"signal": "useful"});
652 let state = fold.reduce(state, &e, &ctx);
653
654 assert_eq!(
655 state.total_events, total_before,
656 "no section_signals should be noop"
657 );
658 }
659
660 #[test]
661 fn section_fold_epoch_decrements() {
662 let fold = SectionPosteriorFold::new();
663 let ctx = FoldContext::new();
664 let state = fold.init(&ctx);
665 let epoch_before = state.exploration_epoch;
666
667 let event = make_section_feedback_event(serde_json::json!({
668 "overview": "useful"
669 }));
670 let state = fold.reduce(state, &event, &ctx);
671
672 assert_eq!(state.exploration_epoch, epoch_before - 1);
673 }
674
675 #[test]
676 fn section_fold_epoch_floors_at_zero() {
677 let fold = SectionPosteriorFold::new();
678 let ctx = FoldContext::new();
679 let mut state = fold.init(&ctx);
680 state.exploration_epoch = 0;
681
682 let event = make_section_feedback_event(serde_json::json!({
683 "overview": "useful"
684 }));
685 let state = fold.reduce(state, &event, &ctx);
686
687 assert_eq!(state.exploration_epoch, 0, "epoch must floor at 0");
688 }
689
690 #[test]
691 fn section_fold_deterministic_replay() {
692 let fold = SectionPosteriorFold::new();
693 let ctx = FoldContext::new();
694
695 let events = vec![
696 make_section_feedback_event(
697 serde_json::json!({"overview": "useful", "formalism": "not_useful"}),
698 ),
699 make_section_feedback_event(serde_json::json!({"examples": "wrong"})),
700 make_section_feedback_event(serde_json::json!({"overview": "useful"})),
701 ];
702
703 let mut s1 = fold.init(&ctx);
704 for e in &events {
705 s1 = fold.reduce(s1, e, &ctx);
706 }
707
708 let mut s2 = fold.init(&ctx);
709 for e in &events {
710 s2 = fold.reduce(s2, e, &ctx);
711 }
712
713 let snap1 = s1.to_snapshot();
714 let snap2 = s2.to_snapshot();
715 assert_eq!(snap1.total_events, snap2.total_events);
716 for st in ST::all() {
717 assert!(
718 (snap1.posteriors[st].alpha - snap2.posteriors[st].alpha).abs() < 1e-12,
719 "replay alpha mismatch for {:?}",
720 st
721 );
722 assert!(
723 (snap1.posteriors[st].beta - snap2.posteriors[st].beta).abs() < 1e-12,
724 "replay beta mismatch for {:?}",
725 st
726 );
727 }
728 }
729}