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