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 responses.
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 respond to 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 a safe validation message when a converted call box is malformed.
39    ///
40    /// The message must be fixed, safely displayable text supplied by the codec;
41    /// it must not be derived from raw arbitrary tool-call values. A malformed
42    /// box is answered unsuccessfully and is never passed to the launcher.
43    /// Existing codecs retain their prior behavior through the `None` default.
44    fn malformed_tool_call_message(&self, _box_: &Self::Box) -> Option<&'static str> {
45        None
46    }
47
48    /// Returns the complete representation of one box for Codex history.
49    fn box_text<'a>(&self, box_: &'a Self::Box) -> &'a str;
50}
51
52/// One ordered item produced by a completed native Codex turn.
53#[derive(Clone, Debug, PartialEq)]
54pub enum ShimItem<B> {
55    /// Assistant text remaining after the final call wave.
56    Text(String),
57    /// Reserved canonical box output.
58    Box(B),
59}
60
61/// Atomic terminal output from one native Codex turn.
62#[derive(Clone, Debug, PartialEq)]
63pub struct ShimOutput<B> {
64    /// Empty, or one final assistant-text item.
65    pub items: Vec<ShimItem<B>>,
66}
67
68#[derive(Clone, Copy, Debug, PartialEq, Eq)]
69enum Health {
70    Ready,
71    Unusable,
72}
73
74/// A single-owner, per-conversation bridge between K1 boxes and Codex turns.
75///
76/// Invoke a shim sequentially. Different shims may share cloned [`Adapter`]s.
77pub 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    /// Creates a ready shim for one conversation.
88    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    /// Appends one externally produced box in canonical history order.
105    pub fn record_box(&mut self, box_: C::Box) {
106        self.pending_boxes.push_back(box_);
107    }
108
109    /// Appends externally produced boxes without changing their order.
110    pub fn record_boxes(&mut self, boxes: impl IntoIterator<Item = C::Box>) {
111        self.pending_boxes.extend(boxes);
112    }
113
114    /// Returns the number of boxes waiting for the next fresh native turn.
115    pub fn pending_box_count(&self) -> usize {
116        self.pending_boxes.len()
117    }
118
119    /// Closes this shim's runtime conversation while the shim is ready.
120    ///
121    /// Pending external boxes are retained and can prefix a later fresh thread.
122    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    /// Runs one fresh native turn through terminal completion.
132    ///
133    /// Pending boxes clear only after start acceptance. Each immediately
134    /// buffered call wave completes its launcher barrier before ordered responses.
135    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}