#![forbid(unsafe_code)]
pub use kcode_k1_chat_codex_state::{
BoxValue, ChatBox, PreflightItem, PreflightMode, PreparedCall, PreparedMailboxFlush,
PreparedPreflightCall, RestartError, ShimOutput, Start, Status, ToolCallId,
};
pub use kcode_k1_chat_persistence::EventRecord;
use kcode_k1_chat_codex_state::{AGENT_RESPONSE_TYPE, ConversationState};
use kcode_k1_chat_persistence::{Record, Session};
use kcode_k1_chat_thread_recovery::recover as recover_thread;
use serde::{Deserialize, Serialize};
use serde_json::{Value, json};
const PREFLIGHT_HANDLER: &str = "chat_preflight";
const DIAGNOSTIC_HANDLER: &str = "chat_diagnostic";
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum ChatDiagnostic {
ModelUsageSubscribe,
ModelUsageReceive,
ModelUsageFinish,
ModelUsageRestart,
ModelUsageShutdown,
ProviderInference,
MailboxTransport,
CriticalIntegrity,
}
impl ChatDiagnostic {
fn code(self) -> &'static str {
match self {
Self::ModelUsageSubscribe => "model_usage_subscribe",
Self::ModelUsageReceive => "model_usage_receive",
Self::ModelUsageFinish => "model_usage_finish",
Self::ModelUsageRestart => "model_usage_restart",
Self::ModelUsageShutdown => "model_usage_shutdown",
Self::ProviderInference => "provider_inference",
Self::MailboxTransport => "mailbox_transport",
Self::CriticalIntegrity => "critical_integrity",
}
}
}
#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
#[serde(rename_all = "snake_case", deny_unknown_fields)]
pub struct TokenBreakdown {
pub input_tokens: i64,
pub cached_input_tokens: i64,
pub cache_write_input_tokens: i64,
pub output_tokens: i64,
pub reasoning_output_tokens: i64,
pub total_tokens: i64,
}
#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
#[serde(rename_all = "snake_case", deny_unknown_fields)]
pub struct ModelUsage {
pub provider: String,
pub model: String,
pub context_id: String,
pub provider_turn_id: String,
pub usage: TokenBreakdown,
pub cumulative_usage: Option<TokenBreakdown>,
pub context_limit_tokens: Option<i64>,
}
pub struct DurableTurn {
state: ConversationState,
records: Vec<Record>,
mirrored: usize,
durable: usize,
session: Session,
returned: Vec<ToolCallId>,
preflight: Vec<PreparedPreflightCall>,
}
impl DurableTurn {
pub fn recover(session: Session) -> Result<Self, String> {
let recovered = recover_thread(&session)?;
let returned = returned_ids(recovered.state.boxes())?;
let preflight = recovered_preflight(&recovered.records, recovered.state.boxes())?;
Ok(Self {
state: recovered.state,
records: recovered.records,
mirrored: recovered.mirrored,
durable: recovered.durable,
session,
returned,
preflight,
})
}
pub fn boxes(&self) -> &[ChatBox] {
self.state.boxes()
}
pub fn events(&self) -> Vec<EventRecord> {
self.records[..self.durable]
.iter()
.filter_map(|record| match record {
Record::Event(event) => Some(event.clone()),
Record::Box(_) => None,
})
.collect()
}
pub fn status(&self) -> Status {
self.state.status()
}
pub fn preflight_calls(&self) -> &[PreparedPreflightCall] {
&self.preflight
}
pub fn prepare_preflight(
&mut self,
items: Vec<PreflightItem>,
) -> Result<Vec<PreparedPreflightCall>, String> {
let calls = self.state.prepare_preflight(items)?;
self.mirror_boxes()?;
let after_box_id = self.latest_box_id()?;
let event = EventRecord {
after_box_id,
event_index: self.next_event_index(after_box_id)?,
connected_box_id: 0,
handler: PREFLIGHT_HANDLER.into(),
data: preflight_data(&calls),
};
let mut pending = self.records[self.durable..].to_vec();
pending.push(Record::Event(event.clone()));
self.session.persist(pending)?;
self.records.push(Record::Event(event));
self.durable = self.records.len();
self.preflight = calls.clone();
Ok(calls)
}
pub fn accept(
&mut self,
box_type: String,
contents: String,
hidden_type: String,
hidden_contents: String,
) -> Result<(), String> {
let result = self
.state
.accept(box_type, contents, hidden_type, hidden_contents);
self.finish(result)
}
pub fn accept_tool_return(
&mut self,
tool_call_id: ToolCallId,
result: Result<String, String>,
) -> Result<(), String> {
if self.returned.contains(&tool_call_id) {
return self.finish(Ok(()));
}
let accepted = self.state.accept_tool_return(tool_call_id, result);
if accepted.is_ok() {
self.returned.push(tool_call_id);
}
self.finish(accepted)
}
pub fn accept_tool_message(
&mut self,
tool_call_id: ToolCallId,
message: String,
) -> Result<(), String> {
let result = self.state.accept_tool_message(tool_call_id, message);
self.finish(result)
}
pub fn accept_tool_return_v2(
&mut self,
tool_call_id: ToolCallId,
result: Result<String, String>,
metadata_type: String,
metadata_contents: String,
) -> Result<(), String> {
if self.returned.contains(&tool_call_id) {
return self.finish(Ok(()));
}
let accepted = self.state.accept_tool_return_v2(
tool_call_id,
result,
metadata_type,
metadata_contents,
);
if accepted.is_ok() {
self.returned.push(tool_call_id);
}
self.finish(accepted)
}
pub fn begin(&mut self) -> Result<Option<Start>, String> {
self.state.begin()
}
pub fn prepare_stage(
&mut self,
job: u64,
text: String,
values: Vec<BoxValue>,
) -> Result<Vec<PreparedCall>, String> {
let result = self.state.prepare_stage(job, text, values);
self.finish(result)
}
pub fn prepare_mailbox_flush(
&mut self,
job: u64,
) -> Result<Option<PreparedMailboxFlush>, String> {
let result = self.state.prepare_mailbox_flush(job);
self.finish(result)
}
pub fn validate_mailbox_flush(&self, prepared: &PreparedMailboxFlush) -> Result<(), String> {
self.state.validate_mailbox_flush(prepared)
}
pub fn commit_mailbox_flush(&mut self, prepared: PreparedMailboxFlush) -> Result<(), String> {
self.state.commit_mailbox_flush(prepared)
}
pub fn complete(&mut self, job: u64, output: ShimOutput<BoxValue>) -> Result<bool, String> {
if let Err(error) = self.state.complete(job, output) {
return self.finish(Err(error));
}
self.persist_completion()
}
pub fn complete_recoverable_failure(
&mut self,
job: u64,
message: String,
) -> Result<bool, String> {
if let Err(error) = self.state.complete_recoverable_failure(job, message) {
return self.finish(Err(error));
}
self.persist_completion()
}
pub fn reset_provider_context(&mut self) -> Result<(), String> {
self.state.reset_provider_context()
}
pub fn halt_critical(&mut self, message: String) -> Result<(), String> {
self.state.halt_critical(message);
self.mirror_and_persist()
}
pub fn record_diagnostic(&mut self, diagnostic: ChatDiagnostic) -> Result<(), String> {
self.mirror_boxes()?;
let after_box_id = self.latest_box_id()?;
self.persist_event(EventRecord {
after_box_id,
event_index: self.next_event_index(after_box_id)?,
connected_box_id: 0,
handler: DIAGNOSTIC_HANDLER.into(),
data: json!({"version": 1, "code": diagnostic.code()}),
})
}
pub fn record_model_usage(
&mut self,
connected_box_id: u64,
usage: ModelUsage,
) -> Result<(), String> {
self.mirror_boxes()?;
let after_box_id = self.latest_box_id()?;
if connected_box_id != 0
&& !self.state.boxes().iter().any(|box_| {
box_.id().get() == connected_box_id && box_.box_type() == AGENT_RESPONSE_TYPE
})
{
return Err("model usage must connect to a canonical Agent Response box".to_owned());
}
self.persist_event(EventRecord {
after_box_id,
event_index: self.next_event_index(after_box_id)?,
connected_box_id,
handler: "model_usage".into(),
data: model_usage_data(usage)?,
})
}
pub fn fail(&mut self, job: u64, message: String, restartable_before_launch: bool) {
self.state.fail(job, message, restartable_before_launch);
}
pub fn restart(&mut self) -> Result<(), RestartError> {
self.state.restart()
}
fn persist_completion(&mut self) -> Result<bool, String> {
self.mirror_boxes()?;
let resume = matches!(self.state.status(), Status::Running);
let after_box_id = self.latest_box_id()?;
self.persist_event(EventRecord {
after_box_id,
event_index: self.next_event_index(after_box_id)?,
connected_box_id: 0,
handler: "llm_done".into(),
data: json!({"resume": resume}),
})?;
Ok(resume)
}
fn finish<T>(&mut self, operation: Result<T, String>) -> Result<T, String> {
let persistence = self.mirror_and_persist();
match (operation, persistence) {
(Ok(value), Ok(())) => Ok(value),
(Err(error), Ok(())) | (Ok(_), Err(error)) => Err(error),
(Err(operation), Err(persistence)) => Err(format!(
"{operation}; additionally failed to persist canonical history: {persistence}"
)),
}
}
fn mirror_and_persist(&mut self) -> Result<(), String> {
self.mirror_boxes()?;
self.persist_pending()
}
fn mirror_boxes(&mut self) -> Result<(), String> {
let boxes = self.state.boxes();
let additions = boxes
.get(self.mirrored..)
.ok_or_else(|| "canonical box frontier moved backwards".to_owned())?;
self.records
.extend(additions.iter().cloned().map(Record::Box));
self.mirrored = boxes.len();
Ok(())
}
fn latest_box_id(&self) -> Result<u64, String> {
self.state
.boxes()
.last()
.map(|box_| box_.id().get())
.ok_or_else(|| "durable event requires a canonical box".to_owned())
}
fn next_event_index(&self, after_box_id: u64) -> Result<u64, String> {
match self.records.last() {
Some(Record::Event(event)) if event.after_box_id == after_box_id => event
.event_index
.checked_add(1)
.ok_or_else(|| "durable event index space was exhausted".to_owned()),
Some(Record::Event(_)) => {
Err("durable event frontier diverged from canonical boxes".to_owned())
}
_ => Ok(1),
}
}
fn persist_event(&mut self, event: EventRecord) -> Result<(), String> {
let suffix = self
.records
.get(self.durable..)
.ok_or_else(|| "durable record frontier moved past canonical records".to_owned())?;
let mut pending = suffix.to_vec();
pending.push(Record::Event(event.clone()));
self.session.persist(pending)?;
self.records.push(Record::Event(event));
self.durable = self.records.len();
Ok(())
}
fn persist_pending(&mut self) -> Result<(), String> {
let suffix = self
.records
.get(self.durable..)
.ok_or_else(|| "durable record frontier moved past canonical records".to_owned())?;
if suffix.is_empty() {
return Ok(());
}
self.session.persist(suffix.to_vec())?;
self.durable = self.records.len();
Ok(())
}
}
fn preflight_data(calls: &[PreparedPreflightCall]) -> Value {
json!({
"version": 1,
"calls": calls.iter().map(|call| json!({
"sequence": call.tool_call_id.sequence(),
"mode": match call.mode {
PreflightMode::Blocking => "blocking",
PreflightMode::NonBlocking => "non_blocking",
},
})).collect::<Vec<_>>()
})
}
fn recovered_preflight(
records: &[Record],
boxes: &[ChatBox],
) -> Result<Vec<PreparedPreflightCall>, String> {
let Some(event) = records.iter().find_map(|record| match record {
Record::Event(event) if event.handler == PREFLIGHT_HANDLER => Some(event),
_ => None,
}) else {
return Ok(Vec::new());
};
let calls = event
.data
.get("calls")
.and_then(Value::as_array)
.ok_or_else(|| "malformed chat_preflight event data".to_owned())?;
calls
.iter()
.map(|call| {
let sequence = call
.get("sequence")
.and_then(Value::as_u64)
.ok_or_else(|| "malformed chat_preflight event data".to_owned())?;
let mode = match call.get("mode").and_then(Value::as_str) {
Some("blocking") => PreflightMode::Blocking,
Some("non_blocking") => PreflightMode::NonBlocking,
_ => return Err("malformed chat_preflight event data".to_owned()),
};
let metadata = boxes
.iter()
.find_map(|value| {
value
.tool_call_metadata()
.ok()
.flatten()
.filter(|metadata| metadata.tool_call_id.sequence() == sequence)
.map(|metadata| (value.id(), metadata))
})
.ok_or_else(|| "chat_preflight references an unknown ToolCallId".to_owned())?;
Ok(PreparedPreflightCall {
tool_call_id: metadata.1.tool_call_id,
call_box_id: metadata.0,
name: metadata.1.name,
arguments: metadata.1.arguments,
mode,
})
})
.collect()
}
fn model_usage_data(usage: ModelUsage) -> Result<Value, String> {
let mut data = serde_json::to_value(usage).map_err(|error| error.to_string())?;
let Value::Object(fields) = &mut data else {
return Err("model usage data did not serialize to an object".to_owned());
};
fields.insert("version".into(), Value::from(1));
Ok(data)
}
fn returned_ids(boxes: &[ChatBox]) -> Result<Vec<ToolCallId>, String> {
let mut returned = Vec::new();
for value in boxes {
if let Some(result) = value
.tool_result_metadata()
.map_err(|error| format!("{error:?}"))?
{
returned.push(result.tool_call_id);
}
}
Ok(returned)
}
#[cfg(test)]
mod tests {
use super::*;
use kcode_k1_chat_persistence::K1ChatPersistence;
use kcode_k1_peering::K1Peering;
use kcode_k1_txn_ordering::K1TxnOrdering;
use std::fs;
use std::path::PathBuf;
use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering};
static NEXT: AtomicU64 = AtomicU64::new(0);
struct Fixture {
root: PathBuf,
session: Option<Session>,
}
impl Fixture {
fn new(nonce: u8) -> Self {
let root = std::env::temp_dir().join(format!(
"k1-durable-turn-{}-{}",
std::process::id(),
NEXT.fetch_add(1, Ordering::Relaxed)
));
let _ = fs::remove_dir_all(&root);
let ordering = Arc::new(K1TxnOrdering::open(&root.join("ordering")).unwrap());
let peering =
Arc::new(K1Peering::open(&root.join("peering"), Arc::clone(&ordering)).unwrap());
let persistence =
K1ChatPersistence::open(&root.join("persistence"), ordering, peering).unwrap();
let (session, _) = persistence.session([nonce; 12]).unwrap();
Self {
root,
session: Some(session),
}
}
fn session(&self) -> Session {
self.session.as_ref().unwrap().clone()
}
}
impl Drop for Fixture {
fn drop(&mut self) {
drop(self.session.take());
let _ = fs::remove_dir_all(&self.root);
}
}
#[test]
fn preflight_boxes_and_event_are_one_recoverable_history() {
let fixture = Fixture::new(9);
let session = fixture.session();
let mut turn = DurableTurn::recover(session.clone()).unwrap();
let calls = turn
.prepare_preflight(vec![
PreflightItem::SystemMessage {
contents: "context".into(),
},
PreflightItem::KtoolCall {
name: "CurrentTime".into(),
arguments: "{}".into(),
mode: PreflightMode::Blocking,
},
])
.unwrap();
assert_eq!(calls.len(), 1);
assert_eq!(turn.preflight_calls(), calls);
let log = session.load().unwrap();
assert_eq!(log.boxes.len(), 2);
assert_eq!(log.events.len(), 1);
assert_eq!(log.events[0].handler, PREFLIGHT_HANDLER);
drop(turn);
let mut recovered = DurableTurn::recover(session).unwrap();
assert_eq!(recovered.preflight_calls(), calls);
assert!(recovered.begin().unwrap().is_none());
}
#[test]
fn terminal_completion_and_model_usage_remain_ordered() {
let fixture = Fixture::new(1);
let session = fixture.session();
let mut turn = DurableTurn::recover(session.clone()).unwrap();
turn.accept(
"User Message".into(),
"hello".into(),
String::new(),
String::new(),
)
.unwrap();
let start = turn.begin().unwrap().unwrap();
turn.complete(start.job, ShimOutput { items: Vec::new() })
.unwrap();
let terminal = turn.boxes().last().unwrap().id().get();
let usage = TokenBreakdown {
input_tokens: 1,
cached_input_tokens: 0,
cache_write_input_tokens: 0,
output_tokens: 1,
reasoning_output_tokens: 0,
total_tokens: 2,
};
turn.record_model_usage(
terminal,
ModelUsage {
provider: "provider".into(),
model: "model".into(),
context_id: "context".into(),
provider_turn_id: "turn".into(),
usage,
cumulative_usage: None,
context_limit_tokens: None,
},
)
.unwrap();
assert_eq!(turn.events().len(), 2);
drop(turn);
assert_eq!(DurableTurn::recover(session).unwrap().events().len(), 2);
}
#[test]
fn recoverable_failure_and_diagnostic_are_durable() {
let fixture = Fixture::new(2);
let session = fixture.session();
let mut turn = DurableTurn::recover(session.clone()).unwrap();
turn.accept(
"User Message".into(),
"hello".into(),
String::new(),
String::new(),
)
.unwrap();
let job = turn.begin().unwrap().unwrap().job;
assert!(
!turn
.complete_recoverable_failure(job, "Please try again.".into())
.unwrap()
);
turn.reset_provider_context().unwrap();
turn.record_diagnostic(ChatDiagnostic::ProviderInference)
.unwrap();
assert_eq!(turn.boxes().last().unwrap().contents(), "Please try again.");
drop(turn);
let recovered = DurableTurn::recover(session).unwrap();
assert_eq!(
recovered.boxes().last().unwrap().contents(),
"Please try again."
);
assert_eq!(
recovered.events().last().unwrap().data,
json!({"version": 1, "code": "provider_inference"})
);
}
}