kcode_k1_chat_codex_state/
lib.rs1pub 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}