Skip to main content

khive_pack_memory/
tunable.rs

1use 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
9/// `MemoryPack` implements `PackTunable` so that the brain can adjust the
10/// recall scoring pipeline based on observed usage patterns (Issue #159).
11///
12/// Parameter names (`memory::relevance_weight`, `memory::salience_weight`,
13/// `memory::temporal_weight`) correspond to the three Beta posteriors in
14/// `BalancedRecallState` (ADR-032 §5a). Posterior means flow directly into
15/// `RecallConfig`.
16///
17/// `project_config` reads posterior means → `RecallConfig`.
18/// `apply_config` validates and stores the new config; future recall calls
19/// pick it up via `MemoryPack::active_config()`.
20impl PackTunable for MemoryPack {
21    fn parameter_space(&self) -> ParameterSpace {
22        ParameterSpace {
23            parameters: vec![
24                ParameterDef {
25                    name: "memory::relevance_weight".into(),
26                    // Prior: relevance is the dominant signal (7:3), matching
27                    // BalancedRecallState's `relevance` posterior prior.
28                    prior_alpha: 7.0,
29                    prior_beta: 3.0,
30                    bounds: (0.0, 1.0),
31                },
32                ParameterDef {
33                    name: "memory::salience_weight".into(),
34                    // Prior: salience is secondary (2:8).
35                    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: temporal is weakest signal (1:9).
42                    prior_alpha: 1.0,
43                    prior_beta: 9.0,
44                    bounds: (0.0, 1.0),
45                },
46            ],
47        }
48    }
49
50    /// Project the current `BalancedRecallState` posteriors into a `RecallConfig` value.
51    ///
52    /// Reads the three posterior means from the profile state. Falls back to the
53    /// current active config if a parameter is absent (brain not yet warmed up).
54    fn project_config(&self, state: &BalancedRecallState) -> Value {
55        let current = self.active_config();
56
57        let relevance = state.relevance.mean();
58        let salience = state.salience.mean();
59        let temporal = state.temporal.mean();
60
61        let projected = RecallConfig {
62            relevance_weight: relevance,
63            salience_weight: salience,
64            temporal_weight: temporal,
65            ..current
66        };
67
68        serde_json::to_value(projected).unwrap_or_else(|_| serde_json::json!({}))
69    }
70
71    /// Apply a projected config to the pack.
72    ///
73    /// Deserializes the JSON value into a `RecallConfig`, validates it, and
74    /// stores it as the active config. Future recall calls pick up the new
75    /// weights via `MemoryPack::active_config()`.
76    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        salience_mean: f64,
99        temporal_mean: f64,
100    ) -> BalancedRecallState {
101        // Construct Beta posteriors whose means match the supplied values.
102        // Using ESS=10 for each: alpha = mean * 10, beta = (1-mean) * 10.
103        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.salience = to_posterior(salience_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::salience_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.salience_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        // Default BalancedRecallState priors: Beta(7,3)=0.7, Beta(2,8)=0.2, Beta(1,9)=0.1
138        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.salience_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            salience_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.salience_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            salience_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        // BalancedRecallState uses Beta(7,3) for relevance; ParameterDef must match.
190        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}