1#![forbid(unsafe_code)]
2
3pub use kcode_k1_chat_boxes::{
4 AGENT_ATTACHMENT_TYPE, AGENT_MESSAGE_TYPE, ATTACHMENT_TYPE, BoxId, ChatBox, MetadataError,
5 ProviderCall, SYSTEM_MESSAGE_TYPE, TOOL_ATTACHMENT_TYPE, TOOL_CALL_HIDDEN_TYPE, TOOL_CALL_TYPE,
6 TOOL_MESSAGE_HIDDEN_TYPE, TOOL_MESSAGE_TYPE, TOOL_RESULT_HIDDEN_TYPE, TOOL_RESULT_TYPE,
7 TOOL_RESULT_V2_HIDDEN_TYPE, ToolCallId, ToolMessageMetadata, ToolResultMetadata,
8 ToolResultV2Metadata, USER_ATTACHMENT_TYPE, USER_MESSAGE_TYPE, tool_call_box, tool_message_box,
9 tool_result_box, tool_result_v2_box,
10};
11
12#[derive(Clone, Debug, Eq, PartialEq)]
13pub struct DispatchedToolCall {
14 pub tool_call_id: ToolCallId,
15 pub call_box_id: BoxId,
16}
17
18#[derive(Clone, Copy, Debug, Eq, PartialEq)]
19pub enum TransitionError {
20 InvalidPhase,
21 BoxIdOverflow,
22 DuplicateToolCall,
23 UnknownToolCall,
24 DuplicateReturn,
25 MalformedToolConvention,
26 WrongOriginatingCall,
27 NonConsecutiveToolMessage,
28 ToolMessageAfterResult,
29}
30
31#[derive(Clone, Copy, Debug, Eq, PartialEq)]
32pub enum RecoveryError {
33 NonContiguousBoxId,
34 MalformedToolConvention,
35 DuplicateToolCall,
36 UnknownToolCall,
37 DuplicateReturn,
38 WrongOriginatingCall,
39 NonConsecutiveToolMessage,
40 ToolMessageAfterResult,
41}
42
43pub struct Chatend {
44 boxes: Vec<ChatBox>,
45 round_active: bool,
46 active_arrivals: Vec<ChatBox>,
47}
48
49impl Chatend {
50 pub const fn new() -> Self {
51 Self {
52 boxes: Vec::new(),
53 round_active: false,
54 active_arrivals: Vec::new(),
55 }
56 }
57
58 pub fn boxes(&self) -> &[ChatBox] {
59 &self.boxes
60 }
61
62 pub const fn round_active(&self) -> bool {
63 self.round_active
64 }
65
66 pub fn accept_box(
67 &mut self,
68 box_type: String,
69 contents: String,
70 hidden_type: String,
71 hidden_contents: String,
72 ) -> Result<Option<BoxId>, TransitionError> {
73 self.accept_arrival(ChatBox::new(
74 BoxId::new(0),
75 box_type,
76 contents,
77 hidden_type,
78 hidden_contents,
79 ))
80 }
81
82 pub fn accept_system(&mut self, contents: String) -> Result<Option<BoxId>, TransitionError> {
83 self.accept_box(
84 SYSTEM_MESSAGE_TYPE.to_owned(),
85 contents,
86 String::new(),
87 String::new(),
88 )
89 }
90
91 pub fn accept_user(&mut self, contents: String) -> Result<Option<BoxId>, TransitionError> {
92 self.accept_box(
93 USER_MESSAGE_TYPE.to_owned(),
94 contents,
95 String::new(),
96 String::new(),
97 )
98 }
99
100 pub fn accept_attachment(
101 &mut self,
102 contents: String,
103 hidden_type: String,
104 hidden_contents: String,
105 ) -> Result<Option<BoxId>, TransitionError> {
106 self.accept_box(
107 USER_ATTACHMENT_TYPE.to_owned(),
108 contents,
109 hidden_type,
110 hidden_contents,
111 )
112 }
113
114 pub fn start_round(&mut self) -> Result<Option<BoxId>, TransitionError> {
115 if self.round_active {
116 return Err(TransitionError::InvalidPhase);
117 }
118 self.round_active = true;
119 Ok(self.boxes.last().map(ChatBox::id))
120 }
121
122 pub fn append_stage(
123 &mut self,
124 agent_contents: String,
125 calls: Vec<ProviderCall>,
126 ) -> Result<Vec<DispatchedToolCall>, TransitionError> {
127 if !self.round_active {
128 return Err(TransitionError::InvalidPhase);
129 }
130 let existing = self.conventions()?;
131 for (index, call) in calls.iter().enumerate() {
132 if existing
133 .iter()
134 .any(|state| state.tool_call_id == call.tool_call_id)
135 || calls[..index]
136 .iter()
137 .any(|earlier| earlier.tool_call_id == call.tool_call_id)
138 {
139 return Err(TransitionError::DuplicateToolCall);
140 }
141 }
142
143 let has_agent = !agent_contents.is_empty();
144 let mut additions = Vec::with_capacity(calls.len() + usize::from(has_agent));
145 if has_agent {
146 additions.push(ChatBox::new(
147 BoxId::new(0),
148 AGENT_MESSAGE_TYPE.to_owned(),
149 agent_contents,
150 String::new(),
151 String::new(),
152 ));
153 }
154 additions.extend(calls.iter().map(tool_call_box));
155 let appended = self.append_batch(additions)?;
156 let call_offset = usize::from(has_agent);
157 Ok(calls
158 .iter()
159 .zip(&appended[call_offset..])
160 .map(|(call, value)| DispatchedToolCall {
161 tool_call_id: call.tool_call_id,
162 call_box_id: value.id(),
163 })
164 .collect())
165 }
166
167 pub fn accept_tool_message(
168 &mut self,
169 tool_call_id: ToolCallId,
170 message: String,
171 ) -> Result<Option<BoxId>, TransitionError> {
172 let state = self.call_state(tool_call_id)?;
173 if state.terminal {
174 return Err(TransitionError::ToolMessageAfterResult);
175 }
176 let message_index = state
177 .messages
178 .checked_add(1)
179 .ok_or(TransitionError::BoxIdOverflow)?;
180 let value = tool_message_box(&ToolMessageMetadata {
181 tool_call_id,
182 originating_call: state.box_id,
183 message_index,
184 message,
185 })
186 .map_err(|_| TransitionError::MalformedToolConvention)?;
187 self.accept_arrival(value)
188 }
189
190 pub fn accept_async_return(
191 &mut self,
192 tool_call_id: ToolCallId,
193 result: Result<String, String>,
194 ) -> Result<Option<BoxId>, TransitionError> {
195 let state = self.open_call_state(tool_call_id)?;
196 self.accept_arrival(tool_result_box(tool_call_id, state.box_id, result))
197 }
198
199 pub fn accept_async_return_v2(
200 &mut self,
201 tool_call_id: ToolCallId,
202 result: Result<String, String>,
203 metadata_type: String,
204 metadata_contents: String,
205 ) -> Result<Option<BoxId>, TransitionError> {
206 let state = self.open_call_state(tool_call_id)?;
207 self.accept_arrival(tool_result_v2_box(&ToolResultV2Metadata {
208 tool_call_id,
209 originating_call: state.box_id,
210 result,
211 metadata_type,
212 metadata_contents,
213 }))
214 }
215
216 pub fn flush_active_arrivals(&mut self) -> Result<Vec<ChatBox>, TransitionError> {
217 if !self.round_active {
218 return Err(TransitionError::InvalidPhase);
219 }
220 let appended = self.append_batch(self.active_arrivals.clone())?;
221 self.active_arrivals.clear();
222 Ok(appended)
223 }
224
225 pub fn done(&mut self, final_agent_contents: String) -> Result<Vec<ChatBox>, TransitionError> {
226 if !self.round_active {
227 return Err(TransitionError::InvalidPhase);
228 }
229 let has_final = !final_agent_contents.is_empty();
230 let mut additions = Vec::with_capacity(self.active_arrivals.len() + usize::from(has_final));
231 if has_final {
232 additions.push(ChatBox::new(
233 BoxId::new(0),
234 AGENT_MESSAGE_TYPE.to_owned(),
235 final_agent_contents,
236 String::new(),
237 String::new(),
238 ));
239 }
240 additions.extend(self.active_arrivals.iter().cloned());
241 let appended = self.append_batch(additions)?;
242 self.active_arrivals.clear();
243 self.round_active = false;
244 Ok(appended)
245 }
246
247 pub fn recover(boxes: Vec<ChatBox>) -> Result<Self, RecoveryError> {
248 for (index, value) in boxes.iter().enumerate() {
249 let expected = u64::try_from(index)
250 .ok()
251 .and_then(|index| index.checked_add(1))
252 .ok_or(RecoveryError::NonContiguousBoxId)?;
253 if value.id().get() != expected {
254 return Err(RecoveryError::NonContiguousBoxId);
255 }
256 }
257 audit(boxes.iter()).map_err(recovery_error)?;
258 Ok(Self {
259 boxes,
260 round_active: false,
261 active_arrivals: Vec::new(),
262 })
263 }
264
265 fn accept_arrival(&mut self, value: ChatBox) -> Result<Option<BoxId>, TransitionError> {
266 audit(
267 self.boxes
268 .iter()
269 .chain(&self.active_arrivals)
270 .chain(std::iter::once(&value)),
271 )
272 .map_err(transition_error)?;
273 if self.round_active {
274 self.active_arrivals.push(value);
275 Ok(None)
276 } else {
277 let mut appended = self.append_batch(vec![value])?;
278 Ok(appended.pop().map(|value| value.id()))
279 }
280 }
281
282 fn conventions(&self) -> Result<Vec<CallState>, TransitionError> {
283 audit(self.boxes.iter().chain(&self.active_arrivals)).map_err(transition_error)
284 }
285
286 fn call_state(&self, tool_call_id: ToolCallId) -> Result<CallState, TransitionError> {
287 self.conventions()?
288 .into_iter()
289 .find(|state| state.tool_call_id == tool_call_id && state.box_id.get() != 0)
290 .ok_or(TransitionError::UnknownToolCall)
291 }
292
293 fn open_call_state(&self, tool_call_id: ToolCallId) -> Result<CallState, TransitionError> {
294 let state = self.call_state(tool_call_id)?;
295 if state.terminal {
296 return Err(TransitionError::DuplicateReturn);
297 }
298 Ok(state)
299 }
300
301 fn append_batch(&mut self, additions: Vec<ChatBox>) -> Result<Vec<ChatBox>, TransitionError> {
302 self.ensure_capacity(additions.len())?;
303 let mut previous = self.boxes.last().map_or(0, |value| value.id().get());
304 let mut appended = Vec::with_capacity(additions.len());
305 for value in additions {
306 previous = previous
307 .checked_add(1)
308 .ok_or(TransitionError::BoxIdOverflow)?;
309 appended.push(ChatBox::new(
310 BoxId::new(previous),
311 value.box_type().to_owned(),
312 value.contents().to_owned(),
313 value.hidden_type().to_owned(),
314 value.hidden_contents().to_owned(),
315 ));
316 }
317 self.boxes.extend(appended.iter().cloned());
318 Ok(appended)
319 }
320
321 fn ensure_capacity(&self, additional: usize) -> Result<(), TransitionError> {
322 let additional = u64::try_from(additional).map_err(|_| TransitionError::BoxIdOverflow)?;
323 let previous = self.boxes.last().map_or(0, |value| value.id().get());
324 previous
325 .checked_add(additional)
326 .ok_or(TransitionError::BoxIdOverflow)?;
327 Ok(())
328 }
329}
330
331impl Default for Chatend {
332 fn default() -> Self {
333 Self::new()
334 }
335}
336
337#[derive(Clone, Copy)]
338struct CallState {
339 tool_call_id: ToolCallId,
340 box_id: BoxId,
341 messages: u64,
342 terminal: bool,
343}
344
345#[derive(Clone, Copy)]
346enum ConventionError {
347 Malformed,
348 DuplicateCall,
349 UnknownCall,
350 DuplicateReturn,
351 WrongOrigin,
352 NonConsecutiveMessage,
353 MessageAfterResult,
354}
355
356fn audit<'a>(
357 values: impl IntoIterator<Item = &'a ChatBox>,
358) -> Result<Vec<CallState>, ConventionError> {
359 let mut calls = Vec::<CallState>::new();
360 for value in values {
361 if let Some(call) = value
362 .tool_call_metadata()
363 .map_err(|_| ConventionError::Malformed)?
364 {
365 if calls
366 .iter()
367 .any(|state| state.tool_call_id == call.tool_call_id)
368 {
369 return Err(ConventionError::DuplicateCall);
370 }
371 calls.push(CallState {
372 tool_call_id: call.tool_call_id,
373 box_id: value.id(),
374 messages: 0,
375 terminal: false,
376 });
377 }
378 if let Some(message) = value
379 .tool_message_metadata()
380 .map_err(|_| ConventionError::Malformed)?
381 {
382 let state = find_call(&mut calls, message.tool_call_id)?;
383 if state.box_id.get() == 0 {
384 return Err(ConventionError::UnknownCall);
385 }
386 if state.box_id != message.originating_call {
387 return Err(ConventionError::WrongOrigin);
388 }
389 if state.terminal {
390 return Err(ConventionError::MessageAfterResult);
391 }
392 let expected = state
393 .messages
394 .checked_add(1)
395 .ok_or(ConventionError::NonConsecutiveMessage)?;
396 if message.message_index != expected {
397 return Err(ConventionError::NonConsecutiveMessage);
398 }
399 state.messages = expected;
400 }
401 if let Some(result) = value
402 .tool_result_metadata()
403 .map_err(|_| ConventionError::Malformed)?
404 {
405 let state = find_call(&mut calls, result.tool_call_id)?;
406 if state.box_id.get() == 0 {
407 return Err(ConventionError::UnknownCall);
408 }
409 if state.box_id != result.originating_call {
410 return Err(ConventionError::WrongOrigin);
411 }
412 if state.terminal {
413 return Err(ConventionError::DuplicateReturn);
414 }
415 state.terminal = true;
416 }
417 }
418 Ok(calls)
419}
420
421fn find_call(
422 calls: &mut [CallState],
423 tool_call_id: ToolCallId,
424) -> Result<&mut CallState, ConventionError> {
425 calls
426 .iter_mut()
427 .find(|state| state.tool_call_id == tool_call_id)
428 .ok_or(ConventionError::UnknownCall)
429}
430
431fn transition_error(error: ConventionError) -> TransitionError {
432 match error {
433 ConventionError::Malformed => TransitionError::MalformedToolConvention,
434 ConventionError::DuplicateCall => TransitionError::DuplicateToolCall,
435 ConventionError::UnknownCall => TransitionError::UnknownToolCall,
436 ConventionError::DuplicateReturn => TransitionError::DuplicateReturn,
437 ConventionError::WrongOrigin => TransitionError::WrongOriginatingCall,
438 ConventionError::NonConsecutiveMessage => TransitionError::NonConsecutiveToolMessage,
439 ConventionError::MessageAfterResult => TransitionError::ToolMessageAfterResult,
440 }
441}
442
443fn recovery_error(error: ConventionError) -> RecoveryError {
444 match error {
445 ConventionError::Malformed => RecoveryError::MalformedToolConvention,
446 ConventionError::DuplicateCall => RecoveryError::DuplicateToolCall,
447 ConventionError::UnknownCall => RecoveryError::UnknownToolCall,
448 ConventionError::DuplicateReturn => RecoveryError::DuplicateReturn,
449 ConventionError::WrongOrigin => RecoveryError::WrongOriginatingCall,
450 ConventionError::NonConsecutiveMessage => RecoveryError::NonConsecutiveToolMessage,
451 ConventionError::MessageAfterResult => RecoveryError::ToolMessageAfterResult,
452 }
453}
454
455#[cfg(test)]
456mod tests;