Skip to main content

kcode_k1_chat_codex_state/
lib.rs

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}