Skip to main content

kcode_k1_chat_thread_durable_state/
lib.rs

1#![forbid(unsafe_code)]
2
3use kcode_k1_access_kmap::K1AccessKmap;
4use kcode_k1_chat_persistence::Session;
5pub use kcode_k1_chat_state::BoxId;
6use kcode_k1_chat_state::{AGENT_MESSAGE_TYPE, AGENT_RESPONSE_TYPE, USER_MESSAGE_TYPE};
7pub use kcode_k1_chat_thread_durable_turn::{
8    BoxValue, ChatBox, ChatDiagnostic, EventRecord, ModelUsage, PreflightItem, PreflightMode,
9    PreparedCall, PreparedMailboxFlush, PreparedPreflightCall, Status, TokenBreakdown, ToolCallId,
10};
11use kcode_k1_chat_thread_durable_turn::{DurableTurn, RestartError, ShimOutput};
12pub use kcode_k1_chat_thread_ktools::{AccessContext, AccessPolicy, ProfileId, SetLaunchNodeKtool};
13use kcode_k1_chat_thread_ktools::{ChatThreadKtoolExecutor, ChatThreadKtools};
14use serde::Deserialize;
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    turn: DurableTurn,
27    ktools: ChatThreadKtoolExecutor,
28    authorized: bool,
29}
30
31impl DurableThread {
32    pub fn recover(session: Session, kmap: Arc<K1AccessKmap>) -> Result<Self, String> {
33        Self::recover_with_ktools(session, ChatThreadKtools::new(kmap))
34    }
35
36    pub fn recover_with_social(
37        session: Session,
38        kmap: Arc<K1AccessKmap>,
39        social: kcode_k1_ktool_social::SocialKtools,
40    ) -> Result<Self, String> {
41        Self::recover_with_ktools(session, ChatThreadKtools::new_with_social(kmap, social))
42    }
43
44    pub fn recover_with_social_and_set_launch_node(
45        session: Session,
46        kmap: Arc<K1AccessKmap>,
47        social: kcode_k1_ktool_social::SocialKtools,
48        set_launch_node: SetLaunchNodeKtool,
49    ) -> Result<Self, String> {
50        Self::recover_with_ktools(
51            session,
52            ChatThreadKtools::new_with_social_and_set_launch_node(kmap, social, set_launch_node),
53        )
54    }
55
56    pub fn boxes(&self) -> &[ChatBox] {
57        self.turn.boxes()
58    }
59
60    pub fn events(&self) -> Vec<EventRecord> {
61        self.turn.events()
62    }
63
64    pub fn status(&self) -> Status {
65        self.turn.status()
66    }
67
68    pub fn preflight_calls(&self) -> &[PreparedPreflightCall] {
69        self.turn.preflight_calls()
70    }
71
72    pub fn preflight_executor(&self) -> ChatThreadKtoolExecutor {
73        self.ktools.clone()
74    }
75
76    pub fn prepare_preflight(
77        &mut self,
78        context: AccessContext,
79        profile_id: ProfileId,
80        policy: AccessPolicy,
81        items: Vec<PreflightItem>,
82    ) -> Result<Vec<PreparedPreflightCall>, TransitionError> {
83        for item in &items {
84            if let PreflightItem::KtoolCall { name, .. } = item {
85                let supported = self
86                    .ktools
87                    .supports(name)
88                    .map_err(TransitionError::Internal)?;
89                if !supported {
90                    return Err(TransitionError::Internal(
91                        "unsupported preflight Ktool".to_owned(),
92                    ));
93                }
94            }
95        }
96        let installed = self.bind_authorization(context, profile_id, policy)?;
97        match self.turn.prepare_preflight(items) {
98            Ok(calls) => Ok(calls),
99            Err(error) => {
100                if installed {
101                    self.clear_authorization();
102                }
103                Err(TransitionError::Internal(error))
104            }
105        }
106    }
107
108    pub fn authorize_preflight(
109        &mut self,
110        context: AccessContext,
111        profile_id: ProfileId,
112        policy: AccessPolicy,
113    ) -> Result<(), TransitionError> {
114        self.bind_authorization(context, profile_id, policy)
115            .map(|_| ())
116    }
117
118    pub fn accept_box(
119        &mut self,
120        box_type: String,
121        contents: String,
122        hidden_type: String,
123        hidden_contents: String,
124    ) -> Result<(), String> {
125        self.turn
126            .accept(box_type, contents, hidden_type, hidden_contents)
127    }
128
129    pub fn accept_external_box(
130        &mut self,
131        box_type: String,
132        contents: String,
133        hidden_type: String,
134        hidden_contents: String,
135    ) -> Result<(), TransitionError> {
136        if box_type == USER_MESSAGE_TYPE {
137            return Err(TransitionError::Unauthorized);
138        }
139        self.accept_box(box_type, contents, hidden_type, hidden_contents)
140            .map_err(TransitionError::Internal)
141    }
142
143    pub fn accept_user(
144        &mut self,
145        context: AccessContext,
146        profile_id: ProfileId,
147        policy: AccessPolicy,
148        contents: String,
149    ) -> Result<(), TransitionError> {
150        let installed = self.bind_authorization(context, profile_id, policy)?;
151        match self.turn.accept(
152            USER_MESSAGE_TYPE.into(),
153            contents,
154            String::new(),
155            String::new(),
156        ) {
157            Ok(()) => Ok(()),
158            Err(error) => {
159                if installed {
160                    self.clear_authorization();
161                }
162                Err(TransitionError::Internal(error))
163            }
164        }
165    }
166
167    pub fn accept_return(
168        &mut self,
169        id: ToolCallId,
170        result: Result<String, String>,
171    ) -> Result<(), String> {
172        self.turn.accept_tool_return(id, result)
173    }
174
175    pub fn prepare_stage(
176        &mut self,
177        job: u64,
178        text: String,
179        boxes: Vec<BoxValue>,
180    ) -> Result<Vec<PreparedCall>, String> {
181        self.turn.prepare_stage(job, text, boxes)
182    }
183
184    pub fn launch_action(&mut self, name: &str, arguments: &str) -> Result<String, String> {
185        if !kcode_k1_ktool_docs::is_known_ktool(name) {
186            return Err("unknown Ktool".into());
187        }
188        match name {
189            "KtoolDocs" => kcode_k1_ktool_docs::ktool_docs(arguments),
190            "SendMessage" => launch_send_message(&mut self.turn, arguments),
191            _ => self.ktools.launch(name, arguments),
192        }
193    }
194
195    pub fn accept_tool_message(&mut self, id: ToolCallId, contents: String) -> Result<(), String> {
196        self.turn.accept_tool_message(id, contents)
197    }
198
199    pub fn accept_tool_return(
200        &mut self,
201        id: ToolCallId,
202        result: Result<String, String>,
203    ) -> Result<(), String> {
204        self.turn.accept_tool_return(id, result)
205    }
206
207    pub fn accept_tool_return_v2(
208        &mut self,
209        id: ToolCallId,
210        result: Result<String, String>,
211        metadata_type: String,
212        metadata_contents: String,
213    ) -> Result<(), String> {
214        self.turn
215            .accept_tool_return_v2(id, result, metadata_type, metadata_contents)
216    }
217
218    pub fn prepare_mailbox_flush(
219        &mut self,
220        job: u64,
221    ) -> Result<Option<PreparedMailboxFlush>, String> {
222        self.turn.prepare_mailbox_flush(job)
223    }
224
225    pub fn prepared_input(&self, prepared: &PreparedMailboxFlush) -> Result<String, String> {
226        self.turn.validate_mailbox_flush(prepared)?;
227        render_input(prepared.values())
228    }
229
230    pub fn commit_mailbox_flush(&mut self, prepared: PreparedMailboxFlush) -> Result<(), String> {
231        self.turn.commit_mailbox_flush(prepared)
232    }
233
234    pub fn begin_input(&mut self) -> Result<Option<(u64, String)>, String> {
235        let Some(start) = self.turn.begin()? else {
236            return Ok(None);
237        };
238        Ok(Some((start.job, render_input(&start.values)?)))
239    }
240
241    pub fn complete(&mut self, job: u64, output: ShimOutput<BoxValue>) -> Result<bool, String> {
242        self.turn.complete(job, output)
243    }
244
245    pub fn complete_with_terminal_response(
246        &mut self,
247        job: u64,
248        output: ShimOutput<BoxValue>,
249    ) -> Result<(bool, u64), String> {
250        let terminal_index = self.turn.boxes().len();
251        let resume = self.complete(job, output)?;
252        let terminal =
253            self.turn.boxes().get(terminal_index).ok_or_else(|| {
254                "completion did not append a terminal Agent Response box".to_owned()
255            })?;
256        if terminal.box_type() != AGENT_RESPONSE_TYPE {
257            return Err("completion terminal box was not an Agent Response".to_owned());
258        }
259        Ok((resume, terminal.id().get()))
260    }
261
262    pub fn complete_recoverable_failure_with_terminal_response(
263        &mut self,
264        job: u64,
265        message: String,
266    ) -> Result<(bool, u64), String> {
267        let terminal_index = self.turn.boxes().len();
268        let resume = self.turn.complete_recoverable_failure(job, message)?;
269        let terminal = self.turn.boxes().get(terminal_index).ok_or_else(|| {
270            "recoverable failure did not append a terminal Agent Response box".to_owned()
271        })?;
272        if terminal.box_type() != AGENT_RESPONSE_TYPE {
273            return Err("recoverable failure terminal box was not an Agent Response".to_owned());
274        }
275        Ok((resume, terminal.id().get()))
276    }
277
278    pub fn reset_provider_context(&mut self) -> Result<(), String> {
279        self.turn.reset_provider_context()
280    }
281
282    pub fn record_diagnostic(&mut self, diagnostic: ChatDiagnostic) -> Result<(), String> {
283        self.turn.record_diagnostic(diagnostic)
284    }
285
286    pub fn halt_critical(&mut self, message: String) -> Result<(), String> {
287        let result = self.turn.halt_critical(message);
288        let _ = self.ktools.clear_authorization();
289        self.authorized = false;
290        result
291    }
292
293    pub fn record_model_usage(
294        &mut self,
295        connected_box_id: u64,
296        usage: ModelUsage,
297    ) -> Result<(), String> {
298        self.turn.record_model_usage(connected_box_id, usage)
299    }
300
301    pub fn fail(&mut self, job: u64, error: String, restartable: bool) {
302        self.turn.fail(job, error, restartable);
303        self.clear_authorization();
304    }
305
306    pub fn restart(
307        &mut self,
308        context: AccessContext,
309        profile_id: ProfileId,
310        policy: AccessPolicy,
311    ) -> Result<(), TransitionError> {
312        let installed = self.bind_authorization(context, profile_id, policy)?;
313        if let Err(error) = self.turn.restart().map_err(|error| match error {
314            RestartError::NotStalled => TransitionError::NotStalled,
315            RestartError::ProviderActionAccepted => TransitionError::NotRestartable,
316        }) {
317            if installed {
318                self.clear_authorization();
319            }
320            return Err(error);
321        }
322        Ok(())
323    }
324
325    pub fn clear_authorization(&mut self) {
326        if self.has_unresolved_preflight() {
327            return;
328        }
329        let _ = self.ktools.clear_authorization();
330        self.authorized = false;
331    }
332
333    fn has_unresolved_preflight(&self) -> bool {
334        self.turn.preflight_calls().iter().any(|call| {
335            !self.turn.boxes().iter().any(|value| {
336                value
337                    .tool_result_metadata()
338                    .ok()
339                    .flatten()
340                    .is_some_and(|result| result.tool_call_id == call.tool_call_id)
341            })
342        })
343    }
344
345    fn recover_with_ktools(session: Session, ktools: ChatThreadKtools) -> Result<Self, String> {
346        Ok(Self {
347            turn: DurableTurn::recover(session)?,
348            ktools: ChatThreadKtoolExecutor::new(ktools),
349            authorized: false,
350        })
351    }
352
353    fn bind_authorization(
354        &mut self,
355        context: AccessContext,
356        profile_id: ProfileId,
357        policy: AccessPolicy,
358    ) -> Result<bool, TransitionError> {
359        let installed = !self.authorized;
360        if self
361            .ktools
362            .bind_authorization(context, profile_id, policy)
363            .is_err()
364        {
365            if installed {
366                let _ = self.ktools.clear_authorization();
367            }
368            return Err(TransitionError::Unauthorized);
369        }
370        self.authorized = true;
371        Ok(installed)
372    }
373}
374
375#[derive(Deserialize)]
376#[serde(deny_unknown_fields)]
377struct SendMessageArguments {
378    message: String,
379}
380
381fn launch_send_message(turn: &mut DurableTurn, arguments: &str) -> Result<String, String> {
382    let parsed: SendMessageArguments =
383        serde_json::from_str(arguments).map_err(|_| invalid_send_message())?;
384    if parsed.message.is_empty() {
385        return Err(invalid_send_message());
386    }
387    turn.accept(
388        AGENT_MESSAGE_TYPE.into(),
389        parsed.message,
390        String::new(),
391        String::new(),
392    )?;
393    Ok("success".into())
394}
395
396fn invalid_send_message() -> String {
397    "invalid SendMessage arguments".into()
398}
399
400fn render_input(values: &[BoxValue]) -> Result<String, String> {
401    let mut output = String::new();
402    for value in values {
403        let BoxValue::History(section) = value else {
404            return Err("Codex provider input contains a non-history value".into());
405        };
406        if section.is_empty() {
407            continue;
408        }
409        if !output.is_empty() && !output.ends_with('\n') {
410            output.push('\n');
411        }
412        output.push_str(section);
413    }
414    Ok(output)
415}
416
417#[cfg(test)]
418mod tests;