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";
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}