Skip to main content

kcode_k1_chat_thread_durable_turn/
lib.rs

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}