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