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, 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}