Skip to main content

kcode_k1_chat_codex_state/
lib.rs

1#![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}