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