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_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    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
267fn launch_send_message(turn: &mut DurableTurn, arguments: &str) -> Result<String, String> {
268    let parsed: Value = serde_json::from_str(arguments).map_err(|_| invalid_send_message())?;
269    let Value::Object(mut fields) = parsed else {
270        return Err(invalid_send_message());
271    };
272    if fields.len() != 1 {
273        return Err(invalid_send_message());
274    }
275    let Some(Value::String(message)) = fields.remove("message") else {
276        return Err(invalid_send_message());
277    };
278    if message.is_empty() {
279        return Err(invalid_send_message());
280    }
281    turn.accept(
282        AGENT_MESSAGE_TYPE.into(),
283        message,
284        String::new(),
285        String::new(),
286    )?;
287    Ok("success".into())
288}
289
290fn invalid_send_message() -> String {
291    "invalid SendMessage arguments".into()
292}
293
294fn render_input(values: &[BoxValue]) -> Result<String, String> {
295    let mut output = String::new();
296    for value in values {
297        let BoxValue::History(section) = value else {
298            return Err("Codex provider input contains a non-history value".into());
299        };
300        if section.is_empty() {
301            continue;
302        }
303        if !output.is_empty() && !output.ends_with('\n') {
304            output.push('\n');
305        }
306        output.push_str(section);
307    }
308    Ok(output)
309}
310
311#[cfg(test)]
312mod tests {
313    use super::*;
314    use kcode_k1_access::K1Access;
315    use kcode_k1_chat_codex_state::Call;
316    use kcode_k1_chat_persistence::{K1ChatPersistence, Session};
317    use kcode_k1_chat_state::{AGENT_RESPONSE_TYPE, TOOL_CALL_TYPE, TOOL_RESULT_TYPE};
318    use kcode_k1_groups::K1Groups;
319    use kcode_k1_kmap::K1Kmap;
320    use kcode_k1_peering::K1Peering;
321    use kcode_k1_txn_ordering::K1TxnOrdering;
322    use tempfile::TempDir;
323
324    fn fixture() -> (TempDir, Session, Arc<K1AccessKmap>) {
325        let root = TempDir::new().unwrap();
326        let ordering = Arc::new(K1TxnOrdering::open(&root.path().join("ordering")).unwrap());
327        let peering =
328            Arc::new(K1Peering::open(&root.path().join("peering"), Arc::clone(&ordering)).unwrap());
329        let groups = Arc::new(
330            K1Groups::open(
331                &root.path().join("groups"),
332                Arc::clone(&ordering),
333                Arc::clone(&peering),
334            )
335            .unwrap(),
336        );
337        let access = Arc::new(
338            K1Access::open(
339                &root.path().join("access"),
340                Arc::clone(&ordering),
341                Arc::clone(&peering),
342                groups,
343            )
344            .unwrap(),
345        );
346        let kmap = Arc::new(
347            K1Kmap::open(
348                &root.path().join("kmap"),
349                Arc::clone(&ordering),
350                Arc::clone(&peering),
351            )
352            .unwrap(),
353        );
354        let access_kmap = Arc::new(K1AccessKmap::open(access, kmap).unwrap());
355        let persistence =
356            K1ChatPersistence::open(&root.path().join("persistence"), ordering, peering).unwrap();
357        let (session, original) = persistence.session([11; 12]).unwrap();
358        assert!(original.records.is_empty());
359        (root, session, access_kmap)
360    }
361
362    fn accept_user_box(thread: &mut DurableThread, contents: &str) {
363        thread
364            .accept_box(
365                USER_MESSAGE_TYPE.into(),
366                contents.into(),
367                String::new(),
368                String::new(),
369            )
370            .unwrap();
371    }
372
373    #[test]
374    fn send_message_is_durable_ordered_and_recovered_once() {
375        let (_root, session, access_kmap) = fixture();
376        let mut thread = DurableThread::recover(session.clone(), Arc::clone(&access_kmap)).unwrap();
377        accept_user_box(&mut thread, "hello");
378        let (job, _) = thread.begin_input().unwrap().unwrap();
379        let arguments = r#"{"message":"WORKING_MESSAGE"}"#;
380        let calls = thread
381            .prepare_stage(
382                job,
383                String::new(),
384                vec![BoxValue::Call(Ok(Call {
385                    name: "SendMessage".into(),
386                    arguments: arguments.into(),
387                }))],
388            )
389            .unwrap();
390        assert_eq!(calls.len(), 1);
391        let result = thread.launch_action("SendMessage", arguments);
392        assert_eq!(result, Ok("success".into()));
393        thread
394            .accept_tool_return(calls[0].tool_call_id, result)
395            .unwrap();
396
397        let prepared = thread.prepare_mailbox_flush(job).unwrap().unwrap();
398        let input = thread.prepared_input(&prepared).unwrap();
399        let call_at = input.find("| Tool Call]").unwrap();
400        let message_at = input.find("| Agent Message]").unwrap();
401        let result_at = input.find("| Tool Result]").unwrap();
402        assert!(call_at < message_at && message_at < result_at);
403        assert!(input.ends_with("| Agent Response]\n"));
404        assert_eq!(
405            thread
406                .boxes()
407                .iter()
408                .map(ChatBox::box_type)
409                .collect::<Vec<_>>(),
410            [
411                USER_MESSAGE_TYPE,
412                AGENT_RESPONSE_TYPE,
413                TOOL_CALL_TYPE,
414                AGENT_MESSAGE_TYPE,
415                TOOL_RESULT_TYPE,
416            ]
417        );
418        let message = thread
419            .boxes()
420            .iter()
421            .find(|value| value.box_type() == AGENT_MESSAGE_TYPE)
422            .unwrap();
423        assert_eq!(message.contents(), "WORKING_MESSAGE");
424        assert_eq!((message.hidden_type(), message.hidden_contents()), ("", ""));
425
426        thread.commit_mailbox_flush(prepared).unwrap();
427        assert!(
428            !thread
429                .complete(job, ShimOutput { items: Vec::new() })
430                .unwrap()
431        );
432        drop(thread);
433
434        let mut recovered = DurableThread::recover(session, access_kmap).unwrap();
435        assert_eq!(
436            recovered
437                .boxes()
438                .iter()
439                .filter(|value| value.box_type() == AGENT_MESSAGE_TYPE)
440                .count(),
441            1
442        );
443        let before = recovered.boxes().len();
444        assert!(recovered.launch_action("CurrentTime", "{}").is_ok());
445        assert_eq!(
446            recovered.launch_action("NotRegistered", "{}"),
447            Err("unknown Ktool".into())
448        );
449        let social = recovered.launch_action("ListContacts", "{}");
450        assert_eq!(social, Err("social Ktools are unavailable".into()));
451        assert_eq!(recovered.boxes().len(), before);
452    }
453
454    #[test]
455    fn complete_returns_terminal_response_before_queued_arrivals() {
456        let (_root, session, access_kmap) = fixture();
457        let mut thread = DurableThread::recover(session, access_kmap).unwrap();
458        accept_user_box(&mut thread, "first");
459        let (job, _) = thread.begin_input().unwrap().unwrap();
460        let terminal_index = thread.boxes().len();
461        accept_user_box(&mut thread, "queued");
462
463        let (resume, terminal_id) = thread
464            .complete_with_terminal_response(job, ShimOutput { items: Vec::new() })
465            .unwrap();
466        assert!(resume);
467        assert_eq!(terminal_id, thread.boxes()[terminal_index].id().get());
468        assert_eq!(
469            thread.boxes()[terminal_index].box_type(),
470            AGENT_RESPONSE_TYPE
471        );
472        assert_eq!(
473            thread.boxes()[terminal_index + 1].box_type(),
474            USER_MESSAGE_TYPE
475        );
476    }
477
478    #[test]
479    fn send_message_rejects_invalid_arguments_without_a_message() {
480        let (_root, session, access_kmap) = fixture();
481        let mut thread = DurableThread::recover(session, access_kmap).unwrap();
482        for arguments in [
483            "",
484            "{",
485            "null",
486            "[]",
487            "{}",
488            r#"{"message":""}"#,
489            r#"{"message":1}"#,
490            r#"{"message":"x","extra":true}"#,
491        ] {
492            let before = thread.boxes().len();
493            assert_eq!(
494                thread.launch_action("SendMessage", arguments),
495                Err("invalid SendMessage arguments".into())
496            );
497            assert_eq!(thread.boxes().len(), before);
498        }
499    }
500}