#![forbid(unsafe_code)]
use kcode_k1_chat_codex_codec::{BoxValue, project};
use kcode_k1_chat_state::{ActorState, BoxId, ChatBox, ProviderCall};
use kcode_k1_codex_adapter::{ShimItem, ShimOutput};
pub use kcode_k1_chat_codex_codec::Codec;
pub use kcode_k1_chat_state::{
DispatchedToolCall, ResultView, ToolCall, ToolCallId, ToolResult, ToolResultStatus,
};
#[derive(Clone, Debug, PartialEq)]
pub struct Start {
pub job: u64,
pub boxes: Vec<BoxValue>,
}
#[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,
StateRejected,
}
#[derive(Clone, Debug, PartialEq)]
pub struct PreparedSteer {
job: u64,
generation: u64,
values: Vec<BoxValue>,
}
impl PreparedSteer {
pub fn values(&self) -> &[BoxValue] {
&self.values
}
}
struct ActiveTurn {
job: u64,
initial_frontier: Option<BoxId>,
accepted_call_wave: bool,
}
struct PendingSteer {
job: u64,
generation: u64,
frontier: BoxId,
values: Vec<BoxValue>,
}
struct Stall {
message: String,
restartable: bool,
}
pub struct ConversationState {
state: ActorState,
submitted: Option<BoxId>,
active: Option<ActiveTurn>,
pending_steer: Option<PendingSteer>,
next_steer_generation: u64,
stall: Option<Stall>,
}
impl ConversationState {
pub fn new() -> Self {
Self::from_state(ActorState::new(false))
}
pub fn recover(boxes: Vec<ChatBox>, force: bool) -> Result<Self, String> {
let state = ActorState::recover(boxes, force).map_err(state_error)?;
Ok(Self::from_state(state))
}
pub fn boxes(&self) -> &[ChatBox] {
self.state.boxes()
}
pub fn status(&self) -> Status {
if let Some(stall) = &self.stall {
return Status::Stalled {
message: stall.message.clone(),
restartable: stall.restartable,
};
}
if self.active.is_some() {
Status::Running
} else {
Status::Quiet
}
}
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(state_error)
}
pub fn accept_tool_return(&mut self, result: ToolResult) -> Result<(), String> {
self.state.accept_async_return(result).map_err(state_error)
}
pub fn begin(&mut self) -> Result<Option<Start>, String> {
let Some(start) = self.state.begin_inference().map_err(state_error)? else {
return Ok(None);
};
self.state
.flush_active_arrivals(start.job)
.map_err(state_error)?;
let initial_frontier = self.state.boxes().last().map(ChatBox::id);
let boxes = self
.state
.boxes()
.iter()
.filter(|value| self.submitted.is_none_or(|id| value.id() > id))
.map(project)
.collect();
self.active = Some(ActiveTurn {
job: start.job,
initial_frontier,
accepted_call_wave: false,
});
self.pending_steer = None;
Ok(Some(Start {
job: start.job,
boxes,
}))
}
pub fn prepare_stage(
&mut self,
job: u64,
text: String,
values: Vec<BoxValue>,
) -> Result<Vec<DispatchedToolCall>, String> {
self.require_job(job, "stage")?;
if self.pending_steer.is_some() {
return Err("cannot append a stage while a steer is pending".into());
}
let calls = values
.into_iter()
.map(|value| match value {
BoxValue::Call(Ok(call)) => Ok(call),
BoxValue::Call(Err(error)) => Err(error),
BoxValue::History(_) => Err("stage contains a non-call value".into()),
})
.collect::<Result<Vec<ProviderCall>, String>>()?;
let dispatched = self
.state
.append_stage(job, text, calls)
.map_err(state_error)?;
self.submitted = self.state.boxes().last().map(ChatBox::id);
if !dispatched.is_empty() {
self.active
.as_mut()
.expect("validated active turn")
.accepted_call_wave = true;
}
Ok(dispatched)
}
pub fn prepare_steer(&mut self, job: u64) -> Result<Option<PreparedSteer>, String> {
self.require_job(job, "steer")?;
if let Some(pending) = &self.pending_steer {
return Ok(Some(pending.token()));
}
let boxes = self.state.flush_active_arrivals(job).map_err(state_error)?;
let Some(frontier) = boxes.last().map(ChatBox::id) else {
return Ok(None);
};
let generation = self
.next_steer_generation
.checked_add(1)
.ok_or_else(|| "steer generation exhausted".to_string())?;
self.next_steer_generation = generation;
let values = boxes.iter().map(project).collect();
let pending = PendingSteer {
job,
generation,
frontier,
values,
};
let token = pending.token();
self.pending_steer = Some(pending);
Ok(Some(token))
}
pub fn validate_steer(&self, token: &PreparedSteer) -> Result<(), String> {
self.require_job(token.job, "steer")?;
match &self.pending_steer {
Some(pending) if pending.job == token.job && pending.generation == token.generation => {
Ok(())
}
_ => Err("stale Codex inference steer".into()),
}
}
pub fn commit_steer(&mut self, token: PreparedSteer) -> Result<(), String> {
self.validate_steer(&token)?;
let pending = self.pending_steer.take().expect("validated pending steer");
self.submitted = Some(pending.frontier);
Ok(())
}
pub fn complete(&mut self, job: u64, output: ShimOutput<BoxValue>) -> Result<(), String> {
self.require_job(job, "completion")?;
if self.pending_steer.is_some() {
return Err("cannot complete while a steer is pending".into());
}
let text = terminal_text(output)?;
let before = self.state.boxes().last().map(ChatBox::id);
let initial = self
.active
.as_ref()
.expect("validated active turn")
.initial_frontier;
self.state
.complete_inference(job, text.clone())
.map_err(state_error)?;
self.submitted = max_box(self.submitted, initial);
if !text.is_empty() {
let terminal = before.map_or(1, |id| id.get().saturating_add(1));
self.submitted = max_box(self.submitted, Some(BoxId::new(terminal)));
}
self.active = None;
Ok(())
}
pub fn fail(&mut self, job: u64, message: String, restartable: bool) {
let Some(active) = self.active.as_ref() else {
return;
};
if active.job != job {
return;
}
let restartable = restartable && !active.accepted_call_wave;
let _ = self.state.flush_active_arrivals(job);
if self.state.stall_inference(job, message.clone()).is_err() {
return;
}
self.active = None;
self.pending_steer = None;
self.stall = Some(Stall {
message,
restartable,
});
}
pub fn restart(&mut self) -> Result<(), RestartError> {
let Some(stall) = &self.stall else {
return Err(RestartError::NotStalled);
};
if !stall.restartable {
return Err(RestartError::NotRestartable);
}
self.state.take_halt();
self.state
.restart()
.map_err(|_| RestartError::StateRejected)?;
self.stall = None;
Ok(())
}
fn from_state(state: ActorState) -> Self {
Self {
state,
submitted: None,
active: None,
pending_steer: None,
next_steer_generation: 0,
stall: None,
}
}
fn require_job(&self, job: u64, operation: &str) -> Result<(), String> {
if self.active.as_ref().is_some_and(|active| active.job == job) {
Ok(())
} else {
Err(format!("stale Codex inference {operation}"))
}
}
}
impl Default for ConversationState {
fn default() -> Self {
Self::new()
}
}
impl PendingSteer {
fn token(&self) -> PreparedSteer {
PreparedSteer {
job: self.job,
generation: self.generation,
values: self.values.clone(),
}
}
}
fn terminal_text(output: ShimOutput<BoxValue>) -> Result<String, String> {
match output.items.as_slice() {
[] => Ok(String::new()),
[ShimItem::Text(text)] => Ok(text.clone()),
_ => Err("Codex completion must contain zero items or one text item".into()),
}
}
fn max_box(left: Option<BoxId>, right: Option<BoxId>) -> Option<BoxId> {
match (left, right) {
(Some(left), Some(right)) => Some(left.max(right)),
(left, right) => left.or(right),
}
}
fn state_error(error: impl std::fmt::Debug) -> String {
format!("chat state error: {error:?}")
}