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}