Skip to main content

khive_pack_brain/
persist.rs

1use std::collections::{HashMap, HashSet};
2use std::sync::Mutex;
3
4use khive_fold::Fold;
5use khive_runtime::{KhiveRuntime, NamespaceToken, RuntimeError};
6use khive_storage::types::{SqlStatement, SqlValue};
7use khive_storage::SqlAccess;
8use serde_json::Value;
9
10use crate::state::{BrainState, BrainStateSnapshot};
11
12const SNAPSHOT_PROFILE_ID: &str = "__brain__";
13const DEFAULT_SNAPSHOT_BATCH_SIZE: u64 = 5;
14
15pub struct PersistenceTracker {
16    loaded_namespaces: HashSet<String>,
17    dirty_counts: HashMap<String, u64>,
18    snapshot_batch_size: u64,
19}
20
21impl Default for PersistenceTracker {
22    fn default() -> Self {
23        Self::new()
24    }
25}
26
27impl PersistenceTracker {
28    pub fn new() -> Self {
29        Self {
30            loaded_namespaces: HashSet::new(),
31            dirty_counts: HashMap::new(),
32            snapshot_batch_size: DEFAULT_SNAPSHOT_BATCH_SIZE,
33        }
34    }
35
36    pub fn is_loaded(&self, namespace: &str) -> bool {
37        self.loaded_namespaces.contains(namespace)
38    }
39
40    pub fn mark_loaded(&mut self, namespace: String) {
41        self.loaded_namespaces.insert(namespace);
42    }
43
44    pub fn increment_dirty(&mut self, namespace: &str) -> bool {
45        let count = self.dirty_counts.entry(namespace.to_string()).or_insert(0);
46        *count += 1;
47        *count >= self.snapshot_batch_size
48    }
49
50    pub fn reset_dirty(&mut self, namespace: &str) {
51        self.dirty_counts.insert(namespace.to_string(), 0);
52    }
53}
54
55fn sql_err(context: &str, e: impl std::fmt::Display) -> RuntimeError {
56    RuntimeError::Internal(format!("brain persistence {context}: {e}"))
57}
58
59pub async fn append_brain_event(
60    sql: &dyn SqlAccess,
61    namespace: &str,
62    profile_id: &str,
63    event_kind: &str,
64    payload: &Value,
65    created_at_us: i64,
66) -> Result<(), RuntimeError> {
67    let payload_str = serde_json::to_string(payload).map_err(|e| sql_err("serialize event", e))?;
68
69    let mut writer = sql.writer().await.map_err(|e| sql_err("writer", e))?;
70    writer
71        .execute(SqlStatement {
72            sql: "INSERT INTO brain_event_log (profile_id, namespace, event_kind, payload, created_at) VALUES (?1, ?2, ?3, ?4, ?5)".into(),
73            params: vec![
74                SqlValue::Text(profile_id.to_string()),
75                SqlValue::Text(namespace.to_string()),
76                SqlValue::Text(event_kind.to_string()),
77                SqlValue::Text(payload_str),
78                SqlValue::Integer(created_at_us),
79            ],
80            label: Some("brain_event_log_append".into()),
81        })
82        .await
83        .map_err(|e| sql_err("append event", e))?;
84
85    Ok(())
86}
87
88pub async fn upsert_snapshot(
89    sql: &dyn SqlAccess,
90    namespace: &str,
91    snapshot: &BrainStateSnapshot,
92    updated_at_us: i64,
93) -> Result<(), RuntimeError> {
94    let snapshot_json =
95        serde_json::to_string(snapshot).map_err(|e| sql_err("serialize snapshot", e))?;
96
97    let mut writer = sql.writer().await.map_err(|e| sql_err("writer", e))?;
98    writer
99        .execute(SqlStatement {
100            sql: "INSERT INTO brain_profile_snapshots (profile_id, namespace, snapshot_json, updated_at) VALUES (?1, ?2, ?3, ?4) ON CONFLICT(profile_id, namespace) DO UPDATE SET snapshot_json = excluded.snapshot_json, updated_at = excluded.updated_at".into(),
101            params: vec![
102                SqlValue::Text(SNAPSHOT_PROFILE_ID.to_string()),
103                SqlValue::Text(namespace.to_string()),
104                SqlValue::Text(snapshot_json),
105                SqlValue::Integer(updated_at_us),
106            ],
107            label: Some("brain_snapshot_upsert".into()),
108        })
109        .await
110        .map_err(|e| sql_err("upsert snapshot", e))?;
111
112    Ok(())
113}
114
115pub async fn load_latest_snapshot(
116    sql: &dyn SqlAccess,
117    namespace: &str,
118) -> Result<Option<(BrainStateSnapshot, i64)>, RuntimeError> {
119    let mut reader = sql.reader().await.map_err(|e| sql_err("reader", e))?;
120    let row = reader
121        .query_row(SqlStatement {
122            sql: "SELECT snapshot_json, updated_at FROM brain_profile_snapshots WHERE profile_id = ?1 AND namespace = ?2 ORDER BY updated_at DESC LIMIT 1".into(),
123            params: vec![
124                SqlValue::Text(SNAPSHOT_PROFILE_ID.to_string()),
125                SqlValue::Text(namespace.to_string()),
126            ],
127            label: Some("brain_snapshot_load".into()),
128        })
129        .await
130        .map_err(|e| sql_err("load snapshot", e))?;
131
132    match row {
133        None => Ok(None),
134        Some(row) => {
135            let json_str = match row.get("snapshot_json") {
136                Some(SqlValue::Text(s)) => s.clone(),
137                _ => return Err(sql_err("load snapshot", "missing snapshot_json column")),
138            };
139            let updated_at = match row.get("updated_at") {
140                Some(SqlValue::Integer(n)) => *n,
141                _ => return Err(sql_err("load snapshot", "missing updated_at column")),
142            };
143            let snapshot: BrainStateSnapshot =
144                serde_json::from_str(&json_str).map_err(|e| sql_err("deserialize snapshot", e))?;
145            Ok(Some((snapshot, updated_at)))
146        }
147    }
148}
149
150pub async fn load_events_since(
151    sql: &dyn SqlAccess,
152    namespace: &str,
153    since_us: i64,
154) -> Result<Vec<khive_storage::event::Event>, RuntimeError> {
155    let mut reader = sql.reader().await.map_err(|e| sql_err("reader", e))?;
156    let rows = reader
157        .query_all(SqlStatement {
158            sql: "SELECT payload FROM brain_event_log WHERE namespace = ?1 AND created_at > ?2 ORDER BY created_at ASC, id ASC".into(),
159            params: vec![
160                SqlValue::Text(namespace.to_string()),
161                SqlValue::Integer(since_us),
162            ],
163            label: Some("brain_events_replay".into()),
164        })
165        .await
166        .map_err(|e| sql_err("load events", e))?;
167
168    let mut events = Vec::with_capacity(rows.len());
169    for row in &rows {
170        let payload_str = match row.get("payload") {
171            Some(SqlValue::Text(s)) => s,
172            _ => continue,
173        };
174        match serde_json::from_str::<khive_storage::event::Event>(payload_str) {
175            Ok(event) => events.push(event),
176            Err(_) => continue,
177        }
178    }
179    Ok(events)
180}
181
182pub async fn ensure_loaded(
183    runtime: &KhiveRuntime,
184    token: &NamespaceToken,
185    tracker: &Mutex<PersistenceTracker>,
186    state: &Mutex<BrainState>,
187    fold: &crate::fold::BalancedRecallFold,
188    section_fold: &crate::fold::SectionPosteriorFold,
189    entity_capacity: usize,
190) -> Result<(), RuntimeError> {
191    let namespace = token.namespace().as_str().to_string();
192
193    {
194        let t = tracker.lock().unwrap();
195        if t.is_loaded(&namespace) {
196            return Ok(());
197        }
198    }
199
200    let sql = runtime.sql();
201
202    let snapshot_result = load_latest_snapshot(sql.as_ref(), &namespace).await?;
203
204    if let Some((snapshot, updated_at)) = snapshot_result {
205        let replay_events = load_events_since(sql.as_ref(), &namespace, updated_at).await?;
206
207        let ctx = khive_fold::FoldContext::new();
208        let mut brain_state = BrainState::from_snapshot(snapshot, entity_capacity);
209
210        for event in &replay_events {
211            let current = std::mem::replace(
212                &mut brain_state.balanced_recall,
213                crate::state::BalancedRecallState::new(0),
214            );
215            brain_state.balanced_recall = fold.reduce(current, event, &ctx);
216
217            let serving_profile = event
218                .payload
219                .get("served_by_profile_id")
220                .and_then(|v| v.as_str())
221                .unwrap_or("balanced-recall-v1");
222
223            if let Some(section_state) = brain_state.section_states.remove(serving_profile) {
224                let updated = section_fold.reduce(section_state, event, &ctx);
225                brain_state
226                    .section_states
227                    .insert(serving_profile.to_string(), updated);
228            }
229        }
230
231        crate::sync_balanced_recall_record(&mut brain_state);
232
233        {
234            let mut s = state.lock().unwrap();
235            *s = brain_state;
236        }
237    }
238
239    {
240        let mut t = tracker.lock().unwrap();
241        t.mark_loaded(namespace);
242    }
243
244    Ok(())
245}
246
247pub async fn persist_after_feedback(
248    runtime: &KhiveRuntime,
249    token: &NamespaceToken,
250    tracker: &Mutex<PersistenceTracker>,
251    state: &Mutex<BrainState>,
252    event: &khive_storage::event::Event,
253    serving_profile: &str,
254) -> Result<(), RuntimeError> {
255    let namespace = token.namespace().as_str().to_string();
256    let now_us = chrono::Utc::now().timestamp_micros();
257
258    let sql = runtime.sql();
259
260    let event_payload = serde_json::to_value(event).map_err(|e| sql_err("serialize event", e))?;
261
262    append_brain_event(
263        sql.as_ref(),
264        &namespace,
265        serving_profile,
266        &event.verb,
267        &event_payload,
268        now_us,
269    )
270    .await?;
271
272    let should_snapshot = {
273        let mut t = tracker.lock().unwrap();
274        t.increment_dirty(&namespace)
275    };
276
277    if should_snapshot {
278        let snapshot = {
279            let s = state.lock().unwrap();
280            s.to_snapshot()
281        };
282
283        upsert_snapshot(sql.as_ref(), &namespace, &snapshot, now_us).await?;
284
285        let mut t = tracker.lock().unwrap();
286        t.reset_dirty(&namespace);
287    }
288
289    Ok(())
290}