Skip to main content

kcode_k1_codex_runtime/
lib.rs

1pub use kcode_k1_codex_conversations::{
2    Config, DynamicTool, Error, ErrorKind, Event, ToolCall, ToolResult,
3};
4use kcode_k1_codex_runtime_actor::{Handle, Started, WeakHandle};
5pub use kcode_k1_codex_runtime_actor::{RuntimeObservation, UsageObservation};
6use std::{fmt, sync::Arc};
7use tokio::sync::mpsc;
8
9#[derive(Clone)]
10pub struct Adapter {
11    handle: Handle,
12    profile: Arc<Config>,
13}
14
15impl Adapter {
16    pub async fn open(config: Config) -> Result<Self, Error> {
17        config.validate()?;
18        let handle = kcode_k1_codex_runtime_actor::open(&config).await?;
19        Ok(Self {
20            handle,
21            profile: Arc::new(config),
22        })
23    }
24
25    pub fn with_config(&self, config: Config) -> Result<Adapter, Error> {
26        if let Err(mut error) = validate_derived(&self.profile, &config) {
27            error.diagnostics = self.handle.diagnostics();
28            return Err(error);
29        }
30        Ok(Adapter {
31            handle: self.handle.clone(),
32            profile: Arc::new(config),
33        })
34    }
35
36    pub async fn start_turn(
37        &self,
38        conversation_key: impl Into<String>,
39        input: impl Into<String>,
40    ) -> Result<Turn, Error> {
41        let key = conversation_key.into();
42        let Started { serial, events } = self
43            .handle
44            .start(key.clone(), input.into(), Arc::clone(&self.profile))
45            .await?;
46        Ok(Turn {
47            key,
48            serial,
49            handle: self.handle.downgrade(),
50            events: EventStream::new(events),
51            diagnostics: self.handle.diagnostics(),
52        })
53    }
54
55    pub async fn steer(
56        &self,
57        conversation_key: impl Into<String>,
58        input: impl Into<String>,
59    ) -> Result<(), Error> {
60        self.handle
61            .steer(
62                conversation_key.into(),
63                input.into(),
64                Arc::clone(&self.profile),
65            )
66            .await
67    }
68
69    pub async fn subscribe_usage(
70        &self,
71        key: String,
72    ) -> Result<mpsc::UnboundedReceiver<RuntimeObservation>, Error> {
73        self.handle.subscribe_usage(key).await
74    }
75
76    pub async fn close_conversation(&self, key: impl Into<String>) -> Result<(), Error> {
77        self.handle
78            .close(key.into(), Arc::clone(&self.profile))
79            .await
80    }
81
82    pub fn diagnostics(&self) -> Vec<u8> {
83        self.handle.diagnostics()
84    }
85}
86
87pub struct Turn {
88    key: String,
89    serial: u64,
90    handle: WeakHandle,
91    events: EventStream,
92    diagnostics: Vec<u8>,
93}
94
95impl fmt::Debug for Turn {
96    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
97        formatter
98            .debug_struct("Turn")
99            .field("key", &self.key)
100            .field("serial", &self.serial)
101            .finish_non_exhaustive()
102    }
103}
104
105impl Turn {
106    pub async fn next_event(&mut self) -> Option<Event> {
107        self.events.next().await
108    }
109
110    pub fn try_next_event(&mut self) -> Result<Option<Event>, Error> {
111        self.events.try_next().map_err(|()| self.unavailable())
112    }
113
114    pub async fn respond(
115        &self,
116        call_id: impl Into<String>,
117        result: ToolResult,
118    ) -> Result<(), Error> {
119        let handle = self.handle.upgrade().ok_or_else(|| self.unavailable())?;
120        handle
121            .respond(self.key.clone(), self.serial, call_id.into(), result)
122            .await
123    }
124
125    fn unavailable(&self) -> Error {
126        let diagnostics = self
127            .handle
128            .upgrade()
129            .map_or_else(|| self.diagnostics.clone(), |handle| handle.diagnostics());
130        Error {
131            kind: ErrorKind::Unavailable,
132            message: "Codex app-server is unavailable".into(),
133            diagnostics,
134        }
135    }
136}
137
138impl Drop for Turn {
139    fn drop(&mut self) {
140        if !self.events.terminal
141            && let Some(handle) = self.handle.upgrade()
142        {
143            handle.abandon(self.key.clone(), self.serial);
144        }
145    }
146}
147
148struct EventStream {
149    receiver: mpsc::UnboundedReceiver<Event>,
150    terminal: bool,
151}
152
153impl EventStream {
154    fn new(receiver: mpsc::UnboundedReceiver<Event>) -> Self {
155        Self {
156            receiver,
157            terminal: false,
158        }
159    }
160
161    async fn next(&mut self) -> Option<Event> {
162        let event = self.receiver.recv().await;
163        self.observe(event.as_ref());
164        event
165    }
166
167    fn try_next(&mut self) -> Result<Option<Event>, ()> {
168        match self.receiver.try_recv() {
169            Ok(event) => {
170                self.observe(Some(&event));
171                Ok(Some(event))
172            }
173            Err(mpsc::error::TryRecvError::Empty) => Ok(None),
174            Err(mpsc::error::TryRecvError::Disconnected) => {
175                self.terminal = true;
176                Err(())
177            }
178        }
179    }
180
181    fn observe(&mut self, event: Option<&Event>) {
182        if event.is_none_or(|event| matches!(event, Event::Done | Event::Error(_))) {
183            self.terminal = true;
184        }
185    }
186}
187
188fn validate_derived(base: &Config, derived: &Config) -> Result<(), Error> {
189    derived.validate()?;
190    if derived.executable != base.executable || derived.working_directory != base.working_directory
191    {
192        return Err(Error::new(
193            ErrorKind::Protocol,
194            "derived configuration must use the same executable and working directory",
195        ));
196    }
197    Ok(())
198}
199
200#[cfg(test)]
201mod tests {
202    use super::*;
203
204    fn profile() -> Config {
205        Config {
206            executable: "codex".into(),
207            working_directory: "/workspace".into(),
208            model: "model-a".into(),
209            reasoning_effort: Some("high".into()),
210            base_instructions: "base instructions".into(),
211            tools: Vec::new(),
212        }
213    }
214
215    #[test]
216    fn derived_profiles_validate_process_identity_and_retain_all_other_fields() {
217        let base = profile();
218        let mut derived = profile();
219        derived.model = "model-b".into();
220        derived.reasoning_effort = None;
221        derived.base_instructions = "other instructions".into();
222        assert!(validate_derived(&base, &derived).is_ok());
223
224        let mut different_executable = derived.clone();
225        different_executable.executable = "other-codex".into();
226        let error = validate_derived(&base, &different_executable).unwrap_err();
227        assert_eq!(error.kind, ErrorKind::Protocol);
228        assert_eq!(
229            error.message,
230            "derived configuration must use the same executable and working directory"
231        );
232
233        let mut different_directory = derived;
234        different_directory.working_directory = "/other".into();
235        assert!(validate_derived(&base, &different_directory).is_err());
236    }
237
238    #[test]
239    fn try_next_event_distinguishes_empty_terminal_and_disconnected_states() {
240        let (sender, receiver) = mpsc::unbounded_channel();
241        let mut events = EventStream::new(receiver);
242        assert!(matches!(events.try_next(), Ok(None)));
243        assert!(!events.terminal);
244
245        sender.send(Event::TextDelta("hello".into())).unwrap();
246        assert!(matches!(
247            events.try_next(),
248            Ok(Some(Event::TextDelta(text))) if text == "hello"
249        ));
250        assert!(!events.terminal);
251
252        sender.send(Event::Done).unwrap();
253        drop(sender);
254        assert!(matches!(events.try_next(), Ok(Some(Event::Done))));
255        assert!(events.terminal);
256        assert!(events.try_next().is_err());
257    }
258}