1#![forbid(unsafe_code)]
2
3mod recovery;
4
5pub use kcode_k1_chat_codex_codec::{BoxValue, Call};
6pub use kcode_k1_chat_state::{
7 AGENT_ATTACHMENT_TYPE, AGENT_MESSAGE_TYPE, AGENT_RESPONSE_TYPE, ActorState, BoxId, ChatBox,
8 ProviderGenerated, SYSTEM_MESSAGE_TYPE, TOOL_ATTACHMENT_TYPE, TOOL_CALL_TYPE,
9 TOOL_MESSAGE_TYPE, TOOL_RESULT_TYPE, ToolCallId, USER_ATTACHMENT_TYPE, USER_MESSAGE_TYPE,
10};
11pub use kcode_k1_codex_adapter::{ShimItem, ShimOutput};
12
13use std::sync::Arc;
14use std::sync::atomic::AtomicU8;
15
16use kcode_k1_chat_codex_codec::{open_agent_response, project};
17use kcode_k1_chat_state::{ProviderCall, StateError};
18use recovery::recovered_sequence;
19
20pub const MALFORMED_NATIVE_CALL_METADATA_TYPE: &str = "k1.malformed-native-call.v1";
21
22#[derive(Clone, Debug)]
23pub struct Start {
24 pub job: u64,
25 pub values: Vec<BoxValue>,
26 pub attempt: Arc<AtomicU8>,
27}
28
29#[derive(Clone, Debug, Eq, PartialEq)]
30pub struct PreparedCall {
31 pub tool_call_id: ToolCallId,
32 pub name: String,
33 pub arguments: String,
34 disposition: PreparedCallDisposition,
35}
36
37impl PreparedCall {
38 pub fn disposition(&self) -> &PreparedCallDisposition {
39 &self.disposition
40 }
41}
42
43#[derive(Clone, Debug, Eq, PartialEq)]
44pub enum PreparedCallDisposition {
45 External,
46 ImmediateError(Box<ImmediateToolError>),
47}
48
49#[derive(Clone, Debug, Eq, PartialEq)]
50pub struct ImmediateToolError {
51 pub message: String,
52 pub native_tool: String,
53 pub attempted_ktool: Option<String>,
54 pub validation_code: String,
55 pub path: String,
56 pub expected: String,
57 pub received: String,
58 pub native_arguments: String,
59 pub metadata_type: String,
60 pub metadata_contents: String,
61}
62
63#[derive(Clone, Debug)]
64pub struct PreparedMailboxFlush(Arc<Prepared>);
65
66#[derive(Debug)]
67struct Prepared {
68 token: u64,
69 values: Vec<BoxValue>,
70 job: u64,
71 external_count: usize,
72}
73
74impl PreparedMailboxFlush {
75 pub fn values(&self) -> &[BoxValue] {
76 &self.0.values
77 }
78}
79
80#[derive(Clone, Debug, Eq, PartialEq)]
81pub enum Status {
82 Running,
83 Quiet,
84 Stalled { message: String, restartable: bool },
85}
86
87#[derive(Clone, Copy, Debug, Eq, PartialEq)]
88pub enum RestartError {
89 NotStalled,
90 ProviderActionAccepted,
91}
92
93#[derive(Clone, Copy, Debug, Eq, PartialEq)]
94enum Phase {
95 ProviderActive,
96 ChatendBoundary,
97 PendingGeneration,
98}
99
100struct Round {
101 job: u64,
102 accepted_provider_action: bool,
103 phase: Phase,
104 mailbox_flush_needed: bool,
105 restartable: bool,
106}
107
108enum Mode {
109 Idle,
110 Running(Round),
111 Stalled(Status, bool),
112}
113
114pub struct ConversationState {
115 state: ActorState,
116 session: [u8; 12],
117 sequence: u64,
118 unsubmitted: Vec<BoxValue>,
119 queued_trigger: bool,
120 token: u64,
121 prepared: Option<Arc<Prepared>>,
122 mode: Mode,
123}
124
125impl ConversationState {
126 pub fn new(session: [u8; 12]) -> Self {
127 Self {
128 state: ActorState::new(false),
129 session,
130 sequence: 0,
131 unsubmitted: Vec::new(),
132 queued_trigger: false,
133 token: 0,
134 prepared: None,
135 mode: Mode::Idle,
136 }
137 }
138
139 pub fn recover(session: [u8; 12], boxes: Vec<ChatBox>, force: bool) -> Result<Self, String> {
140 let sequence = recovered_sequence(session, &boxes)?;
141 let state = ActorState::recover(boxes, force).map_err(debug)?;
142 Ok(Self {
143 unsubmitted: state.boxes().iter().map(project).collect(),
144 state,
145 sequence,
146 ..Self::new(session)
147 })
148 }
149
150 pub fn boxes(&self) -> &[ChatBox] {
151 self.state.boxes()
152 }
153
154 pub fn status(&self) -> Status {
155 match &self.mode {
156 Mode::Running(_) => Status::Running,
157 Mode::Idle if self.state.quiet() => Status::Quiet,
158 Mode::Idle => Status::Running,
159 Mode::Stalled(status, _) => status.clone(),
160 }
161 }
162
163 pub fn accept(
164 &mut self,
165 box_type: String,
166 contents: String,
167 hidden_type: String,
168 hidden_contents: String,
169 ) -> Result<(), String> {
170 self.accept_arrival(true, |state| {
171 state.accept_box(box_type, contents, hidden_type, hidden_contents)
172 })
173 }
174
175 pub fn accept_tool_message(
176 &mut self,
177 tool_call_id: ToolCallId,
178 message: String,
179 ) -> Result<(), String> {
180 self.accept_arrival(false, |state| {
181 state.accept_tool_message(tool_call_id, message)
182 })
183 }
184
185 pub fn accept_tool_return(
186 &mut self,
187 tool_call_id: ToolCallId,
188 result: Result<String, String>,
189 ) -> Result<(), String> {
190 self.accept_arrival(true, |state| {
191 state.accept_async_return(tool_call_id, result)
192 })
193 }
194
195 pub fn accept_tool_return_v2(
196 &mut self,
197 tool_call_id: ToolCallId,
198 result: Result<String, String>,
199 metadata_type: String,
200 metadata_contents: String,
201 ) -> Result<(), String> {
202 self.accept_arrival(true, |state| {
203 state.accept_async_return_v2(tool_call_id, result, metadata_type, metadata_contents)
204 })
205 }
206
207 pub fn begin(&mut self) -> Result<Option<Start>, String> {
208 if !matches!(self.mode, Mode::Idle) {
209 return Ok(None);
210 }
211 let Some(start) = self.state.begin_inference().map_err(debug)? else {
212 return Ok(None);
213 };
214 let promised_id = self.promised_id()?;
215 let mut values = std::mem::take(&mut self.unsubmitted);
216 values.push(open_agent_response(promised_id));
217 let mailbox_flush_needed = std::mem::take(&mut self.queued_trigger);
218 self.mode = Mode::Running(Round {
219 job: start.job,
220 accepted_provider_action: false,
221 phase: Phase::ProviderActive,
222 mailbox_flush_needed,
223 restartable: true,
224 });
225 Ok(Some(Start {
226 job: start.job,
227 values,
228 attempt: start.attempt,
229 }))
230 }
231
232 pub fn prepare_stage(
233 &mut self,
234 job: u64,
235 text: String,
236 values: Vec<BoxValue>,
237 ) -> Result<Vec<PreparedCall>, String> {
238 if !matches!(
239 &self.mode,
240 Mode::Running(round) if round.job == job && round.phase == Phase::ProviderActive
241 ) {
242 return Err("stale Codex inference stage".to_owned());
243 }
244 if self.prepared.is_some() {
245 return Err("a Codex mailbox flush remains uncommitted".to_owned());
246 }
247 let accepted_provider_action = !values.is_empty();
248 let mut sequence = self.sequence;
249 let mut generated = Vec::with_capacity(values.len());
250 let mut prepared = Vec::new();
251 for value in values {
252 match value {
253 BoxValue::AgentMessage(Ok(contents)) => {
254 generated.push(ProviderGenerated::AgentMessage { contents });
255 }
256 BoxValue::Call(Ok(call)) => {
257 sequence = next_sequence(sequence)?;
258 let tool_call_id = ToolCallId::new(self.session, sequence);
259 generated.push(ProviderGenerated::ToolCall(ProviderCall {
260 tool_call_id,
261 name: call.name.clone(),
262 arguments: call.arguments.clone(),
263 }));
264 prepared.push(PreparedCall {
265 tool_call_id,
266 name: call.name,
267 arguments: call.arguments,
268 disposition: PreparedCallDisposition::External,
269 });
270 }
271 BoxValue::MalformedNativeAction(action) => {
272 sequence = next_sequence(sequence)?;
273 let tool_call_id = ToolCallId::new(self.session, sequence);
274 let name = action
275 .attempted_ktool()
276 .unwrap_or_else(|| action.native_tool())
277 .to_owned();
278 let arguments = action.native_arguments_json();
279 let error = ImmediateToolError {
280 message: action.diagnostic(),
281 native_tool: action.native_tool().to_owned(),
282 attempted_ktool: action.attempted_ktool().map(str::to_owned),
283 validation_code: action.validation_code().to_owned(),
284 path: action.path().to_owned(),
285 expected: action.expected().to_owned(),
286 received: action.received().to_owned(),
287 native_arguments: arguments.clone(),
288 metadata_type: MALFORMED_NATIVE_CALL_METADATA_TYPE.to_owned(),
289 metadata_contents: action.diagnostic_json(),
290 };
291 generated.push(ProviderGenerated::ToolCall(ProviderCall {
292 tool_call_id,
293 name: name.clone(),
294 arguments: arguments.clone(),
295 }));
296 prepared.push(PreparedCall {
297 tool_call_id,
298 name,
299 arguments,
300 disposition: PreparedCallDisposition::ImmediateError(Box::new(error)),
301 });
302 }
303 _ => return Err("stage contains a malformed provider action".to_owned()),
304 }
305 }
306 self.state
307 .append_stage(job, text, generated)
308 .map_err(debug)?;
309 self.sequence = sequence;
310 if let Mode::Running(round) = &mut self.mode {
311 round.accepted_provider_action |= accepted_provider_action;
312 round.phase = Phase::ChatendBoundary;
313 round.mailbox_flush_needed = true;
314 round.restartable = false;
315 }
316 Ok(prepared)
317 }
318
319 pub fn mailbox_flush(&mut self, job: u64) -> Result<Vec<ChatBox>, String> {
320 if !matches!(
321 &self.mode,
322 Mode::Running(round) if round.job == job && round.phase == Phase::ChatendBoundary
323 ) {
324 return Err("stale Codex active-arrival mailbox flush".to_owned());
325 }
326 let boxes = self.state.flush_active_arrivals(job).map_err(debug)?;
327 self.unsubmitted.extend(boxes.iter().map(project));
328 if let Mode::Running(round) = &mut self.mode {
329 round.phase = Phase::PendingGeneration;
330 }
331 Ok(boxes)
332 }
333
334 pub fn prepare_mailbox_flush(
335 &mut self,
336 job: u64,
337 ) -> Result<Option<PreparedMailboxFlush>, String> {
338 let (mailbox_flush_needed, phase) = match &self.mode {
339 Mode::Running(round) if round.job == job => (round.mailbox_flush_needed, round.phase),
340 _ => return Err("stale Codex inference mailbox flush".to_owned()),
341 };
342 if let Some(prepared) = &self.prepared {
343 return Ok(Some(PreparedMailboxFlush(Arc::clone(prepared))));
344 }
345 if !mailbox_flush_needed {
346 return Ok(None);
347 }
348 if phase == Phase::ProviderActive {
349 return Ok(None);
350 }
351 if phase == Phase::ChatendBoundary {
352 self.mailbox_flush(job)?;
353 }
354 let token = self
355 .token
356 .checked_add(1)
357 .ok_or_else(|| "Codex mailbox-flush token space was exhausted".to_owned())?;
358 let external_count = self.unsubmitted.len();
359 let mut values = self.unsubmitted.clone();
360 values.push(open_agent_response(self.promised_id()?));
361 let prepared = Arc::new(Prepared {
362 token,
363 values,
364 job,
365 external_count,
366 });
367 self.token = token;
368 self.prepared = Some(Arc::clone(&prepared));
369 if let Mode::Running(round) = &mut self.mode {
370 round.mailbox_flush_needed = false;
371 }
372 Ok(Some(PreparedMailboxFlush(prepared)))
373 }
374
375 pub fn validate_mailbox_flush(&self, prepared: &PreparedMailboxFlush) -> Result<(), String> {
376 let prepared = &prepared.0;
377 let prefix_matches = self.unsubmitted.get(..prepared.external_count)
378 == Some(&prepared.values[..prepared.external_count]);
379 let valid = matches!(
380 &self.mode,
381 Mode::Running(round)
382 if round.job == prepared.job && round.phase == Phase::PendingGeneration
383 ) && self.token == prepared.token
384 && prefix_matches
385 && self
386 .prepared
387 .as_ref()
388 .is_some_and(|current| Arc::ptr_eq(current, prepared));
389 if valid {
390 Ok(())
391 } else {
392 Err("stale or invalid Codex mailbox flush".to_owned())
393 }
394 }
395
396 pub fn commit_mailbox_flush(&mut self, prepared: PreparedMailboxFlush) -> Result<(), String> {
397 self.validate_mailbox_flush(&prepared)?;
398 self.unsubmitted.drain(..prepared.0.external_count);
399 self.prepared = None;
400 if let Mode::Running(round) = &mut self.mode {
401 round.phase = Phase::ProviderActive;
402 }
403 Ok(())
404 }
405
406 pub fn complete(&mut self, job: u64, output: ShimOutput<BoxValue>) -> Result<(), String> {
407 let mut round = self.take_round(job)?;
408 if round.phase != Phase::ProviderActive {
409 let message = "Codex inference completed outside a provider generation".to_owned();
410 self.preserve(round, message.clone(), false);
411 return Err(message);
412 }
413 if self.prepared.is_some() {
414 let message = "Codex inference completed with an uncommitted mailbox flush".to_owned();
415 self.preserve(round, message.clone(), false);
416 return Err(message);
417 }
418 let mut text = String::new();
419 for item in output.items {
420 match item {
421 ShimItem::Text(value) => text.push_str(&value),
422 ShimItem::Box(_) => {
423 round.accepted_provider_action = true;
424 let message = "terminal Codex output contains a box".to_owned();
425 self.preserve(round, message.clone(), false);
426 return Err(message);
427 }
428 }
429 }
430 let before = self.state.boxes().len();
431 if let Err(error) = self.state.complete_inference(job, text).map_err(debug) {
432 self.preserve(round, error.clone(), false);
433 return Err(error);
434 }
435 self.unsubmitted
436 .extend(self.state.boxes()[before..].iter().skip(1).map(project));
437 self.queued_trigger = false;
438 self.mode = Mode::Idle;
439 Ok(())
440 }
441
442 pub fn fail(&mut self, job: u64, message: String, restartable_before_launch: bool) {
443 if let Ok(round) = self.take_round(job) {
444 self.preserve(round, message, restartable_before_launch);
445 }
446 }
447
448 pub fn restart(&mut self) -> Result<(), RestartError> {
449 match &self.mode {
450 Mode::Stalled(_, true) => return Err(RestartError::ProviderActionAccepted),
451 Mode::Stalled(
452 Status::Stalled {
453 restartable: true, ..
454 },
455 false,
456 ) => {}
457 _ => return Err(RestartError::NotStalled),
458 }
459 self.state.restart().map_err(|_| RestartError::NotStalled)?;
460 self.unsubmitted = self.state.boxes().iter().map(project).collect();
461 self.queued_trigger = false;
462 self.prepared = None;
463 self.mode = Mode::Idle;
464 Ok(())
465 }
466
467 fn accept_arrival<F>(&mut self, triggering: bool, accept: F) -> Result<(), String>
468 where
469 F: FnOnce(&mut ActorState) -> Result<(), StateError>,
470 {
471 let before = self.state.boxes().len();
472 accept(&mut self.state).map_err(debug)?;
473 let appended = &self.state.boxes()[before..];
474 self.unsubmitted.extend(appended.iter().map(project));
475 match &mut self.mode {
476 Mode::Running(round) => round.mailbox_flush_needed |= triggering,
477 Mode::Idle if appended.is_empty() => self.queued_trigger |= triggering,
478 Mode::Idle => self.queued_trigger = false,
479 Mode::Stalled(_, _) => {}
480 }
481 Ok(())
482 }
483
484 fn promised_id(&self) -> Result<BoxId, String> {
485 let previous = self.state.boxes().last().map_or(0, |box_| box_.id().get());
486 let value = previous
487 .checked_add(1)
488 .ok_or_else(|| "BoxId space was exhausted".to_owned())?;
489 Ok(BoxId::new(value))
490 }
491
492 fn take_round(&mut self, job: u64) -> Result<Round, String> {
493 match std::mem::replace(&mut self.mode, Mode::Idle) {
494 Mode::Running(round) if round.job == job => Ok(round),
495 other => {
496 self.mode = other;
497 Err("stale Codex inference completion".to_owned())
498 }
499 }
500 }
501
502 fn preserve(&mut self, round: Round, message: String, restartable_before_launch: bool) {
503 self.prepared = None;
504 let restartable = restartable_before_launch
505 && round.restartable
506 && !round.accepted_provider_action
507 && self
508 .state
509 .stall_inference(round.job, message.clone())
510 .is_ok();
511 if !restartable {
512 let _ = self.state.halt(message.clone());
513 }
514 self.mode = Mode::Stalled(
515 Status::Stalled {
516 message,
517 restartable,
518 },
519 round.accepted_provider_action,
520 );
521 }
522}
523
524fn next_sequence(sequence: u64) -> Result<u64, String> {
525 sequence
526 .checked_add(1)
527 .ok_or_else(|| "ToolCallId space was exhausted".to_owned())
528}
529
530fn debug(error: impl std::fmt::Debug) -> String {
531 format!("{error:?}")
532}
533
534#[cfg(test)]
535mod tests {
536 use super::*;
537 use kcode_k1_chat_codex_codec::Codec;
538 use kcode_k1_codex_adapter::{BoxCodec, ToolCall};
539
540 fn active() -> (ConversationState, u64) {
541 let mut state = ConversationState::new([7; 12]);
542 state
543 .accept(
544 USER_MESSAGE_TYPE.into(),
545 "start".into(),
546 String::new(),
547 String::new(),
548 )
549 .unwrap();
550 let job = state.begin().unwrap().unwrap().job;
551 (state, job)
552 }
553
554 fn native(name: &str, arguments: &str) -> BoxValue {
555 let mut codec = Codec;
556 codec.tool_call_box(&ToolCall {
557 call_id: "native".into(),
558 name: name.into(),
559 arguments: arguments.parse().unwrap(),
560 })
561 }
562
563 fn active_after_mailbox_flush() -> (ConversationState, u64, ToolCallId) {
564 let (mut state, job) = active();
565 let calls = state
566 .prepare_stage(
567 job,
568 String::new(),
569 vec![BoxValue::Call(Ok(Call {
570 name: "tool".into(),
571 arguments: "{}".into(),
572 }))],
573 )
574 .unwrap();
575 let prepared = state.prepare_mailbox_flush(job).unwrap().unwrap();
576 state.commit_mailbox_flush(prepared).unwrap();
577 (state, job, calls[0].tool_call_id)
578 }
579
580 fn empty_output() -> ShimOutput<BoxValue> {
581 ShimOutput { items: Vec::new() }
582 }
583
584 #[test]
585 fn malformed_only_is_persisted_and_prepared_as_deterministic_immediate_error() {
586 let (mut state, job) = active();
587 let malformed = native(
588 "call_ktool",
589 r#"{"z":0,"name":"recoverable","arguments":{"b":2,"a":1}}"#,
590 );
591 let (diagnostic, diagnostic_json) = match &malformed {
592 BoxValue::MalformedNativeAction(action) => {
593 (action.diagnostic(), action.diagnostic_json())
594 }
595 _ => panic!("test value must be malformed"),
596 };
597 let calls = state
598 .prepare_stage(job, String::new(), vec![malformed])
599 .unwrap();
600 assert_eq!(state.status(), Status::Running);
601 assert_eq!(calls.len(), 1);
602 assert_eq!(calls[0].name, "recoverable");
603 assert_eq!(
604 calls[0].arguments,
605 r#"{"arguments":{"a":1,"b":2},"name":"recoverable","z":0}"#
606 );
607 let PreparedCallDisposition::ImmediateError(error) = calls[0].disposition() else {
608 panic!("malformed call must be local");
609 };
610 assert_eq!(error.message, diagnostic);
611 assert_eq!(error.native_tool, "call_ktool");
612 assert_eq!(error.attempted_ktool.as_deref(), Some("recoverable"));
613 assert_eq!(error.validation_code, "invalid_wrapper_fields");
614 assert_eq!(error.path, "$");
615 assert_eq!(error.expected, "object with exactly name and arguments");
616 assert_eq!(error.received, calls[0].arguments);
617 assert_eq!(error.native_arguments, calls[0].arguments);
618 assert_eq!(error.metadata_type, MALFORMED_NATIVE_CALL_METADATA_TYPE);
619 assert_eq!(error.metadata_contents, diagnostic_json);
620 let persisted = state
621 .boxes()
622 .last()
623 .unwrap()
624 .tool_call_metadata()
625 .unwrap()
626 .unwrap();
627 assert_eq!(persisted.name, calls[0].name);
628 assert_eq!(persisted.arguments, calls[0].arguments);
629 }
630
631 #[test]
632 fn mixed_stage_preserves_order_contiguous_ids_name_selection_and_external_calls() {
633 let (mut state, job) = active();
634 let before = state.boxes().len();
635 let values = vec![
636 BoxValue::AgentMessage(Ok("first".into())),
637 native(
638 "call_ktool",
639 r#"{"name":"attempted","arguments":{},"extra":1}"#,
640 ),
641 BoxValue::Call(Ok(Call {
642 name: "valid".into(),
643 arguments: "{\"ok\":true}".into(),
644 })),
645 native("future_native", r#"{"z":0,"a":true}"#),
646 ];
647 let calls = state.prepare_stage(job, String::new(), values).unwrap();
648 assert_eq!(
649 calls
650 .iter()
651 .map(|call| call.tool_call_id.sequence())
652 .collect::<Vec<_>>(),
653 vec![1, 2, 3]
654 );
655 assert_eq!(
656 calls
657 .iter()
658 .map(|call| call.name.as_str())
659 .collect::<Vec<_>>(),
660 vec!["attempted", "valid", "future_native"]
661 );
662 assert!(matches!(
663 calls[0].disposition(),
664 PreparedCallDisposition::ImmediateError(_)
665 ));
666 assert_eq!(calls[1].disposition(), &PreparedCallDisposition::External);
667 assert!(matches!(
668 calls[2].disposition(),
669 PreparedCallDisposition::ImmediateError(_)
670 ));
671 assert_eq!(calls[2].arguments, r#"{"a":true,"z":0}"#);
672 assert_eq!(
673 state.boxes()[before..]
674 .iter()
675 .map(ChatBox::box_type)
676 .collect::<Vec<_>>(),
677 vec![
678 AGENT_RESPONSE_TYPE,
679 AGENT_MESSAGE_TYPE,
680 TOOL_CALL_TYPE,
681 TOOL_CALL_TYPE,
682 TOOL_CALL_TYPE,
683 ]
684 );
685 }
686
687 #[test]
688 fn committed_mailbox_flush_does_not_schedule_a_fresh_turn() {
689 let (mut state, job, _) = active_after_mailbox_flush();
690 state.complete(job, empty_output()).unwrap();
691 assert_eq!(state.status(), Status::Quiet);
692 assert!(state.begin().unwrap().is_none());
693 }
694
695 #[test]
696 fn tool_messages_after_a_mailbox_flush_remain_inert() {
697 let (mut state, job, tool_call_id) = active_after_mailbox_flush();
698 state
699 .accept_tool_message(tool_call_id, "still running".into())
700 .unwrap();
701 state.complete(job, empty_output()).unwrap();
702 assert_eq!(state.status(), Status::Quiet);
703 assert!(state.begin().unwrap().is_none());
704 }
705
706 #[test]
707 fn one_tool_result_schedules_exactly_one_fresh_turn() {
708 let (mut state, job, tool_call_id) = active_after_mailbox_flush();
709 state
710 .accept_tool_return(tool_call_id, Ok("done".into()))
711 .unwrap();
712 state.complete(job, empty_output()).unwrap();
713 let followup = state.begin().unwrap().unwrap();
714 assert!(state.begin().unwrap().is_none());
715 state.complete(followup.job, empty_output()).unwrap();
716 assert_eq!(state.status(), Status::Quiet);
717 assert!(state.begin().unwrap().is_none());
718 }
719}