#![forbid(unsafe_code)]
use kcode_k1_chat_chatend::{
BoxId, ChatBox, Chatend, ProviderCall, RecoveryError, TOOL_CALL_HIDDEN_TYPE, TOOL_CALL_TYPE,
TOOL_MESSAGE_TYPE, TOOL_RESULT_TYPE, ToolCallId, ToolMessageMetadata, ToolResultV2Metadata,
TransitionError, USER_MESSAGE_TYPE, tool_call_box, tool_message_box, tool_result_box,
tool_result_v2_box,
};
pub fn assert_chatend_conformance() {
assert_active_lifecycle();
assert_done_order();
assert_result_without_messages();
assert_recovery_rejections();
assert_convention_boundaries();
}
fn assert_active_lifecycle() {
let call_id = tool_id(1, 7);
let mut chat = Chatend::new();
assert_eq!(
chat.accept_user("question".into()).unwrap(),
Some(BoxId::new(1))
);
chat.start_round().unwrap();
let dispatched = chat
.append_stage("calling".into(), vec![provider_call(call_id)])
.unwrap();
assert_eq!(dispatched[0].call_box_id, BoxId::new(3));
assert_eq!(
chat.accept_tool_message(call_id, "first".into()).unwrap(),
None
);
assert_eq!(
chat.accept_tool_message(call_id, "second".into()).unwrap(),
None
);
assert_eq!(
chat.accept_async_return_v2(
call_id,
Ok("answer".into()),
"k1.web-search-result/v1".into(),
r#"{"sources":[]}"#.into(),
)
.unwrap(),
None
);
let arrivals = chat.flush_active_arrivals().unwrap();
assert_eq!(
arrivals.iter().map(ChatBox::box_type).collect::<Vec<_>>(),
vec![TOOL_MESSAGE_TYPE, TOOL_MESSAGE_TYPE, TOOL_RESULT_TYPE]
);
for (offset, value) in arrivals.iter().enumerate() {
assert_eq!(value.id(), BoxId::new(4 + offset as u64));
}
for (offset, value) in arrivals[..2].iter().enumerate() {
let metadata = value.tool_message_metadata().unwrap().unwrap();
assert_eq!(metadata.originating_call, BoxId::new(3));
assert_eq!(metadata.message_index, 1 + offset as u64);
}
let result = arrivals[2].tool_result_v2_metadata().unwrap().unwrap();
assert_eq!(result.originating_call, BoxId::new(3));
assert_eq!(result.metadata_type, "k1.web-search-result/v1");
assert_eq!(
chat.accept_tool_message(call_id, "late".into()),
Err(TransitionError::ToolMessageAfterResult)
);
assert_eq!(
chat.accept_async_return(call_id, Ok("duplicate".into())),
Err(TransitionError::DuplicateReturn)
);
Chatend::recover(chat.boxes().to_vec()).unwrap();
}
fn assert_done_order() {
let call_id = tool_id(2, 1);
let mut chat = Chatend::new();
chat.start_round().unwrap();
chat.append_stage(String::new(), vec![provider_call(call_id)])
.unwrap();
chat.accept_tool_message(call_id, "interim".into()).unwrap();
chat.accept_async_return(call_id, Ok("result".into()))
.unwrap();
let appended = chat.done("provider-final".into()).unwrap();
assert_eq!(
appended.iter().map(ChatBox::box_type).collect::<Vec<_>>(),
vec!["Agent Message", TOOL_MESSAGE_TYPE, TOOL_RESULT_TYPE]
);
assert!(!chat.round_active());
}
fn assert_result_without_messages() {
let call_id = tool_id(3, 1);
let mut chat = Chatend::new();
chat.start_round().unwrap();
chat.append_stage(String::new(), vec![provider_call(call_id)])
.unwrap();
chat.accept_async_return(call_id, Ok("done".into()))
.unwrap();
let arrivals = chat.flush_active_arrivals().unwrap();
assert_eq!(arrivals.len(), 1);
assert_eq!(arrivals[0].box_type(), TOOL_RESULT_TYPE);
Chatend::recover(chat.boxes().to_vec()).unwrap();
}
fn assert_recovery_rejections() {
let call_id = tool_id(4, 9);
let call = with_id(tool_call_box(&provider_call(call_id)), 1);
let gap = with_id(
tool_message_box(&ToolMessageMetadata {
tool_call_id: call_id,
originating_call: BoxId::new(1),
message_index: 2,
message: "gap".into(),
})
.unwrap(),
2,
);
assert_recovery_error(
vec![call.clone(), gap],
RecoveryError::NonConsecutiveToolMessage,
);
let wrong_origin = with_id(
tool_message_box(&ToolMessageMetadata {
tool_call_id: call_id,
originating_call: BoxId::new(99),
message_index: 1,
message: "wrong".into(),
})
.unwrap(),
2,
);
assert_recovery_error(
vec![call.clone(), wrong_origin],
RecoveryError::WrongOriginatingCall,
);
let orphan = with_id(
tool_message_box(&ToolMessageMetadata {
tool_call_id: tool_id(8, 1),
originating_call: BoxId::new(1),
message_index: 1,
message: "orphan".into(),
})
.unwrap(),
2,
);
assert_recovery_error(vec![call.clone(), orphan], RecoveryError::UnknownToolCall);
let result = with_id(tool_result_box(call_id, BoxId::new(1), Ok("one".into())), 2);
let duplicate = with_id(
tool_result_v2_box(&ToolResultV2Metadata {
tool_call_id: call_id,
originating_call: BoxId::new(1),
result: Ok("two".into()),
metadata_type: "test".into(),
metadata_contents: "{}".into(),
}),
3,
);
assert_recovery_error(
vec![call.clone(), result.clone(), duplicate],
RecoveryError::DuplicateReturn,
);
let late = with_id(
tool_message_box(&ToolMessageMetadata {
tool_call_id: call_id,
originating_call: BoxId::new(1),
message_index: 1,
message: "late".into(),
})
.unwrap(),
3,
);
assert_recovery_error(
vec![call.clone(), result, late],
RecoveryError::ToolMessageAfterResult,
);
assert_recovery_error(vec![with_id(call, 2)], RecoveryError::NonContiguousBoxId);
}
fn assert_convention_boundaries() {
let malformed = ChatBox::new(
BoxId::new(1),
USER_MESSAGE_TYPE.into(),
String::new(),
TOOL_CALL_HIDDEN_TYPE.into(),
String::new(),
);
assert_recovery_error(vec![malformed], RecoveryError::MalformedToolConvention);
let unknown = ChatBox::new(
BoxId::new(1),
TOOL_CALL_TYPE.into(),
"future".into(),
"future/tool-call".into(),
"opaque".into(),
);
let recovered = Chatend::recover(vec![unknown.clone()]).unwrap();
assert_eq!(recovered.boxes(), &[unknown]);
}
fn provider_call(tool_call_id: ToolCallId) -> ProviderCall {
ProviderCall {
tool_call_id,
name: "WebSearch".into(),
arguments: "{}".into(),
}
}
fn tool_id(byte: u8, sequence: u64) -> ToolCallId {
ToolCallId::new([byte; 12], sequence)
}
fn with_id(value: ChatBox, id: u64) -> ChatBox {
ChatBox::new(
BoxId::new(id),
value.box_type().to_owned(),
value.contents().to_owned(),
value.hidden_type().to_owned(),
value.hidden_contents().to_owned(),
)
}
fn assert_recovery_error(boxes: Vec<ChatBox>, expected: RecoveryError) {
assert!(matches!(Chatend::recover(boxes), Err(error) if error == expected));
}
#[cfg(test)]
mod tests {
#[test]
fn complete_chatend_conformance() {
super::assert_chatend_conformance();
}
}