kcode_k1_chat_codex_state/
lib.rs1#![forbid(unsafe_code)]
2
3use kcode_k1_chat_codex_codec::{BoxValue, project};
4use kcode_k1_chat_state::{ActorState, BoxId, ChatBox, ProviderCall};
5use kcode_k1_codex_adapter::{ShimItem, ShimOutput};
6
7pub use kcode_k1_chat_codex_codec::Codec;
8pub use kcode_k1_chat_state::{
9 DispatchedToolCall, ResultView, ToolCall, ToolCallId, ToolResult, ToolResultStatus,
10};
11
12#[derive(Clone, Debug, PartialEq)]
13pub struct Start {
14 pub job: u64,
15 pub boxes: Vec<BoxValue>,
16}
17
18#[derive(Clone, Debug, Eq, PartialEq)]
19pub enum Status {
20 Running,
21 Quiet,
22 Stalled { message: String, restartable: bool },
23}
24
25#[derive(Clone, Copy, Debug, Eq, PartialEq)]
26pub enum RestartError {
27 NotStalled,
28 NotRestartable,
29 StateRejected,
30}
31
32#[derive(Clone, Debug, PartialEq)]
33pub struct PreparedSteer {
34 job: u64,
35 generation: u64,
36 values: Vec<BoxValue>,
37}
38
39impl PreparedSteer {
40 pub fn values(&self) -> &[BoxValue] {
41 &self.values
42 }
43}
44
45struct ActiveTurn {
46 job: u64,
47 initial_frontier: Option<BoxId>,
48 accepted_call_wave: bool,
49}
50
51struct PendingSteer {
52 job: u64,
53 generation: u64,
54 frontier: BoxId,
55 values: Vec<BoxValue>,
56}
57
58struct Stall {
59 message: String,
60 restartable: bool,
61}
62
63pub struct ConversationState {
64 state: ActorState,
65 submitted: Option<BoxId>,
66 active: Option<ActiveTurn>,
67 pending_steer: Option<PendingSteer>,
68 next_steer_generation: u64,
69 stall: Option<Stall>,
70}
71
72impl ConversationState {
73 pub fn new() -> Self {
74 Self::from_state(ActorState::new(false))
75 }
76
77 pub fn recover(boxes: Vec<ChatBox>, force: bool) -> Result<Self, String> {
78 let state = ActorState::recover(boxes, force).map_err(state_error)?;
79 Ok(Self::from_state(state))
80 }
81
82 pub fn boxes(&self) -> &[ChatBox] {
83 self.state.boxes()
84 }
85
86 pub fn status(&self) -> Status {
87 if let Some(stall) = &self.stall {
88 return Status::Stalled {
89 message: stall.message.clone(),
90 restartable: stall.restartable,
91 };
92 }
93 if self.active.is_some() {
94 Status::Running
95 } else {
96 Status::Quiet
97 }
98 }
99
100 pub fn accept(
101 &mut self,
102 box_type: String,
103 contents: String,
104 hidden_type: String,
105 hidden_contents: String,
106 ) -> Result<(), String> {
107 self.state
108 .accept_box(box_type, contents, hidden_type, hidden_contents)
109 .map_err(state_error)
110 }
111
112 pub fn accept_tool_return(&mut self, result: ToolResult) -> Result<(), String> {
113 self.state.accept_async_return(result).map_err(state_error)
114 }
115
116 pub fn begin(&mut self) -> Result<Option<Start>, String> {
117 let Some(start) = self.state.begin_inference().map_err(state_error)? else {
118 return Ok(None);
119 };
120 self.state
121 .flush_active_arrivals(start.job)
122 .map_err(state_error)?;
123 let initial_frontier = self.state.boxes().last().map(ChatBox::id);
124 let boxes = self
125 .state
126 .boxes()
127 .iter()
128 .filter(|value| self.submitted.is_none_or(|id| value.id() > id))
129 .map(project)
130 .collect();
131 self.active = Some(ActiveTurn {
132 job: start.job,
133 initial_frontier,
134 accepted_call_wave: false,
135 });
136 self.pending_steer = None;
137 Ok(Some(Start {
138 job: start.job,
139 boxes,
140 }))
141 }
142
143 pub fn prepare_stage(
144 &mut self,
145 job: u64,
146 text: String,
147 values: Vec<BoxValue>,
148 ) -> Result<Vec<DispatchedToolCall>, String> {
149 self.require_job(job, "stage")?;
150 if self.pending_steer.is_some() {
151 return Err("cannot append a stage while a steer is pending".into());
152 }
153 let calls = values
154 .into_iter()
155 .map(|value| match value {
156 BoxValue::Call(Ok(call)) => Ok(call),
157 BoxValue::Call(Err(error)) => Err(error),
158 BoxValue::History(_) => Err("stage contains a non-call value".into()),
159 })
160 .collect::<Result<Vec<ProviderCall>, String>>()?;
161 let dispatched = self
162 .state
163 .append_stage(job, text, calls)
164 .map_err(state_error)?;
165 self.submitted = self.state.boxes().last().map(ChatBox::id);
166 if !dispatched.is_empty() {
167 self.active
168 .as_mut()
169 .expect("validated active turn")
170 .accepted_call_wave = true;
171 }
172 Ok(dispatched)
173 }
174
175 pub fn prepare_steer(&mut self, job: u64) -> Result<Option<PreparedSteer>, String> {
176 self.require_job(job, "steer")?;
177 if let Some(pending) = &self.pending_steer {
178 return Ok(Some(pending.token()));
179 }
180 let boxes = self.state.flush_active_arrivals(job).map_err(state_error)?;
181 let Some(frontier) = boxes.last().map(ChatBox::id) else {
182 return Ok(None);
183 };
184 let generation = self
185 .next_steer_generation
186 .checked_add(1)
187 .ok_or_else(|| "steer generation exhausted".to_string())?;
188 self.next_steer_generation = generation;
189 let values = boxes.iter().map(project).collect();
190 let pending = PendingSteer {
191 job,
192 generation,
193 frontier,
194 values,
195 };
196 let token = pending.token();
197 self.pending_steer = Some(pending);
198 Ok(Some(token))
199 }
200
201 pub fn validate_steer(&self, token: &PreparedSteer) -> Result<(), String> {
202 self.require_job(token.job, "steer")?;
203 match &self.pending_steer {
204 Some(pending) if pending.job == token.job && pending.generation == token.generation => {
205 Ok(())
206 }
207 _ => Err("stale Codex inference steer".into()),
208 }
209 }
210
211 pub fn commit_steer(&mut self, token: PreparedSteer) -> Result<(), String> {
212 self.validate_steer(&token)?;
213 let pending = self.pending_steer.take().expect("validated pending steer");
214 self.submitted = Some(pending.frontier);
215 Ok(())
216 }
217
218 pub fn complete(&mut self, job: u64, output: ShimOutput<BoxValue>) -> Result<(), String> {
219 self.require_job(job, "completion")?;
220 if self.pending_steer.is_some() {
221 return Err("cannot complete while a steer is pending".into());
222 }
223 let text = terminal_text(output)?;
224 let before = self.state.boxes().last().map(ChatBox::id);
225 let initial = self
226 .active
227 .as_ref()
228 .expect("validated active turn")
229 .initial_frontier;
230 self.state
231 .complete_inference(job, text.clone())
232 .map_err(state_error)?;
233 self.submitted = max_box(self.submitted, initial);
234 if !text.is_empty() {
235 let terminal = before.map_or(1, |id| id.get().saturating_add(1));
236 self.submitted = max_box(self.submitted, Some(BoxId::new(terminal)));
237 }
238 self.active = None;
239 Ok(())
240 }
241
242 pub fn fail(&mut self, job: u64, message: String, restartable: bool) {
243 let Some(active) = self.active.as_ref() else {
244 return;
245 };
246 if active.job != job {
247 return;
248 }
249 let restartable = restartable && !active.accepted_call_wave;
250 let _ = self.state.flush_active_arrivals(job);
251 if self.state.stall_inference(job, message.clone()).is_err() {
252 return;
253 }
254 self.active = None;
255 self.pending_steer = None;
256 self.stall = Some(Stall {
257 message,
258 restartable,
259 });
260 }
261
262 pub fn restart(&mut self) -> Result<(), RestartError> {
263 let Some(stall) = &self.stall else {
264 return Err(RestartError::NotStalled);
265 };
266 if !stall.restartable {
267 return Err(RestartError::NotRestartable);
268 }
269 self.state.take_halt();
270 self.state
271 .restart()
272 .map_err(|_| RestartError::StateRejected)?;
273 self.stall = None;
274 Ok(())
275 }
276
277 fn from_state(state: ActorState) -> Self {
278 Self {
279 state,
280 submitted: None,
281 active: None,
282 pending_steer: None,
283 next_steer_generation: 0,
284 stall: None,
285 }
286 }
287
288 fn require_job(&self, job: u64, operation: &str) -> Result<(), String> {
289 if self.active.as_ref().is_some_and(|active| active.job == job) {
290 Ok(())
291 } else {
292 Err(format!("stale Codex inference {operation}"))
293 }
294 }
295}
296
297impl Default for ConversationState {
298 fn default() -> Self {
299 Self::new()
300 }
301}
302
303impl PendingSteer {
304 fn token(&self) -> PreparedSteer {
305 PreparedSteer {
306 job: self.job,
307 generation: self.generation,
308 values: self.values.clone(),
309 }
310 }
311}
312
313fn terminal_text(output: ShimOutput<BoxValue>) -> Result<String, String> {
314 match output.items.as_slice() {
315 [] => Ok(String::new()),
316 [ShimItem::Text(text)] => Ok(text.clone()),
317 _ => Err("Codex completion must contain zero items or one text item".into()),
318 }
319}
320
321fn max_box(left: Option<BoxId>, right: Option<BoxId>) -> Option<BoxId> {
322 match (left, right) {
323 (Some(left), Some(right)) => Some(left.max(right)),
324 (left, right) => left.or(right),
325 }
326}
327
328fn state_error(error: impl std::fmt::Debug) -> String {
329 format!("chat state error: {error:?}")
330}