kcode-k1-chat-codex-state 0.1.0

Deterministic per-conversation K1 and Codex round state
Documentation
pub use kcode_k1_chat_codex_codec::{BoxValue, Call};
pub use kcode_k1_chat_state::{ActionId, ChatBox};

use kcode_k1_chat_codex_codec::project;
use kcode_k1_chat_state::{ActorState, DispatchCall, DispatchOutcome};
use kcode_k1_codex_adapter::{ShimItem, ShimOutput};

#[derive(Clone, Debug, Eq, PartialEq)]
pub enum Arrival {
    System(String),
    User(String),
    Attachment,
    Return {
        action_id: ActionId,
        result: Result<String, String>,
    },
}

#[derive(Clone, Debug, Eq, PartialEq)]
pub struct Start {
    pub job: u64,
    pub boxes: Vec<BoxValue>,
}

#[derive(Clone, Debug, Eq, PartialEq)]
pub struct Launch {
    pub action_id: ActionId,
    pub result: Result<String, String>,
}

#[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,
}

#[derive(Clone)]
struct Launched {
    action_id: ActionId,
    call: Call,
}

struct Round {
    job: u64,
    ledger: Vec<Launched>,
    arrivals: Vec<Arrival>,
}

enum Mode {
    Idle,
    Running(Round),
    Stalled { message: String, restartable: bool },
}

pub struct ConversationState {
    state: ActorState,
    session: [u8; 12],
    sequence: u64,
    submitted: usize,
    mode: Mode,
}

impl ConversationState {
    pub fn new(session: [u8; 12]) -> Self {
        Self {
            state: ActorState::new(false),
            session,
            sequence: 0,
            submitted: 0,
            mode: Mode::Idle,
        }
    }

    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, arrival: Arrival) -> Result<(), String> {
        if let Mode::Running(round) = &mut self.mode {
            round.arrivals.push(arrival);
            Ok(())
        } else {
            self.apply(arrival)
        }
    }

    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,
            ledger: Vec::new(),
            arrivals: Vec::new(),
        });
        Ok(Some(Start {
            job: start.job,
            boxes: projected,
        }))
    }

    pub fn launch(&mut self, job: u64, call: Call) -> Result<Launch, String> {
        let Mode::Running(round) = &mut self.mode else {
            return Err("no K1 inference is running".to_owned());
        };
        if round.job != job {
            return Err("stale K1 inference launch".to_owned());
        }
        let sequence = self
            .sequence
            .checked_add(1)
            .ok_or_else(|| "K1 ActionId space was exhausted".to_owned())?;
        let action_id = ActionId::new(self.session, sequence);
        self.state
            .collect_provider_call(job, call.name.clone(), call.arguments.clone())
            .map_err(debug)?;
        round.ledger.push(Launched {
            action_id,
            call: call.clone(),
        });
        self.sequence = sequence;
        let result = kcode_k1_chat_thread_actions::launch(&call.name, &call.arguments);
        Ok(Launch { action_id, result })
    }

    pub fn complete(&mut self, job: u64, output: ShimOutput<BoxValue>) -> Result<(), String> {
        let round = self.take_round(job)?;
        let text = match validate(output, &round.ledger) {
            Ok(text) => text,
            Err(error) => {
                self.preserve(round, error.clone(), false);
                return Err(error);
            }
        };
        if let Err(error) = self.commit(&round, Some(&text)) {
            let message = format!("failed to commit Codex output: {error}");
            let _ = self.state.halt(message.clone());
            self.mode = Mode::Stalled {
                message: message.clone(),
                restartable: false,
            };
            return Err(message);
        }
        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.mode = Mode::Idle;
        Ok(())
    }

    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 K1 inference completion".to_owned())
            }
        }
    }

    fn commit(&mut self, round: &Round, text: Option<&str>) -> Result<(), String> {
        if let Some(text) = text {
            self.state
                .append_kennedy_text(round.job, text)
                .map_err(debug)?;
        }
        let action_ids = round.ledger.iter().map(|entry| entry.action_id).collect();
        let dispatches = self
            .state
            .complete_provider_output(round.job, action_ids)
            .map_err(debug)?;
        if !dispatches_match(&round.ledger, &dispatches) {
            return Err("committed calls did not match launch ledger".to_owned());
        }
        if !dispatches.is_empty() {
            self.state
                .complete_dispatch(vec![DispatchOutcome::Pending; dispatches.len()])
                .map_err(debug)?;
        }
        self.apply_all(round.arrivals.clone())
    }

    fn preserve(&mut self, round: Round, message: String, restartable_before_launch: bool) {
        let preserved = if round.ledger.is_empty() {
            self.state
                .stall_inference(round.job, message.clone())
                .map_err(debug)
                .and_then(|_| self.apply_all(round.arrivals))
                .is_ok()
        } else {
            let preserved = self.commit(&round, None).is_ok();
            let _ = self.state.halt(message.clone());
            preserved
        };
        self.mode = Mode::Stalled {
            message,
            restartable: round.ledger.is_empty() && restartable_before_launch && preserved,
        };
    }

    fn apply_all(&mut self, arrivals: Vec<Arrival>) -> Result<(), String> {
        for arrival in arrivals {
            self.apply(arrival)?;
        }
        Ok(())
    }

    fn apply(&mut self, arrival: Arrival) -> Result<(), String> {
        match arrival {
            Arrival::System(text) => self.state.accept_system(text).map_err(debug),
            Arrival::User(text) => self.state.accept_user(text).map_err(debug),
            Arrival::Attachment => self.state.accept_attachment().map_err(debug),
            Arrival::Return { action_id, result } => self
                .state
                .accept_async_return(action_id, result)
                .map_err(debug),
        }
    }
}

fn validate(output: ShimOutput<BoxValue>, ledger: &[Launched]) -> Result<String, String> {
    let mut text = String::new();
    let mut calls = Vec::new();
    for item in output.items {
        match item {
            ShimItem::Text(value) => text.push_str(&value),
            ShimItem::Box(BoxValue::Call(Ok(call))) => calls.push(call),
            ShimItem::Box(_) => return Err("shim returned an invalid call box".to_owned()),
        }
    }
    if calls
        != ledger
            .iter()
            .map(|entry| entry.call.clone())
            .collect::<Vec<_>>()
    {
        return Err("shim output did not match the launch ledger".to_string());
    }
    Ok(text)
}

fn dispatches_match(ledger: &[Launched], dispatches: &[DispatchCall]) -> bool {
    ledger.len() == dispatches.len()
        && ledger.iter().zip(dispatches).all(|(expected, actual)| {
            expected.action_id == actual.action_id
                && expected.call.name == actual.name
                && expected.call.arguments == actual.arguments
        })
}

fn debug(error: impl std::fmt::Debug) -> String {
    format!("{error:?}")
}