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