1#![forbid(unsafe_code)]
2
3pub use kcode_k1_chat_codex_codec::{BoxValue, Call};
4pub use kcode_k1_chat_state::{
5 AGENT_ATTACHMENT_TYPE, AGENT_MESSAGE_TYPE, BoxId, ChatBox, SYSTEM_MESSAGE_TYPE,
6 TOOL_ATTACHMENT_TYPE, TOOL_CALL_TYPE, TOOL_MESSAGE_TYPE, TOOL_RESULT_TYPE, ToolCallId,
7 USER_ATTACHMENT_TYPE, USER_MESSAGE_TYPE,
8};
9
10use std::sync::Arc;
11
12use kcode_k1_chat_codex_codec::project;
13use kcode_k1_chat_state::{ActorState, ProviderCall, StateError};
14use kcode_k1_codex_adapter::{ShimItem, ShimOutput};
15
16#[derive(Clone, Debug, Eq, PartialEq)]
17pub struct Start {
18 pub job: u64,
19 pub boxes: Vec<BoxValue>,
20}
21
22#[derive(Clone, Debug, Eq, PartialEq)]
23pub struct PreparedCall {
24 pub tool_call_id: ToolCallId,
25 pub name: String,
26 pub arguments: String,
27}
28
29#[derive(Clone, Debug)]
30pub struct PreparedSteer(Arc<Prepared>);
31
32#[derive(Debug)]
33struct Prepared {
34 values: Vec<BoxValue>,
35 job: u64,
36 generation: u64,
37}
38
39impl PreparedSteer {
40 pub fn values(&self) -> &[BoxValue] {
41 &self.0.values
42 }
43}
44
45#[derive(Clone, Debug, Eq, PartialEq)]
46pub enum Status {
47 Running,
48 Quiet,
49 Stalled { message: String, restartable: bool },
50}
51
52#[derive(Clone, Copy, Debug, Eq, PartialEq)]
53pub enum RestartError {
54 NotStalled,
55 NotRestartable,
56}
57
58struct Round {
59 job: u64,
60 accepted_call_wave: bool,
61}
62
63enum Mode {
64 Idle,
65 Running(Round),
66 Stalled { message: String, restartable: bool },
67}
68
69pub struct ConversationState {
70 state: ActorState,
71 session: [u8; 12],
72 sequence: u64,
73 unsubmitted: Vec<BoxValue>,
74 queued_trigger: bool,
75 steer_required: bool,
76 generation: u64,
77 prepared: Option<Arc<Prepared>>,
78 mode: Mode,
79}
80
81impl ConversationState {
82 pub fn new(session: [u8; 12]) -> Self {
83 Self::from_actor(ActorState::new(false), session, 0)
84 }
85
86 pub fn recover(session: [u8; 12], boxes: Vec<ChatBox>, force: bool) -> Result<Self, String> {
87 let sequence = recovered_sequence(session, &boxes)?;
88 let state = ActorState::recover(boxes, force).map_err(debug)?;
89 Ok(Self::from_actor(state, session, sequence))
90 }
91
92 pub fn boxes(&self) -> &[ChatBox] {
93 self.state.boxes()
94 }
95
96 pub fn status(&self) -> Status {
97 match &self.mode {
98 Mode::Running(_) => Status::Running,
99 Mode::Idle if self.state.quiet() => Status::Quiet,
100 Mode::Idle => Status::Running,
101 Mode::Stalled {
102 message,
103 restartable,
104 } => Status::Stalled {
105 message: message.clone(),
106 restartable: *restartable,
107 },
108 }
109 }
110
111 pub fn accept(
112 &mut self,
113 box_type: String,
114 contents: String,
115 hidden_type: String,
116 hidden_contents: String,
117 ) -> Result<(), String> {
118 self.accept_arrival(true, |state| {
119 state.accept_box(box_type, contents, hidden_type, hidden_contents)
120 })
121 }
122
123 pub fn accept_tool_message(
124 &mut self,
125 tool_call_id: ToolCallId,
126 message: String,
127 ) -> Result<(), String> {
128 self.accept_arrival(false, |state| {
129 state.accept_tool_message(tool_call_id, message)
130 })
131 }
132
133 pub fn accept_tool_return(
134 &mut self,
135 tool_call_id: ToolCallId,
136 result: Result<String, String>,
137 ) -> Result<(), String> {
138 self.accept_arrival(true, |state| {
139 state.accept_async_return(tool_call_id, result)
140 })
141 }
142
143 pub fn accept_tool_return_v2(
144 &mut self,
145 tool_call_id: ToolCallId,
146 result: Result<String, String>,
147 metadata_type: String,
148 metadata_contents: String,
149 ) -> Result<(), String> {
150 self.accept_arrival(true, |state| {
151 state.accept_async_return_v2(tool_call_id, result, metadata_type, metadata_contents)
152 })
153 }
154
155 pub fn begin(&mut self) -> Result<Option<Start>, String> {
156 if !matches!(self.mode, Mode::Idle) {
157 return Ok(None);
158 }
159 let Some(start) = self.state.begin_inference().map_err(debug)? else {
160 return Ok(None);
161 };
162 let boxes = std::mem::take(&mut self.unsubmitted);
163 self.steer_required = self.queued_trigger;
164 self.mode = Mode::Running(Round {
165 job: start.job,
166 accepted_call_wave: false,
167 });
168 Ok(Some(Start {
169 job: start.job,
170 boxes,
171 }))
172 }
173
174 pub fn prepare_stage(
175 &mut self,
176 job: u64,
177 text: String,
178 values: Vec<BoxValue>,
179 ) -> Result<Vec<PreparedCall>, String> {
180 let calls = values
181 .into_iter()
182 .map(|value| match value {
183 BoxValue::Call(Ok(call)) => Ok(call),
184 _ => Err("stage contains a malformed tool call".to_owned()),
185 })
186 .collect::<Result<Vec<_>, _>>()?;
187 if !matches!(&self.mode, Mode::Running(round) if round.job == job) {
188 return Err("stale Codex inference stage".to_owned());
189 }
190 if self.prepared.is_some() {
191 return Err("a Codex steer remains uncommitted".to_owned());
192 }
193 let mut sequence = self.sequence;
194 let mut prepared = Vec::with_capacity(calls.len());
195 let mut provider_calls = Vec::with_capacity(calls.len());
196 for call in calls {
197 sequence = sequence
198 .checked_add(1)
199 .ok_or_else(|| "ToolCallId space was exhausted".to_owned())?;
200 let tool_call_id = ToolCallId::new(self.session, sequence);
201 provider_calls.push(ProviderCall {
202 tool_call_id,
203 name: call.name.clone(),
204 arguments: call.arguments.clone(),
205 });
206 prepared.push(PreparedCall {
207 tool_call_id,
208 name: call.name,
209 arguments: call.arguments,
210 });
211 }
212 self.state
213 .append_stage(job, text, provider_calls)
214 .map_err(debug)?;
215 self.sequence = sequence;
216 if !prepared.is_empty()
217 && let Mode::Running(round) = &mut self.mode
218 {
219 round.accepted_call_wave = true;
220 }
221 Ok(prepared)
222 }
223
224 pub fn flush_active_arrivals(&mut self, job: u64) -> Result<Vec<ChatBox>, String> {
225 let boxes = self.state.flush_active_arrivals(job).map_err(debug)?;
226 self.unsubmitted.extend(boxes.iter().map(project));
227 if !boxes.is_empty() {
228 self.queued_trigger = false;
229 }
230 Ok(boxes)
231 }
232
233 pub fn prepare_steer(&mut self, job: u64) -> Result<Option<PreparedSteer>, String> {
234 if !matches!(&self.mode, Mode::Running(round) if round.job == job) {
235 return Err("stale Codex inference steer".to_owned());
236 }
237 if let Some(prepared) = &self.prepared {
238 return Ok(Some(PreparedSteer(Arc::clone(prepared))));
239 }
240 self.flush_active_arrivals(job)?;
241 if !self.steer_required {
242 return Ok(None);
243 }
244 if self.unsubmitted.is_empty() {
245 return Err("Codex steer trigger has no pending arrival".to_owned());
246 }
247 let generation = self
248 .generation
249 .checked_add(1)
250 .ok_or_else(|| "Codex steer generation was exhausted".to_owned())?;
251 let prepared = Arc::new(Prepared {
252 values: self.unsubmitted.clone(),
253 job,
254 generation,
255 });
256 self.generation = generation;
257 self.steer_required = false;
258 self.prepared = Some(Arc::clone(&prepared));
259 Ok(Some(PreparedSteer(prepared)))
260 }
261
262 pub fn validate_steer(&self, prepared: &PreparedSteer) -> Result<(), String> {
263 let prepared = &prepared.0;
264 let valid = matches!(&self.mode, Mode::Running(round) if round.job == prepared.job)
265 && self.generation == prepared.generation
266 && self.unsubmitted.starts_with(&prepared.values)
267 && self
268 .prepared
269 .as_ref()
270 .is_some_and(|value| Arc::ptr_eq(value, prepared));
271 if valid {
272 Ok(())
273 } else {
274 Err("stale or invalid Codex steer".to_owned())
275 }
276 }
277
278 pub fn commit_steer(&mut self, prepared: PreparedSteer) -> Result<(), String> {
279 self.validate_steer(&prepared)?;
280 self.unsubmitted.drain(..prepared.0.values.len());
281 self.prepared = None;
282 Ok(())
283 }
284
285 pub fn complete(&mut self, job: u64, output: ShimOutput<BoxValue>) -> Result<(), String> {
286 let round = self.take_round(job)?;
287 if self.prepared.is_some() {
288 let message = "Codex inference completed with an uncommitted steer".to_owned();
289 self.preserve(round, message.clone(), false);
290 return Err(message);
291 }
292 let mut text = String::new();
293 for item in output.items {
294 match item {
295 ShimItem::Text(value) => text.push_str(&value),
296 ShimItem::Box(_) => {
297 let message = "terminal Codex output contains a box".to_owned();
298 self.preserve(round, message.clone(), false);
299 return Err(message);
300 }
301 }
302 }
303 if self.steer_required {
304 self.state.force_inference();
305 }
306 let before = self.state.boxes().len();
307 if let Err(error) = self.state.complete_inference(job, text).map_err(debug) {
308 self.preserve(round, error.clone(), false);
309 return Err(error);
310 }
311 self.unsubmitted
312 .extend(self.state.boxes()[before..].iter().map(project));
313 self.queued_trigger = false;
314 self.steer_required = false;
315 self.mode = Mode::Idle;
316 Ok(())
317 }
318
319 pub fn fail(&mut self, job: u64, message: String, restartable_before_launch: bool) {
320 if let Ok(round) = self.take_round(job) {
321 self.preserve(round, message, restartable_before_launch);
322 }
323 }
324
325 pub fn restart(&mut self) -> Result<(), RestartError> {
326 match self.mode {
327 Mode::Stalled {
328 restartable: true, ..
329 } => {}
330 Mode::Stalled { .. } => return Err(RestartError::NotRestartable),
331 _ => return Err(RestartError::NotStalled),
332 }
333 self.state
334 .restart()
335 .map_err(|_| RestartError::NotRestartable)?;
336 self.unsubmitted = self.state.boxes().iter().map(project).collect();
337 self.steer_required = false;
338 self.prepared = None;
339 self.mode = Mode::Idle;
340 Ok(())
341 }
342
343 fn from_actor(state: ActorState, session: [u8; 12], sequence: u64) -> Self {
344 let unsubmitted = state.boxes().iter().map(project).collect();
345 Self {
346 state,
347 session,
348 sequence,
349 unsubmitted,
350 queued_trigger: false,
351 steer_required: false,
352 generation: 0,
353 prepared: None,
354 mode: Mode::Idle,
355 }
356 }
357
358 fn accept_arrival<F>(&mut self, triggering: bool, accept: F) -> Result<(), String>
359 where
360 F: FnOnce(&mut ActorState) -> Result<(), StateError>,
361 {
362 let before = self.state.boxes().len();
363 accept(&mut self.state).map_err(debug)?;
364 let appended = &self.state.boxes()[before..];
365 self.unsubmitted.extend(appended.iter().map(project));
366 self.queued_trigger |= triggering && appended.is_empty();
367 self.steer_required |= triggering && matches!(self.mode, Mode::Running(_));
368 Ok(())
369 }
370
371 fn take_round(&mut self, job: u64) -> Result<Round, String> {
372 let mode = std::mem::replace(&mut self.mode, Mode::Idle);
373 match mode {
374 Mode::Running(round) if round.job == job => Ok(round),
375 other => {
376 self.mode = other;
377 Err("stale Codex inference completion".to_owned())
378 }
379 }
380 }
381
382 fn preserve(&mut self, round: Round, message: String, restartable_before_launch: bool) {
383 self.prepared = None;
384 if round.accepted_call_wave {
385 let _ = self.state.complete_inference(round.job, String::new());
386 let _ = self.state.halt(message.clone());
387 self.mode = Mode::Stalled {
388 message,
389 restartable: false,
390 };
391 } else {
392 let stalled = self
393 .state
394 .stall_inference(round.job, message.clone())
395 .is_ok();
396 self.mode = Mode::Stalled {
397 message,
398 restartable: stalled && restartable_before_launch,
399 };
400 }
401 }
402}
403
404fn recovered_sequence(session: [u8; 12], boxes: &[ChatBox]) -> Result<u64, String> {
405 let mut maximum = 0;
406 for box_ in boxes {
407 if let Some(call) = box_.tool_call_metadata().map_err(debug)? {
408 record_sequence(session, call.tool_call_id, &mut maximum)?;
409 }
410 if let Some(result) = box_.tool_result_metadata().map_err(debug)? {
411 record_sequence(session, result.tool_call_id, &mut maximum)?;
412 }
413 if let Some(result) = box_.tool_result_v2_metadata().map_err(debug)? {
414 record_sequence(session, result.tool_call_id, &mut maximum)?;
415 }
416 }
417 Ok(maximum)
418}
419
420fn record_sequence(session: [u8; 12], id: ToolCallId, maximum: &mut u64) -> Result<(), String> {
421 if id.nonce() != session {
422 return Err("recovered ToolCallId belongs to another session".to_owned());
423 }
424 *maximum = (*maximum).max(id.sequence());
425 Ok(())
426}
427
428fn debug(error: impl std::fmt::Debug) -> String {
429 format!("{error:?}")
430}