kcode-k1-chat-codex-state 0.4.0

Deterministic open-format Codex K1 conversation state
Documentation
#![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:?}")
}