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