khive_pack_memory/
tunable.rs1use 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
11impl PackTunable for MemoryPack {
23 fn parameter_space(&self) -> ParameterSpace {
24 ParameterSpace {
25 parameters: vec![
26 ParameterDef {
27 name: "memory::relevance_weight".into(),
28 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_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_alpha: 1.0,
45 prior_beta: 9.0,
46 bounds: (0.0, 1.0),
47 },
48 ],
49 }
50 }
51
52 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 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 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 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 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}