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 kcode_k1_chat_codex_codec::project;
5use kcode_k1_chat_state::{ActorState, BoxContent, ProviderCall};
6use kcode_k1_codex_adapter::{ShimItem, ShimOutput};
7
8#[derive(Clone, Debug, Eq, PartialEq)]
9pub enum Arrival {
10    System(String),
11    User(String),
12    Attachment,
13    Return {
14        tool_call_id: ToolCallId,
15        result: Result<String, String>,
16    },
17}
18
19#[derive(Clone, Debug, Eq, PartialEq)]
20pub struct Start {
21    pub job: u64,
22    pub boxes: Vec<BoxValue>,
23}
24
25#[derive(Clone, Debug, Eq, PartialEq)]
26pub struct PreparedCall {
27    pub tool_call_id: ToolCallId,
28    pub name: String,
29    pub arguments: String,
30}
31
32#[derive(Clone, Debug, Eq, PartialEq)]
33pub enum Status {
34    Running,
35    Quiet,
36    Stalled { message: String, restartable: bool },
37}
38
39#[derive(Clone, Copy, Debug, Eq, PartialEq)]
40pub enum RestartError {
41    NotStalled,
42    NotRestartable,
43}
44
45struct Round {
46    job: u64,
47    accepted_call_wave: bool,
48}
49
50enum Mode {
51    Idle,
52    Running(Round),
53    Stalled { message: String, restartable: bool },
54}
55
56pub struct ConversationState {
57    state: ActorState,
58    session: [u8; 12],
59    sequence: u64,
60    submitted: usize,
61    mode: Mode,
62}
63
64impl ConversationState {
65    pub fn new(session: [u8; 12]) -> Self {
66        Self {
67            state: ActorState::new(false),
68            session,
69            sequence: 0,
70            submitted: 0,
71            mode: Mode::Idle,
72        }
73    }
74
75    pub fn recover(session: [u8; 12], boxes: Vec<ChatBox>, force: bool) -> Result<Self, String> {
76        let sequence = recovered_sequence(session, &boxes)?;
77        let state = ActorState::recover(boxes, force).map_err(debug)?;
78        Ok(Self {
79            state,
80            session,
81            sequence,
82            submitted: 0,
83            mode: Mode::Idle,
84        })
85    }
86
87    pub fn boxes(&self) -> &[ChatBox] {
88        self.state.boxes()
89    }
90
91    pub fn status(&self) -> Status {
92        match &self.mode {
93            Mode::Running(_) => Status::Running,
94            Mode::Idle if self.state.quiet() => Status::Quiet,
95            Mode::Idle => Status::Running,
96            Mode::Stalled {
97                message,
98                restartable,
99            } => Status::Stalled {
100                message: message.clone(),
101                restartable: *restartable,
102            },
103        }
104    }
105
106    pub fn accept(&mut self, arrival: Arrival) -> Result<(), String> {
107        match arrival {
108            Arrival::System(text) => self.state.accept_system(text).map_err(debug),
109            Arrival::User(text) => self.state.accept_user(text).map_err(debug),
110            Arrival::Attachment => self.state.accept_attachment().map_err(debug),
111            Arrival::Return {
112                tool_call_id,
113                result,
114            } => self
115                .state
116                .accept_async_return(tool_call_id, result)
117                .map_err(debug),
118        }
119    }
120
121    pub fn begin(&mut self) -> Result<Option<Start>, String> {
122        if !matches!(self.mode, Mode::Idle) {
123            return Ok(None);
124        }
125        let Some(start) = self.state.begin_inference().map_err(debug)? else {
126            return Ok(None);
127        };
128        let boxes = self.state.boxes();
129        let projected = boxes[self.submitted..].iter().map(project).collect();
130        self.submitted = boxes.len();
131        self.mode = Mode::Running(Round {
132            job: start.job,
133            accepted_call_wave: false,
134        });
135        Ok(Some(Start {
136            job: start.job,
137            boxes: projected,
138        }))
139    }
140
141    pub fn prepare_stage(
142        &mut self,
143        job: u64,
144        text: String,
145        values: Vec<BoxValue>,
146    ) -> Result<Vec<PreparedCall>, String> {
147        let calls = values
148            .into_iter()
149            .map(|value| match value {
150                BoxValue::Call(Ok(call)) => Ok(call),
151                _ => Err("stage contains a malformed tool call".to_owned()),
152            })
153            .collect::<Result<Vec<_>, _>>()?;
154        match &self.mode {
155            Mode::Running(round) if round.job == job => {}
156            _ => return Err("stale Codex inference stage".to_owned()),
157        }
158        let mut sequence = self.sequence;
159        let prepared = calls
160            .iter()
161            .map(|call| {
162                sequence = sequence
163                    .checked_add(1)
164                    .ok_or_else(|| "ToolCallId space was exhausted".to_owned())?;
165                Ok(PreparedCall {
166                    tool_call_id: ToolCallId::new(self.session, sequence),
167                    name: call.name.clone(),
168                    arguments: call.arguments.clone(),
169                })
170            })
171            .collect::<Result<Vec<_>, String>>()?;
172        let provider_calls = prepared
173            .iter()
174            .map(|call| ProviderCall {
175                tool_call_id: call.tool_call_id,
176                name: call.name.clone(),
177                arguments: call.arguments.clone(),
178            })
179            .collect();
180        self.state
181            .append_stage(job, text, provider_calls)
182            .map_err(debug)?;
183        self.sequence = sequence;
184        if !prepared.is_empty()
185            && let Mode::Running(round) = &mut self.mode
186            && round.job == job
187        {
188            round.accepted_call_wave = true;
189        }
190        Ok(prepared)
191    }
192
193    pub fn complete(&mut self, job: u64, output: ShimOutput<BoxValue>) -> Result<(), String> {
194        let round = self.take_round(job)?;
195        let mut text = String::new();
196        for item in output.items {
197            match item {
198                ShimItem::Text(value) => text.push_str(&value),
199                ShimItem::Box(_) => {
200                    let message = "terminal Codex output contains a box".to_owned();
201                    self.preserve(round, message.clone(), false);
202                    return Err(message);
203                }
204            }
205        }
206        if let Err(error) = self.state.complete_inference(job, text).map_err(debug) {
207            self.preserve(round, error.clone(), false);
208            return Err(error);
209        }
210        self.mode = Mode::Idle;
211        Ok(())
212    }
213
214    pub fn fail(&mut self, job: u64, message: String, restartable_before_launch: bool) {
215        if let Ok(round) = self.take_round(job) {
216            self.preserve(round, message, restartable_before_launch);
217        }
218    }
219
220    pub fn restart(&mut self) -> Result<(), RestartError> {
221        match self.mode {
222            Mode::Stalled {
223                restartable: true, ..
224            } => {}
225            Mode::Stalled { .. } => return Err(RestartError::NotRestartable),
226            _ => return Err(RestartError::NotStalled),
227        }
228        self.state
229            .restart()
230            .map_err(|_| RestartError::NotRestartable)?;
231        self.submitted = 0;
232        self.mode = Mode::Idle;
233        Ok(())
234    }
235
236    fn take_round(&mut self, job: u64) -> Result<Round, String> {
237        let mode = std::mem::replace(&mut self.mode, Mode::Idle);
238        match mode {
239            Mode::Running(round) if round.job == job => Ok(round),
240            other => {
241                self.mode = other;
242                Err("stale Codex inference completion".to_owned())
243            }
244        }
245    }
246
247    fn preserve(&mut self, round: Round, message: String, restartable_before_launch: bool) {
248        if round.accepted_call_wave {
249            let _ = self.state.complete_inference(round.job, String::new());
250            let _ = self.state.halt(message.clone());
251            self.mode = Mode::Stalled {
252                message,
253                restartable: false,
254            };
255        } else {
256            let stalled = self
257                .state
258                .stall_inference(round.job, message.clone())
259                .is_ok();
260            self.mode = Mode::Stalled {
261                message,
262                restartable: stalled && restartable_before_launch,
263            };
264        }
265    }
266}
267
268fn recovered_sequence(session: [u8; 12], boxes: &[ChatBox]) -> Result<u64, String> {
269    let mut maximum = 0;
270    for box_ in boxes {
271        let id = match box_.content() {
272            BoxContent::KtoolCall { tool_call_id, .. }
273            | BoxContent::KtoolReturn { tool_call_id, .. } => tool_call_id,
274            _ => continue,
275        };
276        if id.session() != session {
277            return Err("recovered ToolCallId belongs to another session".to_owned());
278        }
279        maximum = maximum.max(id.sequence());
280    }
281    Ok(maximum)
282}
283
284fn debug(error: impl std::fmt::Debug) -> String {
285    format!("{error:?}")
286}