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, 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}