Skip to main content

kcode_k1_chat_codex_state/
lib.rs

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}