Skip to main content

kcode_k1_chat_chatend/
lib.rs

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;