#![forbid(unsafe_code)]
pub use kcode_k1_chat_codex_codec::{BoxValue, Call};
pub use kcode_k1_chat_state::{
AGENT_ATTACHMENT_TYPE, AGENT_MESSAGE_TYPE, BoxId, ChatBox, SYSTEM_MESSAGE_TYPE,
TOOL_ATTACHMENT_TYPE, TOOL_CALL_TYPE, TOOL_MESSAGE_TYPE, TOOL_RESULT_TYPE, ToolCallId,
USER_ATTACHMENT_TYPE, USER_MESSAGE_TYPE,
};
use std::sync::Arc;
use kcode_k1_chat_codex_codec::project;
use kcode_k1_chat_state::{ActorState, ProviderCall};
use kcode_k1_codex_adapter::{ShimItem, ShimOutput};
#[derive(Clone, Debug, Eq, PartialEq)]
pub struct Start {
pub job: u64,
pub boxes: Vec<BoxValue>,
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub struct PreparedCall {
pub tool_call_id: ToolCallId,
pub name: String,
pub arguments: String,
}
#[derive(Clone, Debug)]
pub struct PreparedSteer(Arc<Prepared>);
#[derive(Debug)]
struct Prepared {
values: Vec<BoxValue>,
job: u64,
generation: u64,
start: usize,
end: usize,
}
impl PreparedSteer {
pub fn values(&self) -> &[BoxValue] {
&self.0.values
}
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub enum Status {
Running,
Quiet,
Stalled { message: String, restartable: bool },
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum RestartError {
NotStalled,
NotRestartable,
}
struct Round {
job: u64,
accepted_call_wave: bool,
}
enum Mode {
Idle,
Running(Round),
Stalled { message: String, restartable: bool },
}
pub struct ConversationState {
state: ActorState,
session: [u8; 12],
sequence: u64,
submitted: usize,
generation: u64,
prepared: Option<Arc<Prepared>>,
mode: Mode,
}
impl ConversationState {
pub fn new(session: [u8; 12]) -> Self {
Self::from_actor(ActorState::new(false), session, 0)
}
pub fn recover(session: [u8; 12], boxes: Vec<ChatBox>, force: bool) -> Result<Self, String> {
let sequence = recovered_sequence(session, &boxes)?;
let state = ActorState::recover(boxes, force).map_err(debug)?;
Ok(Self::from_actor(state, session, sequence))
}
pub fn boxes(&self) -> &[ChatBox] {
self.state.boxes()
}
pub fn status(&self) -> Status {
match &self.mode {
Mode::Running(_) => Status::Running,
Mode::Idle if self.state.quiet() => Status::Quiet,
Mode::Idle => Status::Running,
Mode::Stalled {
message,
restartable,
} => Status::Stalled {
message: message.clone(),
restartable: *restartable,
},
}
}
pub fn accept(
&mut self,
box_type: String,
contents: String,
hidden_type: String,
hidden_contents: String,
) -> Result<(), String> {
self.state
.accept_box(box_type, contents, hidden_type, hidden_contents)
.map_err(debug)
}
pub fn accept_tool_return(
&mut self,
tool_call_id: ToolCallId,
result: Result<String, String>,
) -> Result<(), String> {
self.state
.accept_async_return(tool_call_id, result)
.map_err(debug)
}
pub fn begin(&mut self) -> Result<Option<Start>, String> {
if !matches!(self.mode, Mode::Idle) {
return Ok(None);
}
let Some(start) = self.state.begin_inference().map_err(debug)? else {
return Ok(None);
};
let boxes = self.state.boxes();
let projected = boxes[self.submitted..].iter().map(project).collect();
self.submitted = boxes.len();
self.mode = Mode::Running(Round {
job: start.job,
accepted_call_wave: false,
});
Ok(Some(Start {
job: start.job,
boxes: projected,
}))
}
pub fn prepare_stage(
&mut self,
job: u64,
text: String,
values: Vec<BoxValue>,
) -> Result<Vec<PreparedCall>, String> {
let calls = values
.into_iter()
.map(|value| match value {
BoxValue::Call(Ok(call)) => Ok(call),
_ => Err("stage contains a malformed tool call".to_owned()),
})
.collect::<Result<Vec<_>, _>>()?;
match &self.mode {
Mode::Running(round) if round.job == job => {}
_ => return Err("stale Codex inference stage".to_owned()),
}
if self.prepared.is_some() {
return Err("a Codex steer remains uncommitted".to_owned());
}
let mut sequence = self.sequence;
let prepared = calls
.iter()
.map(|call| {
sequence = sequence
.checked_add(1)
.ok_or_else(|| "ToolCallId space was exhausted".to_owned())?;
Ok(PreparedCall {
tool_call_id: ToolCallId::new(self.session, sequence),
name: call.name.clone(),
arguments: call.arguments.clone(),
})
})
.collect::<Result<Vec<_>, String>>()?;
let provider_calls = prepared
.iter()
.map(|call| ProviderCall {
tool_call_id: call.tool_call_id,
name: call.name.clone(),
arguments: call.arguments.clone(),
})
.collect();
self.state
.append_stage(job, text, provider_calls)
.map_err(debug)?;
self.sequence = sequence;
self.submitted = self.state.boxes().len();
if !prepared.is_empty()
&& let Mode::Running(round) = &mut self.mode
{
round.accepted_call_wave = true;
}
Ok(prepared)
}
pub fn prepare_steer(&mut self, job: u64) -> Result<Option<PreparedSteer>, String> {
if !matches!(&self.mode, Mode::Running(round) if round.job == job) {
return Err("stale Codex inference steer".to_owned());
}
if let Some(prepared) = &self.prepared {
return Ok(Some(PreparedSteer(Arc::clone(prepared))));
}
if self.submitted != self.state.boxes().len() {
return Err("Codex submitted frontier is inconsistent".to_owned());
}
let generation = self
.generation
.checked_add(1)
.ok_or_else(|| "Codex steer generation was exhausted".to_owned())?;
let boxes = self.state.flush_active_arrivals(job).map_err(debug)?;
if boxes.is_empty() {
return Ok(None);
}
let prepared = Arc::new(Prepared {
values: boxes.iter().map(project).collect(),
job,
generation,
start: self.submitted,
end: self.state.boxes().len(),
});
debug_assert_eq!(prepared.end - prepared.start, boxes.len());
self.generation = generation;
self.prepared = Some(Arc::clone(&prepared));
Ok(Some(PreparedSteer(prepared)))
}
pub fn validate_steer(&self, prepared: &PreparedSteer) -> Result<(), String> {
let prepared = &prepared.0;
let valid = matches!(&self.mode, Mode::Running(round) if round.job == prepared.job)
&& self.generation == prepared.generation
&& self.submitted == prepared.start
&& self.state.boxes().len() == prepared.end
&& prepared.end.checked_sub(prepared.start) == Some(prepared.values.len())
&& self
.prepared
.as_ref()
.is_some_and(|value| Arc::ptr_eq(value, prepared));
if valid {
Ok(())
} else {
Err("stale or invalid Codex steer".to_owned())
}
}
pub fn commit_steer(&mut self, prepared: PreparedSteer) -> Result<(), String> {
self.validate_steer(&prepared)?;
self.submitted = prepared.0.end;
self.prepared = None;
Ok(())
}
pub fn complete(&mut self, job: u64, output: ShimOutput<BoxValue>) -> Result<(), String> {
let round = self.take_round(job)?;
if self.prepared.is_some() {
let message = "Codex inference completed with an uncommitted steer".to_owned();
self.preserve(round, message.clone(), false);
return Err(message);
}
let mut text = String::new();
for item in output.items {
match item {
ShimItem::Text(value) => text.push_str(&value),
ShimItem::Box(_) => {
let message = "terminal Codex output contains a box".to_owned();
self.preserve(round, message.clone(), false);
return Err(message);
}
}
}
if let Err(error) = self.state.complete_inference(job, text).map_err(debug) {
self.preserve(round, error.clone(), false);
return Err(error);
}
self.mode = Mode::Idle;
Ok(())
}
pub fn fail(&mut self, job: u64, message: String, restartable_before_launch: bool) {
if let Ok(round) = self.take_round(job) {
self.preserve(round, message, restartable_before_launch);
}
}
pub fn restart(&mut self) -> Result<(), RestartError> {
match self.mode {
Mode::Stalled {
restartable: true, ..
} => {}
Mode::Stalled { .. } => return Err(RestartError::NotRestartable),
_ => return Err(RestartError::NotStalled),
}
self.state
.restart()
.map_err(|_| RestartError::NotRestartable)?;
self.submitted = 0;
self.prepared = None;
self.mode = Mode::Idle;
Ok(())
}
fn from_actor(state: ActorState, session: [u8; 12], sequence: u64) -> Self {
Self {
state,
session,
sequence,
submitted: 0,
generation: 0,
prepared: None,
mode: Mode::Idle,
}
}
fn take_round(&mut self, job: u64) -> Result<Round, String> {
let mode = std::mem::replace(&mut self.mode, Mode::Idle);
match mode {
Mode::Running(round) if round.job == job => Ok(round),
other => {
self.mode = other;
Err("stale Codex inference completion".to_owned())
}
}
}
fn preserve(&mut self, round: Round, message: String, restartable_before_launch: bool) {
self.prepared = None;
if round.accepted_call_wave {
let _ = self.state.complete_inference(round.job, String::new());
let _ = self.state.halt(message.clone());
self.mode = Mode::Stalled {
message,
restartable: false,
};
} else {
let stalled = self
.state
.stall_inference(round.job, message.clone())
.is_ok();
self.mode = Mode::Stalled {
message,
restartable: stalled && restartable_before_launch,
};
}
}
}
fn recovered_sequence(session: [u8; 12], boxes: &[ChatBox]) -> Result<u64, String> {
let mut maximum = 0;
for box_ in boxes {
if let Some(call) = box_.tool_call_metadata().map_err(debug)? {
record_sequence(session, call.tool_call_id, &mut maximum)?;
}
if let Some(result) = box_.tool_result_metadata().map_err(debug)? {
record_sequence(session, result.tool_call_id, &mut maximum)?;
}
}
Ok(maximum)
}
fn record_sequence(session: [u8; 12], id: ToolCallId, maximum: &mut u64) -> Result<(), String> {
if id.nonce() != session {
return Err("recovered ToolCallId belongs to another session".to_owned());
}
*maximum = (*maximum).max(id.sequence());
Ok(())
}
fn debug(error: impl std::fmt::Debug) -> String {
format!("{error:?}")
}