Skip to main content

kcode_k1_chat_codex_state/
lib.rs

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}