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