#![forbid(unsafe_code)]
pub use kcode_k1_chat_boxes::{
AGENT_ATTACHMENT_TYPE, AGENT_MESSAGE_TYPE, ATTACHMENT_TYPE, BoxId, ChatBox, MetadataError,
ProviderCall, SYSTEM_MESSAGE_TYPE, TOOL_ATTACHMENT_TYPE, TOOL_CALL_HIDDEN_TYPE, TOOL_CALL_TYPE,
TOOL_MESSAGE_HIDDEN_TYPE, TOOL_MESSAGE_TYPE, TOOL_RESULT_HIDDEN_TYPE, TOOL_RESULT_TYPE,
TOOL_RESULT_V2_HIDDEN_TYPE, ToolCallId, ToolMessageMetadata, ToolResultMetadata,
ToolResultV2Metadata, USER_ATTACHMENT_TYPE, USER_MESSAGE_TYPE, tool_call_box, tool_message_box,
tool_result_box, tool_result_v2_box,
};
#[derive(Clone, Debug, Eq, PartialEq)]
pub struct DispatchedToolCall {
pub tool_call_id: ToolCallId,
pub call_box_id: BoxId,
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum TransitionError {
InvalidPhase,
BoxIdOverflow,
DuplicateToolCall,
UnknownToolCall,
DuplicateReturn,
MalformedToolConvention,
WrongOriginatingCall,
NonConsecutiveToolMessage,
ToolMessageAfterResult,
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum RecoveryError {
NonContiguousBoxId,
MalformedToolConvention,
DuplicateToolCall,
UnknownToolCall,
DuplicateReturn,
WrongOriginatingCall,
NonConsecutiveToolMessage,
ToolMessageAfterResult,
}
pub struct Chatend {
boxes: Vec<ChatBox>,
round_active: bool,
active_arrivals: Vec<ChatBox>,
}
impl Chatend {
pub const fn new() -> Self {
Self {
boxes: Vec::new(),
round_active: false,
active_arrivals: Vec::new(),
}
}
pub fn boxes(&self) -> &[ChatBox] {
&self.boxes
}
pub const fn round_active(&self) -> bool {
self.round_active
}
pub fn accept_box(
&mut self,
box_type: String,
contents: String,
hidden_type: String,
hidden_contents: String,
) -> Result<Option<BoxId>, TransitionError> {
self.accept_arrival(ChatBox::new(
BoxId::new(0),
box_type,
contents,
hidden_type,
hidden_contents,
))
}
pub fn accept_system(&mut self, contents: String) -> Result<Option<BoxId>, TransitionError> {
self.accept_box(
SYSTEM_MESSAGE_TYPE.to_owned(),
contents,
String::new(),
String::new(),
)
}
pub fn accept_user(&mut self, contents: String) -> Result<Option<BoxId>, TransitionError> {
self.accept_box(
USER_MESSAGE_TYPE.to_owned(),
contents,
String::new(),
String::new(),
)
}
pub fn accept_attachment(
&mut self,
contents: String,
hidden_type: String,
hidden_contents: String,
) -> Result<Option<BoxId>, TransitionError> {
self.accept_box(
USER_ATTACHMENT_TYPE.to_owned(),
contents,
hidden_type,
hidden_contents,
)
}
pub fn start_round(&mut self) -> Result<Option<BoxId>, TransitionError> {
if self.round_active {
return Err(TransitionError::InvalidPhase);
}
self.round_active = true;
Ok(self.boxes.last().map(ChatBox::id))
}
pub fn append_stage(
&mut self,
agent_contents: String,
calls: Vec<ProviderCall>,
) -> Result<Vec<DispatchedToolCall>, TransitionError> {
if !self.round_active {
return Err(TransitionError::InvalidPhase);
}
let existing = self.conventions()?;
for (index, call) in calls.iter().enumerate() {
if existing
.iter()
.any(|state| state.tool_call_id == call.tool_call_id)
|| calls[..index]
.iter()
.any(|earlier| earlier.tool_call_id == call.tool_call_id)
{
return Err(TransitionError::DuplicateToolCall);
}
}
let has_agent = !agent_contents.is_empty();
let mut additions = Vec::with_capacity(calls.len() + usize::from(has_agent));
if has_agent {
additions.push(ChatBox::new(
BoxId::new(0),
AGENT_MESSAGE_TYPE.to_owned(),
agent_contents,
String::new(),
String::new(),
));
}
additions.extend(calls.iter().map(tool_call_box));
let appended = self.append_batch(additions)?;
let call_offset = usize::from(has_agent);
Ok(calls
.iter()
.zip(&appended[call_offset..])
.map(|(call, value)| DispatchedToolCall {
tool_call_id: call.tool_call_id,
call_box_id: value.id(),
})
.collect())
}
pub fn accept_tool_message(
&mut self,
tool_call_id: ToolCallId,
message: String,
) -> Result<Option<BoxId>, TransitionError> {
let state = self.call_state(tool_call_id)?;
if state.terminal {
return Err(TransitionError::ToolMessageAfterResult);
}
let message_index = state
.messages
.checked_add(1)
.ok_or(TransitionError::BoxIdOverflow)?;
let value = tool_message_box(&ToolMessageMetadata {
tool_call_id,
originating_call: state.box_id,
message_index,
message,
})
.map_err(|_| TransitionError::MalformedToolConvention)?;
self.accept_arrival(value)
}
pub fn accept_async_return(
&mut self,
tool_call_id: ToolCallId,
result: Result<String, String>,
) -> Result<Option<BoxId>, TransitionError> {
let state = self.open_call_state(tool_call_id)?;
self.accept_arrival(tool_result_box(tool_call_id, state.box_id, result))
}
pub fn accept_async_return_v2(
&mut self,
tool_call_id: ToolCallId,
result: Result<String, String>,
metadata_type: String,
metadata_contents: String,
) -> Result<Option<BoxId>, TransitionError> {
let state = self.open_call_state(tool_call_id)?;
self.accept_arrival(tool_result_v2_box(&ToolResultV2Metadata {
tool_call_id,
originating_call: state.box_id,
result,
metadata_type,
metadata_contents,
}))
}
pub fn flush_active_arrivals(&mut self) -> Result<Vec<ChatBox>, TransitionError> {
if !self.round_active {
return Err(TransitionError::InvalidPhase);
}
let appended = self.append_batch(self.active_arrivals.clone())?;
self.active_arrivals.clear();
Ok(appended)
}
pub fn done(&mut self, final_agent_contents: String) -> Result<Vec<ChatBox>, TransitionError> {
if !self.round_active {
return Err(TransitionError::InvalidPhase);
}
let has_final = !final_agent_contents.is_empty();
let mut additions = Vec::with_capacity(self.active_arrivals.len() + usize::from(has_final));
if has_final {
additions.push(ChatBox::new(
BoxId::new(0),
AGENT_MESSAGE_TYPE.to_owned(),
final_agent_contents,
String::new(),
String::new(),
));
}
additions.extend(self.active_arrivals.iter().cloned());
let appended = self.append_batch(additions)?;
self.active_arrivals.clear();
self.round_active = false;
Ok(appended)
}
pub fn recover(boxes: Vec<ChatBox>) -> Result<Self, RecoveryError> {
for (index, value) in boxes.iter().enumerate() {
let expected = u64::try_from(index)
.ok()
.and_then(|index| index.checked_add(1))
.ok_or(RecoveryError::NonContiguousBoxId)?;
if value.id().get() != expected {
return Err(RecoveryError::NonContiguousBoxId);
}
}
audit(boxes.iter()).map_err(recovery_error)?;
Ok(Self {
boxes,
round_active: false,
active_arrivals: Vec::new(),
})
}
fn accept_arrival(&mut self, value: ChatBox) -> Result<Option<BoxId>, TransitionError> {
audit(
self.boxes
.iter()
.chain(&self.active_arrivals)
.chain(std::iter::once(&value)),
)
.map_err(transition_error)?;
if self.round_active {
self.active_arrivals.push(value);
Ok(None)
} else {
let mut appended = self.append_batch(vec![value])?;
Ok(appended.pop().map(|value| value.id()))
}
}
fn conventions(&self) -> Result<Vec<CallState>, TransitionError> {
audit(self.boxes.iter().chain(&self.active_arrivals)).map_err(transition_error)
}
fn call_state(&self, tool_call_id: ToolCallId) -> Result<CallState, TransitionError> {
self.conventions()?
.into_iter()
.find(|state| state.tool_call_id == tool_call_id && state.box_id.get() != 0)
.ok_or(TransitionError::UnknownToolCall)
}
fn open_call_state(&self, tool_call_id: ToolCallId) -> Result<CallState, TransitionError> {
let state = self.call_state(tool_call_id)?;
if state.terminal {
return Err(TransitionError::DuplicateReturn);
}
Ok(state)
}
fn append_batch(&mut self, additions: Vec<ChatBox>) -> Result<Vec<ChatBox>, TransitionError> {
self.ensure_capacity(additions.len())?;
let mut previous = self.boxes.last().map_or(0, |value| value.id().get());
let mut appended = Vec::with_capacity(additions.len());
for value in additions {
previous = previous
.checked_add(1)
.ok_or(TransitionError::BoxIdOverflow)?;
appended.push(ChatBox::new(
BoxId::new(previous),
value.box_type().to_owned(),
value.contents().to_owned(),
value.hidden_type().to_owned(),
value.hidden_contents().to_owned(),
));
}
self.boxes.extend(appended.iter().cloned());
Ok(appended)
}
fn ensure_capacity(&self, additional: usize) -> Result<(), TransitionError> {
let additional = u64::try_from(additional).map_err(|_| TransitionError::BoxIdOverflow)?;
let previous = self.boxes.last().map_or(0, |value| value.id().get());
previous
.checked_add(additional)
.ok_or(TransitionError::BoxIdOverflow)?;
Ok(())
}
}
impl Default for Chatend {
fn default() -> Self {
Self::new()
}
}
#[derive(Clone, Copy)]
struct CallState {
tool_call_id: ToolCallId,
box_id: BoxId,
messages: u64,
terminal: bool,
}
#[derive(Clone, Copy)]
enum ConventionError {
Malformed,
DuplicateCall,
UnknownCall,
DuplicateReturn,
WrongOrigin,
NonConsecutiveMessage,
MessageAfterResult,
}
fn audit<'a>(
values: impl IntoIterator<Item = &'a ChatBox>,
) -> Result<Vec<CallState>, ConventionError> {
let mut calls = Vec::<CallState>::new();
for value in values {
if let Some(call) = value
.tool_call_metadata()
.map_err(|_| ConventionError::Malformed)?
{
if calls
.iter()
.any(|state| state.tool_call_id == call.tool_call_id)
{
return Err(ConventionError::DuplicateCall);
}
calls.push(CallState {
tool_call_id: call.tool_call_id,
box_id: value.id(),
messages: 0,
terminal: false,
});
}
if let Some(message) = value
.tool_message_metadata()
.map_err(|_| ConventionError::Malformed)?
{
let state = find_call(&mut calls, message.tool_call_id)?;
if state.box_id.get() == 0 {
return Err(ConventionError::UnknownCall);
}
if state.box_id != message.originating_call {
return Err(ConventionError::WrongOrigin);
}
if state.terminal {
return Err(ConventionError::MessageAfterResult);
}
let expected = state
.messages
.checked_add(1)
.ok_or(ConventionError::NonConsecutiveMessage)?;
if message.message_index != expected {
return Err(ConventionError::NonConsecutiveMessage);
}
state.messages = expected;
}
if let Some(result) = value
.tool_result_metadata()
.map_err(|_| ConventionError::Malformed)?
{
let state = find_call(&mut calls, result.tool_call_id)?;
if state.box_id.get() == 0 {
return Err(ConventionError::UnknownCall);
}
if state.box_id != result.originating_call {
return Err(ConventionError::WrongOrigin);
}
if state.terminal {
return Err(ConventionError::DuplicateReturn);
}
state.terminal = true;
}
}
Ok(calls)
}
fn find_call(
calls: &mut [CallState],
tool_call_id: ToolCallId,
) -> Result<&mut CallState, ConventionError> {
calls
.iter_mut()
.find(|state| state.tool_call_id == tool_call_id)
.ok_or(ConventionError::UnknownCall)
}
fn transition_error(error: ConventionError) -> TransitionError {
match error {
ConventionError::Malformed => TransitionError::MalformedToolConvention,
ConventionError::DuplicateCall => TransitionError::DuplicateToolCall,
ConventionError::UnknownCall => TransitionError::UnknownToolCall,
ConventionError::DuplicateReturn => TransitionError::DuplicateReturn,
ConventionError::WrongOrigin => TransitionError::WrongOriginatingCall,
ConventionError::NonConsecutiveMessage => TransitionError::NonConsecutiveToolMessage,
ConventionError::MessageAfterResult => TransitionError::ToolMessageAfterResult,
}
}
fn recovery_error(error: ConventionError) -> RecoveryError {
match error {
ConventionError::Malformed => RecoveryError::MalformedToolConvention,
ConventionError::DuplicateCall => RecoveryError::DuplicateToolCall,
ConventionError::UnknownCall => RecoveryError::UnknownToolCall,
ConventionError::DuplicateReturn => RecoveryError::DuplicateReturn,
ConventionError::WrongOrigin => RecoveryError::WrongOriginatingCall,
ConventionError::NonConsecutiveMessage => RecoveryError::NonConsecutiveToolMessage,
ConventionError::MessageAfterResult => RecoveryError::ToolMessageAfterResult,
}
}
#[cfg(test)]
mod tests;