Skip to main content

kcode_k1_chat_persistence/
lib.rs

1use kcode_k1_chat_persistence_store::Projection;
2pub use kcode_k1_chat_persistence_store::{
3    Batch, EventRecord, Record, SessionId, SessionLog, TxId,
4};
5use kcode_k1_peering::K1Peering;
6use kcode_k1_txn_ordering::{K1TxnOrdering, Subsystem, SubsystemId};
7use std::collections::HashMap;
8use std::fs::{self, File, OpenOptions};
9use std::io::Write;
10use std::path::{Path, PathBuf};
11use std::sync::{Arc, Mutex, MutexGuard};
12
13const CURSOR: &str = "cursor";
14const CURSOR_TEMP: &str = "cursor.tmp";
15const SUBSYSTEM: &str = "k1-chat-persist";
16
17pub struct K1ChatPersistence {
18    driver: Arc<Driver>,
19    peering: Arc<K1Peering>,
20}
21
22#[derive(Clone)]
23pub struct Session {
24    id: SessionId,
25    driver: Arc<Driver>,
26    peering: Arc<K1Peering>,
27}
28
29struct Driver {
30    state: Mutex<State>,
31}
32
33struct State {
34    root: PathBuf,
35    projection: Projection,
36    cursor: Option<TxId>,
37    first_callback: bool,
38    pending: HashMap<SessionId, Pending>,
39    poison: Option<String>,
40}
41
42struct Pending {
43    payload: Vec<u8>,
44    evidence: Option<TxId>,
45}
46
47impl K1ChatPersistence {
48    pub fn open(
49        root: &Path,
50        ordering: Arc<K1TxnOrdering>,
51        peering: Arc<K1Peering>,
52    ) -> Result<Self, String> {
53        fs::create_dir_all(root).map_err(|e| format!("cannot create persistence root: {e}"))?;
54        let temp = root.join(CURSOR_TEMP);
55        if remove_optional(&temp).map_err(|e| format!("cannot remove stale cursor temp: {e}"))? {
56            sync_directory(root).map_err(|e| format!("cannot sync stale-temp removal: {e}"))?;
57        }
58        let projection =
59            Projection::new(root).map_err(|e| format!("cannot open session projection: {e}"))?;
60        let cursor = read_cursor(&root.join(CURSOR))?;
61        let driver = Arc::new(Driver {
62            state: Mutex::new(State {
63                root: root.to_path_buf(),
64                projection,
65                cursor,
66                first_callback: true,
67                pending: HashMap::new(),
68                poison: None,
69            }),
70        });
71        ordering
72            .register_subsystem(subsystem_id()?, cursor, driver.clone())
73            .map_err(|e| format!("cannot register {SUBSYSTEM}: {e}"))?;
74        Ok(Self { driver, peering })
75    }
76
77    pub fn session(&self, id: SessionId) -> Result<(Session, SessionLog), String> {
78        let session = Session {
79            id,
80            driver: self.driver.clone(),
81            peering: self.peering.clone(),
82        };
83        let log = session.load()?;
84        Ok((session, log))
85    }
86}
87
88impl Session {
89    pub fn id(&self) -> SessionId {
90        self.id
91    }
92
93    pub fn load(&self) -> Result<SessionLog, String> {
94        let state = self.driver.lock()?;
95        ensure_ready(&state)?;
96        state
97            .projection
98            .load(self.id)
99            .map_err(|e| format!("cannot load session: {e}"))
100    }
101
102    pub fn persist(&self, records: Vec<Record>) -> Result<TxId, String> {
103        if records.is_empty() {
104            return Err("cannot persist an empty record batch".to_owned());
105        }
106        let payload = {
107            let mut state = self.driver.lock()?;
108            ensure_ready(&state)?;
109            if state.pending.contains_key(&self.id) {
110                return Err("this session already has a local persist in flight".to_owned());
111            }
112            let predecessor = state
113                .projection
114                .load(self.id)
115                .map_err(|e| format!("cannot load session tail: {e}"))?
116                .records
117                .last()
118                .cloned();
119            let payload = Batch::new(self.id, predecessor, records)
120                .and_then(|batch| batch.encode())
121                .map_err(|e| format!("cannot encode persistence batch: {e}"))?;
122            state.pending.insert(
123                self.id,
124                Pending {
125                    payload: payload.clone(),
126                    evidence: None,
127                },
128            );
129            payload
130        };
131
132        let submitted = self.peering.submit_txn(subsystem_id()?, &payload);
133        self.finish_submit(submitted)
134    }
135
136    fn finish_submit(&self, submitted: Result<TxId, String>) -> Result<TxId, String> {
137        let mut state = self.driver.lock()?;
138        ensure_ready(&state)?;
139        let pending = match state.pending.remove(&self.id) {
140            Some(value) => value,
141            None => {
142                return Err(poison(
143                    &mut state,
144                    "local callback evidence disappeared".into(),
145                ));
146            }
147        };
148        match (submitted, pending.evidence) {
149            (Ok(returned), Some(seen)) if returned == seen => Ok(returned),
150            (Err(_), Some(seen)) => Ok(seen),
151            (Err(error), None) => Err(error),
152            (Ok(returned), Some(seen)) => Err(poison(
153                &mut state,
154                format!(
155                    "callback transaction mismatch: submission returned {returned:?}, callback saw {seen:?}"
156                ),
157            )),
158            (Ok(returned), None) => Err(poison(
159                &mut state,
160                format!("missing callback evidence for returned {returned:?}"),
161            )),
162        }
163    }
164}
165
166impl Driver {
167    fn lock(&self) -> Result<MutexGuard<'_, State>, String> {
168        self.state
169            .lock()
170            .map_err(|_| "persistence state mutex is poisoned".to_owned())
171    }
172
173    fn fault(&self, message: String) -> String {
174        match self.lock() {
175            Ok(mut state) => poison(&mut state, message),
176            Err(lock_error) => format!("{message}; {lock_error}"),
177        }
178    }
179}
180
181impl Subsystem for Driver {
182    fn submit_txn(&self, id: TxId, payload: &[u8]) -> Result<(), String> {
183        let batch = Batch::decode(payload)
184            .map_err(|e| self.fault(format!("cannot decode callback {id:?}: {e}")))?;
185        let mut state = self.lock()?;
186        ensure_ready(&state)?;
187        let reconcile = state.first_callback;
188        if let Err(error) = state.projection.apply(id, &batch, reconcile) {
189            return Err(poison(
190                &mut state,
191                format!("cannot apply callback {id:?}: {error}"),
192            ));
193        }
194        if let Err(error) = replace_cursor(&state.root, id) {
195            return Err(poison(
196                &mut state,
197                format!("cannot persist callback cursor {id:?}: {error}"),
198            ));
199        }
200        state.cursor = Some(id);
201        state.first_callback = false;
202        if let Some(pending) = state.pending.get_mut(&batch.session_id)
203            && pending.payload == payload
204            && pending.evidence.replace(id).is_some()
205        {
206            return Err(poison(
207                &mut state,
208                format!("duplicate local callback evidence for {id:?}"),
209            ));
210        }
211        Ok(())
212    }
213
214    fn reorg(&self) -> Result<(), String> {
215        let mut state = self.lock()?;
216        let mut failures = Vec::new();
217        if let Err(error) = state.projection.discard_all() {
218            failures.push(format!("discard projection: {error}"));
219        }
220        for (label, path) in [
221            (CURSOR, state.root.join(CURSOR)),
222            (CURSOR_TEMP, state.root.join(CURSOR_TEMP)),
223        ] {
224            if let Err(error) = remove_optional(&path) {
225                failures.push(format!("remove {label}: {error}"));
226            }
227        }
228        if let Err(error) = sync_directory(&state.root) {
229            failures.push(format!("sync persistence root: {error}"));
230        }
231        state.cursor.take();
232        state.pending.clear();
233        let mut message = "canonical reorganization faulted the persistence driver".to_owned();
234        if !failures.is_empty() {
235            message.push_str("; cleanup failures: ");
236            message.push_str(&failures.join("; "));
237        }
238        state.poison = Some(message.clone());
239        Err(message)
240    }
241}
242
243fn subsystem_id() -> Result<SubsystemId, String> {
244    SubsystemId::from_str(SUBSYSTEM)
245}
246
247fn ensure_ready(state: &State) -> Result<(), String> {
248    match &state.poison {
249        Some(reason) => Err(format!("persistence driver is faulted: {reason}")),
250        None => Ok(()),
251    }
252}
253
254fn poison(state: &mut State, message: String) -> String {
255    state.pending.clear();
256    if state.poison.is_none() {
257        state.poison = Some(message);
258    }
259    state.poison.clone().expect("poison was just installed")
260}
261
262fn read_cursor(path: &Path) -> Result<Option<TxId>, String> {
263    let bytes = match fs::read(path) {
264        Ok(bytes) => bytes,
265        Err(error) if error.kind() == std::io::ErrorKind::NotFound => return Ok(None),
266        Err(error) => return Err(format!("cannot read cursor: {error}")),
267    };
268    let exact: [u8; 12] = bytes
269        .try_into()
270        .map_err(|_| "cursor must contain exactly 12 bytes".to_owned())?;
271    Ok(Some(TxId::from_bytes(exact)))
272}
273
274fn replace_cursor(root: &Path, id: TxId) -> Result<(), String> {
275    let temp = root.join(CURSOR_TEMP);
276    let mut file = OpenOptions::new()
277        .create(true)
278        .truncate(true)
279        .write(true)
280        .open(&temp)
281        .map_err(|e| format!("open cursor temp: {e}"))?;
282    file.write_all(id.as_bytes())
283        .map_err(|e| format!("write cursor temp: {e}"))?;
284    file.sync_all()
285        .map_err(|e| format!("sync cursor temp: {e}"))?;
286    fs::rename(&temp, root.join(CURSOR)).map_err(|e| format!("rename cursor temp: {e}"))?;
287    sync_directory(root).map_err(|e| format!("sync cursor parent: {e}"))
288}
289
290fn remove_optional(path: &Path) -> std::io::Result<bool> {
291    match fs::remove_file(path) {
292        Ok(()) => Ok(true),
293        Err(error) if error.kind() == std::io::ErrorKind::NotFound => Ok(false),
294        Err(error) => Err(error),
295    }
296}
297
298fn sync_directory(path: &Path) -> std::io::Result<()> {
299    File::open(path)?.sync_all()
300}
301
302#[cfg(test)]
303mod tests {
304    use super::*;
305    use kcode_k1_chat_chatend::{
306        BoxId, ChatBox, ProviderCall, ToolCallId, ToolMessageMetadata, ToolResultV2Metadata,
307        tool_call_box, tool_message_box, tool_result_v2_box,
308    };
309    use std::sync::atomic::{AtomicU64, Ordering};
310
311    static NEXT: AtomicU64 = AtomicU64::new(0);
312
313    struct Roots(PathBuf);
314
315    impl Roots {
316        fn new() -> Self {
317            let path = std::env::temp_dir().join(format!(
318                "k1-chat-persistence-{}-{}",
319                std::process::id(),
320                NEXT.fetch_add(1, Ordering::Relaxed)
321            ));
322            let _ = fs::remove_dir_all(&path);
323            Self(path)
324        }
325    }
326
327    impl Drop for Roots {
328        fn drop(&mut self) {
329            let _ = fs::remove_dir_all(&self.0);
330        }
331    }
332
333    fn generic_history() -> Vec<Record> {
334        Batch::decode(
335            br#"{"version":2,"session_id":"090909090909090909090909","predecessor":null,"records":[{"kind":"box","id":1,"type":"Future Kind","contents":"opaque","hidden_type":"future/v9","hidden_contents":"hidden bytes"},{"kind":"box","id":2,"type":"Future Kind","contents":"next","hidden_type":"","hidden_contents":""}]}"#,
336        )
337        .unwrap()
338        .records
339    }
340
341    fn numbered_box(id: u64, value: ChatBox) -> Record {
342        Record::Box(ChatBox::new(
343            BoxId::new(id),
344            value.box_type().to_owned(),
345            value.contents().to_owned(),
346            value.hidden_type().to_owned(),
347            value.hidden_contents().to_owned(),
348        ))
349    }
350
351    fn stack(roots: &Roots) -> (Arc<K1TxnOrdering>, Arc<K1Peering>) {
352        let ordering = Arc::new(K1TxnOrdering::open(&roots.0.join("ordering")).unwrap());
353        let peering =
354            Arc::new(K1Peering::open(&roots.0.join("peering"), ordering.clone()).unwrap());
355        (ordering, peering)
356    }
357
358    #[test]
359    fn store_v3_generic_box_round_trips_through_kto() {
360        let roots = Roots::new();
361        let (ordering, peering) = stack(&roots);
362        let app = K1ChatPersistence::open(&roots.0.join("persistence"), ordering, peering).unwrap();
363        let (session, _) = app.session([9; 12]).unwrap();
364        session.persist(vec![generic_history()[0].clone()]).unwrap();
365        let Record::Box(value) = &session.load().unwrap().records[0] else {
366            panic!("expected generic box");
367        };
368        assert_eq!(
369            (
370                value.box_type(),
371                value.contents(),
372                value.hidden_type(),
373                value.hidden_contents()
374            ),
375            ("Future Kind", "opaque", "future/v9", "hidden bytes")
376        );
377    }
378
379    #[test]
380    fn two_sessions_cursor_correlation_and_cold_restart() {
381        let roots = Roots::new();
382        let (ordering, peering) = stack(&roots);
383        let app = K1ChatPersistence::open(
384            &roots.0.join("persistence"),
385            ordering.clone(),
386            peering.clone(),
387        )
388        .unwrap();
389        assert!(
390            K1ChatPersistence::open(
391                &roots.0.join("persistence"),
392                ordering.clone(),
393                peering.clone()
394            )
395            .is_err()
396        );
397        let (one, log) = app.session([1; 12]).unwrap();
398        let (two, _) = app.session([2; 12]).unwrap();
399        assert!(log.records.is_empty());
400        let history = generic_history();
401        let first = one.persist(vec![history[0].clone()]).unwrap();
402        two.persist(vec![history[0].clone()]).unwrap();
403        assert_eq!(one.load().unwrap().records.len(), 1);
404        assert_eq!(two.load().unwrap().records.len(), 1);
405        let cursor = fs::read(roots.0.join("persistence/cursor")).unwrap();
406        assert_eq!(cursor.len(), 12);
407        assert_ne!(cursor.as_slice(), first.as_bytes());
408
409        drop(one);
410        drop(two);
411        drop(app);
412        drop(peering);
413        drop(ordering);
414
415        let (ordering, peering) = stack(&roots);
416        let app = K1ChatPersistence::open(&roots.0.join("persistence"), ordering, peering).unwrap();
417        let (one, log) = app.session([1; 12]).unwrap();
418        assert_eq!(log.records, vec![history[0].clone()]);
419        let id = one.persist(vec![history[1].clone()]).unwrap();
420        assert_eq!(
421            fs::read(roots.0.join("persistence/cursor")).unwrap(),
422            id.as_bytes()
423        );
424        assert_eq!(one.load().unwrap().records, history);
425    }
426
427    #[test]
428    fn overlap_rejection_and_reorg_discard_fault() {
429        let roots = Roots::new();
430        let (ordering, peering) = stack(&roots);
431        let app = K1ChatPersistence::open(&roots.0.join("persistence"), ordering, peering).unwrap();
432        let (session, _) = app.session([3; 12]).unwrap();
433        let history = generic_history();
434        session.persist(vec![history[0].clone()]).unwrap();
435        app.driver.state.lock().unwrap().pending.insert(
436            session.id(),
437            Pending {
438                payload: Vec::new(),
439                evidence: None,
440            },
441        );
442        assert!(
443            session
444                .persist(vec![history[1].clone()])
445                .unwrap_err()
446                .contains("in flight")
447        );
448        app.driver.state.lock().unwrap().pending.clear();
449
450        let error = app.driver.reorg().unwrap_err();
451        assert!(error.contains("reorganization"));
452        assert!(session.load().unwrap_err().contains("faulted"));
453        assert!(!roots.0.join("persistence/cursor").exists());
454        assert!(!roots.0.join("persistence/sessions").exists());
455    }
456
457    #[test]
458    fn tool_message_and_v2_result_commit_together() {
459        let roots = Roots::new();
460        let (ordering, peering) = stack(&roots);
461        let session_id = [4; 12];
462        let app = K1ChatPersistence::open(&roots.0.join("persistence"), ordering, peering).unwrap();
463        let (session, _) = app.session(session_id).unwrap();
464        let tool_call_id = ToolCallId::new(session_id, 1);
465        let call = numbered_box(
466            1,
467            tool_call_box(&ProviderCall {
468                tool_call_id,
469                name: "lookup".into(),
470                arguments: r#"{"key":"value"}"#.into(),
471            }),
472        );
473        session.persist(vec![call.clone()]).unwrap();
474
475        let message_metadata = ToolMessageMetadata {
476            tool_call_id,
477            originating_call: BoxId::new(1),
478            message_index: 1,
479            message: "working".into(),
480        };
481        let message = numbered_box(2, tool_message_box(&message_metadata).unwrap());
482        let result_metadata = ToolResultV2Metadata {
483            tool_call_id,
484            originating_call: BoxId::new(1),
485            result: Ok("done".into()),
486            metadata_type: "application/x-k1-test".into(),
487            metadata_contents: "opaque metadata".into(),
488        };
489        let result = numbered_box(3, tool_result_v2_box(&result_metadata));
490        session
491            .persist(vec![message.clone(), result.clone()])
492            .unwrap();
493
494        let loaded = session.load().unwrap().records;
495        assert_eq!(loaded, vec![call, message, result]);
496        let Record::Box(loaded_message) = &loaded[1] else {
497            panic!("expected Tool Message");
498        };
499        assert_eq!(
500            loaded_message.tool_message_metadata().unwrap(),
501            Some(message_metadata)
502        );
503        let Record::Box(loaded_result) = &loaded[2] else {
504            panic!("expected Tool Result");
505        };
506        assert_eq!(
507            loaded_result.tool_result_v2_metadata().unwrap(),
508            Some(result_metadata)
509        );
510    }
511}