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 malformed_tool_call_message(&self, _box_: &Self::Box) -> Option<&'static str> {
45 None
46 }
47
48 fn box_text<'a>(&self, box_: &'a Self::Box) -> &'a str;
50}
51
52#[derive(Clone, Debug, PartialEq)]
54pub enum ShimItem<B> {
55 Text(String),
57 Box(B),
59}
60
61#[derive(Clone, Debug, PartialEq)]
63pub struct ShimOutput<B> {
64 pub items: Vec<ShimItem<B>>,
66}
67
68#[derive(Clone, Copy, Debug, PartialEq, Eq)]
69enum Health {
70 Ready,
71 Unusable,
72}
73
74pub struct Shim<C: BoxCodec> {
78 adapter: Adapter,
79 conversation_key: String,
80 codec: C,
81 launcher: Box<dyn ToolCallLauncher<C::Box>>,
82 health: Health,
83 pending_boxes: VecDeque<C::Box>,
84}
85
86impl<C: BoxCodec> Shim<C> {
87 pub fn new(
89 adapter: Adapter,
90 conversation_key: impl Into<String>,
91 codec: C,
92 launcher: Box<dyn ToolCallLauncher<C::Box>>,
93 ) -> Self {
94 Self {
95 adapter,
96 conversation_key: conversation_key.into(),
97 codec,
98 launcher,
99 health: Health::Ready,
100 pending_boxes: VecDeque::new(),
101 }
102 }
103
104 pub fn record_box(&mut self, box_: C::Box) {
106 self.pending_boxes.push_back(box_);
107 }
108
109 pub fn record_boxes(&mut self, boxes: impl IntoIterator<Item = C::Box>) {
111 self.pending_boxes.extend(boxes);
112 }
113
114 pub fn pending_box_count(&self) -> usize {
116 self.pending_boxes.len()
117 }
118
119 pub async fn close_conversation(&mut self) -> Result<(), Error> {
123 if self.health == Health::Unusable {
124 return Err(self.unusable());
125 }
126 self.adapter
127 .close_conversation(self.conversation_key.clone())
128 .await
129 }
130
131 pub async fn infer(&mut self, input: impl Into<String>) -> Result<ShimOutput<C::Box>, Error> {
136 if self.health == Health::Unusable {
137 return Err(self.unusable());
138 }
139
140 let submitted_box_count = self.pending_boxes.len();
141 let input = append_section(self.render_pending_boxes(), &input.into());
142
143 self.health = Health::Unusable;
144 let mut turn = match self
145 .adapter
146 .start_turn(self.conversation_key.clone(), input)
147 .await
148 {
149 Ok(turn) => turn,
150 Err(error) => {
151 self.health = Health::Ready;
152 return Err(error);
153 }
154 };
155
156 for _ in 0..submitted_box_count {
157 debug_assert!(self.pending_boxes.pop_front().is_some());
158 }
159
160 let diagnostics = self.adapter.clone();
161 let result = {
162 let mut turn = ConvertedTurn {
163 turn: &mut turn,
164 codec: &mut self.codec,
165 };
166 drive_turn(self.launcher.as_mut(), &mut turn, || {
167 diagnostics.diagnostics()
168 })
169 .await
170 };
171 if result.is_ok() {
172 self.health = Health::Ready;
173 }
174 result
175 }
176
177 fn render_pending_boxes(&self) -> String {
178 let mut output = String::new();
179 for box_ in &self.pending_boxes {
180 output = append_section(output, self.codec.box_text(box_));
181 }
182 output
183 }
184
185 fn unusable(&self) -> Error {
186 self.error("Codex shim cannot be reused after an active turn failed or was cancelled")
187 }
188
189 fn error(&self, message: impl Into<String>) -> Error {
190 Error {
191 kind: ErrorKind::Unavailable,
192 message: message.into(),
193 diagnostics: self.adapter.diagnostics(),
194 }
195 }
196}
197
198enum ActiveEvent<B> {
199 TextDelta(String),
200 Call {
201 call_id: String,
202 box_: B,
203 malformed_message: Option<String>,
204 },
205 Done,
206 Error(Error),
207}
208
209trait ActiveTurn<B> {
210 async fn next_event(&mut self) -> Option<ActiveEvent<B>>;
211 fn try_next_event(&mut self) -> Result<Option<ActiveEvent<B>>, Error>;
212 async fn respond(&mut self, call_id: String, result: ToolResult) -> Result<(), Error>;
213}
214
215struct ConvertedTurn<'a, C> {
216 turn: &'a mut Turn,
217 codec: &'a mut C,
218}
219
220impl<C: BoxCodec> ConvertedTurn<'_, C> {
221 fn convert(&mut self, event: Event) -> ActiveEvent<C::Box> {
222 match event {
223 Event::TextDelta(delta) => ActiveEvent::TextDelta(delta),
224 Event::ToolCall(call) => {
225 let box_ = self.codec.tool_call_box(&call);
226 let malformed_message = self
227 .codec
228 .malformed_tool_call_message(&box_)
229 .map(str::to_owned);
230 ActiveEvent::Call {
231 call_id: call.call_id,
232 box_,
233 malformed_message,
234 }
235 }
236 Event::Done => ActiveEvent::Done,
237 Event::Error(error) => ActiveEvent::Error(error),
238 }
239 }
240}
241
242impl<C: BoxCodec> ActiveTurn<C::Box> for ConvertedTurn<'_, C> {
243 async fn next_event(&mut self) -> Option<ActiveEvent<C::Box>> {
244 let event = self.turn.next_event().await?;
245 Some(self.convert(event))
246 }
247
248 fn try_next_event(&mut self) -> Result<Option<ActiveEvent<C::Box>>, Error> {
249 let event = self.turn.try_next_event()?;
250 Ok(event.map(|event| self.convert(event)))
251 }
252
253 async fn respond(&mut self, call_id: String, result: ToolResult) -> Result<(), Error> {
254 self.turn.respond(call_id, result).await
255 }
256}
257
258async fn drive_turn<B, T, D>(
259 launcher: &mut dyn ToolCallLauncher<B>,
260 turn: &mut T,
261 diagnostics: D,
262) -> Result<ShimOutput<B>, Error>
263where
264 T: ActiveTurn<B>,
265 D: Fn() -> Vec<u8>,
266{
267 let mut text = String::new();
268 let mut lookahead = None;
269 loop {
270 let event = match lookahead.take() {
271 Some(event) => Some(event),
272 None => turn.next_event().await,
273 };
274 match event {
275 Some(ActiveEvent::TextDelta(delta)) => text.push_str(&delta),
276 Some(ActiveEvent::Call {
277 call_id,
278 box_,
279 malformed_message,
280 }) => {
281 let mut calls = vec![(call_id, malformed_message)];
282 let mut boxes = Vec::new();
283 if calls[0].1.is_none() {
284 boxes.push(box_);
285 }
286 let mut drain_error = None;
287
288 loop {
289 match turn.try_next_event() {
290 Ok(Some(ActiveEvent::Call {
291 call_id,
292 box_,
293 malformed_message,
294 })) => {
295 if malformed_message.is_none() {
296 boxes.push(box_);
297 }
298 calls.push((call_id, malformed_message));
299 }
300 Ok(Some(event)) => {
301 lookahead = Some(event);
302 break;
303 }
304 Ok(None) => break,
305 Err(error) => {
306 drain_error = Some(error);
307 break;
308 }
309 }
310 }
311
312 if let Err(message) = launcher
313 .launch_stage(std::mem::take(&mut text), boxes)
314 .await
315 {
316 return Err(Error {
317 kind: ErrorKind::LaunchRejected,
318 message,
319 diagnostics: diagnostics(),
320 });
321 }
322
323 for (call_id, malformed_message) in calls {
324 let result = match malformed_message {
325 Some(output) => ToolResult {
326 success: false,
327 output,
328 },
329 None => ToolResult {
330 success: true,
331 output: ASYNC_TOOL_ACKNOWLEDGEMENT.to_owned(),
332 },
333 };
334 turn.respond(call_id, result).await?;
335 }
336
337 if let Some(error) = drain_error {
338 return Err(error);
339 }
340 }
341 Some(ActiveEvent::Done) => {
342 let items = if text.is_empty() {
343 Vec::new()
344 } else {
345 vec![ShimItem::Text(text)]
346 };
347 return Ok(ShimOutput { items });
348 }
349 Some(ActiveEvent::Error(error)) => return Err(error),
350 None => {
351 return Err(Error {
352 kind: ErrorKind::Unavailable,
353 message: "Codex app-server closed before the active turn completed".into(),
354 diagnostics: diagnostics(),
355 });
356 }
357 }
358 }
359}
360
361fn append_section(mut output: String, section: &str) -> String {
362 if section.is_empty() {
363 return output;
364 }
365 if !output.is_empty() && !output.ends_with('\n') {
366 output.push('\n');
367 }
368 output.push_str(section);
369 output
370}
371
372#[cfg(test)]
373mod tests {
374 use super::*;
375 use std::sync::{
376 Arc, Mutex,
377 atomic::{AtomicUsize, Ordering},
378 };
379 use std::task::{Context, Poll, Waker};
380
381 const MALFORMED_MESSAGE: &str = "The tool call did not match the required schema.";
382
383 type Stages = Arc<Mutex<Vec<(String, Vec<String>)>>>;
384
385 struct RecordingLauncher {
386 stages: Stages,
387 completed: Arc<AtomicUsize>,
388 reject: bool,
389 }
390
391 impl ToolCallLauncher<String> for RecordingLauncher {
392 fn launch_stage<'a>(
393 &'a mut self,
394 text: String,
395 boxes: Vec<String>,
396 ) -> ToolLaunchFuture<'a> {
397 let stages = Arc::clone(&self.stages);
398 let completed = Arc::clone(&self.completed);
399 let reject = self.reject;
400 Box::pin(async move {
401 stages.lock().unwrap().push((text, boxes));
402 if reject {
403 Err("consumer barrier failed".into())
404 } else {
405 completed.fetch_add(1, Ordering::SeqCst);
406 Ok(())
407 }
408 })
409 }
410 }
411
412 struct Acknowledgement {
413 call_id: String,
414 result: ToolResult,
415 completed_stages: usize,
416 }
417
418 struct ScriptedTurn {
419 events: VecDeque<ActiveEvent<String>>,
420 acknowledgements: Arc<Mutex<Vec<Acknowledgement>>>,
421 completed: Arc<AtomicUsize>,
422 }
423
424 impl ActiveTurn<String> for ScriptedTurn {
425 async fn next_event(&mut self) -> Option<ActiveEvent<String>> {
426 self.events.pop_front()
427 }
428
429 fn try_next_event(&mut self) -> Result<Option<ActiveEvent<String>>, Error> {
430 Ok(self.events.pop_front())
431 }
432
433 async fn respond(&mut self, call_id: String, result: ToolResult) -> Result<(), Error> {
434 self.acknowledgements.lock().unwrap().push(Acknowledgement {
435 call_id,
436 result,
437 completed_stages: self.completed.load(Ordering::SeqCst),
438 });
439 Ok(())
440 }
441 }
442
443 #[test]
444 fn grouped_valid_waves_wait_for_callback_and_acknowledge_in_provider_order() {
445 let stages = Arc::new(Mutex::new(Vec::new()));
446 let acknowledgements = Arc::new(Mutex::new(Vec::new()));
447 let completed = Arc::new(AtomicUsize::new(0));
448 let mut launcher = RecordingLauncher {
449 stages: Arc::clone(&stages),
450 completed: Arc::clone(&completed),
451 reject: false,
452 };
453 let mut turn = ScriptedTurn {
454 events: VecDeque::from([
455 ActiveEvent::TextDelta("first stage".into()),
456 ActiveEvent::Call {
457 call_id: "call-1".into(),
458 box_: "box-1".into(),
459 malformed_message: None,
460 },
461 ActiveEvent::Call {
462 call_id: "call-2".into(),
463 box_: "box-2".into(),
464 malformed_message: None,
465 },
466 ActiveEvent::TextDelta("second stage".into()),
467 ActiveEvent::Call {
468 call_id: "call-3".into(),
469 box_: "box-3".into(),
470 malformed_message: None,
471 },
472 ActiveEvent::TextDelta("final text".into()),
473 ActiveEvent::Done,
474 ]),
475 acknowledgements: Arc::clone(&acknowledgements),
476 completed: Arc::clone(&completed),
477 };
478
479 let output = run_ready(drive_turn(&mut launcher, &mut turn, Vec::new)).unwrap();
480
481 assert_eq!(
482 *stages.lock().unwrap(),
483 vec![
484 ("first stage".into(), vec!["box-1".into(), "box-2".into()]),
485 ("second stage".into(), vec!["box-3".into()]),
486 ]
487 );
488 let acknowledgements = acknowledgements.lock().unwrap();
489 assert_eq!(
490 acknowledgements
491 .iter()
492 .map(|ack| ack.call_id.as_str())
493 .collect::<Vec<_>>(),
494 vec!["call-1", "call-2", "call-3"]
495 );
496 assert_eq!(
497 acknowledgements
498 .iter()
499 .map(|ack| ack.completed_stages)
500 .collect::<Vec<_>>(),
501 vec![1, 1, 2]
502 );
503 assert!(acknowledgements.iter().all(|ack| ack.result.success));
504 assert!(
505 acknowledgements
506 .iter()
507 .all(|ack| ack.result.output == ASYNC_TOOL_ACKNOWLEDGEMENT)
508 );
509 assert_eq!(
510 ASYNC_TOOL_ACKNOWLEDGEMENT,
511 "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."
512 );
513 assert_eq!(
514 output,
515 ShimOutput {
516 items: vec![ShimItem::Text("final text".into())]
517 }
518 );
519 }
520
521 #[test]
522 fn malformed_only_wave_retains_text_launches_no_boxes_and_continues() {
523 let stages = Arc::new(Mutex::new(Vec::new()));
524 let acknowledgements = Arc::new(Mutex::new(Vec::new()));
525 let completed = Arc::new(AtomicUsize::new(0));
526 let mut launcher = RecordingLauncher {
527 stages: Arc::clone(&stages),
528 completed: Arc::clone(&completed),
529 reject: false,
530 };
531 let mut turn = ScriptedTurn {
532 events: VecDeque::from([
533 ActiveEvent::TextDelta("malformed stage".into()),
534 ActiveEvent::Call {
535 call_id: "bad-call".into(),
536 box_: "bad-box".into(),
537 malformed_message: Some(MALFORMED_MESSAGE.into()),
538 },
539 ActiveEvent::TextDelta("terminal text".into()),
540 ActiveEvent::Done,
541 ]),
542 acknowledgements: Arc::clone(&acknowledgements),
543 completed: Arc::clone(&completed),
544 };
545
546 let output = run_ready(drive_turn(&mut launcher, &mut turn, Vec::new)).unwrap();
547
548 assert_eq!(
549 *stages.lock().unwrap(),
550 vec![("malformed stage".into(), Vec::new())]
551 );
552 let acknowledgements = acknowledgements.lock().unwrap();
553 assert_eq!(acknowledgements.len(), 1);
554 assert_eq!(acknowledgements[0].call_id, "bad-call");
555 assert!(!acknowledgements[0].result.success);
556 assert_eq!(acknowledgements[0].result.output, MALFORMED_MESSAGE);
557 assert_eq!(acknowledgements[0].completed_stages, 1);
558 assert_eq!(
559 output,
560 ShimOutput {
561 items: vec![ShimItem::Text("terminal text".into())]
562 }
563 );
564 }
565
566 #[test]
567 fn mixed_wave_launches_only_valid_boxes_and_responds_in_provider_order() {
568 let stages = Arc::new(Mutex::new(Vec::new()));
569 let acknowledgements = Arc::new(Mutex::new(Vec::new()));
570 let completed = Arc::new(AtomicUsize::new(0));
571 let mut launcher = RecordingLauncher {
572 stages: Arc::clone(&stages),
573 completed: Arc::clone(&completed),
574 reject: false,
575 };
576 let mut turn = ScriptedTurn {
577 events: VecDeque::from([
578 ActiveEvent::TextDelta("mixed stage".into()),
579 ActiveEvent::Call {
580 call_id: "valid-1".into(),
581 box_: "box-1".into(),
582 malformed_message: None,
583 },
584 ActiveEvent::Call {
585 call_id: "malformed-2".into(),
586 box_: "bad-box".into(),
587 malformed_message: Some(MALFORMED_MESSAGE.into()),
588 },
589 ActiveEvent::Call {
590 call_id: "valid-3".into(),
591 box_: "box-3".into(),
592 malformed_message: None,
593 },
594 ActiveEvent::TextDelta("continued terminal text".into()),
595 ActiveEvent::Done,
596 ]),
597 acknowledgements: Arc::clone(&acknowledgements),
598 completed: Arc::clone(&completed),
599 };
600
601 let output = run_ready(drive_turn(&mut launcher, &mut turn, Vec::new)).unwrap();
602
603 assert_eq!(
604 *stages.lock().unwrap(),
605 vec![("mixed stage".into(), vec!["box-1".into(), "box-3".into()])]
606 );
607 let acknowledgements = acknowledgements.lock().unwrap();
608 assert_eq!(
609 acknowledgements
610 .iter()
611 .map(|ack| ack.call_id.as_str())
612 .collect::<Vec<_>>(),
613 vec!["valid-1", "malformed-2", "valid-3"]
614 );
615 assert!(acknowledgements[0].result.success);
616 assert_eq!(
617 acknowledgements[0].result.output,
618 ASYNC_TOOL_ACKNOWLEDGEMENT
619 );
620 assert!(!acknowledgements[1].result.success);
621 assert_eq!(acknowledgements[1].result.output, MALFORMED_MESSAGE);
622 assert!(acknowledgements[2].result.success);
623 assert_eq!(
624 acknowledgements[2].result.output,
625 ASYNC_TOOL_ACKNOWLEDGEMENT
626 );
627 assert!(acknowledgements.iter().all(|ack| ack.completed_stages == 1));
628 assert_eq!(
629 output,
630 ShimOutput {
631 items: vec![ShimItem::Text("continued terminal text".into())]
632 }
633 );
634 }
635
636 #[test]
637 fn callback_failure_responds_to_none_of_a_mixed_wave() {
638 let stages = Arc::new(Mutex::new(Vec::new()));
639 let acknowledgements = Arc::new(Mutex::new(Vec::new()));
640 let completed = Arc::new(AtomicUsize::new(0));
641 let mut launcher = RecordingLauncher {
642 stages: Arc::clone(&stages),
643 completed: Arc::clone(&completed),
644 reject: true,
645 };
646 let mut turn = ScriptedTurn {
647 events: VecDeque::from([
648 ActiveEvent::TextDelta("accepted text".into()),
649 ActiveEvent::Call {
650 call_id: "valid-call".into(),
651 box_: "valid-box".into(),
652 malformed_message: None,
653 },
654 ActiveEvent::Call {
655 call_id: "bad-call".into(),
656 box_: "bad-box".into(),
657 malformed_message: Some(MALFORMED_MESSAGE.into()),
658 },
659 ActiveEvent::Done,
660 ]),
661 acknowledgements: Arc::clone(&acknowledgements),
662 completed,
663 };
664
665 let error = run_ready(drive_turn(&mut launcher, &mut turn, Vec::new)).unwrap_err();
666
667 assert_eq!(error.kind, ErrorKind::LaunchRejected);
668 assert_eq!(error.message, "consumer barrier failed");
669 assert!(acknowledgements.lock().unwrap().is_empty());
670 assert_eq!(
671 *stages.lock().unwrap(),
672 vec![("accepted text".into(), vec!["valid-box".into()])]
673 );
674 }
675
676 struct LegacyCodec;
677
678 impl BoxCodec for LegacyCodec {
679 type Box = String;
680
681 fn tool_call_box(&mut self, _call: &ToolCall) -> Self::Box {
682 "legacy box".into()
683 }
684
685 fn box_text<'a>(&self, box_: &'a Self::Box) -> &'a str {
686 box_
687 }
688 }
689
690 #[test]
691 fn codec_without_classifier_override_retains_valid_behavior() {
692 let codec = LegacyCodec;
693 let box_ = "legacy box".to_owned();
694
695 assert_eq!(codec.malformed_tool_call_message(&box_), None);
696 }
697
698 fn run_ready<F: Future>(future: F) -> F::Output {
699 let mut context = Context::from_waker(Waker::noop());
700 let mut future = Box::pin(future);
701 match future.as_mut().poll(&mut context) {
702 Poll::Ready(output) => output,
703 Poll::Pending => panic!("bounded scripted future unexpectedly pending"),
704 }
705 }
706}