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