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