Skip to main content

kcode_k1_chat_chatend_testkit/
lib.rs

1#![forbid(unsafe_code)]
2
3use kcode_k1_chat_chatend::{
4    BoxId, ChatBox, Chatend, ProviderCall, RecoveryError, TOOL_CALL_HIDDEN_TYPE, TOOL_CALL_TYPE,
5    TOOL_MESSAGE_TYPE, TOOL_RESULT_TYPE, ToolCallId, ToolMessageMetadata, ToolResultV2Metadata,
6    TransitionError, USER_MESSAGE_TYPE, tool_call_box, tool_message_box, tool_result_box,
7    tool_result_v2_box,
8};
9
10pub fn assert_chatend_conformance() {
11    assert_active_lifecycle();
12    assert_done_order();
13    assert_result_without_messages();
14    assert_recovery_rejections();
15    assert_convention_boundaries();
16}
17
18fn assert_active_lifecycle() {
19    let call_id = tool_id(1, 7);
20    let mut chat = Chatend::new();
21    assert_eq!(
22        chat.accept_user("question".into()).unwrap(),
23        Some(BoxId::new(1))
24    );
25    chat.start_round().unwrap();
26    let dispatched = chat
27        .append_stage("calling".into(), vec![provider_call(call_id)])
28        .unwrap();
29    assert_eq!(dispatched[0].call_box_id, BoxId::new(3));
30    assert_eq!(
31        chat.accept_tool_message(call_id, "first".into()).unwrap(),
32        None
33    );
34    assert_eq!(
35        chat.accept_tool_message(call_id, "second".into()).unwrap(),
36        None
37    );
38    assert_eq!(
39        chat.accept_async_return_v2(
40            call_id,
41            Ok("answer".into()),
42            "k1.web-search-result/v1".into(),
43            r#"{"sources":[]}"#.into(),
44        )
45        .unwrap(),
46        None
47    );
48
49    let arrivals = chat.flush_active_arrivals().unwrap();
50    assert_eq!(
51        arrivals.iter().map(ChatBox::box_type).collect::<Vec<_>>(),
52        vec![TOOL_MESSAGE_TYPE, TOOL_MESSAGE_TYPE, TOOL_RESULT_TYPE]
53    );
54    for (offset, value) in arrivals.iter().enumerate() {
55        assert_eq!(value.id(), BoxId::new(4 + offset as u64));
56    }
57    for (offset, value) in arrivals[..2].iter().enumerate() {
58        let metadata = value.tool_message_metadata().unwrap().unwrap();
59        assert_eq!(metadata.originating_call, BoxId::new(3));
60        assert_eq!(metadata.message_index, 1 + offset as u64);
61    }
62    let result = arrivals[2].tool_result_v2_metadata().unwrap().unwrap();
63    assert_eq!(result.originating_call, BoxId::new(3));
64    assert_eq!(result.metadata_type, "k1.web-search-result/v1");
65    assert_eq!(
66        chat.accept_tool_message(call_id, "late".into()),
67        Err(TransitionError::ToolMessageAfterResult)
68    );
69    assert_eq!(
70        chat.accept_async_return(call_id, Ok("duplicate".into())),
71        Err(TransitionError::DuplicateReturn)
72    );
73    Chatend::recover(chat.boxes().to_vec()).unwrap();
74}
75
76fn assert_done_order() {
77    let call_id = tool_id(2, 1);
78    let mut chat = Chatend::new();
79    chat.start_round().unwrap();
80    chat.append_stage(String::new(), vec![provider_call(call_id)])
81        .unwrap();
82    chat.accept_tool_message(call_id, "interim".into()).unwrap();
83    chat.accept_async_return(call_id, Ok("result".into()))
84        .unwrap();
85    let appended = chat.done("provider-final".into()).unwrap();
86    assert_eq!(
87        appended.iter().map(ChatBox::box_type).collect::<Vec<_>>(),
88        vec!["Agent Message", TOOL_MESSAGE_TYPE, TOOL_RESULT_TYPE]
89    );
90    assert!(!chat.round_active());
91}
92
93fn assert_result_without_messages() {
94    let call_id = tool_id(3, 1);
95    let mut chat = Chatend::new();
96    chat.start_round().unwrap();
97    chat.append_stage(String::new(), vec![provider_call(call_id)])
98        .unwrap();
99    chat.accept_async_return(call_id, Ok("done".into()))
100        .unwrap();
101    let arrivals = chat.flush_active_arrivals().unwrap();
102    assert_eq!(arrivals.len(), 1);
103    assert_eq!(arrivals[0].box_type(), TOOL_RESULT_TYPE);
104    Chatend::recover(chat.boxes().to_vec()).unwrap();
105}
106
107fn assert_recovery_rejections() {
108    let call_id = tool_id(4, 9);
109    let call = with_id(tool_call_box(&provider_call(call_id)), 1);
110
111    let gap = with_id(
112        tool_message_box(&ToolMessageMetadata {
113            tool_call_id: call_id,
114            originating_call: BoxId::new(1),
115            message_index: 2,
116            message: "gap".into(),
117        })
118        .unwrap(),
119        2,
120    );
121    assert_recovery_error(
122        vec![call.clone(), gap],
123        RecoveryError::NonConsecutiveToolMessage,
124    );
125
126    let wrong_origin = with_id(
127        tool_message_box(&ToolMessageMetadata {
128            tool_call_id: call_id,
129            originating_call: BoxId::new(99),
130            message_index: 1,
131            message: "wrong".into(),
132        })
133        .unwrap(),
134        2,
135    );
136    assert_recovery_error(
137        vec![call.clone(), wrong_origin],
138        RecoveryError::WrongOriginatingCall,
139    );
140
141    let orphan = with_id(
142        tool_message_box(&ToolMessageMetadata {
143            tool_call_id: tool_id(8, 1),
144            originating_call: BoxId::new(1),
145            message_index: 1,
146            message: "orphan".into(),
147        })
148        .unwrap(),
149        2,
150    );
151    assert_recovery_error(vec![call.clone(), orphan], RecoveryError::UnknownToolCall);
152
153    let result = with_id(tool_result_box(call_id, BoxId::new(1), Ok("one".into())), 2);
154    let duplicate = with_id(
155        tool_result_v2_box(&ToolResultV2Metadata {
156            tool_call_id: call_id,
157            originating_call: BoxId::new(1),
158            result: Ok("two".into()),
159            metadata_type: "test".into(),
160            metadata_contents: "{}".into(),
161        }),
162        3,
163    );
164    assert_recovery_error(
165        vec![call.clone(), result.clone(), duplicate],
166        RecoveryError::DuplicateReturn,
167    );
168
169    let late = with_id(
170        tool_message_box(&ToolMessageMetadata {
171            tool_call_id: call_id,
172            originating_call: BoxId::new(1),
173            message_index: 1,
174            message: "late".into(),
175        })
176        .unwrap(),
177        3,
178    );
179    assert_recovery_error(
180        vec![call.clone(), result, late],
181        RecoveryError::ToolMessageAfterResult,
182    );
183
184    assert_recovery_error(vec![with_id(call, 2)], RecoveryError::NonContiguousBoxId);
185}
186
187fn assert_convention_boundaries() {
188    let malformed = ChatBox::new(
189        BoxId::new(1),
190        USER_MESSAGE_TYPE.into(),
191        String::new(),
192        TOOL_CALL_HIDDEN_TYPE.into(),
193        String::new(),
194    );
195    assert_recovery_error(vec![malformed], RecoveryError::MalformedToolConvention);
196
197    let unknown = ChatBox::new(
198        BoxId::new(1),
199        TOOL_CALL_TYPE.into(),
200        "future".into(),
201        "future/tool-call".into(),
202        "opaque".into(),
203    );
204    let recovered = Chatend::recover(vec![unknown.clone()]).unwrap();
205    assert_eq!(recovered.boxes(), &[unknown]);
206}
207
208fn provider_call(tool_call_id: ToolCallId) -> ProviderCall {
209    ProviderCall {
210        tool_call_id,
211        name: "WebSearch".into(),
212        arguments: "{}".into(),
213    }
214}
215
216fn tool_id(byte: u8, sequence: u64) -> ToolCallId {
217    ToolCallId::new([byte; 12], sequence)
218}
219
220fn with_id(value: ChatBox, id: u64) -> ChatBox {
221    ChatBox::new(
222        BoxId::new(id),
223        value.box_type().to_owned(),
224        value.contents().to_owned(),
225        value.hidden_type().to_owned(),
226        value.hidden_contents().to_owned(),
227    )
228}
229
230fn assert_recovery_error(boxes: Vec<ChatBox>, expected: RecoveryError) {
231    assert!(matches!(Chatend::recover(boxes), Err(error) if error == expected));
232}
233
234#[cfg(test)]
235mod tests {
236    #[test]
237    fn complete_chatend_conformance() {
238        super::assert_chatend_conformance();
239    }
240}