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