Skip to main content

kcode_k1_codex_shim/
lib.rs

1//! Per-conversation K1 bridge for native Codex turns.
2//!
3//! One [`Shim::infer`] drives one fresh runtime turn through [`Event::Done`],
4//! completing each consumer-owned wave barrier before ordered acknowledgement.
5
6use 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
12/// Fixed successful response used to release every Codex dynamic-tool request.
13pub 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
15/// Future returned while handing one output stage to K1.
16pub type ToolLaunchFuture<'a> = Pin<Box<dyn Future<Output = Result<(), String>> + Send + 'a>>;
17
18/// Accepts one assistant-text and tool-call wave for durable launch.
19pub trait ToolCallLauncher<B>: Send {
20    /// Completes every consumer-required step before this wave may be acknowledged.
21    ///
22    /// Success includes optional same-turn steering. Failure can follow durable
23    /// acceptance or launch, does not certify replayability, and causes the shim
24    /// to acknowledge none of the wave and remain unusable.
25    fn launch_stage<'a>(&'a mut self, text: String, boxes: Vec<B>) -> ToolLaunchFuture<'a>;
26}
27
28/// Converts between K1's canonical box type and text visible to Codex.
29pub trait BoxCodec {
30    /// Canonical box type owned by K1.
31    type Box: Clone;
32
33    /// Converts one typed Codex dynamic-tool request into a canonical box.
34    ///
35    /// A shim invokes this exactly once for each call it receives.
36    fn tool_call_box(&mut self, call: &ToolCall) -> Self::Box;
37
38    /// Returns the complete representation of one box for Codex history.
39    fn box_text<'a>(&self, box_: &'a Self::Box) -> &'a str;
40}
41
42/// One ordered item produced by a completed native Codex turn.
43#[derive(Clone, Debug, PartialEq)]
44pub enum ShimItem<B> {
45    /// Assistant text remaining after the final call wave.
46    Text(String),
47    /// Reserved canonical box output.
48    Box(B),
49}
50
51/// Atomic terminal output from one native Codex turn.
52#[derive(Clone, Debug, PartialEq)]
53pub struct ShimOutput<B> {
54    /// Empty, or one final assistant-text item.
55    pub items: Vec<ShimItem<B>>,
56}
57
58#[derive(Clone, Copy, Debug, PartialEq, Eq)]
59enum Health {
60    Ready,
61    Unusable,
62}
63
64/// A single-owner, per-conversation bridge between K1 boxes and Codex turns.
65///
66/// Invoke a shim sequentially. Different shims may share cloned [`Adapter`]s.
67pub 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    /// Creates a ready shim for one conversation.
78    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    /// Appends one externally produced box in canonical history order.
95    pub fn record_box(&mut self, box_: C::Box) {
96        self.pending_boxes.push_back(box_);
97    }
98
99    /// Appends externally produced boxes without changing their order.
100    pub fn record_boxes(&mut self, boxes: impl IntoIterator<Item = C::Box>) {
101        self.pending_boxes.extend(boxes);
102    }
103
104    /// Returns the number of boxes waiting for the next fresh native turn.
105    pub fn pending_box_count(&self) -> usize {
106        self.pending_boxes.len()
107    }
108
109    /// Closes this shim's runtime conversation while the shim is ready.
110    ///
111    /// Pending external boxes are retained and can prefix a later fresh thread.
112    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    /// Runs one fresh native turn through terminal completion.
122    ///
123    /// Pending boxes clear only after start acceptance. Each immediately
124    /// buffered call wave completes its launcher barrier before acknowledgement.
125    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}