#![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)
}