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 std::sync::atomic::{AtomicU64, Ordering};
306
307    static NEXT: AtomicU64 = AtomicU64::new(0);
308
309    struct Roots(PathBuf);
310
311    impl Roots {
312        fn new() -> Self {
313            let path = std::env::temp_dir().join(format!(
314                "k1-chat-persistence-{}-{}",
315                std::process::id(),
316                NEXT.fetch_add(1, Ordering::Relaxed)
317            ));
318            let _ = fs::remove_dir_all(&path);
319            Self(path)
320        }
321    }
322
323    impl Drop for Roots {
324        fn drop(&mut self) {
325            let _ = fs::remove_dir_all(&self.0);
326        }
327    }
328
329    fn event(index: u64) -> Record {
330        Record::Event(EventRecord::new(0, index, 0, "test".into(), "value".into()).unwrap())
331    }
332
333    fn generic_box() -> 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"}]}"#,
336        )
337        .unwrap()
338        .records
339        .into_iter()
340        .next()
341        .unwrap()
342    }
343
344    fn stack(roots: &Roots) -> (Arc<K1TxnOrdering>, Arc<K1Peering>) {
345        let ordering = Arc::new(K1TxnOrdering::open(&roots.0.join("ordering")).unwrap());
346        let peering =
347            Arc::new(K1Peering::open(&roots.0.join("peering"), ordering.clone()).unwrap());
348        (ordering, peering)
349    }
350
351    #[test]
352    fn store_v3_generic_box_round_trips_through_kto() {
353        let roots = Roots::new();
354        let (ordering, peering) = stack(&roots);
355        let app = K1ChatPersistence::open(&roots.0.join("persistence"), ordering, peering).unwrap();
356        let (session, _) = app.session([9; 12]).unwrap();
357        session.persist(vec![generic_box()]).unwrap();
358        let Record::Box(value) = &session.load().unwrap().records[0] else {
359            panic!("expected generic box");
360        };
361        assert_eq!(
362            (
363                value.box_type(),
364                value.contents(),
365                value.hidden_type(),
366                value.hidden_contents()
367            ),
368            ("Future Kind", "opaque", "future/v9", "hidden bytes")
369        );
370    }
371
372    #[test]
373    fn two_sessions_cursor_correlation_and_cold_restart() {
374        let roots = Roots::new();
375        let (ordering, peering) = stack(&roots);
376        let app = K1ChatPersistence::open(
377            &roots.0.join("persistence"),
378            ordering.clone(),
379            peering.clone(),
380        )
381        .unwrap();
382        assert!(
383            K1ChatPersistence::open(
384                &roots.0.join("persistence"),
385                ordering.clone(),
386                peering.clone()
387            )
388            .is_err()
389        );
390        let (one, log) = app.session([1; 12]).unwrap();
391        let (two, _) = app.session([2; 12]).unwrap();
392        assert!(log.records.is_empty());
393        let first = one.persist(vec![event(1)]).unwrap();
394        two.persist(vec![event(1)]).unwrap();
395        assert_eq!(one.load().unwrap().records.len(), 1);
396        assert_eq!(two.load().unwrap().records.len(), 1);
397        let cursor = fs::read(roots.0.join("persistence/cursor")).unwrap();
398        assert_eq!(cursor.len(), 12);
399        assert_ne!(cursor.as_slice(), first.as_bytes());
400
401        drop(one);
402        drop(two);
403        drop(app);
404        drop(peering);
405        drop(ordering);
406
407        let (ordering, peering) = stack(&roots);
408        let app = K1ChatPersistence::open(&roots.0.join("persistence"), ordering, peering).unwrap();
409        let (one, log) = app.session([1; 12]).unwrap();
410        assert_eq!(log.records, vec![event(1)]);
411        let id = one.persist(vec![event(2)]).unwrap();
412        assert_eq!(
413            fs::read(roots.0.join("persistence/cursor")).unwrap(),
414            id.as_bytes()
415        );
416        assert_eq!(one.load().unwrap().records, vec![event(1), event(2)]);
417    }
418
419    #[test]
420    fn overlap_rejection_and_reorg_discard_fault() {
421        let roots = Roots::new();
422        let (ordering, peering) = stack(&roots);
423        let app = K1ChatPersistence::open(&roots.0.join("persistence"), ordering, peering).unwrap();
424        let (session, _) = app.session([3; 12]).unwrap();
425        session.persist(vec![event(1)]).unwrap();
426        app.driver.state.lock().unwrap().pending.insert(
427            session.id(),
428            Pending {
429                payload: Vec::new(),
430                evidence: None,
431            },
432        );
433        assert!(
434            session
435                .persist(vec![event(2)])
436                .unwrap_err()
437                .contains("in flight")
438        );
439        app.driver.state.lock().unwrap().pending.clear();
440
441        let error = app.driver.reorg().unwrap_err();
442        assert!(error.contains("reorganization"));
443        assert!(session.load().unwrap_err().contains("faulted"));
444        assert!(!roots.0.join("persistence/cursor").exists());
445        assert!(!roots.0.join("persistence/sessions").exists());
446    }
447}