1#![forbid(unsafe_code)]
2
3pub use kcode_k1_chat_codex_state::{
4 BoxValue, ChatBox, PreflightItem, PreflightMode, PreparedCall, PreparedMailboxFlush,
5 PreparedPreflightCall, RestartError, ShimOutput, Start, Status, ToolCallId,
6};
7pub use kcode_k1_chat_persistence::EventRecord;
8
9use kcode_k1_chat_codex_state::{AGENT_RESPONSE_TYPE, ConversationState};
10use kcode_k1_chat_persistence::{Record, Session};
11use kcode_k1_chat_thread_recovery::recover as recover_thread;
12use serde::{Deserialize, Serialize};
13use serde_json::{Value, json};
14
15const PREFLIGHT_HANDLER: &str = "chat_preflight";
16const DIAGNOSTIC_HANDLER: &str = "chat_diagnostic";
17
18#[derive(Clone, Copy, Debug, Eq, PartialEq)]
19pub enum ChatDiagnostic {
20 ModelUsageSubscribe,
21 ModelUsageReceive,
22 ModelUsageFinish,
23 ModelUsageRestart,
24 ModelUsageShutdown,
25 ProviderInference,
26 MailboxTransport,
27 CriticalIntegrity,
28}
29
30impl ChatDiagnostic {
31 fn code(self) -> &'static str {
32 match self {
33 Self::ModelUsageSubscribe => "model_usage_subscribe",
34 Self::ModelUsageReceive => "model_usage_receive",
35 Self::ModelUsageFinish => "model_usage_finish",
36 Self::ModelUsageRestart => "model_usage_restart",
37 Self::ModelUsageShutdown => "model_usage_shutdown",
38 Self::ProviderInference => "provider_inference",
39 Self::MailboxTransport => "mailbox_transport",
40 Self::CriticalIntegrity => "critical_integrity",
41 }
42 }
43}
44
45#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
46#[serde(rename_all = "snake_case", deny_unknown_fields)]
47pub struct TokenBreakdown {
48 pub input_tokens: i64,
49 pub cached_input_tokens: i64,
50 pub cache_write_input_tokens: i64,
51 pub output_tokens: i64,
52 pub reasoning_output_tokens: i64,
53 pub total_tokens: i64,
54}
55
56#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
57#[serde(rename_all = "snake_case", deny_unknown_fields)]
58pub struct ModelUsage {
59 pub provider: String,
60 pub model: String,
61 pub context_id: String,
62 pub provider_turn_id: String,
63 pub usage: TokenBreakdown,
64 pub cumulative_usage: Option<TokenBreakdown>,
65 pub context_limit_tokens: Option<i64>,
66}
67
68pub struct DurableTurn {
69 state: ConversationState,
70 records: Vec<Record>,
71 mirrored: usize,
72 durable: usize,
73 session: Session,
74 returned: Vec<ToolCallId>,
75 preflight: Vec<PreparedPreflightCall>,
76}
77
78impl DurableTurn {
79 pub fn recover(session: Session) -> Result<Self, String> {
80 let recovered = recover_thread(&session)?;
81 let returned = returned_ids(recovered.state.boxes())?;
82 let preflight = recovered_preflight(&recovered.records, recovered.state.boxes())?;
83 Ok(Self {
84 state: recovered.state,
85 records: recovered.records,
86 mirrored: recovered.mirrored,
87 durable: recovered.durable,
88 session,
89 returned,
90 preflight,
91 })
92 }
93
94 pub fn boxes(&self) -> &[ChatBox] {
95 self.state.boxes()
96 }
97
98 pub fn events(&self) -> Vec<EventRecord> {
99 self.records[..self.durable]
100 .iter()
101 .filter_map(|record| match record {
102 Record::Event(event) => Some(event.clone()),
103 Record::Box(_) => None,
104 })
105 .collect()
106 }
107
108 pub fn status(&self) -> Status {
109 self.state.status()
110 }
111
112 pub fn preflight_calls(&self) -> &[PreparedPreflightCall] {
113 &self.preflight
114 }
115
116 pub fn prepare_preflight(
117 &mut self,
118 items: Vec<PreflightItem>,
119 ) -> Result<Vec<PreparedPreflightCall>, String> {
120 let calls = self.state.prepare_preflight(items)?;
121 self.mirror_boxes()?;
122 let after_box_id = self.latest_box_id()?;
123 let event = EventRecord {
124 after_box_id,
125 event_index: self.next_event_index(after_box_id)?,
126 connected_box_id: 0,
127 handler: PREFLIGHT_HANDLER.into(),
128 data: preflight_data(&calls),
129 };
130 let mut pending = self.records[self.durable..].to_vec();
131 pending.push(Record::Event(event.clone()));
132 self.session.persist(pending)?;
133 self.records.push(Record::Event(event));
134 self.durable = self.records.len();
135 self.preflight = calls.clone();
136 Ok(calls)
137 }
138
139 pub fn accept(
140 &mut self,
141 box_type: String,
142 contents: String,
143 hidden_type: String,
144 hidden_contents: String,
145 ) -> Result<(), String> {
146 let result = self
147 .state
148 .accept(box_type, contents, hidden_type, hidden_contents);
149 self.finish(result)
150 }
151
152 pub fn accept_tool_return(
153 &mut self,
154 tool_call_id: ToolCallId,
155 result: Result<String, String>,
156 ) -> Result<(), String> {
157 if self.returned.contains(&tool_call_id) {
158 return self.finish(Ok(()));
159 }
160 let accepted = self.state.accept_tool_return(tool_call_id, result);
161 if accepted.is_ok() {
162 self.returned.push(tool_call_id);
163 }
164 self.finish(accepted)
165 }
166
167 pub fn accept_tool_message(
168 &mut self,
169 tool_call_id: ToolCallId,
170 message: String,
171 ) -> Result<(), String> {
172 let result = self.state.accept_tool_message(tool_call_id, message);
173 self.finish(result)
174 }
175
176 pub fn accept_tool_return_v2(
177 &mut self,
178 tool_call_id: ToolCallId,
179 result: Result<String, String>,
180 metadata_type: String,
181 metadata_contents: String,
182 ) -> Result<(), String> {
183 if self.returned.contains(&tool_call_id) {
184 return self.finish(Ok(()));
185 }
186 let accepted = self.state.accept_tool_return_v2(
187 tool_call_id,
188 result,
189 metadata_type,
190 metadata_contents,
191 );
192 if accepted.is_ok() {
193 self.returned.push(tool_call_id);
194 }
195 self.finish(accepted)
196 }
197
198 pub fn begin(&mut self) -> Result<Option<Start>, String> {
199 self.state.begin()
200 }
201
202 pub fn prepare_stage(
203 &mut self,
204 job: u64,
205 text: String,
206 values: Vec<BoxValue>,
207 ) -> Result<Vec<PreparedCall>, String> {
208 let result = self.state.prepare_stage(job, text, values);
209 self.finish(result)
210 }
211
212 pub fn prepare_mailbox_flush(
213 &mut self,
214 job: u64,
215 ) -> Result<Option<PreparedMailboxFlush>, String> {
216 let result = self.state.prepare_mailbox_flush(job);
217 self.finish(result)
218 }
219
220 pub fn validate_mailbox_flush(&self, prepared: &PreparedMailboxFlush) -> Result<(), String> {
221 self.state.validate_mailbox_flush(prepared)
222 }
223
224 pub fn commit_mailbox_flush(&mut self, prepared: PreparedMailboxFlush) -> Result<(), String> {
225 self.state.commit_mailbox_flush(prepared)
226 }
227
228 pub fn complete(&mut self, job: u64, output: ShimOutput<BoxValue>) -> Result<bool, String> {
229 if let Err(error) = self.state.complete(job, output) {
230 return self.finish(Err(error));
231 }
232 self.persist_completion()
233 }
234
235 pub fn complete_recoverable_failure(
236 &mut self,
237 job: u64,
238 message: String,
239 ) -> Result<bool, String> {
240 if let Err(error) = self.state.complete_recoverable_failure(job, message) {
241 return self.finish(Err(error));
242 }
243 self.persist_completion()
244 }
245
246 pub fn reset_provider_context(&mut self) -> Result<(), String> {
247 self.state.reset_provider_context()
248 }
249
250 pub fn halt_critical(&mut self, message: String) -> Result<(), String> {
251 self.state.halt_critical(message);
252 self.mirror_and_persist()
253 }
254
255 pub fn record_diagnostic(&mut self, diagnostic: ChatDiagnostic) -> Result<(), String> {
256 self.mirror_boxes()?;
257 let after_box_id = self.latest_box_id()?;
258 self.persist_event(EventRecord {
259 after_box_id,
260 event_index: self.next_event_index(after_box_id)?,
261 connected_box_id: 0,
262 handler: DIAGNOSTIC_HANDLER.into(),
263 data: json!({"version": 1, "code": diagnostic.code()}),
264 })
265 }
266
267 pub fn record_model_usage(
268 &mut self,
269 connected_box_id: u64,
270 usage: ModelUsage,
271 ) -> Result<(), String> {
272 self.mirror_boxes()?;
273 let after_box_id = self.latest_box_id()?;
274 if connected_box_id != 0
275 && !self.state.boxes().iter().any(|box_| {
276 box_.id().get() == connected_box_id && box_.box_type() == AGENT_RESPONSE_TYPE
277 })
278 {
279 return Err("model usage must connect to a canonical Agent Response box".to_owned());
280 }
281 self.persist_event(EventRecord {
282 after_box_id,
283 event_index: self.next_event_index(after_box_id)?,
284 connected_box_id,
285 handler: "model_usage".into(),
286 data: model_usage_data(usage)?,
287 })
288 }
289
290 pub fn fail(&mut self, job: u64, message: String, restartable_before_launch: bool) {
291 self.state.fail(job, message, restartable_before_launch);
292 }
293
294 pub fn restart(&mut self) -> Result<(), RestartError> {
295 self.state.restart()
296 }
297
298 fn persist_completion(&mut self) -> Result<bool, String> {
299 self.mirror_boxes()?;
300 let resume = matches!(self.state.status(), Status::Running);
301 let after_box_id = self.latest_box_id()?;
302 self.persist_event(EventRecord {
303 after_box_id,
304 event_index: self.next_event_index(after_box_id)?,
305 connected_box_id: 0,
306 handler: "llm_done".into(),
307 data: json!({"resume": resume}),
308 })?;
309 Ok(resume)
310 }
311
312 fn finish<T>(&mut self, operation: Result<T, String>) -> Result<T, String> {
313 let persistence = self.mirror_and_persist();
314 match (operation, persistence) {
315 (Ok(value), Ok(())) => Ok(value),
316 (Err(error), Ok(())) | (Ok(_), Err(error)) => Err(error),
317 (Err(operation), Err(persistence)) => Err(format!(
318 "{operation}; additionally failed to persist canonical history: {persistence}"
319 )),
320 }
321 }
322
323 fn mirror_and_persist(&mut self) -> Result<(), String> {
324 self.mirror_boxes()?;
325 self.persist_pending()
326 }
327
328 fn mirror_boxes(&mut self) -> Result<(), String> {
329 let boxes = self.state.boxes();
330 let additions = boxes
331 .get(self.mirrored..)
332 .ok_or_else(|| "canonical box frontier moved backwards".to_owned())?;
333 self.records
334 .extend(additions.iter().cloned().map(Record::Box));
335 self.mirrored = boxes.len();
336 Ok(())
337 }
338
339 fn latest_box_id(&self) -> Result<u64, String> {
340 self.state
341 .boxes()
342 .last()
343 .map(|box_| box_.id().get())
344 .ok_or_else(|| "durable event requires a canonical box".to_owned())
345 }
346
347 fn next_event_index(&self, after_box_id: u64) -> Result<u64, String> {
348 match self.records.last() {
349 Some(Record::Event(event)) if event.after_box_id == after_box_id => event
350 .event_index
351 .checked_add(1)
352 .ok_or_else(|| "durable event index space was exhausted".to_owned()),
353 Some(Record::Event(_)) => {
354 Err("durable event frontier diverged from canonical boxes".to_owned())
355 }
356 _ => Ok(1),
357 }
358 }
359
360 fn persist_event(&mut self, event: EventRecord) -> Result<(), String> {
361 let suffix = self
362 .records
363 .get(self.durable..)
364 .ok_or_else(|| "durable record frontier moved past canonical records".to_owned())?;
365 let mut pending = suffix.to_vec();
366 pending.push(Record::Event(event.clone()));
367 self.session.persist(pending)?;
368 self.records.push(Record::Event(event));
369 self.durable = self.records.len();
370 Ok(())
371 }
372
373 fn persist_pending(&mut self) -> Result<(), String> {
374 let suffix = self
375 .records
376 .get(self.durable..)
377 .ok_or_else(|| "durable record frontier moved past canonical records".to_owned())?;
378 if suffix.is_empty() {
379 return Ok(());
380 }
381 self.session.persist(suffix.to_vec())?;
382 self.durable = self.records.len();
383 Ok(())
384 }
385}
386
387fn preflight_data(calls: &[PreparedPreflightCall]) -> Value {
388 json!({
389 "version": 1,
390 "calls": calls.iter().map(|call| json!({
391 "sequence": call.tool_call_id.sequence(),
392 "mode": match call.mode {
393 PreflightMode::Blocking => "blocking",
394 PreflightMode::NonBlocking => "non_blocking",
395 },
396 })).collect::<Vec<_>>()
397 })
398}
399
400fn recovered_preflight(
401 records: &[Record],
402 boxes: &[ChatBox],
403) -> Result<Vec<PreparedPreflightCall>, String> {
404 let Some(event) = records.iter().find_map(|record| match record {
405 Record::Event(event) if event.handler == PREFLIGHT_HANDLER => Some(event),
406 _ => None,
407 }) else {
408 return Ok(Vec::new());
409 };
410 let calls = event
411 .data
412 .get("calls")
413 .and_then(Value::as_array)
414 .ok_or_else(|| "malformed chat_preflight event data".to_owned())?;
415 calls
416 .iter()
417 .map(|call| {
418 let sequence = call
419 .get("sequence")
420 .and_then(Value::as_u64)
421 .ok_or_else(|| "malformed chat_preflight event data".to_owned())?;
422 let mode = match call.get("mode").and_then(Value::as_str) {
423 Some("blocking") => PreflightMode::Blocking,
424 Some("non_blocking") => PreflightMode::NonBlocking,
425 _ => return Err("malformed chat_preflight event data".to_owned()),
426 };
427 let metadata = boxes
428 .iter()
429 .find_map(|value| {
430 value
431 .tool_call_metadata()
432 .ok()
433 .flatten()
434 .filter(|metadata| metadata.tool_call_id.sequence() == sequence)
435 .map(|metadata| (value.id(), metadata))
436 })
437 .ok_or_else(|| "chat_preflight references an unknown ToolCallId".to_owned())?;
438 Ok(PreparedPreflightCall {
439 tool_call_id: metadata.1.tool_call_id,
440 call_box_id: metadata.0,
441 name: metadata.1.name,
442 arguments: metadata.1.arguments,
443 mode,
444 })
445 })
446 .collect()
447}
448
449fn model_usage_data(usage: ModelUsage) -> Result<Value, String> {
450 let mut data = serde_json::to_value(usage).map_err(|error| error.to_string())?;
451 let Value::Object(fields) = &mut data else {
452 return Err("model usage data did not serialize to an object".to_owned());
453 };
454 fields.insert("version".into(), Value::from(1));
455 Ok(data)
456}
457
458fn returned_ids(boxes: &[ChatBox]) -> Result<Vec<ToolCallId>, String> {
459 let mut returned = Vec::new();
460 for value in boxes {
461 if let Some(result) = value
462 .tool_result_metadata()
463 .map_err(|error| format!("{error:?}"))?
464 {
465 returned.push(result.tool_call_id);
466 }
467 }
468 Ok(returned)
469}
470
471#[cfg(test)]
472mod tests {
473 use super::*;
474 use kcode_k1_chat_persistence::K1ChatPersistence;
475 use kcode_k1_peering::K1Peering;
476 use kcode_k1_txn_ordering::K1TxnOrdering;
477 use std::fs;
478 use std::path::PathBuf;
479 use std::sync::Arc;
480 use std::sync::atomic::{AtomicU64, Ordering};
481
482 static NEXT: AtomicU64 = AtomicU64::new(0);
483
484 struct Fixture {
485 root: PathBuf,
486 session: Option<Session>,
487 }
488
489 impl Fixture {
490 fn new(nonce: u8) -> Self {
491 let root = std::env::temp_dir().join(format!(
492 "k1-durable-turn-{}-{}",
493 std::process::id(),
494 NEXT.fetch_add(1, Ordering::Relaxed)
495 ));
496 let _ = fs::remove_dir_all(&root);
497 let ordering = Arc::new(K1TxnOrdering::open(&root.join("ordering")).unwrap());
498 let peering =
499 Arc::new(K1Peering::open(&root.join("peering"), Arc::clone(&ordering)).unwrap());
500 let persistence =
501 K1ChatPersistence::open(&root.join("persistence"), ordering, peering).unwrap();
502 let (session, _) = persistence.session([nonce; 12]).unwrap();
503 Self {
504 root,
505 session: Some(session),
506 }
507 }
508
509 fn session(&self) -> Session {
510 self.session.as_ref().unwrap().clone()
511 }
512 }
513
514 impl Drop for Fixture {
515 fn drop(&mut self) {
516 drop(self.session.take());
517 let _ = fs::remove_dir_all(&self.root);
518 }
519 }
520
521 #[test]
522 fn preflight_boxes_and_event_are_one_recoverable_history() {
523 let fixture = Fixture::new(9);
524 let session = fixture.session();
525 let mut turn = DurableTurn::recover(session.clone()).unwrap();
526 let calls = turn
527 .prepare_preflight(vec![
528 PreflightItem::SystemMessage {
529 contents: "context".into(),
530 },
531 PreflightItem::KtoolCall {
532 name: "CurrentTime".into(),
533 arguments: "{}".into(),
534 mode: PreflightMode::Blocking,
535 },
536 ])
537 .unwrap();
538 assert_eq!(calls.len(), 1);
539 assert_eq!(turn.preflight_calls(), calls);
540 let log = session.load().unwrap();
541 assert_eq!(log.boxes.len(), 2);
542 assert_eq!(log.events.len(), 1);
543 assert_eq!(log.events[0].handler, PREFLIGHT_HANDLER);
544 drop(turn);
545 let mut recovered = DurableTurn::recover(session).unwrap();
546 assert_eq!(recovered.preflight_calls(), calls);
547 assert!(recovered.begin().unwrap().is_none());
548 }
549
550 #[test]
551 fn terminal_completion_and_model_usage_remain_ordered() {
552 let fixture = Fixture::new(1);
553 let session = fixture.session();
554 let mut turn = DurableTurn::recover(session.clone()).unwrap();
555 turn.accept(
556 "User Message".into(),
557 "hello".into(),
558 String::new(),
559 String::new(),
560 )
561 .unwrap();
562 let start = turn.begin().unwrap().unwrap();
563 turn.complete(start.job, ShimOutput { items: Vec::new() })
564 .unwrap();
565 let terminal = turn.boxes().last().unwrap().id().get();
566 let usage = TokenBreakdown {
567 input_tokens: 1,
568 cached_input_tokens: 0,
569 cache_write_input_tokens: 0,
570 output_tokens: 1,
571 reasoning_output_tokens: 0,
572 total_tokens: 2,
573 };
574 turn.record_model_usage(
575 terminal,
576 ModelUsage {
577 provider: "provider".into(),
578 model: "model".into(),
579 context_id: "context".into(),
580 provider_turn_id: "turn".into(),
581 usage,
582 cumulative_usage: None,
583 context_limit_tokens: None,
584 },
585 )
586 .unwrap();
587 assert_eq!(turn.events().len(), 2);
588 drop(turn);
589 assert_eq!(DurableTurn::recover(session).unwrap().events().len(), 2);
590 }
591
592 #[test]
593 fn recoverable_failure_and_diagnostic_are_durable() {
594 let fixture = Fixture::new(2);
595 let session = fixture.session();
596 let mut turn = DurableTurn::recover(session.clone()).unwrap();
597 turn.accept(
598 "User Message".into(),
599 "hello".into(),
600 String::new(),
601 String::new(),
602 )
603 .unwrap();
604 let job = turn.begin().unwrap().unwrap().job;
605 assert!(
606 !turn
607 .complete_recoverable_failure(job, "Please try again.".into())
608 .unwrap()
609 );
610 turn.reset_provider_context().unwrap();
611 turn.record_diagnostic(ChatDiagnostic::ProviderInference)
612 .unwrap();
613 assert_eq!(turn.boxes().last().unwrap().contents(), "Please try again.");
614 drop(turn);
615 let recovered = DurableTurn::recover(session).unwrap();
616 assert_eq!(
617 recovered.boxes().last().unwrap().contents(),
618 "Please try again."
619 );
620 assert_eq!(
621 recovered.events().last().unwrap().data,
622 json!({"version": 1, "code": "provider_inference"})
623 );
624 }
625}