1use std::{collections::VecDeque, future::Future, pin::Pin};
7
8pub use kcode_k1_codex_runtime::{
9 Adapter, Config, DynamicTool, Error, ErrorKind, Event, ToolCall, ToolResult, Turn,
10};
11
12pub const ASYNC_TOOL_ACKNOWLEDGEMENT: &str = "The tool was launched asynchronously. Its result is not included in this acknowledgement. Continue without waiting or polling; available results will be provided at a later inference boundary, which may be within this same turn.";
14
15pub type ToolLaunchFuture<'a> = Pin<Box<dyn Future<Output = Result<(), String>> + Send + 'a>>;
17
18pub trait ToolCallLauncher<B>: Send {
20 fn launch_stage<'a>(&'a mut self, text: String, boxes: Vec<B>) -> ToolLaunchFuture<'a>;
26}
27
28pub trait BoxCodec {
30 type Box: Clone;
32
33 fn tool_call_box(&mut self, call: &ToolCall) -> Self::Box;
37
38 fn box_text<'a>(&self, box_: &'a Self::Box) -> &'a str;
40}
41
42#[derive(Clone, Debug, PartialEq)]
44pub enum ShimItem<B> {
45 Text(String),
47 Box(B),
49}
50
51#[derive(Clone, Debug, PartialEq)]
53pub struct ShimOutput<B> {
54 pub items: Vec<ShimItem<B>>,
56}
57
58#[derive(Clone, Copy, Debug, PartialEq, Eq)]
59enum Health {
60 Ready,
61 Unusable,
62}
63
64pub struct Shim<C: BoxCodec> {
68 adapter: Adapter,
69 conversation_key: String,
70 codec: C,
71 launcher: Box<dyn ToolCallLauncher<C::Box>>,
72 health: Health,
73 pending_boxes: VecDeque<C::Box>,
74}
75
76impl<C: BoxCodec> Shim<C> {
77 pub fn new(
79 adapter: Adapter,
80 conversation_key: impl Into<String>,
81 codec: C,
82 launcher: Box<dyn ToolCallLauncher<C::Box>>,
83 ) -> Self {
84 Self {
85 adapter,
86 conversation_key: conversation_key.into(),
87 codec,
88 launcher,
89 health: Health::Ready,
90 pending_boxes: VecDeque::new(),
91 }
92 }
93
94 pub fn record_box(&mut self, box_: C::Box) {
96 self.pending_boxes.push_back(box_);
97 }
98
99 pub fn record_boxes(&mut self, boxes: impl IntoIterator<Item = C::Box>) {
101 self.pending_boxes.extend(boxes);
102 }
103
104 pub fn pending_box_count(&self) -> usize {
106 self.pending_boxes.len()
107 }
108
109 pub async fn close_conversation(&mut self) -> Result<(), Error> {
113 if self.health == Health::Unusable {
114 return Err(self.unusable());
115 }
116 self.adapter
117 .close_conversation(self.conversation_key.clone())
118 .await
119 }
120
121 pub async fn infer(&mut self, input: impl Into<String>) -> Result<ShimOutput<C::Box>, Error> {
126 if self.health == Health::Unusable {
127 return Err(self.unusable());
128 }
129
130 let submitted_box_count = self.pending_boxes.len();
131 let input = append_section(self.render_pending_boxes(), &input.into());
132
133 self.health = Health::Unusable;
134 let mut turn = match self
135 .adapter
136 .start_turn(self.conversation_key.clone(), input)
137 .await
138 {
139 Ok(turn) => turn,
140 Err(error) => {
141 self.health = Health::Ready;
142 return Err(error);
143 }
144 };
145
146 for _ in 0..submitted_box_count {
147 debug_assert!(self.pending_boxes.pop_front().is_some());
148 }
149
150 let diagnostics = self.adapter.clone();
151 let result = {
152 let mut turn = ConvertedTurn {
153 turn: &mut turn,
154 codec: &mut self.codec,
155 };
156 drive_turn(self.launcher.as_mut(), &mut turn, || {
157 diagnostics.diagnostics()
158 })
159 .await
160 };
161 if result.is_ok() {
162 self.health = Health::Ready;
163 }
164 result
165 }
166
167 fn render_pending_boxes(&self) -> String {
168 let mut output = String::new();
169 for box_ in &self.pending_boxes {
170 output = append_section(output, self.codec.box_text(box_));
171 }
172 output
173 }
174
175 fn unusable(&self) -> Error {
176 self.error("Codex shim cannot be reused after an active turn failed or was cancelled")
177 }
178
179 fn error(&self, message: impl Into<String>) -> Error {
180 Error {
181 kind: ErrorKind::Unavailable,
182 message: message.into(),
183 diagnostics: self.adapter.diagnostics(),
184 }
185 }
186}
187
188enum ActiveEvent<B> {
189 TextDelta(String),
190 Call { call_id: String, box_: B },
191 Done,
192 Error(Error),
193}
194
195trait ActiveTurn<B> {
196 async fn next_event(&mut self) -> Option<ActiveEvent<B>>;
197 fn try_next_event(&mut self) -> Result<Option<ActiveEvent<B>>, Error>;
198 async fn respond(&mut self, call_id: String, result: ToolResult) -> Result<(), Error>;
199}
200
201struct ConvertedTurn<'a, C> {
202 turn: &'a mut Turn,
203 codec: &'a mut C,
204}
205
206impl<C: BoxCodec> ConvertedTurn<'_, C> {
207 fn convert(&mut self, event: Event) -> ActiveEvent<C::Box> {
208 match event {
209 Event::TextDelta(delta) => ActiveEvent::TextDelta(delta),
210 Event::ToolCall(call) => ActiveEvent::Call {
211 box_: self.codec.tool_call_box(&call),
212 call_id: call.call_id,
213 },
214 Event::Done => ActiveEvent::Done,
215 Event::Error(error) => ActiveEvent::Error(error),
216 }
217 }
218}
219
220impl<C: BoxCodec> ActiveTurn<C::Box> for ConvertedTurn<'_, C> {
221 async fn next_event(&mut self) -> Option<ActiveEvent<C::Box>> {
222 let event = self.turn.next_event().await?;
223 Some(self.convert(event))
224 }
225
226 fn try_next_event(&mut self) -> Result<Option<ActiveEvent<C::Box>>, Error> {
227 let event = self.turn.try_next_event()?;
228 Ok(event.map(|event| self.convert(event)))
229 }
230
231 async fn respond(&mut self, call_id: String, result: ToolResult) -> Result<(), Error> {
232 self.turn.respond(call_id, result).await
233 }
234}
235
236async fn drive_turn<B, T, D>(
237 launcher: &mut dyn ToolCallLauncher<B>,
238 turn: &mut T,
239 diagnostics: D,
240) -> Result<ShimOutput<B>, Error>
241where
242 T: ActiveTurn<B>,
243 D: Fn() -> Vec<u8>,
244{
245 let mut text = String::new();
246 let mut lookahead = None;
247 loop {
248 let event = match lookahead.take() {
249 Some(event) => Some(event),
250 None => turn.next_event().await,
251 };
252 match event {
253 Some(ActiveEvent::TextDelta(delta)) => text.push_str(&delta),
254 Some(ActiveEvent::Call { call_id, box_ }) => {
255 let mut call_ids = vec![call_id];
256 let mut boxes = vec![box_];
257 let mut drain_error = None;
258
259 loop {
260 match turn.try_next_event() {
261 Ok(Some(ActiveEvent::Call { call_id, box_ })) => {
262 call_ids.push(call_id);
263 boxes.push(box_);
264 }
265 Ok(Some(event)) => {
266 lookahead = Some(event);
267 break;
268 }
269 Ok(None) => break,
270 Err(error) => {
271 drain_error = Some(error);
272 break;
273 }
274 }
275 }
276
277 if let Err(message) = launcher
278 .launch_stage(std::mem::take(&mut text), boxes)
279 .await
280 {
281 return Err(Error {
282 kind: ErrorKind::LaunchRejected,
283 message,
284 diagnostics: diagnostics(),
285 });
286 }
287
288 for call_id in call_ids {
289 turn.respond(
290 call_id,
291 ToolResult {
292 success: true,
293 output: ASYNC_TOOL_ACKNOWLEDGEMENT.to_owned(),
294 },
295 )
296 .await?;
297 }
298
299 if let Some(error) = drain_error {
300 return Err(error);
301 }
302 }
303 Some(ActiveEvent::Done) => {
304 let items = if text.is_empty() {
305 Vec::new()
306 } else {
307 vec![ShimItem::Text(text)]
308 };
309 return Ok(ShimOutput { items });
310 }
311 Some(ActiveEvent::Error(error)) => return Err(error),
312 None => {
313 return Err(Error {
314 kind: ErrorKind::Unavailable,
315 message: "Codex app-server closed before the active turn completed".into(),
316 diagnostics: diagnostics(),
317 });
318 }
319 }
320 }
321}
322
323fn append_section(mut output: String, section: &str) -> String {
324 if section.is_empty() {
325 return output;
326 }
327 if !output.is_empty() && !output.ends_with('\n') {
328 output.push('\n');
329 }
330 output.push_str(section);
331 output
332}
333
334#[cfg(test)]
335mod tests {
336 use super::*;
337 use std::sync::{
338 Arc, Mutex,
339 atomic::{AtomicUsize, Ordering},
340 };
341 use std::task::{Context, Poll, Waker};
342
343 type Stages = Arc<Mutex<Vec<(String, Vec<String>)>>>;
344
345 struct RecordingLauncher {
346 stages: Stages,
347 completed: Arc<AtomicUsize>,
348 reject: bool,
349 }
350
351 impl ToolCallLauncher<String> for RecordingLauncher {
352 fn launch_stage<'a>(
353 &'a mut self,
354 text: String,
355 boxes: Vec<String>,
356 ) -> ToolLaunchFuture<'a> {
357 let stages = Arc::clone(&self.stages);
358 let completed = Arc::clone(&self.completed);
359 let reject = self.reject;
360 Box::pin(async move {
361 stages.lock().unwrap().push((text, boxes));
362 if reject {
363 Err("consumer barrier failed".into())
364 } else {
365 completed.fetch_add(1, Ordering::SeqCst);
366 Ok(())
367 }
368 })
369 }
370 }
371
372 struct Acknowledgement {
373 call_id: String,
374 result: ToolResult,
375 completed_stages: usize,
376 }
377
378 struct ScriptedTurn {
379 events: VecDeque<ActiveEvent<String>>,
380 acknowledgements: Arc<Mutex<Vec<Acknowledgement>>>,
381 completed: Arc<AtomicUsize>,
382 }
383
384 impl ActiveTurn<String> for ScriptedTurn {
385 async fn next_event(&mut self) -> Option<ActiveEvent<String>> {
386 self.events.pop_front()
387 }
388
389 fn try_next_event(&mut self) -> Result<Option<ActiveEvent<String>>, Error> {
390 Ok(self.events.pop_front())
391 }
392
393 async fn respond(&mut self, call_id: String, result: ToolResult) -> Result<(), Error> {
394 self.acknowledgements.lock().unwrap().push(Acknowledgement {
395 call_id,
396 result,
397 completed_stages: self.completed.load(Ordering::SeqCst),
398 });
399 Ok(())
400 }
401 }
402
403 #[test]
404 fn grouped_waves_wait_for_callback_and_acknowledge_in_provider_order() {
405 let stages = Arc::new(Mutex::new(Vec::new()));
406 let acknowledgements = Arc::new(Mutex::new(Vec::new()));
407 let completed = Arc::new(AtomicUsize::new(0));
408 let mut launcher = RecordingLauncher {
409 stages: Arc::clone(&stages),
410 completed: Arc::clone(&completed),
411 reject: false,
412 };
413 let mut turn = ScriptedTurn {
414 events: VecDeque::from([
415 ActiveEvent::TextDelta("first stage".into()),
416 ActiveEvent::Call {
417 call_id: "call-1".into(),
418 box_: "box-1".into(),
419 },
420 ActiveEvent::Call {
421 call_id: "call-2".into(),
422 box_: "box-2".into(),
423 },
424 ActiveEvent::TextDelta("second stage".into()),
425 ActiveEvent::Call {
426 call_id: "call-3".into(),
427 box_: "box-3".into(),
428 },
429 ActiveEvent::TextDelta("final text".into()),
430 ActiveEvent::Done,
431 ]),
432 acknowledgements: Arc::clone(&acknowledgements),
433 completed: Arc::clone(&completed),
434 };
435
436 let output = run_ready(drive_turn(&mut launcher, &mut turn, Vec::new)).unwrap();
437
438 assert_eq!(
439 *stages.lock().unwrap(),
440 vec![
441 ("first stage".into(), vec!["box-1".into(), "box-2".into()]),
442 ("second stage".into(), vec!["box-3".into()]),
443 ]
444 );
445 let acknowledgements = acknowledgements.lock().unwrap();
446 assert_eq!(
447 acknowledgements
448 .iter()
449 .map(|ack| ack.call_id.as_str())
450 .collect::<Vec<_>>(),
451 vec!["call-1", "call-2", "call-3"]
452 );
453 assert_eq!(
454 acknowledgements
455 .iter()
456 .map(|ack| ack.completed_stages)
457 .collect::<Vec<_>>(),
458 vec![1, 1, 2]
459 );
460 assert!(acknowledgements.iter().all(|ack| ack.result.success));
461 assert!(
462 acknowledgements
463 .iter()
464 .all(|ack| ack.result.output == ASYNC_TOOL_ACKNOWLEDGEMENT)
465 );
466 assert_eq!(
467 ASYNC_TOOL_ACKNOWLEDGEMENT,
468 "The tool was launched asynchronously. Its result is not included in this acknowledgement. Continue without waiting or polling; available results will be provided at a later inference boundary, which may be within this same turn."
469 );
470 assert_eq!(
471 output,
472 ShimOutput {
473 items: vec![ShimItem::Text("final text".into())]
474 }
475 );
476 }
477
478 #[test]
479 fn callback_failure_acknowledges_nothing() {
480 let stages = Arc::new(Mutex::new(Vec::new()));
481 let acknowledgements = Arc::new(Mutex::new(Vec::new()));
482 let completed = Arc::new(AtomicUsize::new(0));
483 let mut launcher = RecordingLauncher {
484 stages: Arc::clone(&stages),
485 completed: Arc::clone(&completed),
486 reject: true,
487 };
488 let mut turn = ScriptedTurn {
489 events: VecDeque::from([
490 ActiveEvent::TextDelta("accepted text".into()),
491 ActiveEvent::Call {
492 call_id: "call-1".into(),
493 box_: "box-1".into(),
494 },
495 ActiveEvent::Done,
496 ]),
497 acknowledgements: Arc::clone(&acknowledgements),
498 completed,
499 };
500
501 let error = run_ready(drive_turn(&mut launcher, &mut turn, Vec::new)).unwrap_err();
502
503 assert_eq!(error.kind, ErrorKind::LaunchRejected);
504 assert_eq!(error.message, "consumer barrier failed");
505 assert!(acknowledgements.lock().unwrap().is_empty());
506 assert_eq!(
507 *stages.lock().unwrap(),
508 vec![("accepted text".into(), vec!["box-1".into()])]
509 );
510 }
511
512 fn run_ready<F: Future>(future: F) -> F::Output {
513 let mut context = Context::from_waker(Waker::noop());
514 let mut future = Box::pin(future);
515 match future.as_mut().poll(&mut context) {
516 Poll::Ready(output) => output,
517 Poll::Pending => panic!("bounded scripted future unexpectedly pending"),
518 }
519 }
520}