kcode-k1-chat-thread-durable-turn 0.1.0

Synchronous durable turn owner for one K1 chat thread
Documentation
#![forbid(unsafe_code)]

pub use kcode_k1_chat_codex_state::{
    BoxValue, ChatBox, PreparedCall, PreparedSteer, RestartError, ShimOutput, Start, Status,
    ToolCallId,
};

use kcode_k1_chat_codex_state::ConversationState;
use kcode_k1_chat_persistence::{EventRecord, Record, Session};
use kcode_k1_chat_thread_recovery::recover as recover_thread;
use serde_json::json;

pub struct DurableTurn {
    state: ConversationState,
    records: Vec<Record>,
    mirrored: usize,
    durable: usize,
    session: Session,
    returned: Vec<ToolCallId>,
}

impl DurableTurn {
    pub fn recover(session: Session) -> Result<Self, String> {
        let recovered = recover_thread(&session)?;
        let returned = returned_ids(recovered.state.boxes())?;
        Ok(Self {
            state: recovered.state,
            records: recovered.records,
            mirrored: recovered.mirrored,
            durable: recovered.durable,
            session,
            returned,
        })
    }

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

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

    pub fn accept(
        &mut self,
        box_type: String,
        contents: String,
        hidden_type: String,
        hidden_contents: String,
    ) -> Result<(), String> {
        let result = self
            .state
            .accept(box_type, contents, hidden_type, hidden_contents);
        self.finish(result)
    }

    pub fn accept_tool_return(
        &mut self,
        tool_call_id: ToolCallId,
        result: Result<String, String>,
    ) -> Result<(), String> {
        if self.returned.contains(&tool_call_id) {
            return self.finish(Ok(()));
        }
        let accepted = self.state.accept_tool_return(tool_call_id, result);
        if accepted.is_ok() {
            self.returned.push(tool_call_id);
        }
        self.finish(accepted)
    }

    pub fn accept_tool_message(
        &mut self,
        tool_call_id: ToolCallId,
        message: String,
    ) -> Result<(), String> {
        let result = self.state.accept_tool_message(tool_call_id, message);
        self.finish(result)
    }

    pub fn accept_tool_return_v2(
        &mut self,
        tool_call_id: ToolCallId,
        result: Result<String, String>,
        metadata_type: String,
        metadata_contents: String,
    ) -> Result<(), String> {
        let accepted = self.state.accept_tool_return_v2(
            tool_call_id,
            result,
            metadata_type,
            metadata_contents,
        );
        if accepted.is_ok() {
            self.returned.push(tool_call_id);
        }
        self.finish(accepted)
    }

    pub fn begin(&mut self) -> Result<Option<Start>, String> {
        self.state.begin()
    }

    pub fn prepare_stage(
        &mut self,
        job: u64,
        text: String,
        values: Vec<BoxValue>,
    ) -> Result<Vec<PreparedCall>, String> {
        let result = self.state.prepare_stage(job, text, values);
        self.finish(result)
    }

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

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

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

    pub fn complete(&mut self, job: u64, output: ShimOutput<BoxValue>) -> Result<bool, String> {
        if let Err(error) = self.state.complete(job, output) {
            return self.finish(Err(error));
        }
        self.mirror_boxes()?;
        let resume = matches!(self.state.status(), Status::Running);
        let after_box_id = self
            .state
            .boxes()
            .last()
            .ok_or_else(|| "completed turn has no terminal box".to_owned())?
            .id()
            .get();
        self.records.push(Record::Event(EventRecord {
            after_box_id,
            event_index: 1,
            connected_box_id: 0,
            handler: "llm_done".into(),
            data: json!({"resume": resume}),
        }));
        self.persist_pending()?;
        Ok(resume)
    }

    pub fn fail(&mut self, job: u64, message: String, restartable_before_launch: bool) {
        self.state.fail(job, message, restartable_before_launch);
    }

    pub fn restart(&mut self) -> Result<(), RestartError> {
        self.state.restart()
    }

    fn finish<T>(&mut self, operation: Result<T, String>) -> Result<T, String> {
        let persistence = self.mirror_and_persist();
        match (operation, persistence) {
            (Ok(value), Ok(())) => Ok(value),
            (Err(error), Ok(())) | (Ok(_), Err(error)) => Err(error),
            (Err(operation), Err(persistence)) => Err(format!(
                "{operation}; additionally failed to persist canonical history: {persistence}"
            )),
        }
    }

    fn mirror_and_persist(&mut self) -> Result<(), String> {
        self.mirror_boxes()?;
        self.persist_pending()
    }

    fn mirror_boxes(&mut self) -> Result<(), String> {
        let boxes = self.state.boxes();
        let additions = boxes
            .get(self.mirrored..)
            .ok_or_else(|| "canonical box frontier moved backwards".to_owned())?;
        self.records
            .extend(additions.iter().cloned().map(Record::Box));
        self.mirrored = boxes.len();
        Ok(())
    }

    fn persist_pending(&mut self) -> Result<(), String> {
        let suffix = self
            .records
            .get(self.durable..)
            .ok_or_else(|| "durable record frontier moved past canonical records".to_owned())?;
        if suffix.is_empty() {
            return Ok(());
        }
        self.session.persist(suffix.to_vec())?;
        self.durable = self.records.len();
        Ok(())
    }
}

fn returned_ids(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)
}

#[cfg(test)]
mod tests {
    use super::*;
    use kcode_k1_chat_codex_state::Call;
    use kcode_k1_chat_persistence::K1ChatPersistence;
    use kcode_k1_peering::K1Peering;
    use kcode_k1_txn_ordering::K1TxnOrdering;
    use std::fs;
    use std::path::PathBuf;
    use std::sync::Arc;
    use std::sync::atomic::{AtomicU64, Ordering};

    static NEXT: AtomicU64 = AtomicU64::new(0);

    struct Fixture {
        root: PathBuf,
        session: Option<Session>,
    }

    impl Fixture {
        fn new(nonce: u8) -> Self {
            let root = std::env::temp_dir().join(format!(
                "k1-durable-turn-{}-{}",
                std::process::id(),
                NEXT.fetch_add(1, Ordering::Relaxed)
            ));
            let _ = fs::remove_dir_all(&root);
            let ordering = Arc::new(K1TxnOrdering::open(&root.join("ordering")).unwrap());
            let peering =
                Arc::new(K1Peering::open(&root.join("peering"), Arc::clone(&ordering)).unwrap());
            let persistence =
                K1ChatPersistence::open(&root.join("persistence"), ordering, peering).unwrap();
            let (session, _) = persistence.session([nonce; 12]).unwrap();
            Self {
                root,
                session: Some(session),
            }
        }

        fn session(&self) -> Session {
            self.session.as_ref().unwrap().clone()
        }
    }

    impl Drop for Fixture {
        fn drop(&mut self) {
            drop(self.session.take());
            let _ = fs::remove_dir_all(&self.root);
        }
    }

    #[test]
    fn completion_persists_terminal_box_event_and_recovery_frontiers() {
        let fixture = Fixture::new(1);
        let session = fixture.session();
        let mut turn = DurableTurn::recover(session.clone()).unwrap();
        turn.accept(
            "User Message".into(),
            "hello".into(),
            String::new(),
            String::new(),
        )
        .unwrap();
        let start = turn.begin().unwrap().unwrap();
        let resume = turn
            .complete(start.job, ShimOutput { items: Vec::new() })
            .unwrap();

        assert!(!resume);
        assert_eq!(turn.status(), Status::Quiet);
        assert_eq!((turn.mirrored, turn.durable, turn.records.len()), (2, 3, 3));
        let log = session.load().unwrap();
        let Record::Event(event) = &log.records[2] else {
            panic!("expected llm_done event");
        };
        assert_eq!(event.after_box_id, 2);
        assert_eq!(event.event_index, 1);
        assert_eq!(event.connected_box_id, 0);
        assert_eq!(event.handler, "llm_done");
        assert_eq!(event.data, json!({"resume": false}));

        drop(turn);
        let recovered = DurableTurn::recover(session).unwrap();
        assert_eq!(recovered.boxes().len(), 2);
        assert_eq!(recovered.status(), Status::Quiet);
        assert_eq!(
            (recovered.mirrored, recovered.durable),
            (recovered.boxes().len(), recovered.records.len())
        );
    }

    #[test]
    fn active_fifo_is_hidden_then_persisted_and_v1_return_is_idempotent() {
        let fixture = Fixture::new(2);
        let session = fixture.session();
        let mut turn = DurableTurn::recover(session.clone()).unwrap();
        turn.accept(
            "User Message".into(),
            "search".into(),
            String::new(),
            String::new(),
        )
        .unwrap();
        let start = turn.begin().unwrap().unwrap();
        let calls = turn
            .prepare_stage(
                start.job,
                "working".into(),
                vec![BoxValue::Call(Ok(Call {
                    name: "WebSearch".into(),
                    arguments: "{}".into(),
                }))],
            )
            .unwrap();
        let tool_call_id = calls[0].tool_call_id;
        assert_eq!(session.load().unwrap().records.len(), 3);

        turn.accept_tool_message(tool_call_id, "searching".into())
            .unwrap();
        turn.accept_tool_return_v2(
            tool_call_id,
            Ok("found".into()),
            "k1.web-search-result/v1".into(),
            "opaque".into(),
        )
        .unwrap();
        turn.accept_tool_return(tool_call_id, Ok("duplicate".into()))
            .unwrap();
        assert_eq!(turn.boxes().len(), 3);
        assert_eq!(session.load().unwrap().records.len(), 3);

        let prepared = turn.prepare_steer(start.job).unwrap().unwrap();
        turn.validate_steer(&prepared).unwrap();
        assert_eq!(turn.boxes().len(), 5);
        assert_eq!(session.load().unwrap().records.len(), 5);
        assert_eq!((turn.mirrored, turn.durable, turn.records.len()), (5, 5, 5));
        assert!(turn.boxes()[3].tool_message_metadata().unwrap().is_some());
        assert!(turn.boxes()[4].tool_result_v2_metadata().unwrap().is_some());
        turn.commit_steer(prepared).unwrap();
    }
}