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}