1#![forbid(unsafe_code)]
2
3pub use kcode_k1_chat_codex_state::{
4 BoxValue, ChatBox, PreparedCall, PreparedSteer, RestartError, ShimOutput, Start, Status,
5 ToolCallId,
6};
7
8use kcode_k1_chat_codex_state::ConversationState;
9use kcode_k1_chat_persistence::{EventRecord, Record, Session};
10use kcode_k1_chat_thread_recovery::recover as recover_thread;
11use serde_json::json;
12
13pub struct DurableTurn {
14 state: ConversationState,
15 records: Vec<Record>,
16 mirrored: usize,
17 durable: usize,
18 session: Session,
19 returned: Vec<ToolCallId>,
20}
21
22impl DurableTurn {
23 pub fn recover(session: Session) -> Result<Self, String> {
24 let recovered = recover_thread(&session)?;
25 let returned = returned_ids(recovered.state.boxes())?;
26 Ok(Self {
27 state: recovered.state,
28 records: recovered.records,
29 mirrored: recovered.mirrored,
30 durable: recovered.durable,
31 session,
32 returned,
33 })
34 }
35
36 pub fn boxes(&self) -> &[ChatBox] {
37 self.state.boxes()
38 }
39
40 pub fn status(&self) -> Status {
41 self.state.status()
42 }
43
44 pub fn accept(
45 &mut self,
46 box_type: String,
47 contents: String,
48 hidden_type: String,
49 hidden_contents: String,
50 ) -> Result<(), String> {
51 let result = self
52 .state
53 .accept(box_type, contents, hidden_type, hidden_contents);
54 self.finish(result)
55 }
56
57 pub fn accept_tool_return(
58 &mut self,
59 tool_call_id: ToolCallId,
60 result: Result<String, String>,
61 ) -> Result<(), String> {
62 if self.returned.contains(&tool_call_id) {
63 return self.finish(Ok(()));
64 }
65 let accepted = self.state.accept_tool_return(tool_call_id, result);
66 if accepted.is_ok() {
67 self.returned.push(tool_call_id);
68 }
69 self.finish(accepted)
70 }
71
72 pub fn accept_tool_message(
73 &mut self,
74 tool_call_id: ToolCallId,
75 message: String,
76 ) -> Result<(), String> {
77 let result = self.state.accept_tool_message(tool_call_id, message);
78 self.finish(result)
79 }
80
81 pub fn accept_tool_return_v2(
82 &mut self,
83 tool_call_id: ToolCallId,
84 result: Result<String, String>,
85 metadata_type: String,
86 metadata_contents: String,
87 ) -> Result<(), String> {
88 let accepted = self.state.accept_tool_return_v2(
89 tool_call_id,
90 result,
91 metadata_type,
92 metadata_contents,
93 );
94 if accepted.is_ok() {
95 self.returned.push(tool_call_id);
96 }
97 self.finish(accepted)
98 }
99
100 pub fn begin(&mut self) -> Result<Option<Start>, String> {
101 self.state.begin()
102 }
103
104 pub fn prepare_stage(
105 &mut self,
106 job: u64,
107 text: String,
108 values: Vec<BoxValue>,
109 ) -> Result<Vec<PreparedCall>, String> {
110 let result = self.state.prepare_stage(job, text, values);
111 self.finish(result)
112 }
113
114 pub fn prepare_steer(&mut self, job: u64) -> Result<Option<PreparedSteer>, String> {
115 let result = self.state.prepare_steer(job);
116 self.finish(result)
117 }
118
119 pub fn validate_steer(&self, prepared: &PreparedSteer) -> Result<(), String> {
120 self.state.validate_steer(prepared)
121 }
122
123 pub fn commit_steer(&mut self, prepared: PreparedSteer) -> Result<(), String> {
124 self.state.commit_steer(prepared)
125 }
126
127 pub fn complete(&mut self, job: u64, output: ShimOutput<BoxValue>) -> Result<bool, String> {
128 if let Err(error) = self.state.complete(job, output) {
129 return self.finish(Err(error));
130 }
131 self.mirror_boxes()?;
132 let resume = matches!(self.state.status(), Status::Running);
133 let after_box_id = self
134 .state
135 .boxes()
136 .last()
137 .ok_or_else(|| "completed turn has no terminal box".to_owned())?
138 .id()
139 .get();
140 self.records.push(Record::Event(EventRecord {
141 after_box_id,
142 event_index: 1,
143 connected_box_id: 0,
144 handler: "llm_done".into(),
145 data: json!({"resume": resume}),
146 }));
147 self.persist_pending()?;
148 Ok(resume)
149 }
150
151 pub fn fail(&mut self, job: u64, message: String, restartable_before_launch: bool) {
152 self.state.fail(job, message, restartable_before_launch);
153 }
154
155 pub fn restart(&mut self) -> Result<(), RestartError> {
156 self.state.restart()
157 }
158
159 fn finish<T>(&mut self, operation: Result<T, String>) -> Result<T, String> {
160 let persistence = self.mirror_and_persist();
161 match (operation, persistence) {
162 (Ok(value), Ok(())) => Ok(value),
163 (Err(error), Ok(())) | (Ok(_), Err(error)) => Err(error),
164 (Err(operation), Err(persistence)) => Err(format!(
165 "{operation}; additionally failed to persist canonical history: {persistence}"
166 )),
167 }
168 }
169
170 fn mirror_and_persist(&mut self) -> Result<(), String> {
171 self.mirror_boxes()?;
172 self.persist_pending()
173 }
174
175 fn mirror_boxes(&mut self) -> Result<(), String> {
176 let boxes = self.state.boxes();
177 let additions = boxes
178 .get(self.mirrored..)
179 .ok_or_else(|| "canonical box frontier moved backwards".to_owned())?;
180 self.records
181 .extend(additions.iter().cloned().map(Record::Box));
182 self.mirrored = boxes.len();
183 Ok(())
184 }
185
186 fn persist_pending(&mut self) -> Result<(), String> {
187 let suffix = self
188 .records
189 .get(self.durable..)
190 .ok_or_else(|| "durable record frontier moved past canonical records".to_owned())?;
191 if suffix.is_empty() {
192 return Ok(());
193 }
194 self.session.persist(suffix.to_vec())?;
195 self.durable = self.records.len();
196 Ok(())
197 }
198}
199
200fn returned_ids(boxes: &[ChatBox]) -> Result<Vec<ToolCallId>, String> {
201 let mut returned = Vec::new();
202 for value in boxes {
203 if let Some(result) = value
204 .tool_result_metadata()
205 .map_err(|error| format!("{error:?}"))?
206 {
207 returned.push(result.tool_call_id);
208 }
209 }
210 Ok(returned)
211}
212
213#[cfg(test)]
214mod tests {
215 use super::*;
216 use kcode_k1_chat_codex_state::Call;
217 use kcode_k1_chat_persistence::K1ChatPersistence;
218 use kcode_k1_peering::K1Peering;
219 use kcode_k1_txn_ordering::K1TxnOrdering;
220 use std::fs;
221 use std::path::PathBuf;
222 use std::sync::Arc;
223 use std::sync::atomic::{AtomicU64, Ordering};
224
225 static NEXT: AtomicU64 = AtomicU64::new(0);
226
227 struct Fixture {
228 root: PathBuf,
229 session: Option<Session>,
230 }
231
232 impl Fixture {
233 fn new(nonce: u8) -> Self {
234 let root = std::env::temp_dir().join(format!(
235 "k1-durable-turn-{}-{}",
236 std::process::id(),
237 NEXT.fetch_add(1, Ordering::Relaxed)
238 ));
239 let _ = fs::remove_dir_all(&root);
240 let ordering = Arc::new(K1TxnOrdering::open(&root.join("ordering")).unwrap());
241 let peering =
242 Arc::new(K1Peering::open(&root.join("peering"), Arc::clone(&ordering)).unwrap());
243 let persistence =
244 K1ChatPersistence::open(&root.join("persistence"), ordering, peering).unwrap();
245 let (session, _) = persistence.session([nonce; 12]).unwrap();
246 Self {
247 root,
248 session: Some(session),
249 }
250 }
251
252 fn session(&self) -> Session {
253 self.session.as_ref().unwrap().clone()
254 }
255 }
256
257 impl Drop for Fixture {
258 fn drop(&mut self) {
259 drop(self.session.take());
260 let _ = fs::remove_dir_all(&self.root);
261 }
262 }
263
264 #[test]
265 fn completion_persists_terminal_box_event_and_recovery_frontiers() {
266 let fixture = Fixture::new(1);
267 let session = fixture.session();
268 let mut turn = DurableTurn::recover(session.clone()).unwrap();
269 turn.accept(
270 "User Message".into(),
271 "hello".into(),
272 String::new(),
273 String::new(),
274 )
275 .unwrap();
276 let start = turn.begin().unwrap().unwrap();
277 let resume = turn
278 .complete(start.job, ShimOutput { items: Vec::new() })
279 .unwrap();
280
281 assert!(!resume);
282 assert_eq!(turn.status(), Status::Quiet);
283 assert_eq!((turn.mirrored, turn.durable, turn.records.len()), (2, 3, 3));
284 let log = session.load().unwrap();
285 let Record::Event(event) = &log.records[2] else {
286 panic!("expected llm_done event");
287 };
288 assert_eq!(event.after_box_id, 2);
289 assert_eq!(event.event_index, 1);
290 assert_eq!(event.connected_box_id, 0);
291 assert_eq!(event.handler, "llm_done");
292 assert_eq!(event.data, json!({"resume": false}));
293
294 drop(turn);
295 let recovered = DurableTurn::recover(session).unwrap();
296 assert_eq!(recovered.boxes().len(), 2);
297 assert_eq!(recovered.status(), Status::Quiet);
298 assert_eq!(
299 (recovered.mirrored, recovered.durable),
300 (recovered.boxes().len(), recovered.records.len())
301 );
302 }
303
304 #[test]
305 fn active_turn_mailbox_flush_continues_generation_without_resume() {
306 let fixture = Fixture::new(4);
307 let session = fixture.session();
308 let mut turn = DurableTurn::recover(session.clone()).unwrap();
309 turn.accept(
310 "User Message".into(),
311 "first".into(),
312 String::new(),
313 String::new(),
314 )
315 .unwrap();
316
317 let start = turn.begin().unwrap().unwrap();
318 assert_eq!(
319 start.values.last(),
320 Some(&BoxValue::History("[Box 2 | Agent Response]\n".into()))
321 );
322
323 turn.prepare_stage(start.job, "working".into(), Vec::new())
324 .unwrap();
325 turn.accept(
326 "User Message".into(),
327 "second".into(),
328 String::new(),
329 String::new(),
330 )
331 .unwrap();
332
333 let prepared = turn.prepare_steer(start.job).unwrap().unwrap();
334 assert_eq!(
335 prepared.values().last(),
336 Some(&BoxValue::History("[Box 4 | Agent Response]\n".into()))
337 );
338 turn.commit_steer(prepared).unwrap();
339
340 let resume = turn
341 .complete(start.job, ShimOutput { items: Vec::new() })
342 .unwrap();
343 assert!(!resume);
344 assert_eq!(turn.status(), Status::Quiet);
345 assert!(turn.begin().unwrap().is_none());
346
347 let log = session.load().unwrap();
348 let Record::Event(event) = log.records.last().unwrap() else {
349 panic!("expected final llm_done event");
350 };
351 assert_eq!(event.handler, "llm_done");
352 assert_eq!(event.data, json!({"resume": false}));
353 }
354
355 #[test]
356 fn active_fifo_is_hidden_then_persisted_and_v1_return_is_idempotent() {
357 let fixture = Fixture::new(2);
358 let session = fixture.session();
359 let mut turn = DurableTurn::recover(session.clone()).unwrap();
360 turn.accept(
361 "User Message".into(),
362 "search".into(),
363 String::new(),
364 String::new(),
365 )
366 .unwrap();
367 let start = turn.begin().unwrap().unwrap();
368 let calls = turn
369 .prepare_stage(
370 start.job,
371 "working".into(),
372 vec![BoxValue::Call(Ok(Call {
373 name: "WebSearch".into(),
374 arguments: "{}".into(),
375 }))],
376 )
377 .unwrap();
378 let tool_call_id = calls[0].tool_call_id;
379 assert_eq!(session.load().unwrap().records.len(), 3);
380
381 turn.accept_tool_message(tool_call_id, "searching".into())
382 .unwrap();
383 turn.accept_tool_return_v2(
384 tool_call_id,
385 Ok("found".into()),
386 "k1.web-search-result/v1".into(),
387 "opaque".into(),
388 )
389 .unwrap();
390 turn.accept_tool_return(tool_call_id, Ok("duplicate".into()))
391 .unwrap();
392 assert_eq!(turn.boxes().len(), 3);
393 assert_eq!(session.load().unwrap().records.len(), 3);
394
395 let prepared = turn.prepare_steer(start.job).unwrap().unwrap();
396 turn.validate_steer(&prepared).unwrap();
397 assert_eq!(turn.boxes().len(), 5);
398 assert_eq!(session.load().unwrap().records.len(), 5);
399 assert_eq!((turn.mirrored, turn.durable, turn.records.len()), (5, 5, 5));
400 assert!(turn.boxes()[3].tool_message_metadata().unwrap().is_some());
401 assert!(turn.boxes()[4].tool_result_v2_metadata().unwrap().is_some());
402 turn.commit_steer(prepared).unwrap();
403 }
404
405 #[test]
406 fn unresolved_search_messages_are_inert_and_result_v2_begins_once() {
407 let fixture = Fixture::new(3);
408 let session = fixture.session();
409 let mut turn = DurableTurn::recover(session.clone()).unwrap();
410 turn.accept(
411 "User Message".into(),
412 "search".into(),
413 String::new(),
414 String::new(),
415 )
416 .unwrap();
417 let start = turn.begin().unwrap().unwrap();
418 let calls = turn
419 .prepare_stage(
420 start.job,
421 "working".into(),
422 vec![BoxValue::Call(Ok(Call {
423 name: "WebSearch".into(),
424 arguments: "{}".into(),
425 }))],
426 )
427 .unwrap();
428 let tool_call_id = calls[0].tool_call_id;
429
430 let prepared = turn.prepare_steer(start.job).unwrap().unwrap();
431 turn.commit_steer(prepared).unwrap();
432
433 let resume = turn
434 .complete(start.job, ShimOutput { items: Vec::new() })
435 .unwrap();
436 assert!(!resume);
437 assert_eq!(turn.status(), Status::Quiet);
438 let log = session.load().unwrap();
439 let Record::Event(event) = log.records.last().unwrap() else {
440 panic!("expected llm_done event");
441 };
442 assert_eq!(event.data, json!({"resume": false}));
443
444 turn.accept_tool_message(tool_call_id, "still searching".into())
445 .unwrap();
446 assert_eq!(turn.status(), Status::Quiet);
447 assert!(turn.begin().unwrap().is_none());
448
449 turn.accept_tool_return_v2(
450 tool_call_id,
451 Ok("found".into()),
452 "k1.web-search-result/v1".into(),
453 "opaque".into(),
454 )
455 .unwrap();
456 assert_eq!(turn.status(), Status::Running);
457 assert!(turn.begin().unwrap().is_some());
458 assert!(turn.begin().unwrap().is_none());
459 }
460}