kcode-k1-chat-thread-durable-state 0.4.15

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

use kcode_k1_access_kmap::K1AccessKmap;
use kcode_k1_chat_persistence::Session;
pub use kcode_k1_chat_state::BoxId;
use kcode_k1_chat_state::{AGENT_MESSAGE_TYPE, AGENT_RESPONSE_TYPE, USER_MESSAGE_TYPE};
pub use kcode_k1_chat_thread_durable_turn::{
    BoxValue, ChatBox, EventRecord, ModelUsage, PreflightItem, PreflightMode, PreparedCall,
    PreparedMailboxFlush, PreparedPreflightCall, Status, TokenBreakdown, ToolCallId,
};
use kcode_k1_chat_thread_durable_turn::{DurableTurn, RestartError, ShimOutput};
pub use kcode_k1_chat_thread_ktools::{AccessContext, AccessPolicy, ProfileId, SetLaunchNodeKtool};
use kcode_k1_chat_thread_ktools::{ChatThreadKtoolExecutor, ChatThreadKtools};
use serde::Deserialize;
use std::sync::Arc;

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

pub struct DurableThread {
    turn: DurableTurn,
    ktools: ChatThreadKtoolExecutor,
    authorized: bool,
}

impl DurableThread {
    pub fn recover(session: Session, kmap: Arc<K1AccessKmap>) -> Result<Self, String> {
        Self::recover_with_ktools(session, ChatThreadKtools::new(kmap))
    }

    pub fn recover_with_social(
        session: Session,
        kmap: Arc<K1AccessKmap>,
        social: kcode_k1_ktool_social::SocialKtools,
    ) -> Result<Self, String> {
        Self::recover_with_ktools(session, ChatThreadKtools::new_with_social(kmap, social))
    }

    pub fn recover_with_social_and_set_launch_node(
        session: Session,
        kmap: Arc<K1AccessKmap>,
        social: kcode_k1_ktool_social::SocialKtools,
        set_launch_node: SetLaunchNodeKtool,
    ) -> Result<Self, String> {
        Self::recover_with_ktools(
            session,
            ChatThreadKtools::new_with_social_and_set_launch_node(kmap, social, set_launch_node),
        )
    }

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

    pub fn events(&self) -> Vec<EventRecord> {
        self.turn.events()
    }

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

    pub fn preflight_calls(&self) -> &[PreparedPreflightCall] {
        self.turn.preflight_calls()
    }

    pub fn preflight_executor(&self) -> ChatThreadKtoolExecutor {
        self.ktools.clone()
    }

    pub fn prepare_preflight(
        &mut self,
        context: AccessContext,
        profile_id: ProfileId,
        policy: AccessPolicy,
        items: Vec<PreflightItem>,
    ) -> Result<Vec<PreparedPreflightCall>, TransitionError> {
        for item in &items {
            if let PreflightItem::KtoolCall { name, .. } = item {
                let supported = self
                    .ktools
                    .supports(name)
                    .map_err(TransitionError::Internal)?;
                if !supported {
                    return Err(TransitionError::Internal(
                        "unsupported preflight Ktool".to_owned(),
                    ));
                }
            }
        }
        let installed = self.bind_authorization(context, profile_id, policy)?;
        match self.turn.prepare_preflight(items) {
            Ok(calls) => Ok(calls),
            Err(error) => {
                if installed {
                    self.clear_authorization();
                }
                Err(TransitionError::Internal(error))
            }
        }
    }

    pub fn authorize_preflight(
        &mut self,
        context: AccessContext,
        profile_id: ProfileId,
        policy: AccessPolicy,
    ) -> Result<(), TransitionError> {
        self.bind_authorization(context, profile_id, policy)
            .map(|_| ())
    }

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

    pub fn accept_external_box(
        &mut self,
        box_type: String,
        contents: String,
        hidden_type: String,
        hidden_contents: String,
    ) -> Result<(), TransitionError> {
        if box_type == USER_MESSAGE_TYPE {
            return Err(TransitionError::Unauthorized);
        }
        self.accept_box(box_type, contents, hidden_type, hidden_contents)
            .map_err(TransitionError::Internal)
    }

    pub fn accept_user(
        &mut self,
        context: AccessContext,
        profile_id: ProfileId,
        policy: AccessPolicy,
        contents: String,
    ) -> Result<(), TransitionError> {
        let installed = self.bind_authorization(context, profile_id, policy)?;
        match self.turn.accept(
            USER_MESSAGE_TYPE.into(),
            contents,
            String::new(),
            String::new(),
        ) {
            Ok(()) => Ok(()),
            Err(error) => {
                if installed {
                    self.clear_authorization();
                }
                Err(TransitionError::Internal(error))
            }
        }
    }

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

    pub fn prepare_stage(
        &mut self,
        job: u64,
        text: String,
        boxes: Vec<BoxValue>,
    ) -> Result<Vec<PreparedCall>, String> {
        self.turn.prepare_stage(job, text, boxes)
    }

    pub fn launch_action(&mut self, name: &str, arguments: &str) -> Result<String, String> {
        if !kcode_k1_ktool_docs::is_known_ktool(name) {
            return Err("unknown Ktool".into());
        }
        match name {
            "KtoolDocs" => kcode_k1_ktool_docs::ktool_docs(arguments),
            "SendMessage" => launch_send_message(&mut self.turn, arguments),
            _ => self.ktools.launch(name, arguments),
        }
    }

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

    pub fn accept_tool_return(
        &mut self,
        id: ToolCallId,
        result: Result<String, String>,
    ) -> Result<(), String> {
        self.turn.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.turn
            .accept_tool_return_v2(id, result, metadata_type, metadata_contents)
    }

    pub fn prepare_mailbox_flush(
        &mut self,
        job: u64,
    ) -> Result<Option<PreparedMailboxFlush>, String> {
        self.turn.prepare_mailbox_flush(job)
    }

    pub fn prepared_input(&self, prepared: &PreparedMailboxFlush) -> Result<String, String> {
        self.turn.validate_mailbox_flush(prepared)?;
        render_input(prepared.values())
    }

    pub fn commit_mailbox_flush(&mut self, prepared: PreparedMailboxFlush) -> Result<(), String> {
        self.turn.commit_mailbox_flush(prepared)
    }

    pub fn begin_input(&mut self) -> Result<Option<(u64, String)>, String> {
        let Some(start) = self.turn.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.turn.complete(job, output)
    }

    pub fn complete_with_terminal_response(
        &mut self,
        job: u64,
        output: ShimOutput<BoxValue>,
    ) -> Result<(bool, u64), String> {
        let terminal_index = self.turn.boxes().len();
        let resume = self.complete(job, output)?;
        let terminal =
            self.turn.boxes().get(terminal_index).ok_or_else(|| {
                "completion did not append a terminal Agent Response box".to_owned()
            })?;
        if terminal.box_type() != AGENT_RESPONSE_TYPE {
            return Err("completion terminal box was not an Agent Response".to_owned());
        }
        Ok((resume, terminal.id().get()))
    }

    pub fn record_model_usage(
        &mut self,
        connected_box_id: u64,
        usage: ModelUsage,
    ) -> Result<(), String> {
        self.turn.record_model_usage(connected_box_id, usage)
    }

    pub fn fail(&mut self, job: u64, error: String, restartable: bool) {
        self.turn.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.turn.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) {
        let _ = self.ktools.clear_authorization();
        self.authorized = false;
    }

    fn recover_with_ktools(session: Session, ktools: ChatThreadKtools) -> Result<Self, String> {
        Ok(Self {
            turn: DurableTurn::recover(session)?,
            ktools: ChatThreadKtoolExecutor::new(ktools),
            authorized: false,
        })
    }

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

#[derive(Deserialize)]
#[serde(deny_unknown_fields)]
struct SendMessageArguments {
    message: String,
}

fn launch_send_message(turn: &mut DurableTurn, arguments: &str) -> Result<String, String> {
    let parsed: SendMessageArguments =
        serde_json::from_str(arguments).map_err(|_| invalid_send_message())?;
    if parsed.message.is_empty() {
        return Err(invalid_send_message());
    }
    turn.accept(
        AGENT_MESSAGE_TYPE.into(),
        parsed.message,
        String::new(),
        String::new(),
    )?;
    Ok("success".into())
}

fn invalid_send_message() -> String {
    "invalid SendMessage arguments".into()
}

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 provider input 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)
}

#[cfg(test)]
mod tests;