Skip to main content

kcode_k1_chat_thread_durable_state/
lib.rs

1#![forbid(unsafe_code)]
2
3use kcode_k1_access_kmap::K1AccessKmap;
4pub use kcode_k1_chat_codex_codec::BoxValue;
5use kcode_k1_chat_codex_state::{ConversationState, RestartError};
6pub use kcode_k1_chat_codex_state::{PreparedCall, PreparedSteer, Status};
7use kcode_k1_chat_persistence::{EventRecord, Record, Session};
8use kcode_k1_chat_state::USER_MESSAGE_TYPE;
9pub use kcode_k1_chat_state::{BoxId, ChatBox, ToolCallId};
10use kcode_k1_chat_thread_actions::ChatThreadActions;
11pub use kcode_k1_chat_thread_actions::{AccessContext, AccessPolicy, ProfileId};
12use kcode_k1_chat_thread_recovery::recover as recover_thread;
13use kcode_k1_codex_adapter::ShimOutput;
14use serde_json::json;
15use std::sync::Arc;
16
17#[derive(Clone, Debug, Eq, PartialEq)]
18pub enum TransitionError {
19    Unauthorized,
20    NotStalled,
21    NotRestartable,
22    Internal(String),
23}
24
25pub struct DurableThread {
26    state: ConversationState,
27    records: Vec<Record>,
28    mirrored: usize,
29    durable: usize,
30    session: Session,
31    actions: ChatThreadActions,
32    authorized: bool,
33}
34
35impl DurableThread {
36    pub fn recover(session: Session, kmap: Arc<K1AccessKmap>) -> Result<Self, String> {
37        let recovered = recover_thread(&session)?;
38        Ok(Self {
39            state: recovered.state,
40            records: recovered.records,
41            mirrored: recovered.mirrored,
42            durable: recovered.durable,
43            session,
44            actions: ChatThreadActions::new(kmap),
45            authorized: false,
46        })
47    }
48
49    pub fn boxes(&self) -> &[ChatBox] {
50        self.state.boxes()
51    }
52
53    pub fn status(&self) -> Status {
54        self.state.status()
55    }
56
57    pub fn accept_box(
58        &mut self,
59        box_type: String,
60        contents: String,
61        hidden_type: String,
62        hidden_contents: String,
63    ) -> Result<(), String> {
64        self.state
65            .accept(box_type, contents, hidden_type, hidden_contents)
66    }
67
68    pub fn accept_user(
69        &mut self,
70        context: AccessContext,
71        profile_id: ProfileId,
72        policy: AccessPolicy,
73        contents: String,
74    ) -> Result<(), TransitionError> {
75        self.bind_authorization(context, profile_id, policy)?;
76        self.state
77            .accept(
78                USER_MESSAGE_TYPE.into(),
79                contents,
80                String::new(),
81                String::new(),
82            )
83            .and_then(|()| self.checkpoint())
84            .map_err(TransitionError::Internal)
85    }
86
87    pub fn accept_return(
88        &mut self,
89        id: ToolCallId,
90        result: Result<String, String>,
91    ) -> Result<(), String> {
92        if returned_calls(self.state.boxes())?.contains(&id) {
93            Ok(())
94        } else {
95            self.state.accept_tool_return(id, result)
96        }
97    }
98
99    pub fn prepare_stage(
100        &mut self,
101        job: u64,
102        text: String,
103        boxes: Vec<BoxValue>,
104    ) -> Result<Vec<PreparedCall>, String> {
105        let calls = self.state.prepare_stage(job, text, boxes)?;
106        self.checkpoint()?;
107        Ok(calls)
108    }
109
110    pub fn launch_action(&mut self, name: &str, arguments: &str) -> Result<String, String> {
111        self.actions.launch(name, arguments)
112    }
113
114    pub fn accept_tool_message(&mut self, id: ToolCallId, contents: String) -> Result<(), String> {
115        self.state.accept_tool_message(id, contents)
116    }
117
118    pub fn accept_tool_return(
119        &mut self,
120        id: ToolCallId,
121        result: Result<String, String>,
122    ) -> Result<(), String> {
123        self.state.accept_tool_return(id, result)
124    }
125
126    pub fn accept_tool_return_v2(
127        &mut self,
128        id: ToolCallId,
129        result: Result<String, String>,
130        metadata_type: String,
131        metadata_contents: String,
132    ) -> Result<(), String> {
133        self.state
134            .accept_tool_return_v2(id, result, metadata_type, metadata_contents)
135    }
136
137    pub fn flush_active_arrivals(&mut self, job: u64) -> Result<(), String> {
138        self.state.flush_active_arrivals(job).map(|_| ())
139    }
140
141    pub fn checkpoint(&mut self) -> Result<(), String> {
142        self.mirror();
143        self.flush()
144    }
145
146    pub fn prepare_steer(&mut self, job: u64) -> Result<Option<PreparedSteer>, String> {
147        let prepared = self.state.prepare_steer(job)?;
148        self.checkpoint()?;
149        Ok(prepared)
150    }
151
152    pub fn prepared_input(&self, prepared: &PreparedSteer) -> Result<String, String> {
153        self.state.validate_steer(prepared)?;
154        render_input(prepared.values())
155    }
156
157    pub fn commit_steer(&mut self, prepared: PreparedSteer) -> Result<(), String> {
158        self.state.commit_steer(prepared)
159    }
160
161    pub fn begin_input(&mut self) -> Result<Option<(u64, String)>, String> {
162        let Some(start) = self.state.begin()? else {
163            return Ok(None);
164        };
165        Ok(Some((start.job, render_input(&start.values)?)))
166    }
167
168    pub fn complete(&mut self, job: u64, output: ShimOutput<BoxValue>) -> Result<bool, String> {
169        self.state.complete(job, output)?;
170        let resume = matches!(self.state.status(), Status::Running);
171        self.persist_done(resume)?;
172        if !resume {
173            self.clear_authorization();
174        }
175        Ok(resume)
176    }
177
178    pub fn fail(&mut self, job: u64, error: String, restartable: bool) {
179        self.state.fail(job, error, restartable);
180        self.clear_authorization();
181    }
182
183    pub fn restart(
184        &mut self,
185        context: AccessContext,
186        profile_id: ProfileId,
187        policy: AccessPolicy,
188    ) -> Result<(), TransitionError> {
189        let installed = self.bind_authorization(context, profile_id, policy)?;
190        if let Err(error) = self.state.restart().map_err(|error| match error {
191            RestartError::NotStalled => TransitionError::NotStalled,
192            RestartError::ProviderActionAccepted => TransitionError::NotRestartable,
193        }) {
194            if installed {
195                self.clear_authorization();
196            }
197            return Err(error);
198        }
199        Ok(())
200    }
201
202    fn bind_authorization(
203        &mut self,
204        context: AccessContext,
205        profile_id: ProfileId,
206        policy: AccessPolicy,
207    ) -> Result<bool, TransitionError> {
208        let installed = !self.authorized;
209        if self
210            .actions
211            .bind_authorization(context, profile_id, policy)
212            .is_err()
213        {
214            if installed {
215                self.actions.clear_authorization();
216            }
217            return Err(TransitionError::Unauthorized);
218        }
219        self.authorized = true;
220        Ok(installed)
221    }
222
223    fn clear_authorization(&mut self) {
224        self.actions.clear_authorization();
225        self.authorized = false;
226    }
227
228    fn mirror(&mut self) {
229        self.records.extend(
230            self.state.boxes()[self.mirrored..]
231                .iter()
232                .cloned()
233                .map(Record::Box),
234        );
235        self.mirrored = self.state.boxes().len();
236    }
237
238    fn flush(&mut self) -> Result<(), String> {
239        if self.durable < self.records.len() {
240            self.session
241                .persist(self.records[self.durable..].to_vec())?;
242            self.durable = self.records.len();
243        }
244        Ok(())
245    }
246
247    fn persist_done(&mut self, resume: bool) -> Result<(), String> {
248        self.mirror();
249        let anchor = self
250            .state
251            .boxes()
252            .last()
253            .map_or(0, |value| value.id().get());
254        let index = match self.records.last() {
255            Some(Record::Event(event)) if event.after_box_id == anchor => {
256                event.event_index.checked_add(1)
257            }
258            _ => Some(1),
259        }
260        .ok_or_else(|| "event index space was exhausted".to_owned())?;
261        self.records.push(Record::Event(
262            EventRecord::new(
263                anchor,
264                index,
265                0,
266                "llm_done".into(),
267                json!({"resume": resume}),
268            )
269            .map_err(|error| error.to_string())?,
270        ));
271        self.flush()
272    }
273}
274
275fn render_input(values: &[BoxValue]) -> Result<String, String> {
276    let mut output = String::new();
277    for value in values {
278        let BoxValue::History(section) = value else {
279            return Err("Codex steer contains a non-history value".into());
280        };
281        if section.is_empty() {
282            continue;
283        }
284        if !output.is_empty() && !output.ends_with('\n') {
285            output.push('\n');
286        }
287        output.push_str(section);
288    }
289    Ok(output)
290}
291
292fn returned_calls(boxes: &[ChatBox]) -> Result<Vec<ToolCallId>, String> {
293    let mut returned = Vec::new();
294    for value in boxes {
295        if let Some(result) = value
296            .tool_result_metadata()
297            .map_err(|error| format!("{error:?}"))?
298        {
299            returned.push(result.tool_call_id);
300        }
301    }
302    Ok(returned)
303}