1#![forbid(unsafe_code)]
2
3use std::sync::{Arc, atomic::AtomicU8};
4
5pub use kcode_k1_chat_chatend::{
6 AGENT_ATTACHMENT_TYPE, AGENT_MESSAGE_TYPE, AGENT_RESPONSE_TYPE, ATTACHMENT_TYPE, BoxId,
7 ChatBox, DispatchedToolCall, ProviderCall, ProviderGenerated, RecoveryError,
8 SYSTEM_MESSAGE_TYPE, TOOL_ATTACHMENT_TYPE, TOOL_CALL_TYPE, TOOL_MESSAGE_TYPE, TOOL_RESULT_TYPE,
9 ToolCallId, ToolMessageMetadata, ToolResultMetadata, ToolResultV2Metadata, TransitionError,
10 USER_ATTACHMENT_TYPE, USER_MESSAGE_TYPE,
11};
12
13#[derive(Clone, Debug)]
14pub struct InferenceStart {
15 pub job: u64,
16 pub frontier: Option<BoxId>,
17 pub attempt: Arc<AtomicU8>,
18}
19
20#[derive(Debug)]
21pub enum StateError {
22 Transition(TransitionError),
23 Recovery(RecoveryError),
24 WrongInference { expected: Option<u64>, actual: u64 },
25 JobIdExhausted,
26 NotStalled,
27 Busy,
28}
29
30impl From<TransitionError> for StateError {
31 fn from(error: TransitionError) -> Self {
32 Self::Transition(error)
33 }
34}
35
36struct ActiveInference {
37 job: u64,
38 frontier: Option<BoxId>,
39 _attempt: Arc<AtomicU8>,
40}
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
87 pub fn halted(&self) -> bool {
88 self.halt.is_some() || self.retry.is_some()
89 }
90
91 pub fn halt(&mut self, text: String) -> bool {
92 if self.halted() {
93 false
94 } else {
95 self.halt = Some(text);
96 true
97 }
98 }
99
100 pub fn take_halt(&mut self) -> Option<String> {
101 self.halt.take()
102 }
103
104 pub fn restart(&mut self) -> Result<(), StateError> {
105 if !self.halted() {
106 return Err(StateError::NotStalled);
107 }
108 if self.active.is_some() {
109 return Err(StateError::Busy);
110 }
111 self.halt = None;
112 if let Some(retry) = &mut self.retry {
113 retry.ready = true;
114 } else {
115 self.scheduled = true;
116 }
117 Ok(())
118 }
119
120 pub fn accept_box(
121 &mut self,
122 box_type: String,
123 contents: String,
124 hidden_type: String,
125 hidden_contents: String,
126 ) -> Result<(), StateError> {
127 self.chatend
128 .accept_box(box_type, contents, hidden_type, hidden_contents)?;
129 self.arrival();
130 Ok(())
131 }
132
133 pub fn accept_system(&mut self, contents: String) -> Result<(), StateError> {
134 self.chatend.accept_system(contents)?;
135 self.arrival();
136 Ok(())
137 }
138
139 pub fn accept_user(&mut self, contents: String) -> Result<(), StateError> {
140 self.chatend.accept_user(contents)?;
141 self.arrival();
142 Ok(())
143 }
144
145 pub fn accept_attachment(
146 &mut self,
147 contents: String,
148 hidden_type: String,
149 hidden_contents: String,
150 ) -> Result<(), StateError> {
151 self.chatend
152 .accept_attachment(contents, hidden_type, hidden_contents)?;
153 self.arrival();
154 Ok(())
155 }
156
157 pub fn accept_tool_message(
158 &mut self,
159 tool_call_id: ToolCallId,
160 message: String,
161 ) -> Result<(), StateError> {
162 self.chatend.accept_tool_message(tool_call_id, message)?;
163 Ok(())
164 }
165
166 pub fn accept_async_return(
167 &mut self,
168 tool_call_id: ToolCallId,
169 result: Result<String, String>,
170 ) -> Result<(), StateError> {
171 self.chatend.accept_async_return(tool_call_id, result)?;
172 self.arrival();
173 Ok(())
174 }
175
176 pub fn accept_async_return_v2(
177 &mut self,
178 tool_call_id: ToolCallId,
179 result: Result<String, String>,
180 metadata_type: String,
181 metadata_contents: String,
182 ) -> Result<(), StateError> {
183 self.chatend.accept_async_return_v2(
184 tool_call_id,
185 result,
186 metadata_type,
187 metadata_contents,
188 )?;
189 self.arrival();
190 Ok(())
191 }
192
193 pub fn force_inference(&mut self) {
194 self.scheduled = true;
195 }
196
197 pub fn begin_inference(&mut self) -> Result<Option<InferenceStart>, StateError> {
198 if self.halt.is_some() || self.active.is_some() {
199 return Ok(None);
200 }
201 if let Some(retry) = &self.retry {
202 if !retry.ready {
203 return Ok(None);
204 }
205 let frontier = retry.frontier;
206 let start = self.activate(frontier)?;
207 self.retry = None;
208 return Ok(Some(start));
209 }
210 if !self.scheduled {
211 return Ok(None);
212 }
213 let job = self.next_job()?;
214 self.chatend.start_round()?;
215 let frontier = self.boxes().last().map(ChatBox::id);
216 self.next_job = job;
217 self.scheduled = false;
218 Ok(Some(self.install_active(job, frontier)))
219 }
220
221 pub fn append_stage(
222 &mut self,
223 job: u64,
224 contents: String,
225 generated: Vec<ProviderGenerated>,
226 ) -> Result<Vec<DispatchedToolCall>, StateError> {
227 self.require_job(job)?;
228 Ok(self.chatend.append_stage(contents, generated)?)
229 }
230
231 pub fn flush_active_arrivals(&mut self, job: u64) -> Result<Vec<ChatBox>, StateError> {
232 self.require_job(job)?;
233 let arrivals = self.chatend.flush_active_arrivals()?;
234 if !arrivals.is_empty() {
235 self.arrival_during_active = false;
236 }
237 Ok(arrivals)
238 }
239
240 pub fn complete_inference(&mut self, job: u64, contents: String) -> Result<(), StateError> {
241 self.require_job(job)?;
242 self.chatend.done(contents)?;
243 self.active = None;
244 if self.arrival_during_active {
245 self.scheduled = true;
246 self.arrival_during_active = false;
247 }
248 Ok(())
249 }
250
251 pub fn stall_inference(&mut self, job: u64, text: String) -> Result<(), StateError> {
252 self.require_job(job)?;
253 let active = self.active.take().expect("validated active inference");
254 self.halt = Some(text);
255 self.retry = Some(RetryRound {
256 frontier: active.frontier,
257 ready: false,
258 });
259 Ok(())
260 }
261
262 pub fn quiet(&self) -> bool {
263 self.active.is_none() && self.retry.is_none() && !self.scheduled
264 }
265
266 fn arrival(&mut self) {
267 if self.active.is_some() || self.retry.is_some() {
268 self.arrival_during_active = true;
269 } else {
270 self.scheduled = true;
271 }
272 }
273
274 fn activate(&mut self, frontier: Option<BoxId>) -> Result<InferenceStart, StateError> {
275 let job = self.next_job()?;
276 self.next_job = job;
277 Ok(self.install_active(job, frontier))
278 }
279
280 fn install_active(&mut self, job: u64, frontier: Option<BoxId>) -> InferenceStart {
281 let attempt = Arc::new(AtomicU8::new(1));
282 self.active = Some(ActiveInference {
283 job,
284 frontier,
285 _attempt: attempt.clone(),
286 });
287 InferenceStart {
288 job,
289 frontier,
290 attempt,
291 }
292 }
293
294 fn next_job(&self) -> Result<u64, StateError> {
295 self.next_job
296 .checked_add(1)
297 .ok_or(StateError::JobIdExhausted)
298 }
299
300 fn require_job(&self, job: u64) -> Result<(), StateError> {
301 let expected = self.active.as_ref().map(|active| active.job);
302 if expected == Some(job) {
303 Ok(())
304 } else {
305 Err(StateError::WrongInference {
306 expected,
307 actual: job,
308 })
309 }
310 }
311}
312
313impl Default for ActorState {
314 fn default() -> Self {
315 Self::new(false)
316 }
317}
318
319#[cfg(test)]
320mod tests;