kcode-k1-chat-thread-durable-state 0.2.0

Durable synchronous state transitions for one K1 chat thread
Documentation
#![forbid(unsafe_code)]

use kcode_k1_access_kmap::K1AccessKmap;
pub use kcode_k1_chat_codex_codec::BoxValue;
use kcode_k1_chat_codex_state::{ConversationState, RestartError};
pub use kcode_k1_chat_codex_state::{PreparedCall, PreparedSteer, Status};
use kcode_k1_chat_persistence::{EventRecord, Record, Session};
use kcode_k1_chat_state::USER_MESSAGE_TYPE;
pub use kcode_k1_chat_state::{BoxId, ChatBox, ToolCallId};
use kcode_k1_chat_thread_actions::ChatThreadActions;
pub use kcode_k1_chat_thread_actions::{AccessContext, AccessPolicy, ProfileId};
use kcode_k1_chat_thread_recovery::recover as recover_thread;
use kcode_k1_codex_adapter::ShimOutput;
use serde_json::json;
use std::sync::Arc;

#[derive(Clone, Debug, Eq, PartialEq)]
pub enum TransitionError {
    Unauthorized,
    NotStalled,
    NotRestartable,
    Internal(String),
}

pub struct DurableThread {
    state: ConversationState,
    records: Vec<Record>,
    mirrored: usize,
    durable: usize,
    session: Session,
    actions: ChatThreadActions,
    authorized: bool,
}

impl DurableThread {
    pub fn recover(session: Session, kmap: Arc<K1AccessKmap>) -> Result<Self, String> {
        let recovered = recover_thread(&session)?;
        Ok(Self {
            state: recovered.state,
            records: recovered.records,
            mirrored: recovered.mirrored,
            durable: recovered.durable,
            session,
            actions: ChatThreadActions::new(kmap),
            authorized: false,
        })
    }

    pub fn boxes(&self) -> &[ChatBox] {
        self.state.boxes()
    }

    pub fn status(&self) -> Status {
        self.state.status()
    }

    pub fn accept_box(
        &mut self,
        box_type: String,
        contents: String,
        hidden_type: String,
        hidden_contents: String,
    ) -> Result<(), String> {
        self.state
            .accept(box_type, contents, hidden_type, hidden_contents)
    }

    pub fn accept_user(
        &mut self,
        context: AccessContext,
        profile_id: ProfileId,
        policy: AccessPolicy,
        contents: String,
    ) -> Result<(), TransitionError> {
        self.bind_authorization(context, profile_id, policy)?;
        self.state
            .accept(
                USER_MESSAGE_TYPE.into(),
                contents,
                String::new(),
                String::new(),
            )
            .and_then(|()| self.checkpoint())
            .map_err(TransitionError::Internal)
    }

    pub fn accept_return(
        &mut self,
        id: ToolCallId,
        result: Result<String, String>,
    ) -> Result<(), String> {
        if returned_calls(self.state.boxes())?.contains(&id) {
            Ok(())
        } else {
            self.state.accept_tool_return(id, result)
        }
    }

    pub fn prepare_stage(
        &mut self,
        job: u64,
        text: String,
        boxes: Vec<BoxValue>,
    ) -> Result<Vec<PreparedCall>, String> {
        let calls = self.state.prepare_stage(job, text, boxes)?;
        self.checkpoint()?;
        Ok(calls)
    }

    pub fn launch_action(&mut self, name: &str, arguments: &str) -> Result<String, String> {
        self.actions.launch(name, arguments)
    }

    pub fn accept_tool_message(&mut self, id: ToolCallId, contents: String) -> Result<(), String> {
        self.state.accept_tool_message(id, contents)
    }

    pub fn accept_tool_return(
        &mut self,
        id: ToolCallId,
        result: Result<String, String>,
    ) -> Result<(), String> {
        self.state.accept_tool_return(id, result)
    }

    pub fn accept_tool_return_v2(
        &mut self,
        id: ToolCallId,
        result: Result<String, String>,
        metadata_type: String,
        metadata_contents: String,
    ) -> Result<(), String> {
        self.state
            .accept_tool_return_v2(id, result, metadata_type, metadata_contents)
    }

    pub fn flush_active_arrivals(&mut self, job: u64) -> Result<(), String> {
        self.state.flush_active_arrivals(job).map(|_| ())
    }

    pub fn checkpoint(&mut self) -> Result<(), String> {
        self.mirror();
        self.flush()
    }

    pub fn prepare_steer(&mut self, job: u64) -> Result<Option<PreparedSteer>, String> {
        let prepared = self.state.prepare_steer(job)?;
        self.checkpoint()?;
        Ok(prepared)
    }

    pub fn prepared_input(&self, prepared: &PreparedSteer) -> Result<String, String> {
        self.state.validate_steer(prepared)?;
        render_input(prepared.values())
    }

    pub fn commit_steer(&mut self, prepared: PreparedSteer) -> Result<(), String> {
        self.state.commit_steer(prepared)
    }

    pub fn begin_input(&mut self) -> Result<Option<(u64, String)>, String> {
        let Some(start) = self.state.begin()? else {
            return Ok(None);
        };
        Ok(Some((start.job, render_input(&start.values)?)))
    }

    pub fn complete(&mut self, job: u64, output: ShimOutput<BoxValue>) -> Result<bool, String> {
        self.state.complete(job, output)?;
        let resume = matches!(self.state.status(), Status::Running);
        self.persist_done(resume)?;
        Ok(resume)
    }

    pub fn fail(&mut self, job: u64, error: String, restartable: bool) {
        self.state.fail(job, error, restartable);
        self.clear_authorization();
    }

    pub fn restart(
        &mut self,
        context: AccessContext,
        profile_id: ProfileId,
        policy: AccessPolicy,
    ) -> Result<(), TransitionError> {
        let installed = self.bind_authorization(context, profile_id, policy)?;
        if let Err(error) = self.state.restart().map_err(|error| match error {
            RestartError::NotStalled => TransitionError::NotStalled,
            RestartError::ProviderActionAccepted => TransitionError::NotRestartable,
        }) {
            if installed {
                self.clear_authorization();
            }
            return Err(error);
        }
        Ok(())
    }

    pub fn clear_authorization(&mut self) {
        self.actions.clear_authorization();
        self.authorized = false;
    }

    fn bind_authorization(
        &mut self,
        context: AccessContext,
        profile_id: ProfileId,
        policy: AccessPolicy,
    ) -> Result<bool, TransitionError> {
        let installed = !self.authorized;
        if self
            .actions
            .bind_authorization(context, profile_id, policy)
            .is_err()
        {
            if installed {
                self.actions.clear_authorization();
            }
            return Err(TransitionError::Unauthorized);
        }
        self.authorized = true;
        Ok(installed)
    }

    fn mirror(&mut self) {
        self.records.extend(
            self.state.boxes()[self.mirrored..]
                .iter()
                .cloned()
                .map(Record::Box),
        );
        self.mirrored = self.state.boxes().len();
    }

    fn flush(&mut self) -> Result<(), String> {
        if self.durable < self.records.len() {
            self.session
                .persist(self.records[self.durable..].to_vec())?;
            self.durable = self.records.len();
        }
        Ok(())
    }

    fn persist_done(&mut self, resume: bool) -> Result<(), String> {
        self.mirror();
        let anchor = self
            .state
            .boxes()
            .last()
            .map_or(0, |value| value.id().get());
        let index = match self.records.last() {
            Some(Record::Event(event)) if event.after_box_id == anchor => {
                event.event_index.checked_add(1)
            }
            _ => Some(1),
        }
        .ok_or_else(|| "event index space was exhausted".to_owned())?;
        self.records.push(Record::Event(
            EventRecord::new(
                anchor,
                index,
                0,
                "llm_done".into(),
                json!({"resume": resume}),
            )
            .map_err(|error| error.to_string())?,
        ));
        self.flush()
    }
}

fn render_input(values: &[BoxValue]) -> Result<String, String> {
    let mut output = String::new();
    for value in values {
        let BoxValue::History(section) = value else {
            return Err("Codex steer contains a non-history value".into());
        };
        if section.is_empty() {
            continue;
        }
        if !output.is_empty() && !output.ends_with('\n') {
            output.push('\n');
        }
        output.push_str(section);
    }
    Ok(output)
}

fn returned_calls(boxes: &[ChatBox]) -> Result<Vec<ToolCallId>, String> {
    let mut returned = Vec::new();
    for value in boxes {
        if let Some(result) = value
            .tool_result_metadata()
            .map_err(|error| format!("{error:?}"))?
        {
            returned.push(result.tool_call_id);
        }
    }
    Ok(returned)
}