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}