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 AGENT_RESPONSE_TYPE, BoxId, ChatBox, ProviderCall, ToolCallId, ToolMessageMetadata,
307 ToolResultV2Metadata, 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 agent_response_hidden_fields_survive_cold_restart() {
381 let roots = Roots::new();
382 let persistence = roots.0.join("persistence");
383 let record = Record::Box(ChatBox::new(
384 BoxId::new(1),
385 AGENT_RESPONSE_TYPE.to_owned(),
386 String::new(),
387 "agent/hidden".to_owned(),
388 "hidden response".to_owned(),
389 ));
390 {
391 let (ordering, peering) = stack(&roots);
392 let app = K1ChatPersistence::open(&persistence, ordering, peering).unwrap();
393 let (session, _) = app.session([5; 12]).unwrap();
394 session.persist(vec![record.clone()]).unwrap();
395 }
396 let (ordering, peering) = stack(&roots);
397 let app = K1ChatPersistence::open(&persistence, ordering, peering).unwrap();
398 let (_, loaded) = app.session([5; 12]).unwrap();
399 assert_eq!(loaded.records, vec![record]);
400 }
401
402 #[test]
403 fn two_sessions_cursor_correlation_and_cold_restart() {
404 let roots = Roots::new();
405 let (ordering, peering) = stack(&roots);
406 let app = K1ChatPersistence::open(
407 &roots.0.join("persistence"),
408 ordering.clone(),
409 peering.clone(),
410 )
411 .unwrap();
412 assert!(
413 K1ChatPersistence::open(
414 &roots.0.join("persistence"),
415 ordering.clone(),
416 peering.clone()
417 )
418 .is_err()
419 );
420 let (one, log) = app.session([1; 12]).unwrap();
421 let (two, _) = app.session([2; 12]).unwrap();
422 assert!(log.records.is_empty());
423 let history = generic_history();
424 let first = one.persist(vec![history[0].clone()]).unwrap();
425 two.persist(vec![history[0].clone()]).unwrap();
426 assert_eq!(one.load().unwrap().records.len(), 1);
427 assert_eq!(two.load().unwrap().records.len(), 1);
428 let cursor = fs::read(roots.0.join("persistence/cursor")).unwrap();
429 assert_eq!(cursor.len(), 12);
430 assert_ne!(cursor.as_slice(), first.as_bytes());
431
432 drop(one);
433 drop(two);
434 drop(app);
435 drop(peering);
436 drop(ordering);
437
438 let (ordering, peering) = stack(&roots);
439 let app = K1ChatPersistence::open(&roots.0.join("persistence"), ordering, peering).unwrap();
440 let (one, log) = app.session([1; 12]).unwrap();
441 assert_eq!(log.records, vec![history[0].clone()]);
442 let id = one.persist(vec![history[1].clone()]).unwrap();
443 assert_eq!(
444 fs::read(roots.0.join("persistence/cursor")).unwrap(),
445 id.as_bytes()
446 );
447 assert_eq!(one.load().unwrap().records, history);
448 }
449
450 #[test]
451 fn overlap_rejection_and_reorg_discard_fault() {
452 let roots = Roots::new();
453 let (ordering, peering) = stack(&roots);
454 let app = K1ChatPersistence::open(&roots.0.join("persistence"), ordering, peering).unwrap();
455 let (session, _) = app.session([3; 12]).unwrap();
456 let history = generic_history();
457 session.persist(vec![history[0].clone()]).unwrap();
458 app.driver.state.lock().unwrap().pending.insert(
459 session.id(),
460 Pending {
461 payload: Vec::new(),
462 evidence: None,
463 },
464 );
465 assert!(
466 session
467 .persist(vec![history[1].clone()])
468 .unwrap_err()
469 .contains("in flight")
470 );
471 app.driver.state.lock().unwrap().pending.clear();
472
473 let error = app.driver.reorg().unwrap_err();
474 assert!(error.contains("reorganization"));
475 assert!(session.load().unwrap_err().contains("faulted"));
476 assert!(!roots.0.join("persistence/cursor").exists());
477 assert!(!roots.0.join("persistence/sessions").exists());
478 }
479
480 #[test]
481 fn tool_message_and_v2_result_commit_together() {
482 let roots = Roots::new();
483 let (ordering, peering) = stack(&roots);
484 let session_id = [4; 12];
485 let app = K1ChatPersistence::open(&roots.0.join("persistence"), ordering, peering).unwrap();
486 let (session, _) = app.session(session_id).unwrap();
487 let tool_call_id = ToolCallId::new(session_id, 1);
488 let call = numbered_box(
489 1,
490 tool_call_box(&ProviderCall {
491 tool_call_id,
492 name: "lookup".into(),
493 arguments: r#"{"key":"value"}"#.into(),
494 }),
495 );
496 session.persist(vec![call.clone()]).unwrap();
497
498 let message_metadata = ToolMessageMetadata {
499 tool_call_id,
500 originating_call: BoxId::new(1),
501 message_index: 1,
502 message: "working".into(),
503 };
504 let message = numbered_box(2, tool_message_box(&message_metadata).unwrap());
505 let result_metadata = ToolResultV2Metadata {
506 tool_call_id,
507 originating_call: BoxId::new(1),
508 result: Ok("done".into()),
509 metadata_type: "application/x-k1-test".into(),
510 metadata_contents: "opaque metadata".into(),
511 };
512 let result = numbered_box(3, tool_result_v2_box(&result_metadata));
513 session
514 .persist(vec![message.clone(), result.clone()])
515 .unwrap();
516
517 let loaded = session.load().unwrap().records;
518 assert_eq!(loaded, vec![call, message, result]);
519 let Record::Box(loaded_message) = &loaded[1] else {
520 panic!("expected Tool Message");
521 };
522 assert_eq!(
523 loaded_message.tool_message_metadata().unwrap(),
524 Some(message_metadata)
525 );
526 let Record::Box(loaded_result) = &loaded[2] else {
527 panic!("expected Tool Result");
528 };
529 assert_eq!(
530 loaded_result.tool_result_v2_metadata().unwrap(),
531 Some(result_metadata)
532 );
533 }
534}