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.unsubmitted.drain(..prepared.0.external_count);
336        self.prepared = None;
337        if let Mode::Running(round) = &mut self.mode {
338            round.phase = Phase::ProviderActive;
339        }
340        Ok(())
341    }
342
343    pub fn complete(&mut self, job: u64, output: ShimOutput<BoxValue>) -> Result<(), String> {
344        let mut round = self.take_round(job)?;
345        if round.phase != Phase::ProviderActive {
346            let message = "Codex inference completed outside a provider generation".to_owned();
347            self.preserve(round, message.clone(), false);
348            return Err(message);
349        }
350        if self.prepared.is_some() {
351            let message = "Codex inference completed with an uncommitted steer".to_owned();
352            self.preserve(round, message.clone(), false);
353            return Err(message);
354        }
355        let mut text = String::new();
356        for item in output.items {
357            match item {
358                ShimItem::Text(value) => text.push_str(&value),
359                ShimItem::Box(_) => {
360                    round.accepted_provider_action = true;
361                    let message = "terminal Codex output contains a box".to_owned();
362                    self.preserve(round, message.clone(), false);
363                    return Err(message);
364                }
365            }
366        }
367        let before = self.state.boxes().len();
368        if let Err(error) = self.state.complete_inference(job, text).map_err(debug) {
369            self.preserve(round, error.clone(), false);
370            return Err(error);
371        }
372        self.unsubmitted
373            .extend(self.state.boxes()[before..].iter().skip(1).map(project));
374        self.queued_trigger = false;
375        self.mode = Mode::Idle;
376        Ok(())
377    }
378
379    pub fn fail(&mut self, job: u64, message: String, restartable_before_launch: bool) {
380        if let Ok(round) = self.take_round(job) {
381            self.preserve(round, message, restartable_before_launch);
382        }
383    }
384
385    pub fn restart(&mut self) -> Result<(), RestartError> {
386        match &self.mode {
387            Mode::Stalled(_, true) => return Err(RestartError::ProviderActionAccepted),
388            Mode::Stalled(
389                Status::Stalled {
390                    restartable: true, ..
391                },
392                false,
393            ) => {}
394            _ => return Err(RestartError::NotStalled),
395        }
396        self.state.restart().map_err(|_| RestartError::NotStalled)?;
397        self.unsubmitted = self.state.boxes().iter().map(project).collect();
398        self.queued_trigger = false;
399        self.prepared = None;
400        self.mode = Mode::Idle;
401        Ok(())
402    }
403
404    fn accept_arrival<F>(&mut self, triggering: bool, accept: F) -> Result<(), String>
405    where
406        F: FnOnce(&mut ActorState) -> Result<(), StateError>,
407    {
408        let before = self.state.boxes().len();
409        accept(&mut self.state).map_err(debug)?;
410        let appended = &self.state.boxes()[before..];
411        self.unsubmitted.extend(appended.iter().map(project));
412        match &mut self.mode {
413            Mode::Running(round) => round.steer_needed |= triggering,
414            Mode::Idle if appended.is_empty() => self.queued_trigger |= triggering,
415            Mode::Idle => self.queued_trigger = false,
416            Mode::Stalled(_, _) => {}
417        }
418        Ok(())
419    }
420
421    fn promised_id(&self) -> Result<BoxId, String> {
422        let previous = self.state.boxes().last().map_or(0, |box_| box_.id().get());
423        let value = previous
424            .checked_add(1)
425            .ok_or_else(|| "BoxId space was exhausted".to_owned())?;
426        Ok(BoxId::new(value))
427    }
428
429    fn take_round(&mut self, job: u64) -> Result<Round, String> {
430        match std::mem::replace(&mut self.mode, Mode::Idle) {
431            Mode::Running(round) if round.job == job => Ok(round),
432            other => {
433                self.mode = other;
434                Err("stale Codex inference completion".to_owned())
435            }
436        }
437    }
438
439    fn preserve(&mut self, round: Round, message: String, restartable_before_launch: bool) {
440        self.prepared = None;
441        let restartable = restartable_before_launch
442            && round.restartable
443            && !round.accepted_provider_action
444            && self
445                .state
446                .stall_inference(round.job, message.clone())
447                .is_ok();
448        if !restartable {
449            let _ = self.state.halt(message.clone());
450        }
451        self.mode = Mode::Stalled(
452            Status::Stalled {
453                message,
454                restartable,
455            },
456            round.accepted_provider_action,
457        );
458    }
459}
460
461fn debug(error: impl std::fmt::Debug) -> String {
462    format!("{error:?}")
463}
464
465#[cfg(test)]
466mod tests {
467    use super::*;
468
469    fn active_after_steer() -> (ConversationState, u64, ToolCallId) {
470        let mut state = ConversationState::new([7; 12]);
471        state
472            .accept(
473                USER_MESSAGE_TYPE.into(),
474                "start".into(),
475                String::new(),
476                String::new(),
477            )
478            .unwrap();
479        let start = state.begin().unwrap().unwrap();
480        let calls = state
481            .prepare_stage(
482                start.job,
483                String::new(),
484                vec![BoxValue::Call(Ok(Call {
485                    name: "tool".into(),
486                    arguments: "{}".into(),
487                }))],
488            )
489            .unwrap();
490        let prepared = state.prepare_steer(start.job).unwrap().unwrap();
491        state.commit_steer(prepared).unwrap();
492        (state, start.job, calls[0].tool_call_id)
493    }
494
495    fn empty_output() -> ShimOutput<BoxValue> {
496        ShimOutput { items: Vec::new() }
497    }
498
499    #[test]
500    fn committed_steer_does_not_schedule_a_fresh_turn() {
501        let (mut state, job, _) = active_after_steer();
502        state.complete(job, empty_output()).unwrap();
503
504        assert_eq!(state.status(), Status::Quiet);
505        assert!(state.begin().unwrap().is_none());
506    }
507
508    #[test]
509    fn tool_messages_after_a_steer_remain_inert() {
510        let (mut state, job, tool_call_id) = active_after_steer();
511        state
512            .accept_tool_message(tool_call_id, "still running".into())
513            .unwrap();
514        state.complete(job, empty_output()).unwrap();
515
516        assert_eq!(state.status(), Status::Quiet);
517        assert!(state.begin().unwrap().is_none());
518    }
519
520    #[test]
521    fn one_tool_result_schedules_exactly_one_fresh_turn() {
522        let (mut state, job, tool_call_id) = active_after_steer();
523        state
524            .accept_tool_return(tool_call_id, Ok("done".into()))
525            .unwrap();
526        state.complete(job, empty_output()).unwrap();
527
528        let followup = state.begin().unwrap().unwrap();
529        assert!(state.begin().unwrap().is_none());
530        state.complete(followup.job, empty_output()).unwrap();
531        assert_eq!(state.status(), Status::Quiet);
532        assert!(state.begin().unwrap().is_none());
533    }
534}