Skip to main content

kcode_k1_chat_state/
lib.rs

1#![forbid(unsafe_code)]
2#![doc = include_str!("../Documentation.md")]
3
4use std::sync::{Arc, atomic::AtomicU8};
5
6pub use kcode_k1_chat_chatend::{
7    AGENT_ATTACHMENT_TYPE, AGENT_MESSAGE_TYPE, AGENT_RESPONSE_TYPE, ATTACHMENT_TYPE, BoxId,
8    ChatBox, DispatchedToolCall, ProviderCall, ProviderGenerated, RecoveryError,
9    SYSTEM_MESSAGE_TYPE, TOOL_ATTACHMENT_TYPE, TOOL_CALL_TYPE, TOOL_MESSAGE_TYPE, TOOL_RESULT_TYPE,
10    ToolCallId, ToolMessageMetadata, ToolResultMetadata, ToolResultV2Metadata, TransitionError,
11    USER_ATTACHMENT_TYPE, USER_MESSAGE_TYPE,
12};
13
14#[derive(Clone, Debug)]
15pub struct InferenceStart {
16    pub job: u64,
17    pub frontier: Option<BoxId>,
18    pub attempt: Arc<AtomicU8>,
19}
20
21#[derive(Debug)]
22pub enum StateError {
23    Transition(TransitionError),
24    Recovery(RecoveryError),
25    WrongInference { expected: Option<u64>, actual: u64 },
26    JobIdExhausted,
27    NotStalled,
28    Busy,
29}
30
31impl From<TransitionError> for StateError {
32    fn from(error: TransitionError) -> Self {
33        Self::Transition(error)
34    }
35}
36
37struct ActiveInference {
38    job: u64,
39    frontier: Option<BoxId>,
40    _attempt: Arc<AtomicU8>,
41}
42struct RetryRound {
43    frontier: Option<BoxId>,
44    ready: bool,
45}
46
47pub struct ActorState {
48    chatend: kcode_k1_chat_chatend::Chatend,
49    next_job: u64,
50    active: Option<ActiveInference>,
51    retry: Option<RetryRound>,
52    scheduled: bool,
53    arrival_during_active: bool,
54    halt: Option<String>,
55}
56
57impl ActorState {
58    pub fn new(force: bool) -> Self {
59        Self {
60            chatend: kcode_k1_chat_chatend::Chatend::new(),
61            next_job: 0,
62            active: None,
63            retry: None,
64            scheduled: force,
65            arrival_during_active: false,
66            halt: None,
67        }
68    }
69
70    pub fn recover(boxes: Vec<ChatBox>, force: bool) -> Result<Self, StateError> {
71        Ok(Self {
72            chatend: kcode_k1_chat_chatend::Chatend::recover(boxes)
73                .map_err(StateError::Recovery)?,
74            next_job: 0,
75            active: None,
76            retry: None,
77            scheduled: force,
78            arrival_during_active: false,
79            halt: None,
80        })
81    }
82
83    pub fn boxes(&self) -> &[ChatBox] {
84        self.chatend.boxes()
85    }
86    pub fn halted(&self) -> bool {
87        self.halt.is_some() || self.retry.is_some()
88    }
89    pub fn halt(&mut self, text: String) -> bool {
90        if self.halted() {
91            false
92        } else {
93            self.halt = Some(text);
94            true
95        }
96    }
97    pub fn take_halt(&mut self) -> Option<String> {
98        self.halt.take()
99    }
100
101    pub fn restart(&mut self) -> Result<(), StateError> {
102        if !self.halted() {
103            return Err(StateError::NotStalled);
104        }
105        if self.active.is_some() {
106            return Err(StateError::Busy);
107        }
108        self.halt = None;
109        if let Some(retry) = &mut self.retry {
110            retry.ready = true;
111        } else {
112            self.scheduled = true;
113        }
114        Ok(())
115    }
116
117    pub fn accept_box(
118        &mut self,
119        box_type: String,
120        contents: String,
121        hidden_type: String,
122        hidden_contents: String,
123    ) -> Result<(), StateError> {
124        self.chatend
125            .accept_box(box_type, contents, hidden_type, hidden_contents)?;
126        self.arrival();
127        Ok(())
128    }
129    pub fn accept_system(&mut self, contents: String) -> Result<(), StateError> {
130        self.chatend.accept_system(contents)?;
131        self.arrival();
132        Ok(())
133    }
134    pub fn accept_user(&mut self, contents: String) -> Result<(), StateError> {
135        self.chatend.accept_user(contents)?;
136        self.arrival();
137        Ok(())
138    }
139    pub fn accept_attachment(
140        &mut self,
141        contents: String,
142        hidden_type: String,
143        hidden_contents: String,
144    ) -> Result<(), StateError> {
145        self.chatend
146            .accept_attachment(contents, hidden_type, hidden_contents)?;
147        self.arrival();
148        Ok(())
149    }
150    pub fn accept_tool_message(
151        &mut self,
152        tool_call_id: ToolCallId,
153        message: String,
154    ) -> Result<(), StateError> {
155        self.chatend.accept_tool_message(tool_call_id, message)?;
156        Ok(())
157    }
158    pub fn accept_async_return(
159        &mut self,
160        tool_call_id: ToolCallId,
161        result: Result<String, String>,
162    ) -> Result<(), StateError> {
163        self.chatend.accept_async_return(tool_call_id, result)?;
164        self.arrival();
165        Ok(())
166    }
167    pub fn accept_async_return_v2(
168        &mut self,
169        tool_call_id: ToolCallId,
170        result: Result<String, String>,
171        metadata_type: String,
172        metadata_contents: String,
173    ) -> Result<(), StateError> {
174        self.chatend.accept_async_return_v2(
175            tool_call_id,
176            result,
177            metadata_type,
178            metadata_contents,
179        )?;
180        self.arrival();
181        Ok(())
182    }
183
184    pub fn accept_idle_context_box(
185        &mut self,
186        box_type: String,
187        contents: String,
188        hidden_type: String,
189        hidden_contents: String,
190    ) -> Result<BoxId, StateError> {
191        self.require_idle_context()?;
192        self.chatend
193            .accept_box(box_type, contents, hidden_type, hidden_contents)?
194            .ok_or(StateError::Busy)
195    }
196
197    pub fn accept_idle_context_tool_call(
198        &mut self,
199        call: ProviderCall,
200    ) -> Result<BoxId, StateError> {
201        self.require_idle_context()?;
202        let value = kcode_k1_chat_chatend::tool_call_box(&call);
203        self.chatend
204            .accept_box(
205                value.box_type().to_owned(),
206                value.contents().to_owned(),
207                value.hidden_type().to_owned(),
208                value.hidden_contents().to_owned(),
209            )?
210            .ok_or(StateError::Busy)
211    }
212
213    pub fn accept_idle_context_tool_return(
214        &mut self,
215        tool_call_id: ToolCallId,
216        result: Result<String, String>,
217    ) -> Result<BoxId, StateError> {
218        self.require_idle_context()?;
219        self.chatend
220            .accept_async_return(tool_call_id, result)?
221            .ok_or(StateError::Busy)
222    }
223
224    pub fn force_inference(&mut self) {
225        self.scheduled = true;
226    }
227
228    pub fn begin_inference(&mut self) -> Result<Option<InferenceStart>, StateError> {
229        if self.halt.is_some() || self.active.is_some() {
230            return Ok(None);
231        }
232        if let Some(retry) = &self.retry {
233            if !retry.ready {
234                return Ok(None);
235            }
236            let frontier = retry.frontier;
237            let start = self.activate(frontier)?;
238            self.retry = None;
239            return Ok(Some(start));
240        }
241        if !self.scheduled {
242            return Ok(None);
243        }
244        let job = self.next_job()?;
245        self.chatend.start_round()?;
246        let frontier = self.boxes().last().map(ChatBox::id);
247        self.next_job = job;
248        self.scheduled = false;
249        Ok(Some(self.install_active(job, frontier)))
250    }
251
252    pub fn append_stage(
253        &mut self,
254        job: u64,
255        contents: String,
256        generated: Vec<ProviderGenerated>,
257    ) -> Result<Vec<DispatchedToolCall>, StateError> {
258        self.require_job(job)?;
259        Ok(self.chatend.append_stage(contents, generated)?)
260    }
261    pub fn flush_active_arrivals(&mut self, job: u64) -> Result<Vec<ChatBox>, StateError> {
262        self.require_job(job)?;
263        let arrivals = self.chatend.flush_active_arrivals()?;
264        if !arrivals.is_empty() {
265            self.arrival_during_active = false;
266        }
267        Ok(arrivals)
268    }
269    pub fn complete_inference(&mut self, job: u64, contents: String) -> Result<(), StateError> {
270        self.require_job(job)?;
271        self.chatend.done(contents)?;
272        self.active = None;
273        if self.arrival_during_active {
274            self.scheduled = true;
275            self.arrival_during_active = false;
276        }
277        Ok(())
278    }
279    pub fn stall_inference(&mut self, job: u64, text: String) -> Result<(), StateError> {
280        self.require_job(job)?;
281        let active = self.active.take().expect("validated active inference");
282        self.halt = Some(text);
283        self.retry = Some(RetryRound {
284            frontier: active.frontier,
285            ready: false,
286        });
287        Ok(())
288    }
289    pub fn quiet(&self) -> bool {
290        self.active.is_none() && self.retry.is_none() && !self.scheduled
291    }
292
293    fn require_idle_context(&self) -> Result<(), StateError> {
294        if self.quiet() && !self.halted() && !self.arrival_during_active {
295            Ok(())
296        } else {
297            Err(StateError::Busy)
298        }
299    }
300    fn arrival(&mut self) {
301        if self.active.is_some() || self.retry.is_some() {
302            self.arrival_during_active = true;
303        } else {
304            self.scheduled = true;
305        }
306    }
307    fn activate(&mut self, frontier: Option<BoxId>) -> Result<InferenceStart, StateError> {
308        let job = self.next_job()?;
309        self.next_job = job;
310        Ok(self.install_active(job, frontier))
311    }
312    fn install_active(&mut self, job: u64, frontier: Option<BoxId>) -> InferenceStart {
313        let attempt = Arc::new(AtomicU8::new(1));
314        self.active = Some(ActiveInference {
315            job,
316            frontier,
317            _attempt: attempt.clone(),
318        });
319        InferenceStart {
320            job,
321            frontier,
322            attempt,
323        }
324    }
325    fn next_job(&self) -> Result<u64, StateError> {
326        self.next_job
327            .checked_add(1)
328            .ok_or(StateError::JobIdExhausted)
329    }
330    fn require_job(&self, job: u64) -> Result<(), StateError> {
331        let expected = self.active.as_ref().map(|a| a.job);
332        if expected == Some(job) {
333            Ok(())
334        } else {
335            Err(StateError::WrongInference {
336                expected,
337                actual: job,
338            })
339        }
340    }
341}
342
343impl Default for ActorState {
344    fn default() -> Self {
345        Self::new(false)
346    }
347}
348
349#[cfg(test)]
350mod tests {
351    use super::*;
352    #[test]
353    fn idle_context_is_canonical_and_nontriggering() {
354        let mut state = ActorState::new(false);
355        state
356            .accept_idle_context_box(
357                SYSTEM_MESSAGE_TYPE.into(),
358                "prefix".into(),
359                String::new(),
360                String::new(),
361            )
362            .unwrap();
363        let id = ToolCallId::new([7; 12], 1);
364        state
365            .accept_idle_context_tool_call(ProviderCall {
366                tool_call_id: id,
367                name: "KmapOpenNode".into(),
368                arguments: "{}".into(),
369            })
370            .unwrap();
371        state
372            .accept_idle_context_tool_return(id, Ok("loaded".into()))
373            .unwrap();
374        assert!(state.quiet());
375        assert!(state.begin_inference().unwrap().is_none());
376        state.accept_user("kickoff".into()).unwrap();
377        assert!(state.begin_inference().unwrap().is_some());
378        assert_eq!(
379            state
380                .boxes()
381                .iter()
382                .map(ChatBox::box_type)
383                .collect::<Vec<_>>(),
384            [
385                SYSTEM_MESSAGE_TYPE,
386                TOOL_CALL_TYPE,
387                TOOL_RESULT_TYPE,
388                USER_MESSAGE_TYPE
389            ]
390        );
391    }
392}