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