Skip to main content

khive_pack_memory/
tunable.rs

1//! Brain-tunable parameter surface for the memory pack's recall scoring pipeline.
2
3use khive_brain_core::BalancedRecallState;
4use khive_brain_core::{PackTunable, ParameterDef, ParameterSpace};
5use khive_runtime::RuntimeError;
6use serde_json::Value;
7
8use crate::config::RecallConfig;
9use crate::MemoryPack;
10
11/// `MemoryPack` implements `PackTunable` so that the brain can adjust the
12/// recall scoring pipeline based on observed usage patterns (Issue #159).
13///
14/// Parameter names (`memory::relevance_weight`, `memory::salience_weight`,
15/// `memory::temporal_weight`) correspond to the three Beta posteriors in
16/// `BalancedRecallState`. Posterior means flow directly into
17/// `RecallConfig`.
18///
19/// `project_config` reads posterior means → `RecallConfig`.
20/// `apply_config` validates and stores the new config; future recall calls
21/// pick it up via `MemoryPack::active_config()`.
22impl PackTunable for MemoryPack {
23    fn parameter_space(&self) -> ParameterSpace {
24        ParameterSpace {
25            parameters: vec![
26                ParameterDef {
27                    name: "memory::relevance_weight".into(),
28                    // Prior: relevance is the dominant signal (7:3), matching
29                    // BalancedRecallState's `relevance` posterior prior.
30                    prior_alpha: 7.0,
31                    prior_beta: 3.0,
32                    bounds: (0.0, 1.0),
33                },
34                ParameterDef {
35                    name: "memory::salience_weight".into(),
36                    // Prior: salience is secondary (2:8).
37                    prior_alpha: 2.0,
38                    prior_beta: 8.0,
39                    bounds: (0.0, 1.0),
40                },
41                ParameterDef {
42                    name: "memory::temporal_weight".into(),
43                    // Prior: temporal is weakest signal (1:9).
44                    prior_alpha: 1.0,
45                    prior_beta: 9.0,
46                    bounds: (0.0, 1.0),
47                },
48            ],
49        }
50    }
51
52    /// Project the current `BalancedRecallState` posteriors into a `RecallConfig` value.
53    ///
54    /// Reads the three posterior means from the profile state. Falls back to the
55    /// current active config if a parameter is absent (brain not yet warmed up).
56    fn project_config(&self, state: &BalancedRecallState) -> Value {
57        let current = self.active_config();
58
59        let relevance = state.relevance.mean();
60        let salience = state.salience.mean();
61        let temporal = state.temporal.mean();
62
63        let projected = RecallConfig {
64            relevance_weight: relevance,
65            salience_weight: salience,
66            temporal_weight: temporal,
67            ..current
68        };
69
70        serde_json::to_value(projected).unwrap_or_else(|_| serde_json::json!({}))
71    }
72
73    /// Apply a projected config to the pack.
74    ///
75    /// Deserializes the JSON value into a `RecallConfig`, validates it, and
76    /// stores it as the active config. Future recall calls pick up the new
77    /// weights via `MemoryPack::active_config()`.
78    fn apply_config(&self, config: Value) -> Result<(), RuntimeError> {
79        let new_cfg: RecallConfig = serde_json::from_value(config)
80            .map_err(|e| RuntimeError::InvalidInput(format!("invalid RecallConfig: {e}")))?;
81        new_cfg.validate()?;
82        *self.config.lock().unwrap() = new_cfg;
83        Ok(())
84    }
85}
86
87#[cfg(test)]
88mod tests {
89    use super::*;
90    use khive_brain_core::{BalancedRecallState, BetaPosterior};
91    use khive_runtime::KhiveRuntime;
92
93    fn make_pack() -> MemoryPack {
94        let rt = KhiveRuntime::memory().expect("in-memory runtime");
95        MemoryPack::new(rt)
96    }
97
98    fn balanced_state_with_means(
99        relevance_mean: f64,
100        salience_mean: f64,
101        temporal_mean: f64,
102    ) -> BalancedRecallState {
103        // Construct Beta posteriors whose means match the supplied values.
104        // Using ESS=10 for each: alpha = mean * 10, beta = (1-mean) * 10.
105        let to_posterior =
106            |mean: f64| -> BetaPosterior { BetaPosterior::new(mean * 10.0, (1.0 - mean) * 10.0) };
107        let mut state = BalancedRecallState::new(100);
108        state.relevance = to_posterior(relevance_mean);
109        state.salience = to_posterior(salience_mean);
110        state.temporal = to_posterior(temporal_mean);
111        state
112    }
113
114    #[test]
115    fn parameter_space_has_three_params() {
116        let pack = make_pack();
117        let space = pack.parameter_space();
118        assert_eq!(space.parameters.len(), 3);
119        let names: Vec<&str> = space.parameters.iter().map(|p| p.name.as_str()).collect();
120        assert!(names.contains(&"memory::relevance_weight"));
121        assert!(names.contains(&"memory::salience_weight"));
122        assert!(names.contains(&"memory::temporal_weight"));
123    }
124
125    #[test]
126    fn project_config_reads_posterior_means() {
127        let pack = make_pack();
128        let state = balanced_state_with_means(0.6, 0.3, 0.1);
129        let projected = pack.project_config(&state);
130
131        let cfg: RecallConfig = serde_json::from_value(projected).unwrap();
132        assert!((cfg.relevance_weight - 0.6).abs() < 1e-10);
133        assert!((cfg.salience_weight - 0.3).abs() < 1e-10);
134        assert!((cfg.temporal_weight - 0.1).abs() < 1e-10);
135    }
136
137    #[test]
138    fn project_config_with_default_priors_matches_expected_defaults() {
139        // Default BalancedRecallState priors: Beta(7,3)=0.7, Beta(2,8)=0.2, Beta(1,9)=0.1
140        let pack = make_pack();
141        let state = BalancedRecallState::new(100);
142        let projected = pack.project_config(&state);
143
144        let cfg: RecallConfig = serde_json::from_value(projected).unwrap();
145        assert!((cfg.relevance_weight - 0.70).abs() < 1e-10);
146        assert!((cfg.salience_weight - 0.20).abs() < 1e-10);
147        assert!((cfg.temporal_weight - 0.10).abs() < 1e-10);
148    }
149
150    #[test]
151    fn apply_config_updates_active_config() {
152        let pack = make_pack();
153        let new_cfg = RecallConfig {
154            relevance_weight: 0.5,
155            salience_weight: 0.3,
156            temporal_weight: 0.2,
157            ..RecallConfig::default()
158        };
159        let config_value = serde_json::to_value(&new_cfg).unwrap();
160        pack.apply_config(config_value)
161            .expect("apply_config succeeds");
162
163        let active = pack.active_config();
164        assert!((active.relevance_weight - 0.5).abs() < 1e-10);
165        assert!((active.salience_weight - 0.3).abs() < 1e-10);
166        assert!((active.temporal_weight - 0.2).abs() < 1e-10);
167    }
168
169    #[test]
170    fn apply_config_rejects_all_zero_weights() {
171        let pack = make_pack();
172        let bad_cfg = RecallConfig {
173            relevance_weight: 0.0,
174            salience_weight: 0.0,
175            temporal_weight: 0.0,
176            ..RecallConfig::default()
177        };
178        let config_value = serde_json::to_value(&bad_cfg).unwrap();
179        assert!(pack.apply_config(config_value).is_err());
180    }
181
182    #[test]
183    fn apply_config_rejects_malformed_json() {
184        let pack = make_pack();
185        let bad = serde_json::json!({ "relevance_weight": "not_a_number" });
186        assert!(pack.apply_config(bad).is_err());
187    }
188
189    #[test]
190    fn prior_for_relevance_weight_matches_balanced_recall_state_prior() {
191        // BalancedRecallState uses Beta(7,3) for relevance; ParameterDef must match.
192        let pack = make_pack();
193        let space = pack.parameter_space();
194        let def = space
195            .parameters
196            .iter()
197            .find(|p| p.name == "memory::relevance_weight")
198            .unwrap();
199        assert!((def.prior_alpha - 7.0).abs() < 1e-12);
200        assert!((def.prior_beta - 3.0).abs() < 1e-12);
201    }
202}