kcode-k1-chat-codex-state 0.5.0

Codex turn state over current K1 Chat State
Documentation
#![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:?}")
}