Skip to main content

kcode_k1_chat_thread_durable_turn/
lib.rs

1#![forbid(unsafe_code)]
2#![doc = include_str!("../Documentation.md")]
3
4use kcode_k1_chat_codex_state::{AGENT_RESPONSE_TYPE, ConversationState};
5pub use kcode_k1_chat_codex_state::{
6    BoxValue, ChatBox, PreparedCall, PreparedMailboxFlush, RestartError, ShimOutput, Start, Status,
7    ToolCallId,
8};
9pub use kcode_k1_chat_persistence::EventRecord;
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#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
26#[serde(rename_all = "snake_case", deny_unknown_fields)]
27pub struct ModelUsage {
28    pub provider: String,
29    pub model: String,
30    pub context_id: String,
31    pub provider_turn_id: String,
32    pub usage: TokenBreakdown,
33    pub cumulative_usage: Option<TokenBreakdown>,
34    pub context_limit_tokens: Option<i64>,
35}
36
37pub struct DurableTurn {
38    state: ConversationState,
39    records: Vec<Record>,
40    mirrored: usize,
41    durable: usize,
42    session: Session,
43    returned: Vec<ToolCallId>,
44}
45impl DurableTurn {
46    pub fn recover(session: Session) -> Result<Self, String> {
47        let recovered = recover_thread(&session)?;
48        let returned = returned_ids(recovered.state.boxes())?;
49        Ok(Self {
50            state: recovered.state,
51            records: recovered.records,
52            mirrored: recovered.mirrored,
53            durable: recovered.durable,
54            session,
55            returned,
56        })
57    }
58    pub fn boxes(&self) -> &[ChatBox] {
59        self.state.boxes()
60    }
61    pub fn events(&self) -> Vec<EventRecord> {
62        self.records[..self.durable]
63            .iter()
64            .filter_map(|r| {
65                if let Record::Event(e) = r {
66                    Some(e.clone())
67                } else {
68                    None
69                }
70            })
71            .collect()
72    }
73    pub fn status(&self) -> Status {
74        self.state.status()
75    }
76    pub fn accept(
77        &mut self,
78        box_type: String,
79        contents: String,
80        hidden_type: String,
81        hidden_contents: String,
82    ) -> Result<(), String> {
83        let r = self
84            .state
85            .accept(box_type, contents, hidden_type, hidden_contents);
86        self.finish(r)
87    }
88    pub fn accept_idle_context_box(
89        &mut self,
90        box_type: String,
91        contents: String,
92        hidden_type: String,
93        hidden_contents: String,
94    ) -> Result<(), String> {
95        let r =
96            self.state
97                .accept_idle_context_box(box_type, contents, hidden_type, hidden_contents);
98        self.finish(r)
99    }
100    pub fn accept_idle_context_tool_call(
101        &mut self,
102        name: String,
103        arguments: String,
104    ) -> Result<ToolCallId, String> {
105        let r = self.state.accept_idle_context_tool_call(name, arguments);
106        self.finish(r)
107    }
108    pub fn accept_idle_context_tool_return(
109        &mut self,
110        id: ToolCallId,
111        result: Result<String, String>,
112    ) -> Result<(), String> {
113        let r = self.state.accept_idle_context_tool_return(id, result);
114        if r.is_ok() {
115            self.returned.push(id);
116        }
117        self.finish(r)
118    }
119    pub fn accept_tool_return(
120        &mut self,
121        id: ToolCallId,
122        result: Result<String, String>,
123    ) -> Result<(), String> {
124        if self.returned.contains(&id) {
125            return self.finish(Ok(()));
126        }
127        let r = self.state.accept_tool_return(id, result);
128        if r.is_ok() {
129            self.returned.push(id);
130        }
131        self.finish(r)
132    }
133    pub fn accept_tool_message(&mut self, id: ToolCallId, message: String) -> Result<(), String> {
134        let r = self.state.accept_tool_message(id, message);
135        self.finish(r)
136    }
137    pub fn accept_tool_return_v2(
138        &mut self,
139        id: ToolCallId,
140        result: Result<String, String>,
141        metadata_type: String,
142        metadata_contents: String,
143    ) -> Result<(), String> {
144        let r = self
145            .state
146            .accept_tool_return_v2(id, result, metadata_type, metadata_contents);
147        if r.is_ok() {
148            self.returned.push(id);
149        }
150        self.finish(r)
151    }
152    pub fn begin(&mut self) -> Result<Option<Start>, String> {
153        self.state.begin()
154    }
155    pub fn prepare_stage(
156        &mut self,
157        job: u64,
158        text: String,
159        values: Vec<BoxValue>,
160    ) -> Result<Vec<PreparedCall>, String> {
161        let r = self.state.prepare_stage(job, text, values);
162        self.finish(r)
163    }
164    pub fn prepare_mailbox_flush(
165        &mut self,
166        job: u64,
167    ) -> Result<Option<PreparedMailboxFlush>, String> {
168        let r = self.state.prepare_mailbox_flush(job);
169        self.finish(r)
170    }
171    pub fn validate_mailbox_flush(&self, p: &PreparedMailboxFlush) -> Result<(), String> {
172        self.state.validate_mailbox_flush(p)
173    }
174    pub fn commit_mailbox_flush(&mut self, p: PreparedMailboxFlush) -> Result<(), String> {
175        self.state.commit_mailbox_flush(p)
176    }
177    pub fn complete(&mut self, job: u64, output: ShimOutput<BoxValue>) -> Result<bool, String> {
178        if let Err(e) = self.state.complete(job, output) {
179            return self.finish(Err(e));
180        }
181        self.mirror_boxes()?;
182        let resume = matches!(self.state.status(), Status::Running);
183        let after = self.latest_box_id()?;
184        self.persist_event(EventRecord {
185            after_box_id: after,
186            event_index: self.next_event_index(after)?,
187            connected_box_id: 0,
188            handler: "llm_done".into(),
189            data: json!({"resume":resume}),
190        })?;
191        Ok(resume)
192    }
193    pub fn record_model_usage(
194        &mut self,
195        connected_box_id: u64,
196        usage: ModelUsage,
197    ) -> Result<(), String> {
198        self.mirror_boxes()?;
199        let after = self.latest_box_id()?;
200        if connected_box_id != 0
201            && !self
202                .state
203                .boxes()
204                .iter()
205                .any(|b| b.id().get() == connected_box_id && b.box_type() == AGENT_RESPONSE_TYPE)
206        {
207            return Err("model usage must connect to a canonical Agent Response box".into());
208        }
209        self.persist_event(EventRecord {
210            after_box_id: after,
211            event_index: self.next_event_index(after)?,
212            connected_box_id,
213            handler: "model_usage".into(),
214            data: model_usage_data(usage)?,
215        })
216    }
217    pub fn fail(&mut self, job: u64, message: String, restartable: bool) {
218        self.state.fail(job, message, restartable)
219    }
220    pub fn restart(&mut self) -> Result<(), RestartError> {
221        self.state.restart()
222    }
223    fn finish<T>(&mut self, operation: Result<T, String>) -> Result<T, String> {
224        let persistence = self.mirror_and_persist();
225        match (operation, persistence) {
226            (Ok(v), Ok(())) => Ok(v),
227            (Err(e), Ok(())) | (Ok(_), Err(e)) => Err(e),
228            (Err(a), Err(b)) => Err(format!(
229                "{a}; additionally failed to persist canonical history: {b}"
230            )),
231        }
232    }
233    fn mirror_and_persist(&mut self) -> Result<(), String> {
234        self.mirror_boxes()?;
235        self.persist_pending()
236    }
237    fn mirror_boxes(&mut self) -> Result<(), String> {
238        let additions = self
239            .state
240            .boxes()
241            .get(self.mirrored..)
242            .ok_or_else(|| "canonical box frontier moved backwards".to_owned())?;
243        self.records
244            .extend(additions.iter().cloned().map(Record::Box));
245        self.mirrored = self.state.boxes().len();
246        Ok(())
247    }
248    fn latest_box_id(&self) -> Result<u64, String> {
249        self.state
250            .boxes()
251            .last()
252            .map(|b| b.id().get())
253            .ok_or_else(|| "durable event requires a canonical box".into())
254    }
255    fn next_event_index(&self, after: u64) -> Result<u64, String> {
256        match self.records.last() {
257            Some(Record::Event(e)) if e.after_box_id == after => e
258                .event_index
259                .checked_add(1)
260                .ok_or_else(|| "durable event index space was exhausted".into()),
261            Some(Record::Event(_)) => {
262                Err("durable event frontier diverged from canonical boxes".into())
263            }
264            _ => Ok(1),
265        }
266    }
267    fn persist_event(&mut self, event: EventRecord) -> Result<(), String> {
268        let mut pending = self
269            .records
270            .get(self.durable..)
271            .ok_or_else(|| "durable record frontier moved past canonical records".to_owned())?
272            .to_vec();
273        pending.push(Record::Event(event.clone()));
274        self.session.persist(pending)?;
275        self.records.push(Record::Event(event));
276        self.durable = self.records.len();
277        Ok(())
278    }
279    fn persist_pending(&mut self) -> Result<(), String> {
280        let suffix = self
281            .records
282            .get(self.durable..)
283            .ok_or_else(|| "durable record frontier moved past canonical records".to_owned())?;
284        if suffix.is_empty() {
285            return Ok(());
286        }
287        self.session.persist(suffix.to_vec())?;
288        self.durable = self.records.len();
289        Ok(())
290    }
291}
292fn model_usage_data(usage: ModelUsage) -> Result<Value, String> {
293    let mut data = serde_json::to_value(usage).map_err(|e| e.to_string())?;
294    let Value::Object(fields) = &mut data else {
295        return Err("model usage data did not serialize to an object".into());
296    };
297    fields.insert("version".into(), Value::from(1));
298    Ok(data)
299}
300fn returned_ids(boxes: &[ChatBox]) -> Result<Vec<ToolCallId>, String> {
301    let mut out = Vec::new();
302    for b in boxes {
303        if let Some(r) = b.tool_result_metadata().map_err(|e| format!("{e:?}"))? {
304            out.push(r.tool_call_id);
305        }
306        if let Some(r) = b.tool_result_v2_metadata().map_err(|e| format!("{e:?}"))? {
307            out.push(r.tool_call_id);
308        }
309    }
310    Ok(out)
311}