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