use std::{collections::VecDeque, future::Future, pin::Pin};
pub use kcode_k1_codex_runtime::{
Adapter, Config, DynamicTool, Error, ErrorKind, Event, ToolCall, ToolResult, Turn,
};
pub 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.";
pub type ToolLaunchFuture<'a> = Pin<Box<dyn Future<Output = Result<(), String>> + Send + 'a>>;
pub trait ToolCallLauncher<B>: Send {
fn launch_stage<'a>(&'a mut self, text: String, boxes: Vec<B>) -> ToolLaunchFuture<'a>;
}
pub trait BoxCodec {
type Box: Clone;
fn tool_call_box(&mut self, call: &ToolCall) -> Self::Box;
fn box_text<'a>(&self, box_: &'a Self::Box) -> &'a str;
}
#[derive(Clone, Debug, PartialEq)]
pub enum ShimItem<B> {
Text(String),
Box(B),
}
#[derive(Clone, Debug, PartialEq)]
pub struct ShimOutput<B> {
pub items: Vec<ShimItem<B>>,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum Health {
Ready,
Unusable,
}
pub struct Shim<C: BoxCodec> {
adapter: Adapter,
conversation_key: String,
codec: C,
launcher: Box<dyn ToolCallLauncher<C::Box>>,
health: Health,
pending_boxes: VecDeque<C::Box>,
}
impl<C: BoxCodec> Shim<C> {
pub fn new(
adapter: Adapter,
conversation_key: impl Into<String>,
codec: C,
launcher: Box<dyn ToolCallLauncher<C::Box>>,
) -> Self {
Self {
adapter,
conversation_key: conversation_key.into(),
codec,
launcher,
health: Health::Ready,
pending_boxes: VecDeque::new(),
}
}
pub fn record_box(&mut self, box_: C::Box) {
self.pending_boxes.push_back(box_);
}
pub fn record_boxes(&mut self, boxes: impl IntoIterator<Item = C::Box>) {
self.pending_boxes.extend(boxes);
}
pub fn pending_box_count(&self) -> usize {
self.pending_boxes.len()
}
pub async fn close_conversation(&mut self) -> Result<(), Error> {
if self.health == Health::Unusable {
return Err(self.unusable());
}
self.adapter
.close_conversation(self.conversation_key.clone())
.await
}
pub async fn infer(&mut self, input: impl Into<String>) -> Result<ShimOutput<C::Box>, Error> {
if self.health == Health::Unusable {
return Err(self.unusable());
}
let submitted_box_count = self.pending_boxes.len();
let input = append_section(self.render_pending_boxes(), &input.into());
self.health = Health::Unusable;
let mut turn = match self
.adapter
.start_turn(self.conversation_key.clone(), input)
.await
{
Ok(turn) => turn,
Err(error) => {
self.health = Health::Ready;
return Err(error);
}
};
for _ in 0..submitted_box_count {
debug_assert!(self.pending_boxes.pop_front().is_some());
}
let diagnostics = self.adapter.clone();
let result = {
let mut turn = ConvertedTurn {
turn: &mut turn,
codec: &mut self.codec,
};
drive_turn(self.launcher.as_mut(), &mut turn, || {
diagnostics.diagnostics()
})
.await
};
if result.is_ok() {
self.health = Health::Ready;
}
result
}
fn render_pending_boxes(&self) -> String {
let mut output = String::new();
for box_ in &self.pending_boxes {
output = append_section(output, self.codec.box_text(box_));
}
output
}
fn unusable(&self) -> Error {
self.error("Codex shim cannot be reused after an active turn failed or was cancelled")
}
fn error(&self, message: impl Into<String>) -> Error {
Error {
kind: ErrorKind::Unavailable,
message: message.into(),
diagnostics: self.adapter.diagnostics(),
}
}
}
enum ActiveEvent<B> {
TextDelta(String),
Call { call_id: String, box_: B },
Done,
Error(Error),
}
trait ActiveTurn<B> {
async fn next_event(&mut self) -> Option<ActiveEvent<B>>;
fn try_next_event(&mut self) -> Result<Option<ActiveEvent<B>>, Error>;
async fn respond(&mut self, call_id: String, result: ToolResult) -> Result<(), Error>;
}
struct ConvertedTurn<'a, C> {
turn: &'a mut Turn,
codec: &'a mut C,
}
impl<C: BoxCodec> ConvertedTurn<'_, C> {
fn convert(&mut self, event: Event) -> ActiveEvent<C::Box> {
match event {
Event::TextDelta(delta) => ActiveEvent::TextDelta(delta),
Event::ToolCall(call) => ActiveEvent::Call {
box_: self.codec.tool_call_box(&call),
call_id: call.call_id,
},
Event::Done => ActiveEvent::Done,
Event::Error(error) => ActiveEvent::Error(error),
}
}
}
impl<C: BoxCodec> ActiveTurn<C::Box> for ConvertedTurn<'_, C> {
async fn next_event(&mut self) -> Option<ActiveEvent<C::Box>> {
let event = self.turn.next_event().await?;
Some(self.convert(event))
}
fn try_next_event(&mut self) -> Result<Option<ActiveEvent<C::Box>>, Error> {
let event = self.turn.try_next_event()?;
Ok(event.map(|event| self.convert(event)))
}
async fn respond(&mut self, call_id: String, result: ToolResult) -> Result<(), Error> {
self.turn.respond(call_id, result).await
}
}
async fn drive_turn<B, T, D>(
launcher: &mut dyn ToolCallLauncher<B>,
turn: &mut T,
diagnostics: D,
) -> Result<ShimOutput<B>, Error>
where
T: ActiveTurn<B>,
D: Fn() -> Vec<u8>,
{
let mut text = String::new();
let mut lookahead = None;
loop {
let event = match lookahead.take() {
Some(event) => Some(event),
None => turn.next_event().await,
};
match event {
Some(ActiveEvent::TextDelta(delta)) => text.push_str(&delta),
Some(ActiveEvent::Call { call_id, box_ }) => {
let mut call_ids = vec![call_id];
let mut boxes = vec![box_];
let mut drain_error = None;
loop {
match turn.try_next_event() {
Ok(Some(ActiveEvent::Call { call_id, box_ })) => {
call_ids.push(call_id);
boxes.push(box_);
}
Ok(Some(event)) => {
lookahead = Some(event);
break;
}
Ok(None) => break,
Err(error) => {
drain_error = Some(error);
break;
}
}
}
if let Err(message) = launcher
.launch_stage(std::mem::take(&mut text), boxes)
.await
{
return Err(Error {
kind: ErrorKind::LaunchRejected,
message,
diagnostics: diagnostics(),
});
}
for call_id in call_ids {
turn.respond(
call_id,
ToolResult {
success: true,
output: ASYNC_TOOL_ACKNOWLEDGEMENT.to_owned(),
},
)
.await?;
}
if let Some(error) = drain_error {
return Err(error);
}
}
Some(ActiveEvent::Done) => {
let items = if text.is_empty() {
Vec::new()
} else {
vec![ShimItem::Text(text)]
};
return Ok(ShimOutput { items });
}
Some(ActiveEvent::Error(error)) => return Err(error),
None => {
return Err(Error {
kind: ErrorKind::Unavailable,
message: "Codex app-server closed before the active turn completed".into(),
diagnostics: diagnostics(),
});
}
}
}
}
fn append_section(mut output: String, section: &str) -> String {
if section.is_empty() {
return output;
}
if !output.is_empty() && !output.ends_with('\n') {
output.push('\n');
}
output.push_str(section);
output
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::{
Arc, Mutex,
atomic::{AtomicUsize, Ordering},
};
use std::task::{Context, Poll, Waker};
type Stages = Arc<Mutex<Vec<(String, Vec<String>)>>>;
struct RecordingLauncher {
stages: Stages,
completed: Arc<AtomicUsize>,
reject: bool,
}
impl ToolCallLauncher<String> for RecordingLauncher {
fn launch_stage<'a>(
&'a mut self,
text: String,
boxes: Vec<String>,
) -> ToolLaunchFuture<'a> {
let stages = Arc::clone(&self.stages);
let completed = Arc::clone(&self.completed);
let reject = self.reject;
Box::pin(async move {
stages.lock().unwrap().push((text, boxes));
if reject {
Err("consumer barrier failed".into())
} else {
completed.fetch_add(1, Ordering::SeqCst);
Ok(())
}
})
}
}
struct Acknowledgement {
call_id: String,
result: ToolResult,
completed_stages: usize,
}
struct ScriptedTurn {
events: VecDeque<ActiveEvent<String>>,
acknowledgements: Arc<Mutex<Vec<Acknowledgement>>>,
completed: Arc<AtomicUsize>,
}
impl ActiveTurn<String> for ScriptedTurn {
async fn next_event(&mut self) -> Option<ActiveEvent<String>> {
self.events.pop_front()
}
fn try_next_event(&mut self) -> Result<Option<ActiveEvent<String>>, Error> {
Ok(self.events.pop_front())
}
async fn respond(&mut self, call_id: String, result: ToolResult) -> Result<(), Error> {
self.acknowledgements.lock().unwrap().push(Acknowledgement {
call_id,
result,
completed_stages: self.completed.load(Ordering::SeqCst),
});
Ok(())
}
}
#[test]
fn grouped_waves_wait_for_callback_and_acknowledge_in_provider_order() {
let stages = Arc::new(Mutex::new(Vec::new()));
let acknowledgements = Arc::new(Mutex::new(Vec::new()));
let completed = Arc::new(AtomicUsize::new(0));
let mut launcher = RecordingLauncher {
stages: Arc::clone(&stages),
completed: Arc::clone(&completed),
reject: false,
};
let mut turn = ScriptedTurn {
events: VecDeque::from([
ActiveEvent::TextDelta("first stage".into()),
ActiveEvent::Call {
call_id: "call-1".into(),
box_: "box-1".into(),
},
ActiveEvent::Call {
call_id: "call-2".into(),
box_: "box-2".into(),
},
ActiveEvent::TextDelta("second stage".into()),
ActiveEvent::Call {
call_id: "call-3".into(),
box_: "box-3".into(),
},
ActiveEvent::TextDelta("final text".into()),
ActiveEvent::Done,
]),
acknowledgements: Arc::clone(&acknowledgements),
completed: Arc::clone(&completed),
};
let output = run_ready(drive_turn(&mut launcher, &mut turn, Vec::new)).unwrap();
assert_eq!(
*stages.lock().unwrap(),
vec![
("first stage".into(), vec!["box-1".into(), "box-2".into()]),
("second stage".into(), vec!["box-3".into()]),
]
);
let acknowledgements = acknowledgements.lock().unwrap();
assert_eq!(
acknowledgements
.iter()
.map(|ack| ack.call_id.as_str())
.collect::<Vec<_>>(),
vec!["call-1", "call-2", "call-3"]
);
assert_eq!(
acknowledgements
.iter()
.map(|ack| ack.completed_stages)
.collect::<Vec<_>>(),
vec![1, 1, 2]
);
assert!(acknowledgements.iter().all(|ack| ack.result.success));
assert!(
acknowledgements
.iter()
.all(|ack| ack.result.output == ASYNC_TOOL_ACKNOWLEDGEMENT)
);
assert_eq!(
ASYNC_TOOL_ACKNOWLEDGEMENT,
"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."
);
assert_eq!(
output,
ShimOutput {
items: vec![ShimItem::Text("final text".into())]
}
);
}
#[test]
fn callback_failure_acknowledges_nothing() {
let stages = Arc::new(Mutex::new(Vec::new()));
let acknowledgements = Arc::new(Mutex::new(Vec::new()));
let completed = Arc::new(AtomicUsize::new(0));
let mut launcher = RecordingLauncher {
stages: Arc::clone(&stages),
completed: Arc::clone(&completed),
reject: true,
};
let mut turn = ScriptedTurn {
events: VecDeque::from([
ActiveEvent::TextDelta("accepted text".into()),
ActiveEvent::Call {
call_id: "call-1".into(),
box_: "box-1".into(),
},
ActiveEvent::Done,
]),
acknowledgements: Arc::clone(&acknowledgements),
completed,
};
let error = run_ready(drive_turn(&mut launcher, &mut turn, Vec::new)).unwrap_err();
assert_eq!(error.kind, ErrorKind::LaunchRejected);
assert_eq!(error.message, "consumer barrier failed");
assert!(acknowledgements.lock().unwrap().is_empty());
assert_eq!(
*stages.lock().unwrap(),
vec![("accepted text".into(), vec!["box-1".into()])]
);
}
fn run_ready<F: Future>(future: F) -> F::Output {
let mut context = Context::from_waker(Waker::noop());
let mut future = Box::pin(future);
match future.as_mut().poll(&mut context) {
Poll::Ready(output) => output,
Poll::Pending => panic!("bounded scripted future unexpectedly pending"),
}
}
}