#![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};
use kcode_k1_chat_thread_actions::ChatThreadActions;
pub use kcode_k1_chat_thread_actions::{AccessContext, AccessPolicy, ProfileId};
pub use kcode_k1_chat_thread_durable_turn::{
BoxValue, ChatBox, EventRecord, ModelUsage, PreparedCall, PreparedMailboxFlush, Status,
TokenBreakdown, ToolCallId,
};
use kcode_k1_chat_thread_durable_turn::{DurableTurn, RestartError, ShimOutput};
use serde_json::Value;
use std::sync::Arc;
#[derive(Clone, Debug, Eq, PartialEq)]
pub enum TransitionError {
Unauthorized,
NotStalled,
NotRestartable,
Internal(String),
}
pub struct DurableThread {
turn: DurableTurn,
actions: ChatThreadActions,
authorized: bool,
}
impl DurableThread {
pub fn recover(session: Session, kmap: Arc<K1AccessKmap>) -> Result<Self, String> {
Ok(Self {
turn: DurableTurn::recover(session)?,
actions: ChatThreadActions::new(kmap),
authorized: false,
})
}
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 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> {
match name {
"KtoolDocs" => kcode_k1_ktool_docs::ktool_docs(arguments),
"SendMessage" => launch_send_message(&mut self.turn, arguments),
_ => self.actions.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) {
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 launch_send_message(turn: &mut DurableTurn, arguments: &str) -> Result<String, String> {
let parsed: Value = serde_json::from_str(arguments).map_err(|_| invalid_send_message())?;
let Value::Object(mut fields) = parsed else {
return Err(invalid_send_message());
};
if fields.len() != 1 {
return Err(invalid_send_message());
}
let Some(Value::String(message)) = fields.remove("message") else {
return Err(invalid_send_message());
};
if message.is_empty() {
return Err(invalid_send_message());
}
turn.accept(
AGENT_MESSAGE_TYPE.into(),
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 {
use super::*;
use kcode_k1_access::K1Access;
use kcode_k1_chat_codex_state::Call;
use kcode_k1_chat_persistence::{K1ChatPersistence, Session};
use kcode_k1_chat_state::{AGENT_RESPONSE_TYPE, TOOL_CALL_TYPE, TOOL_RESULT_TYPE};
use kcode_k1_groups::K1Groups;
use kcode_k1_kmap::K1Kmap;
use kcode_k1_peering::K1Peering;
use kcode_k1_txn_ordering::K1TxnOrdering;
use tempfile::TempDir;
fn fixture() -> (TempDir, Session, Arc<K1AccessKmap>) {
let root = TempDir::new().unwrap();
let ordering = Arc::new(K1TxnOrdering::open(&root.path().join("ordering")).unwrap());
let peering =
Arc::new(K1Peering::open(&root.path().join("peering"), Arc::clone(&ordering)).unwrap());
let groups = Arc::new(
K1Groups::open(
&root.path().join("groups"),
Arc::clone(&ordering),
Arc::clone(&peering),
)
.unwrap(),
);
let access = Arc::new(
K1Access::open(
&root.path().join("access"),
Arc::clone(&ordering),
Arc::clone(&peering),
groups,
)
.unwrap(),
);
let kmap = Arc::new(
K1Kmap::open(
&root.path().join("kmap"),
Arc::clone(&ordering),
Arc::clone(&peering),
)
.unwrap(),
);
let access_kmap = Arc::new(K1AccessKmap::open(access, kmap).unwrap());
let persistence =
K1ChatPersistence::open(&root.path().join("persistence"), ordering, peering).unwrap();
let (session, original) = persistence.session([11; 12]).unwrap();
assert!(original.records.is_empty());
(root, session, access_kmap)
}
#[test]
fn send_message_is_durable_ordered_and_recovered_once() {
let (_root, session, access_kmap) = fixture();
let mut thread = DurableThread::recover(session.clone(), Arc::clone(&access_kmap)).unwrap();
thread
.accept_box(
USER_MESSAGE_TYPE.into(),
"hello".into(),
String::new(),
String::new(),
)
.unwrap();
let (job, _) = thread.begin_input().unwrap().unwrap();
let arguments = r#"{"message":"WORKING_MESSAGE"}"#;
let calls = thread
.prepare_stage(
job,
String::new(),
vec![BoxValue::Call(Ok(Call {
name: "SendMessage".into(),
arguments: arguments.into(),
}))],
)
.unwrap();
assert_eq!(calls.len(), 1);
let result = thread.launch_action("SendMessage", arguments);
assert_eq!(result, Ok("success".into()));
thread
.accept_tool_return(calls[0].tool_call_id, result)
.unwrap();
let prepared = thread.prepare_mailbox_flush(job).unwrap().unwrap();
let input = thread.prepared_input(&prepared).unwrap();
let call_at = input.find("| Tool Call]").unwrap();
let message_at = input.find("| Agent Message]").unwrap();
let result_at = input.find("| Tool Result]").unwrap();
assert!(call_at < message_at && message_at < result_at);
assert!(input.ends_with("| Agent Response]\n"));
assert_eq!(
thread
.boxes()
.iter()
.map(ChatBox::box_type)
.collect::<Vec<_>>(),
[
USER_MESSAGE_TYPE,
AGENT_RESPONSE_TYPE,
TOOL_CALL_TYPE,
AGENT_MESSAGE_TYPE,
TOOL_RESULT_TYPE,
]
);
let message = thread
.boxes()
.iter()
.find(|value| value.box_type() == AGENT_MESSAGE_TYPE)
.unwrap();
assert_eq!(message.contents(), "WORKING_MESSAGE");
assert_eq!((message.hidden_type(), message.hidden_contents()), ("", ""));
thread.commit_mailbox_flush(prepared).unwrap();
assert!(
!thread
.complete(job, ShimOutput { items: Vec::new() })
.unwrap()
);
drop(thread);
let mut recovered = DurableThread::recover(session, access_kmap).unwrap();
assert_eq!(
recovered
.boxes()
.iter()
.filter(|value| value.box_type() == AGENT_MESSAGE_TYPE)
.count(),
1
);
let before = recovered.boxes().len();
assert!(recovered.launch_action("CurrentTime", "{}").is_ok());
assert_eq!(recovered.boxes().len(), before);
}
#[test]
fn complete_returns_terminal_response_before_queued_arrivals() {
let (_root, session, access_kmap) = fixture();
let mut thread = DurableThread::recover(session, access_kmap).unwrap();
thread
.accept_box(
USER_MESSAGE_TYPE.into(),
"first".into(),
String::new(),
String::new(),
)
.unwrap();
let (job, _) = thread.begin_input().unwrap().unwrap();
let terminal_index = thread.boxes().len();
thread
.accept_box(
USER_MESSAGE_TYPE.into(),
"queued".into(),
String::new(),
String::new(),
)
.unwrap();
let (resume, terminal_id) = thread
.complete_with_terminal_response(job, ShimOutput { items: Vec::new() })
.unwrap();
assert!(resume);
assert_eq!(terminal_id, thread.boxes()[terminal_index].id().get());
assert_eq!(
thread.boxes()[terminal_index].box_type(),
AGENT_RESPONSE_TYPE
);
assert_eq!(
thread.boxes()[terminal_index + 1].box_type(),
USER_MESSAGE_TYPE
);
}
#[test]
fn send_message_rejects_invalid_arguments_without_a_message() {
let (_root, session, access_kmap) = fixture();
let mut thread = DurableThread::recover(session, access_kmap).unwrap();
for arguments in [
"",
"{",
"null",
"[]",
"{}",
r#"{"message":""}"#,
r#"{"message":1}"#,
r#"{"message":"x","extra":true}"#,
] {
let before = thread.boxes().len();
assert_eq!(
thread.launch_action("SendMessage", arguments),
Err("invalid SendMessage arguments".into())
);
assert_eq!(thread.boxes().len(), before);
}
}
#[test]
fn ktool_docs_is_stateless_without_authorization_and_preserves_dispatch() {
let (_root, session, access_kmap) = fixture();
let mut thread = DurableThread::recover(session, access_kmap).unwrap();
let initial_status = thread.status();
let document = thread
.launch_action("KtoolDocs", r#"{"name":"KtoolDocs"}"#)
.unwrap();
let document: Value = serde_json::from_str(&document).unwrap();
assert_eq!(
document.get("name").and_then(Value::as_str),
Some("KtoolDocs")
);
assert_eq!(
document.get("latest_version").and_then(Value::as_str),
Some("1.0.0")
);
assert_thread_state_unchanged(&thread, &initial_status);
assert_eq!(
thread.launch_action("KtoolDocs", "{}"),
Err("invalid KtoolDocs arguments".into())
);
assert_thread_state_unchanged(&thread, &initial_status);
assert_eq!(
thread.launch_action("KtoolDocs", r#"{"name":"MissingKtool"}"#),
Err("unknown Ktool".into())
);
assert_thread_state_unchanged(&thread, &initial_status);
assert!(thread.launch_action("CurrentTime", "{}").is_ok());
assert_thread_state_unchanged(&thread, &initial_status);
}
fn assert_thread_state_unchanged(thread: &DurableThread, initial_status: &Status) {
assert!(thread.boxes().is_empty());
assert!(thread.events().is_empty());
assert_eq!(&thread.status(), initial_status);
}
}