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, USER_MESSAGE_TYPE};
7use kcode_k1_chat_thread_actions::ChatThreadActions;
8pub use kcode_k1_chat_thread_actions::{AccessContext, AccessPolicy, ProfileId};
9pub use kcode_k1_chat_thread_durable_turn::{
10    BoxValue, ChatBox, PreparedCall, PreparedMailboxFlush, Status, ToolCallId,
11};
12use kcode_k1_chat_thread_durable_turn::{DurableTurn, RestartError, ShimOutput};
13use serde_json::Value;
14use std::sync::Arc;
15
16#[derive(Clone, Debug, Eq, PartialEq)]
17pub enum TransitionError {
18    Unauthorized,
19    NotStalled,
20    NotRestartable,
21    Internal(String),
22}
23
24pub struct DurableThread {
25    turn: DurableTurn,
26    actions: ChatThreadActions,
27    authorized: bool,
28}
29
30impl DurableThread {
31    pub fn recover(session: Session, kmap: Arc<K1AccessKmap>) -> Result<Self, String> {
32        Ok(Self {
33            turn: DurableTurn::recover(session)?,
34            actions: ChatThreadActions::new(kmap),
35            authorized: false,
36        })
37    }
38
39    pub fn boxes(&self) -> &[ChatBox] {
40        self.turn.boxes()
41    }
42
43    pub fn status(&self) -> Status {
44        self.turn.status()
45    }
46
47    pub fn accept_box(
48        &mut self,
49        box_type: String,
50        contents: String,
51        hidden_type: String,
52        hidden_contents: String,
53    ) -> Result<(), String> {
54        self.turn
55            .accept(box_type, contents, hidden_type, hidden_contents)
56    }
57
58    pub fn accept_user(
59        &mut self,
60        context: AccessContext,
61        profile_id: ProfileId,
62        policy: AccessPolicy,
63        contents: String,
64    ) -> Result<(), TransitionError> {
65        let installed = self.bind_authorization(context, profile_id, policy)?;
66        match self.turn.accept(
67            USER_MESSAGE_TYPE.into(),
68            contents,
69            String::new(),
70            String::new(),
71        ) {
72            Ok(()) => Ok(()),
73            Err(error) => {
74                if installed {
75                    self.clear_authorization();
76                }
77                Err(TransitionError::Internal(error))
78            }
79        }
80    }
81
82    pub fn accept_return(
83        &mut self,
84        id: ToolCallId,
85        result: Result<String, String>,
86    ) -> Result<(), String> {
87        self.turn.accept_tool_return(id, result)
88    }
89
90    pub fn prepare_stage(
91        &mut self,
92        job: u64,
93        text: String,
94        boxes: Vec<BoxValue>,
95    ) -> Result<Vec<PreparedCall>, String> {
96        self.turn.prepare_stage(job, text, boxes)
97    }
98
99    pub fn launch_action(&mut self, name: &str, arguments: &str) -> Result<String, String> {
100        if name == "SendMessage" {
101            launch_send_message(&mut self.turn, arguments)
102        } else {
103            self.actions.launch(name, arguments)
104        }
105    }
106
107    pub fn accept_tool_message(&mut self, id: ToolCallId, contents: String) -> Result<(), String> {
108        self.turn.accept_tool_message(id, contents)
109    }
110
111    pub fn accept_tool_return(
112        &mut self,
113        id: ToolCallId,
114        result: Result<String, String>,
115    ) -> Result<(), String> {
116        self.turn.accept_tool_return(id, result)
117    }
118
119    pub fn accept_tool_return_v2(
120        &mut self,
121        id: ToolCallId,
122        result: Result<String, String>,
123        metadata_type: String,
124        metadata_contents: String,
125    ) -> Result<(), String> {
126        self.turn
127            .accept_tool_return_v2(id, result, metadata_type, metadata_contents)
128    }
129
130    pub fn prepare_mailbox_flush(
131        &mut self,
132        job: u64,
133    ) -> Result<Option<PreparedMailboxFlush>, String> {
134        self.turn.prepare_mailbox_flush(job)
135    }
136
137    pub fn prepared_input(&self, prepared: &PreparedMailboxFlush) -> Result<String, String> {
138        self.turn.validate_mailbox_flush(prepared)?;
139        render_input(prepared.values())
140    }
141
142    pub fn commit_mailbox_flush(&mut self, prepared: PreparedMailboxFlush) -> Result<(), String> {
143        self.turn.commit_mailbox_flush(prepared)
144    }
145
146    pub fn begin_input(&mut self) -> Result<Option<(u64, String)>, String> {
147        let Some(start) = self.turn.begin()? else {
148            return Ok(None);
149        };
150        Ok(Some((start.job, render_input(&start.values)?)))
151    }
152
153    pub fn complete(&mut self, job: u64, output: ShimOutput<BoxValue>) -> Result<bool, String> {
154        self.turn.complete(job, output)
155    }
156
157    pub fn fail(&mut self, job: u64, error: String, restartable: bool) {
158        self.turn.fail(job, error, restartable);
159        self.clear_authorization();
160    }
161
162    pub fn restart(
163        &mut self,
164        context: AccessContext,
165        profile_id: ProfileId,
166        policy: AccessPolicy,
167    ) -> Result<(), TransitionError> {
168        let installed = self.bind_authorization(context, profile_id, policy)?;
169        if let Err(error) = self.turn.restart().map_err(|error| match error {
170            RestartError::NotStalled => TransitionError::NotStalled,
171            RestartError::ProviderActionAccepted => TransitionError::NotRestartable,
172        }) {
173            if installed {
174                self.clear_authorization();
175            }
176            return Err(error);
177        }
178        Ok(())
179    }
180
181    pub fn clear_authorization(&mut self) {
182        self.actions.clear_authorization();
183        self.authorized = false;
184    }
185
186    fn bind_authorization(
187        &mut self,
188        context: AccessContext,
189        profile_id: ProfileId,
190        policy: AccessPolicy,
191    ) -> Result<bool, TransitionError> {
192        let installed = !self.authorized;
193        if self
194            .actions
195            .bind_authorization(context, profile_id, policy)
196            .is_err()
197        {
198            if installed {
199                self.actions.clear_authorization();
200            }
201            return Err(TransitionError::Unauthorized);
202        }
203        self.authorized = true;
204        Ok(installed)
205    }
206}
207
208fn launch_send_message(turn: &mut DurableTurn, arguments: &str) -> Result<String, String> {
209    let parsed: Value = serde_json::from_str(arguments).map_err(|_| invalid_send_message())?;
210    let Value::Object(mut fields) = parsed else {
211        return Err(invalid_send_message());
212    };
213    if fields.len() != 1 {
214        return Err(invalid_send_message());
215    }
216    let Some(Value::String(message)) = fields.remove("message") else {
217        return Err(invalid_send_message());
218    };
219    if message.is_empty() {
220        return Err(invalid_send_message());
221    }
222    turn.accept(
223        AGENT_MESSAGE_TYPE.into(),
224        message,
225        String::new(),
226        String::new(),
227    )?;
228    Ok("success".into())
229}
230
231fn invalid_send_message() -> String {
232    "invalid SendMessage arguments".into()
233}
234
235fn render_input(values: &[BoxValue]) -> Result<String, String> {
236    let mut output = String::new();
237    for value in values {
238        let BoxValue::History(section) = value else {
239            return Err("Codex provider input contains a non-history value".into());
240        };
241        if section.is_empty() {
242            continue;
243        }
244        if !output.is_empty() && !output.ends_with('\n') {
245            output.push('\n');
246        }
247        output.push_str(section);
248    }
249    Ok(output)
250}
251
252#[cfg(test)]
253mod tests {
254    use super::*;
255    use kcode_k1_access::K1Access;
256    use kcode_k1_chat_codex_state::Call;
257    use kcode_k1_chat_persistence::{K1ChatPersistence, Session};
258    use kcode_k1_chat_state::{AGENT_RESPONSE_TYPE, TOOL_CALL_TYPE, TOOL_RESULT_TYPE};
259    use kcode_k1_groups::K1Groups;
260    use kcode_k1_kmap::K1Kmap;
261    use kcode_k1_peering::K1Peering;
262    use kcode_k1_txn_ordering::K1TxnOrdering;
263    use tempfile::TempDir;
264
265    fn fixture() -> (TempDir, Session, Arc<K1AccessKmap>) {
266        let root = TempDir::new().unwrap();
267        let ordering = Arc::new(K1TxnOrdering::open(&root.path().join("ordering")).unwrap());
268        let peering =
269            Arc::new(K1Peering::open(&root.path().join("peering"), Arc::clone(&ordering)).unwrap());
270        let groups = Arc::new(
271            K1Groups::open(
272                &root.path().join("groups"),
273                Arc::clone(&ordering),
274                Arc::clone(&peering),
275            )
276            .unwrap(),
277        );
278        let access = Arc::new(
279            K1Access::open(
280                &root.path().join("access"),
281                Arc::clone(&ordering),
282                Arc::clone(&peering),
283                groups,
284            )
285            .unwrap(),
286        );
287        let kmap = Arc::new(
288            K1Kmap::open(
289                &root.path().join("kmap"),
290                Arc::clone(&ordering),
291                Arc::clone(&peering),
292            )
293            .unwrap(),
294        );
295        let access_kmap = Arc::new(K1AccessKmap::open(access, kmap).unwrap());
296        let persistence =
297            K1ChatPersistence::open(&root.path().join("persistence"), ordering, peering).unwrap();
298        let (session, original) = persistence.session([11; 12]).unwrap();
299        assert!(original.records.is_empty());
300        (root, session, access_kmap)
301    }
302
303    #[test]
304    fn send_message_is_durable_ordered_and_recovered_once() {
305        let (_root, session, access_kmap) = fixture();
306        let mut thread = DurableThread::recover(session.clone(), Arc::clone(&access_kmap)).unwrap();
307        thread
308            .accept_box(
309                USER_MESSAGE_TYPE.into(),
310                "hello".into(),
311                String::new(),
312                String::new(),
313            )
314            .unwrap();
315        let (job, _) = thread.begin_input().unwrap().unwrap();
316        let arguments = r#"{"message":"WORKING_MESSAGE"}"#;
317        let calls = thread
318            .prepare_stage(
319                job,
320                String::new(),
321                vec![BoxValue::Call(Ok(Call {
322                    name: "SendMessage".into(),
323                    arguments: arguments.into(),
324                }))],
325            )
326            .unwrap();
327        assert_eq!(calls.len(), 1);
328        let result = thread.launch_action("SendMessage", arguments);
329        assert_eq!(result, Ok("success".into()));
330        thread
331            .accept_tool_return(calls[0].tool_call_id, result)
332            .unwrap();
333
334        let prepared = thread.prepare_mailbox_flush(job).unwrap().unwrap();
335        let input = thread.prepared_input(&prepared).unwrap();
336        let call_at = input.find("| Tool Call]").unwrap();
337        let message_at = input.find("| Agent Message]").unwrap();
338        let result_at = input.find("| Tool Result]").unwrap();
339        assert!(call_at < message_at && message_at < result_at);
340        assert!(input.ends_with("| Agent Response]\n"));
341        assert_eq!(
342            thread
343                .boxes()
344                .iter()
345                .map(ChatBox::box_type)
346                .collect::<Vec<_>>(),
347            [
348                USER_MESSAGE_TYPE,
349                AGENT_RESPONSE_TYPE,
350                TOOL_CALL_TYPE,
351                AGENT_MESSAGE_TYPE,
352                TOOL_RESULT_TYPE,
353            ]
354        );
355        let message = thread
356            .boxes()
357            .iter()
358            .find(|value| value.box_type() == AGENT_MESSAGE_TYPE)
359            .unwrap();
360        assert_eq!(message.contents(), "WORKING_MESSAGE");
361        assert_eq!((message.hidden_type(), message.hidden_contents()), ("", ""));
362
363        thread.commit_mailbox_flush(prepared).unwrap();
364        assert!(
365            !thread
366                .complete(job, ShimOutput { items: Vec::new() })
367                .unwrap()
368        );
369        drop(thread);
370
371        let mut recovered = DurableThread::recover(session, access_kmap).unwrap();
372        assert_eq!(
373            recovered
374                .boxes()
375                .iter()
376                .filter(|value| value.box_type() == AGENT_MESSAGE_TYPE)
377                .count(),
378            1
379        );
380        let before = recovered.boxes().len();
381        assert!(recovered.launch_action("CurrentTime", "{}").is_ok());
382        assert_eq!(recovered.boxes().len(), before);
383    }
384
385    #[test]
386    fn send_message_rejects_invalid_arguments_without_a_message() {
387        let (_root, session, access_kmap) = fixture();
388        let mut thread = DurableThread::recover(session, access_kmap).unwrap();
389        for arguments in [
390            "",
391            "{",
392            "null",
393            "[]",
394            "{}",
395            r#"{"message":""}"#,
396            r#"{"message":1}"#,
397            r#"{"message":"x","extra":true}"#,
398        ] {
399            let before = thread.boxes().len();
400            assert_eq!(
401                thread.launch_action("SendMessage", arguments),
402                Err("invalid SendMessage arguments".into())
403            );
404            assert_eq!(thread.boxes().len(), before);
405        }
406    }
407}