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}