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}