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, PreparedSteer, RestartError, ShimOutput, Start, Status,
5    ToolCallId,
6};
7
8use kcode_k1_chat_codex_state::ConversationState;
9use kcode_k1_chat_persistence::{EventRecord, Record, Session};
10use kcode_k1_chat_thread_recovery::recover as recover_thread;
11use serde_json::json;
12
13pub struct DurableTurn {
14    state: ConversationState,
15    records: Vec<Record>,
16    mirrored: usize,
17    durable: usize,
18    session: Session,
19    returned: Vec<ToolCallId>,
20}
21
22impl DurableTurn {
23    pub fn recover(session: Session) -> Result<Self, String> {
24        let recovered = recover_thread(&session)?;
25        let returned = returned_ids(recovered.state.boxes())?;
26        Ok(Self {
27            state: recovered.state,
28            records: recovered.records,
29            mirrored: recovered.mirrored,
30            durable: recovered.durable,
31            session,
32            returned,
33        })
34    }
35
36    pub fn boxes(&self) -> &[ChatBox] {
37        self.state.boxes()
38    }
39
40    pub fn status(&self) -> Status {
41        self.state.status()
42    }
43
44    pub fn accept(
45        &mut self,
46        box_type: String,
47        contents: String,
48        hidden_type: String,
49        hidden_contents: String,
50    ) -> Result<(), String> {
51        let result = self
52            .state
53            .accept(box_type, contents, hidden_type, hidden_contents);
54        self.finish(result)
55    }
56
57    pub fn accept_tool_return(
58        &mut self,
59        tool_call_id: ToolCallId,
60        result: Result<String, String>,
61    ) -> Result<(), String> {
62        if self.returned.contains(&tool_call_id) {
63            return self.finish(Ok(()));
64        }
65        let accepted = self.state.accept_tool_return(tool_call_id, result);
66        if accepted.is_ok() {
67            self.returned.push(tool_call_id);
68        }
69        self.finish(accepted)
70    }
71
72    pub fn accept_tool_message(
73        &mut self,
74        tool_call_id: ToolCallId,
75        message: String,
76    ) -> Result<(), String> {
77        let result = self.state.accept_tool_message(tool_call_id, message);
78        self.finish(result)
79    }
80
81    pub fn accept_tool_return_v2(
82        &mut self,
83        tool_call_id: ToolCallId,
84        result: Result<String, String>,
85        metadata_type: String,
86        metadata_contents: String,
87    ) -> Result<(), String> {
88        let accepted = self.state.accept_tool_return_v2(
89            tool_call_id,
90            result,
91            metadata_type,
92            metadata_contents,
93        );
94        if accepted.is_ok() {
95            self.returned.push(tool_call_id);
96        }
97        self.finish(accepted)
98    }
99
100    pub fn begin(&mut self) -> Result<Option<Start>, String> {
101        self.state.begin()
102    }
103
104    pub fn prepare_stage(
105        &mut self,
106        job: u64,
107        text: String,
108        values: Vec<BoxValue>,
109    ) -> Result<Vec<PreparedCall>, String> {
110        let result = self.state.prepare_stage(job, text, values);
111        self.finish(result)
112    }
113
114    pub fn prepare_steer(&mut self, job: u64) -> Result<Option<PreparedSteer>, String> {
115        let result = self.state.prepare_steer(job);
116        self.finish(result)
117    }
118
119    pub fn validate_steer(&self, prepared: &PreparedSteer) -> Result<(), String> {
120        self.state.validate_steer(prepared)
121    }
122
123    pub fn commit_steer(&mut self, prepared: PreparedSteer) -> Result<(), String> {
124        self.state.commit_steer(prepared)
125    }
126
127    pub fn complete(&mut self, job: u64, output: ShimOutput<BoxValue>) -> Result<bool, String> {
128        if let Err(error) = self.state.complete(job, output) {
129            return self.finish(Err(error));
130        }
131        self.mirror_boxes()?;
132        let resume = matches!(self.state.status(), Status::Running);
133        let after_box_id = self
134            .state
135            .boxes()
136            .last()
137            .ok_or_else(|| "completed turn has no terminal box".to_owned())?
138            .id()
139            .get();
140        self.records.push(Record::Event(EventRecord {
141            after_box_id,
142            event_index: 1,
143            connected_box_id: 0,
144            handler: "llm_done".into(),
145            data: json!({"resume": resume}),
146        }));
147        self.persist_pending()?;
148        Ok(resume)
149    }
150
151    pub fn fail(&mut self, job: u64, message: String, restartable_before_launch: bool) {
152        self.state.fail(job, message, restartable_before_launch);
153    }
154
155    pub fn restart(&mut self) -> Result<(), RestartError> {
156        self.state.restart()
157    }
158
159    fn finish<T>(&mut self, operation: Result<T, String>) -> Result<T, String> {
160        let persistence = self.mirror_and_persist();
161        match (operation, persistence) {
162            (Ok(value), Ok(())) => Ok(value),
163            (Err(error), Ok(())) | (Ok(_), Err(error)) => Err(error),
164            (Err(operation), Err(persistence)) => Err(format!(
165                "{operation}; additionally failed to persist canonical history: {persistence}"
166            )),
167        }
168    }
169
170    fn mirror_and_persist(&mut self) -> Result<(), String> {
171        self.mirror_boxes()?;
172        self.persist_pending()
173    }
174
175    fn mirror_boxes(&mut self) -> Result<(), String> {
176        let boxes = self.state.boxes();
177        let additions = boxes
178            .get(self.mirrored..)
179            .ok_or_else(|| "canonical box frontier moved backwards".to_owned())?;
180        self.records
181            .extend(additions.iter().cloned().map(Record::Box));
182        self.mirrored = boxes.len();
183        Ok(())
184    }
185
186    fn persist_pending(&mut self) -> Result<(), String> {
187        let suffix = self
188            .records
189            .get(self.durable..)
190            .ok_or_else(|| "durable record frontier moved past canonical records".to_owned())?;
191        if suffix.is_empty() {
192            return Ok(());
193        }
194        self.session.persist(suffix.to_vec())?;
195        self.durable = self.records.len();
196        Ok(())
197    }
198}
199
200fn returned_ids(boxes: &[ChatBox]) -> Result<Vec<ToolCallId>, String> {
201    let mut returned = Vec::new();
202    for value in boxes {
203        if let Some(result) = value
204            .tool_result_metadata()
205            .map_err(|error| format!("{error:?}"))?
206        {
207            returned.push(result.tool_call_id);
208        }
209    }
210    Ok(returned)
211}
212
213#[cfg(test)]
214mod tests {
215    use super::*;
216    use kcode_k1_chat_codex_state::Call;
217    use kcode_k1_chat_persistence::K1ChatPersistence;
218    use kcode_k1_peering::K1Peering;
219    use kcode_k1_txn_ordering::K1TxnOrdering;
220    use std::fs;
221    use std::path::PathBuf;
222    use std::sync::Arc;
223    use std::sync::atomic::{AtomicU64, Ordering};
224
225    static NEXT: AtomicU64 = AtomicU64::new(0);
226
227    struct Fixture {
228        root: PathBuf,
229        session: Option<Session>,
230    }
231
232    impl Fixture {
233        fn new(nonce: u8) -> Self {
234            let root = std::env::temp_dir().join(format!(
235                "k1-durable-turn-{}-{}",
236                std::process::id(),
237                NEXT.fetch_add(1, Ordering::Relaxed)
238            ));
239            let _ = fs::remove_dir_all(&root);
240            let ordering = Arc::new(K1TxnOrdering::open(&root.join("ordering")).unwrap());
241            let peering =
242                Arc::new(K1Peering::open(&root.join("peering"), Arc::clone(&ordering)).unwrap());
243            let persistence =
244                K1ChatPersistence::open(&root.join("persistence"), ordering, peering).unwrap();
245            let (session, _) = persistence.session([nonce; 12]).unwrap();
246            Self {
247                root,
248                session: Some(session),
249            }
250        }
251
252        fn session(&self) -> Session {
253            self.session.as_ref().unwrap().clone()
254        }
255    }
256
257    impl Drop for Fixture {
258        fn drop(&mut self) {
259            drop(self.session.take());
260            let _ = fs::remove_dir_all(&self.root);
261        }
262    }
263
264    #[test]
265    fn completion_persists_terminal_box_event_and_recovery_frontiers() {
266        let fixture = Fixture::new(1);
267        let session = fixture.session();
268        let mut turn = DurableTurn::recover(session.clone()).unwrap();
269        turn.accept(
270            "User Message".into(),
271            "hello".into(),
272            String::new(),
273            String::new(),
274        )
275        .unwrap();
276        let start = turn.begin().unwrap().unwrap();
277        let resume = turn
278            .complete(start.job, ShimOutput { items: Vec::new() })
279            .unwrap();
280
281        assert!(!resume);
282        assert_eq!(turn.status(), Status::Quiet);
283        assert_eq!((turn.mirrored, turn.durable, turn.records.len()), (2, 3, 3));
284        let log = session.load().unwrap();
285        let Record::Event(event) = &log.records[2] else {
286            panic!("expected llm_done event");
287        };
288        assert_eq!(event.after_box_id, 2);
289        assert_eq!(event.event_index, 1);
290        assert_eq!(event.connected_box_id, 0);
291        assert_eq!(event.handler, "llm_done");
292        assert_eq!(event.data, json!({"resume": false}));
293
294        drop(turn);
295        let recovered = DurableTurn::recover(session).unwrap();
296        assert_eq!(recovered.boxes().len(), 2);
297        assert_eq!(recovered.status(), Status::Quiet);
298        assert_eq!(
299            (recovered.mirrored, recovered.durable),
300            (recovered.boxes().len(), recovered.records.len())
301        );
302    }
303
304    #[test]
305    fn active_turn_mailbox_flush_continues_generation_without_resume() {
306        let fixture = Fixture::new(4);
307        let session = fixture.session();
308        let mut turn = DurableTurn::recover(session.clone()).unwrap();
309        turn.accept(
310            "User Message".into(),
311            "first".into(),
312            String::new(),
313            String::new(),
314        )
315        .unwrap();
316
317        let start = turn.begin().unwrap().unwrap();
318        assert_eq!(
319            start.values.last(),
320            Some(&BoxValue::History("[Box 2 | Agent Response]\n".into()))
321        );
322
323        turn.prepare_stage(start.job, "working".into(), Vec::new())
324            .unwrap();
325        turn.accept(
326            "User Message".into(),
327            "second".into(),
328            String::new(),
329            String::new(),
330        )
331        .unwrap();
332
333        let prepared = turn.prepare_steer(start.job).unwrap().unwrap();
334        assert_eq!(
335            prepared.values().last(),
336            Some(&BoxValue::History("[Box 4 | Agent Response]\n".into()))
337        );
338        turn.commit_steer(prepared).unwrap();
339
340        let resume = turn
341            .complete(start.job, ShimOutput { items: Vec::new() })
342            .unwrap();
343        assert!(!resume);
344        assert_eq!(turn.status(), Status::Quiet);
345        assert!(turn.begin().unwrap().is_none());
346
347        let log = session.load().unwrap();
348        let Record::Event(event) = log.records.last().unwrap() else {
349            panic!("expected final llm_done event");
350        };
351        assert_eq!(event.handler, "llm_done");
352        assert_eq!(event.data, json!({"resume": false}));
353    }
354
355    #[test]
356    fn active_fifo_is_hidden_then_persisted_and_v1_return_is_idempotent() {
357        let fixture = Fixture::new(2);
358        let session = fixture.session();
359        let mut turn = DurableTurn::recover(session.clone()).unwrap();
360        turn.accept(
361            "User Message".into(),
362            "search".into(),
363            String::new(),
364            String::new(),
365        )
366        .unwrap();
367        let start = turn.begin().unwrap().unwrap();
368        let calls = turn
369            .prepare_stage(
370                start.job,
371                "working".into(),
372                vec![BoxValue::Call(Ok(Call {
373                    name: "WebSearch".into(),
374                    arguments: "{}".into(),
375                }))],
376            )
377            .unwrap();
378        let tool_call_id = calls[0].tool_call_id;
379        assert_eq!(session.load().unwrap().records.len(), 3);
380
381        turn.accept_tool_message(tool_call_id, "searching".into())
382            .unwrap();
383        turn.accept_tool_return_v2(
384            tool_call_id,
385            Ok("found".into()),
386            "k1.web-search-result/v1".into(),
387            "opaque".into(),
388        )
389        .unwrap();
390        turn.accept_tool_return(tool_call_id, Ok("duplicate".into()))
391            .unwrap();
392        assert_eq!(turn.boxes().len(), 3);
393        assert_eq!(session.load().unwrap().records.len(), 3);
394
395        let prepared = turn.prepare_steer(start.job).unwrap().unwrap();
396        turn.validate_steer(&prepared).unwrap();
397        assert_eq!(turn.boxes().len(), 5);
398        assert_eq!(session.load().unwrap().records.len(), 5);
399        assert_eq!((turn.mirrored, turn.durable, turn.records.len()), (5, 5, 5));
400        assert!(turn.boxes()[3].tool_message_metadata().unwrap().is_some());
401        assert!(turn.boxes()[4].tool_result_v2_metadata().unwrap().is_some());
402        turn.commit_steer(prepared).unwrap();
403    }
404
405    #[test]
406    fn unresolved_search_messages_are_inert_and_result_v2_begins_once() {
407        let fixture = Fixture::new(3);
408        let session = fixture.session();
409        let mut turn = DurableTurn::recover(session.clone()).unwrap();
410        turn.accept(
411            "User Message".into(),
412            "search".into(),
413            String::new(),
414            String::new(),
415        )
416        .unwrap();
417        let start = turn.begin().unwrap().unwrap();
418        let calls = turn
419            .prepare_stage(
420                start.job,
421                "working".into(),
422                vec![BoxValue::Call(Ok(Call {
423                    name: "WebSearch".into(),
424                    arguments: "{}".into(),
425                }))],
426            )
427            .unwrap();
428        let tool_call_id = calls[0].tool_call_id;
429
430        let prepared = turn.prepare_steer(start.job).unwrap().unwrap();
431        turn.commit_steer(prepared).unwrap();
432
433        let resume = turn
434            .complete(start.job, ShimOutput { items: Vec::new() })
435            .unwrap();
436        assert!(!resume);
437        assert_eq!(turn.status(), Status::Quiet);
438        let log = session.load().unwrap();
439        let Record::Event(event) = log.records.last().unwrap() else {
440            panic!("expected llm_done event");
441        };
442        assert_eq!(event.data, json!({"resume": false}));
443
444        turn.accept_tool_message(tool_call_id, "still searching".into())
445            .unwrap();
446        assert_eq!(turn.status(), Status::Quiet);
447        assert!(turn.begin().unwrap().is_none());
448
449        turn.accept_tool_return_v2(
450            tool_call_id,
451            Ok("found".into()),
452            "k1.web-search-result/v1".into(),
453            "opaque".into(),
454        )
455        .unwrap();
456        assert_eq!(turn.status(), Status::Running);
457        assert!(turn.begin().unwrap().is_some());
458        assert!(turn.begin().unwrap().is_none());
459    }
460}