#![forbid(unsafe_code)]
pub use kcode_k1_chat_boxes::{
AGENT_ATTACHMENT_TYPE, AGENT_MESSAGE_TYPE, AGENT_RESPONSE_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 enum ProviderGenerated {
AgentMessage { contents: String },
ToolCall(ProviderCall),
}
#[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,
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
enum Phase {
Idle,
Generating(BoxId),
Boundary,
}
pub struct Chatend {
boxes: Vec<ChatBox>,
phase: Phase,
active_arrivals: Vec<ChatBox>,
}
impl Chatend {
pub const fn new() -> Self {
Self {
boxes: Vec::new(),
phase: Phase::Idle,
active_arrivals: Vec::new(),
}
}
pub fn boxes(&self) -> &[ChatBox] {
&self.boxes
}
pub const fn round_active(&self) -> bool {
!matches!(self.phase, Phase::Idle)
}
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);
}
let promised = self.next_after(0)?;
let anchor = self.boxes.last().map(ChatBox::id);
self.phase = Phase::Generating(promised);
Ok(anchor)
}
pub fn append_stage(
&mut self,
agent_response: String,
generated: Vec<ProviderGenerated>,
) -> Result<Vec<DispatchedToolCall>, TransitionError> {
let promised = match self.phase {
Phase::Generating(promised) => promised,
Phase::Idle | Phase::Boundary => return Err(TransitionError::InvalidPhase),
};
let existing = self.conventions()?;
for (index, value) in generated.iter().enumerate() {
let ProviderGenerated::ToolCall(call) = value else {
continue;
};
if existing
.iter()
.any(|state| state.tool_call_id == call.tool_call_id)
|| generated[..index].iter().any(|earlier| {
matches!(earlier, ProviderGenerated::ToolCall(earlier)
if earlier.tool_call_id == call.tool_call_id)
})
{
return Err(TransitionError::DuplicateToolCall);
}
}
let capacity = generated
.len()
.checked_add(1)
.ok_or(TransitionError::BoxIdOverflow)?;
let mut additions = Vec::with_capacity(capacity);
additions.push(ChatBox::new(
promised,
AGENT_RESPONSE_TYPE.to_owned(),
agent_response,
String::new(),
String::new(),
));
additions.extend(generated.iter().map(|value| match value {
ProviderGenerated::AgentMessage { contents } => ChatBox::new(
BoxId::new(0),
AGENT_MESSAGE_TYPE.to_owned(),
contents.clone(),
String::new(),
String::new(),
),
ProviderGenerated::ToolCall(call) => tool_call_box(call),
}));
let appended = self.append_batch(additions)?;
let dispatched = generated
.iter()
.zip(appended.iter().skip(1))
.filter_map(|(value, appended)| match value {
ProviderGenerated::AgentMessage { .. } => None,
ProviderGenerated::ToolCall(call) => Some(DispatchedToolCall {
tool_call_id: call.tool_call_id,
call_box_id: appended.id(),
}),
})
.collect();
self.phase = Phase::Boundary;
Ok(dispatched)
}
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.phase != Phase::Boundary {
return Err(TransitionError::InvalidPhase);
}
let promised = self.next_after(self.active_arrivals.len())?;
let appended = self.append_batch(self.active_arrivals.clone())?;
self.active_arrivals.clear();
self.phase = Phase::Generating(promised);
Ok(appended)
}
pub fn done(&mut self, agent_response: String) -> Result<Vec<ChatBox>, TransitionError> {
let promised = match self.phase {
Phase::Generating(promised) => promised,
Phase::Idle | Phase::Boundary => return Err(TransitionError::InvalidPhase),
};
let capacity = self
.active_arrivals
.len()
.checked_add(1)
.ok_or(TransitionError::BoxIdOverflow)?;
let mut additions = Vec::with_capacity(capacity);
additions.push(ChatBox::new(
promised,
AGENT_RESPONSE_TYPE.to_owned(),
agent_response,
String::new(),
String::new(),
));
additions.extend(self.active_arrivals.iter().cloned());
let appended = self.append_batch(additions)?;
self.active_arrivals.clear();
self.phase = Phase::Idle;
Ok(appended)
}
pub fn abort(&mut self) -> Result<Vec<ChatBox>, TransitionError> {
if self.phase == Phase::Idle {
return Err(TransitionError::InvalidPhase);
}
let appended = self.append_batch(self.active_arrivals.clone())?;
self.active_arrivals.clear();
self.phase = Phase::Idle;
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,
phase: Phase::Idle,
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 next_after(&self, additional: usize) -> Result<BoxId, TransitionError> {
let additional = u64::try_from(additional).map_err(|_| TransitionError::BoxIdOverflow)?;
let previous = self.boxes.last().map_or(0, |value| value.id().get());
let next = previous
.checked_add(additional)
.and_then(|value| value.checked_add(1))
.ok_or(TransitionError::BoxIdOverflow)?;
Ok(BoxId::new(next))
}
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,
}
}