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