khive_pack_memory/
tunable.rs1use khive_pack_brain::state::BalancedRecallState;
2use khive_pack_brain::tunable::{PackTunable, ParameterDef, ParameterSpace};
3use khive_runtime::RuntimeError;
4use serde_json::Value;
5
6use crate::config::RecallConfig;
7use crate::MemoryPack;
8
9impl PackTunable for MemoryPack {
21 fn parameter_space(&self) -> ParameterSpace {
22 ParameterSpace {
23 parameters: vec![
24 ParameterDef {
25 name: "memory::relevance_weight".into(),
26 prior_alpha: 7.0,
29 prior_beta: 3.0,
30 bounds: (0.0, 1.0),
31 },
32 ParameterDef {
33 name: "memory::importance_weight".into(),
34 prior_alpha: 2.0,
36 prior_beta: 8.0,
37 bounds: (0.0, 1.0),
38 },
39 ParameterDef {
40 name: "memory::temporal_weight".into(),
41 prior_alpha: 1.0,
43 prior_beta: 9.0,
44 bounds: (0.0, 1.0),
45 },
46 ],
47 }
48 }
49
50 fn project_config(&self, state: &BalancedRecallState) -> Value {
55 let current = self.active_config();
56
57 let relevance = state.relevance.mean();
58 let importance = state.importance.mean();
59 let temporal = state.temporal.mean();
60
61 let projected = RecallConfig {
62 relevance_weight: relevance,
63 importance_weight: importance,
64 temporal_weight: temporal,
65 ..current
66 };
67
68 serde_json::to_value(projected).unwrap_or_else(|_| serde_json::json!({}))
69 }
70
71 fn apply_config(&self, config: Value) -> Result<(), RuntimeError> {
77 let new_cfg: RecallConfig = serde_json::from_value(config)
78 .map_err(|e| RuntimeError::InvalidInput(format!("invalid RecallConfig: {e}")))?;
79 new_cfg.validate()?;
80 *self.config.lock().unwrap() = new_cfg;
81 Ok(())
82 }
83}
84
85#[cfg(test)]
86mod tests {
87 use super::*;
88 use khive_pack_brain::state::{BalancedRecallState, BetaPosterior};
89 use khive_runtime::KhiveRuntime;
90
91 fn make_pack() -> MemoryPack {
92 let rt = KhiveRuntime::memory().expect("in-memory runtime");
93 MemoryPack::new(rt)
94 }
95
96 fn balanced_state_with_means(
97 relevance_mean: f64,
98 importance_mean: f64,
99 temporal_mean: f64,
100 ) -> BalancedRecallState {
101 let to_posterior =
104 |mean: f64| -> BetaPosterior { BetaPosterior::new(mean * 10.0, (1.0 - mean) * 10.0) };
105 let mut state = BalancedRecallState::new(100);
106 state.relevance = to_posterior(relevance_mean);
107 state.importance = to_posterior(importance_mean);
108 state.temporal = to_posterior(temporal_mean);
109 state
110 }
111
112 #[test]
113 fn parameter_space_has_three_params() {
114 let pack = make_pack();
115 let space = pack.parameter_space();
116 assert_eq!(space.parameters.len(), 3);
117 let names: Vec<&str> = space.parameters.iter().map(|p| p.name.as_str()).collect();
118 assert!(names.contains(&"memory::relevance_weight"));
119 assert!(names.contains(&"memory::importance_weight"));
120 assert!(names.contains(&"memory::temporal_weight"));
121 }
122
123 #[test]
124 fn project_config_reads_posterior_means() {
125 let pack = make_pack();
126 let state = balanced_state_with_means(0.6, 0.3, 0.1);
127 let projected = pack.project_config(&state);
128
129 let cfg: RecallConfig = serde_json::from_value(projected).unwrap();
130 assert!((cfg.relevance_weight - 0.6).abs() < 1e-10);
131 assert!((cfg.importance_weight - 0.3).abs() < 1e-10);
132 assert!((cfg.temporal_weight - 0.1).abs() < 1e-10);
133 }
134
135 #[test]
136 fn project_config_with_default_priors_matches_expected_defaults() {
137 let pack = make_pack();
139 let state = BalancedRecallState::new(100);
140 let projected = pack.project_config(&state);
141
142 let cfg: RecallConfig = serde_json::from_value(projected).unwrap();
143 assert!((cfg.relevance_weight - 0.70).abs() < 1e-10);
144 assert!((cfg.importance_weight - 0.20).abs() < 1e-10);
145 assert!((cfg.temporal_weight - 0.10).abs() < 1e-10);
146 }
147
148 #[test]
149 fn apply_config_updates_active_config() {
150 let pack = make_pack();
151 let new_cfg = RecallConfig {
152 relevance_weight: 0.5,
153 importance_weight: 0.3,
154 temporal_weight: 0.2,
155 ..RecallConfig::default()
156 };
157 let config_value = serde_json::to_value(&new_cfg).unwrap();
158 pack.apply_config(config_value)
159 .expect("apply_config succeeds");
160
161 let active = pack.active_config();
162 assert!((active.relevance_weight - 0.5).abs() < 1e-10);
163 assert!((active.importance_weight - 0.3).abs() < 1e-10);
164 assert!((active.temporal_weight - 0.2).abs() < 1e-10);
165 }
166
167 #[test]
168 fn apply_config_rejects_all_zero_weights() {
169 let pack = make_pack();
170 let bad_cfg = RecallConfig {
171 relevance_weight: 0.0,
172 importance_weight: 0.0,
173 temporal_weight: 0.0,
174 ..RecallConfig::default()
175 };
176 let config_value = serde_json::to_value(&bad_cfg).unwrap();
177 assert!(pack.apply_config(config_value).is_err());
178 }
179
180 #[test]
181 fn apply_config_rejects_malformed_json() {
182 let pack = make_pack();
183 let bad = serde_json::json!({ "relevance_weight": "not_a_number" });
184 assert!(pack.apply_config(bad).is_err());
185 }
186
187 #[test]
188 fn prior_for_relevance_weight_matches_balanced_recall_state_prior() {
189 let pack = make_pack();
191 let space = pack.parameter_space();
192 let def = space
193 .parameters
194 .iter()
195 .find(|p| p.name == "memory::relevance_weight")
196 .unwrap();
197 assert!((def.prior_alpha - 7.0).abs() < 1e-12);
198 assert!((def.prior_beta - 3.0).abs() < 1e-12);
199 }
200}