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        Ok(resume)
173    }
174
175    pub fn fail(&mut self, job: u64, error: String, restartable: bool) {
176        self.state.fail(job, error, restartable);
177        self.clear_authorization();
178    }
179
180    pub fn restart(
181        &mut self,
182        context: AccessContext,
183        profile_id: ProfileId,
184        policy: AccessPolicy,
185    ) -> Result<(), TransitionError> {
186        let installed = self.bind_authorization(context, profile_id, policy)?;
187        if let Err(error) = self.state.restart().map_err(|error| match error {
188            RestartError::NotStalled => TransitionError::NotStalled,
189            RestartError::ProviderActionAccepted => TransitionError::NotRestartable,
190        }) {
191            if installed {
192                self.clear_authorization();
193            }
194            return Err(error);
195        }
196        Ok(())
197    }
198
199    pub fn clear_authorization(&mut self) {
200        self.actions.clear_authorization();
201        self.authorized = false;
202    }
203
204    fn bind_authorization(
205        &mut self,
206        context: AccessContext,
207        profile_id: ProfileId,
208        policy: AccessPolicy,
209    ) -> Result<bool, TransitionError> {
210        let installed = !self.authorized;
211        if self
212            .actions
213            .bind_authorization(context, profile_id, policy)
214            .is_err()
215        {
216            if installed {
217                self.actions.clear_authorization();
218            }
219            return Err(TransitionError::Unauthorized);
220        }
221        self.authorized = true;
222        Ok(installed)
223    }
224
225    fn mirror(&mut self) {
226        self.records.extend(
227            self.state.boxes()[self.mirrored..]
228                .iter()
229                .cloned()
230                .map(Record::Box),
231        );
232        self.mirrored = self.state.boxes().len();
233    }
234
235    fn flush(&mut self) -> Result<(), String> {
236        if self.durable < self.records.len() {
237            self.session
238                .persist(self.records[self.durable..].to_vec())?;
239            self.durable = self.records.len();
240        }
241        Ok(())
242    }
243
244    fn persist_done(&mut self, resume: bool) -> Result<(), String> {
245        self.mirror();
246        let anchor = self
247            .state
248            .boxes()
249            .last()
250            .map_or(0, |value| value.id().get());
251        let index = match self.records.last() {
252            Some(Record::Event(event)) if event.after_box_id == anchor => {
253                event.event_index.checked_add(1)
254            }
255            _ => Some(1),
256        }
257        .ok_or_else(|| "event index space was exhausted".to_owned())?;
258        self.records.push(Record::Event(
259            EventRecord::new(
260                anchor,
261                index,
262                0,
263                "llm_done".into(),
264                json!({"resume": resume}),
265            )
266            .map_err(|error| error.to_string())?,
267        ));
268        self.flush()
269    }
270}
271
272fn render_input(values: &[BoxValue]) -> Result<String, String> {
273    let mut output = String::new();
274    for value in values {
275        let BoxValue::History(section) = value else {
276            return Err("Codex steer contains a non-history value".into());
277        };
278        if section.is_empty() {
279            continue;
280        }
281        if !output.is_empty() && !output.ends_with('\n') {
282            output.push('\n');
283        }
284        output.push_str(section);
285    }
286    Ok(output)
287}
288
289fn returned_calls(boxes: &[ChatBox]) -> Result<Vec<ToolCallId>, String> {
290    let mut returned = Vec::new();
291    for value in boxes {
292        if let Some(result) = value
293            .tool_result_metadata()
294            .map_err(|error| format!("{error:?}"))?
295        {
296            returned.push(result.tool_call_id);
297        }
298    }
299    Ok(returned)
300}