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