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