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::{ActionId, ChatBox};
3
4use kcode_k1_chat_codex_codec::project;
5use kcode_k1_chat_state::{ActorState, DispatchCall, DispatchOutcome};
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        action_id: ActionId,
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 Launch {
27    pub action_id: ActionId,
28    pub result: Result<String, String>,
29}
30
31#[derive(Clone, Debug, Eq, PartialEq)]
32pub enum Status {
33    Running,
34    Quiet,
35    Stalled { message: String, restartable: bool },
36}
37
38#[derive(Clone, Copy, Debug, Eq, PartialEq)]
39pub enum RestartError {
40    NotStalled,
41    NotRestartable,
42}
43
44#[derive(Clone)]
45struct Launched {
46    action_id: ActionId,
47    call: Call,
48}
49
50struct Round {
51    job: u64,
52    ledger: Vec<Launched>,
53    arrivals: Vec<Arrival>,
54}
55
56enum Mode {
57    Idle,
58    Running(Round),
59    Stalled { message: String, restartable: bool },
60}
61
62pub struct ConversationState {
63    state: ActorState,
64    session: [u8; 12],
65    sequence: u64,
66    submitted: usize,
67    mode: Mode,
68}
69
70impl ConversationState {
71    pub fn new(session: [u8; 12]) -> Self {
72        Self {
73            state: ActorState::new(false),
74            session,
75            sequence: 0,
76            submitted: 0,
77            mode: Mode::Idle,
78        }
79    }
80
81    pub fn boxes(&self) -> &[ChatBox] {
82        self.state.boxes()
83    }
84
85    pub fn status(&self) -> Status {
86        match &self.mode {
87            Mode::Running(_) => Status::Running,
88            Mode::Idle if self.state.quiet() => Status::Quiet,
89            Mode::Idle => Status::Running,
90            Mode::Stalled {
91                message,
92                restartable,
93            } => Status::Stalled {
94                message: message.clone(),
95                restartable: *restartable,
96            },
97        }
98    }
99
100    pub fn accept(&mut self, arrival: Arrival) -> Result<(), String> {
101        if let Mode::Running(round) = &mut self.mode {
102            round.arrivals.push(arrival);
103            Ok(())
104        } else {
105            self.apply(arrival)
106        }
107    }
108
109    pub fn begin(&mut self) -> Result<Option<Start>, String> {
110        if !matches!(self.mode, Mode::Idle) {
111            return Ok(None);
112        }
113        let Some(start) = self.state.begin_inference().map_err(debug)? else {
114            return Ok(None);
115        };
116        let boxes = self.state.boxes();
117        let projected = boxes[self.submitted..].iter().map(project).collect();
118        self.submitted = boxes.len();
119        self.mode = Mode::Running(Round {
120            job: start.job,
121            ledger: Vec::new(),
122            arrivals: Vec::new(),
123        });
124        Ok(Some(Start {
125            job: start.job,
126            boxes: projected,
127        }))
128    }
129
130    pub fn launch(&mut self, job: u64, call: Call) -> Result<Launch, String> {
131        let Mode::Running(round) = &mut self.mode else {
132            return Err("no K1 inference is running".to_owned());
133        };
134        if round.job != job {
135            return Err("stale K1 inference launch".to_owned());
136        }
137        let sequence = self
138            .sequence
139            .checked_add(1)
140            .ok_or_else(|| "K1 ActionId space was exhausted".to_owned())?;
141        let action_id = ActionId::new(self.session, sequence);
142        self.state
143            .collect_provider_call(job, call.name.clone(), call.arguments.clone())
144            .map_err(debug)?;
145        round.ledger.push(Launched {
146            action_id,
147            call: call.clone(),
148        });
149        self.sequence = sequence;
150        let result = kcode_k1_chat_thread_actions::launch(&call.name, &call.arguments);
151        Ok(Launch { action_id, result })
152    }
153
154    pub fn complete(&mut self, job: u64, output: ShimOutput<BoxValue>) -> Result<(), String> {
155        let round = self.take_round(job)?;
156        let text = match validate(output, &round.ledger) {
157            Ok(text) => text,
158            Err(error) => {
159                self.preserve(round, error.clone(), false);
160                return Err(error);
161            }
162        };
163        if let Err(error) = self.commit(&round, Some(&text)) {
164            let message = format!("failed to commit Codex output: {error}");
165            let _ = self.state.halt(message.clone());
166            self.mode = Mode::Stalled {
167                message: message.clone(),
168                restartable: false,
169            };
170            return Err(message);
171        }
172        self.mode = Mode::Idle;
173        Ok(())
174    }
175
176    pub fn fail(&mut self, job: u64, message: String, restartable_before_launch: bool) {
177        if let Ok(round) = self.take_round(job) {
178            self.preserve(round, message, restartable_before_launch);
179        }
180    }
181
182    pub fn restart(&mut self) -> Result<(), RestartError> {
183        match &self.mode {
184            Mode::Stalled {
185                restartable: true, ..
186            } => {}
187            Mode::Stalled { .. } => return Err(RestartError::NotRestartable),
188            _ => return Err(RestartError::NotStalled),
189        }
190        self.state
191            .restart()
192            .map_err(|_| RestartError::NotRestartable)?;
193        self.submitted = 0;
194        self.mode = Mode::Idle;
195        Ok(())
196    }
197
198    fn take_round(&mut self, job: u64) -> Result<Round, String> {
199        let mode = std::mem::replace(&mut self.mode, Mode::Idle);
200        match mode {
201            Mode::Running(round) if round.job == job => Ok(round),
202            other => {
203                self.mode = other;
204                Err("stale K1 inference completion".to_owned())
205            }
206        }
207    }
208
209    fn commit(&mut self, round: &Round, text: Option<&str>) -> Result<(), String> {
210        if let Some(text) = text {
211            self.state
212                .append_kennedy_text(round.job, text)
213                .map_err(debug)?;
214        }
215        let action_ids = round.ledger.iter().map(|entry| entry.action_id).collect();
216        let dispatches = self
217            .state
218            .complete_provider_output(round.job, action_ids)
219            .map_err(debug)?;
220        if !dispatches_match(&round.ledger, &dispatches) {
221            return Err("committed calls did not match launch ledger".to_owned());
222        }
223        if !dispatches.is_empty() {
224            self.state
225                .complete_dispatch(vec![DispatchOutcome::Pending; dispatches.len()])
226                .map_err(debug)?;
227        }
228        self.apply_all(round.arrivals.clone())
229    }
230
231    fn preserve(&mut self, round: Round, message: String, restartable_before_launch: bool) {
232        let preserved = if round.ledger.is_empty() {
233            self.state
234                .stall_inference(round.job, message.clone())
235                .map_err(debug)
236                .and_then(|_| self.apply_all(round.arrivals))
237                .is_ok()
238        } else {
239            let preserved = self.commit(&round, None).is_ok();
240            let _ = self.state.halt(message.clone());
241            preserved
242        };
243        self.mode = Mode::Stalled {
244            message,
245            restartable: round.ledger.is_empty() && restartable_before_launch && preserved,
246        };
247    }
248
249    fn apply_all(&mut self, arrivals: Vec<Arrival>) -> Result<(), String> {
250        for arrival in arrivals {
251            self.apply(arrival)?;
252        }
253        Ok(())
254    }
255
256    fn apply(&mut self, arrival: Arrival) -> Result<(), String> {
257        match arrival {
258            Arrival::System(text) => self.state.accept_system(text).map_err(debug),
259            Arrival::User(text) => self.state.accept_user(text).map_err(debug),
260            Arrival::Attachment => self.state.accept_attachment().map_err(debug),
261            Arrival::Return { action_id, result } => self
262                .state
263                .accept_async_return(action_id, result)
264                .map_err(debug),
265        }
266    }
267}
268
269fn validate(output: ShimOutput<BoxValue>, ledger: &[Launched]) -> Result<String, String> {
270    let mut text = String::new();
271    let mut calls = Vec::new();
272    for item in output.items {
273        match item {
274            ShimItem::Text(value) => text.push_str(&value),
275            ShimItem::Box(BoxValue::Call(Ok(call))) => calls.push(call),
276            ShimItem::Box(_) => return Err("shim returned an invalid call box".to_owned()),
277        }
278    }
279    if calls
280        != ledger
281            .iter()
282            .map(|entry| entry.call.clone())
283            .collect::<Vec<_>>()
284    {
285        return Err("shim output did not match the launch ledger".to_string());
286    }
287    Ok(text)
288}
289
290fn dispatches_match(ledger: &[Launched], dispatches: &[DispatchCall]) -> bool {
291    ledger.len() == dispatches.len()
292        && ledger.iter().zip(dispatches).all(|(expected, actual)| {
293            expected.action_id == actual.action_id
294                && expected.call.name == actual.name
295                && expected.call.arguments == actual.arguments
296        })
297}
298
299fn debug(error: impl std::fmt::Debug) -> String {
300    format!("{error:?}")
301}