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}