Skip to main content

kcode_k1_codex_conversations/
lib.rs

1pub use kcode_k1_codex_conversation_values::{
2    Config, Diagnostics, DynamicTool, Error, ErrorKind, Event, ToolCall, ToolResult,
3    validate_config,
4};
5use serde_json::{Value, json};
6use std::collections::{HashMap, HashSet};
7
8pub type ToolToken = (String, u64, String);
9
10pub struct Active<S> {
11    pub serial: u64,
12    pub turn: Option<String>,
13    pub sink: Option<S>,
14    pub early: Vec<Value>,
15    pub cancelled: bool,
16    pub interrupt_sent: bool,
17    pub failure: Option<Value>,
18}
19
20#[derive(Default)]
21pub struct Conversation<S> {
22    pub thread: Option<String>,
23    pub active: Option<Active<S>>,
24    pub closing: bool,
25}
26
27pub enum Pending<R> {
28    Thread {
29        key: String,
30        serial: u64,
31        input: String,
32    },
33    Turn {
34        key: String,
35        serial: u64,
36    },
37    Steer {
38        key: String,
39        serial: u64,
40        reply: R,
41    },
42    Close {
43        key: String,
44        thread: String,
45        reply: R,
46    },
47    Interrupt,
48}
49
50#[derive(Clone, Debug, PartialEq)]
51pub struct PendingTool {
52    pub id: Value,
53    pub rpc_key: String,
54}
55
56pub struct State<S, R> {
57    pub conversations: HashMap<String, Conversation<S>>,
58    pub by_thread: HashMap<String, String>,
59    pub pending: HashMap<u64, Pending<R>>,
60    pub tools: HashMap<ToolToken, PendingTool>,
61    pub rpc_ids: HashSet<String>,
62    pub next_id: u64,
63    pub next_turn: u64,
64}
65
66impl<S, R> Default for State<S, R> {
67    fn default() -> Self {
68        Self {
69            conversations: HashMap::new(),
70            by_thread: HashMap::new(),
71            pending: HashMap::new(),
72            tools: HashMap::new(),
73            rpc_ids: HashSet::new(),
74            next_id: 1,
75            next_turn: 1,
76        }
77    }
78}
79
80impl<S, R> State<S, R> {
81    pub fn allocate_request_id(&mut self) -> Result<u64, Error> {
82        allocate(&mut self.next_id, "client request id space exhausted")
83    }
84
85    pub fn allocate_turn_id(&mut self) -> Result<u64, Error> {
86        allocate(&mut self.next_turn, "turn serial space exhausted")
87    }
88
89    pub fn begin_turn(&mut self, key: impl Into<String>, sink: S) -> Result<u64, Error> {
90        let key = key.into();
91        ensure(
92            !self
93                .conversations
94                .get(&key)
95                .is_some_and(|value| value.active.is_some() || value.closing),
96            busy("conversation already has an active turn"),
97        )?;
98        let serial = self.allocate_turn_id()?;
99        self.conversations
100            .entry(key)
101            .or_insert_with(empty_conversation)
102            .active = Some(Active {
103            serial,
104            turn: None,
105            sink: Some(sink),
106            early: Vec::new(),
107            cancelled: false,
108            interrupt_sent: false,
109            failure: None,
110        });
111        Ok(serial)
112    }
113
114    pub fn begin_steer(
115        &mut self,
116        key: &str,
117        serial: u64,
118        reply: R,
119    ) -> Result<(u64, String, String), Error> {
120        let target = self
121            .conversations
122            .get(key)
123            .filter(|value| !value.closing)
124            .and_then(|conversation| {
125                conversation
126                    .active
127                    .as_ref()
128                    .filter(|active| active.serial == serial)
129                    .and_then(|active| conversation.thread.clone().zip(active.turn.clone()))
130            })
131            .ok_or_else(|| busy("conversation has no active native turn"))?;
132        ensure(
133            !self
134                .pending
135                .values()
136                .any(|value| matches!(value, Pending::Steer { key: owner, .. } if owner == key)),
137            busy("conversation already has a pending steer"),
138        )?;
139        ensure(
140            !self.pending.contains_key(&self.next_id),
141            protocol("duplicate client request id"),
142        )?;
143        let id = self.allocate_request_id()?;
144        self.pending.insert(
145            id,
146            Pending::Steer {
147                key: key.to_owned(),
148                serial,
149                reply,
150            },
151        );
152        Ok((id, target.0, target.1))
153    }
154
155    pub fn take_active(&mut self, key: &str, serial: u64) -> Option<Active<S>> {
156        let active = &mut self.conversations.get_mut(key)?.active;
157        (active.as_ref()?.serial == serial)
158            .then(|| active.take())
159            .flatten()
160    }
161
162    pub fn set_native_turn(
163        &mut self,
164        key: &str,
165        serial: u64,
166        turn: impl Into<String>,
167    ) -> Option<Vec<Value>> {
168        let active = self
169            .conversations
170            .get_mut(key)?
171            .active
172            .as_mut()
173            .filter(|value| value.serial == serial)?;
174        active.turn = Some(turn.into());
175        Some(std::mem::take(&mut active.early))
176    }
177
178    pub fn interrupt_target(&mut self, key: &str, serial: u64) -> Option<(String, String)> {
179        let conversation = self.conversations.get_mut(key)?;
180        let active = conversation.active.as_mut()?;
181        if active.serial != serial || active.interrupt_sent {
182            return None;
183        }
184        let target = (conversation.thread.clone()?, active.turn.clone()?);
185        active.interrupt_sent = true;
186        Some(target)
187    }
188
189    pub fn take_sinks(&mut self) -> Vec<S> {
190        self.conversations
191            .values_mut()
192            .filter_map(|value| value.active.take()?.sink)
193            .collect()
194    }
195
196    pub fn thread(&self, key: &str) -> Option<&str> {
197        self.conversations.get(key)?.thread.as_deref()
198    }
199
200    pub fn owner(&self, thread: &str) -> Option<&str> {
201        self.by_thread.get(thread).map(String::as_str)
202    }
203
204    pub fn set_thread(&mut self, key: &str, thread: impl Into<String>) -> Result<(), Error> {
205        let thread = thread.into();
206        ensure(
207            !self
208                .by_thread
209                .get(&thread)
210                .is_some_and(|owner| owner != key),
211            protocol("thread/start reused another conversation thread"),
212        )?;
213        let conversation = self
214            .conversations
215            .entry(key.to_owned())
216            .or_insert_with(empty_conversation);
217        if let Some(old) = conversation.thread.replace(thread.clone())
218            && old != thread
219        {
220            self.by_thread.remove(&old);
221        }
222        self.by_thread.insert(thread, key.to_owned());
223        Ok(())
224    }
225
226    pub fn begin_close(&mut self, key: &str) -> Result<Option<String>, Error> {
227        let Some(conversation) = self.conversations.get(key) else {
228            return Ok(None);
229        };
230        ensure(
231            conversation.active.is_none() && !conversation.closing,
232            busy("conversation is active or already closing"),
233        )?;
234        let Some(thread) = conversation.thread.clone() else {
235            self.conversations.remove(key);
236            return Ok(None);
237        };
238        self.conversations.get_mut(key).unwrap().closing = true;
239        Ok(Some(thread))
240    }
241
242    pub fn cancel_close(&mut self, key: &str) {
243        if let Some(conversation) = self.conversations.get_mut(key) {
244            conversation.closing = false;
245        }
246    }
247
248    pub fn finish_close(&mut self, key: &str, thread: &str) -> Result<(), Error> {
249        ensure(
250            self.conversations.get(key).is_some_and(|value| {
251                value.thread.as_deref() == Some(thread) && value.active.is_none() && value.closing
252            }),
253            protocol("thread/unsubscribe response did not match closing conversation"),
254        )?;
255        self.by_thread.remove(thread);
256        self.conversations.remove(key);
257        Ok(())
258    }
259
260    pub fn insert_pending(&mut self, id: u64, pending: Pending<R>) -> Result<(), Error> {
261        ensure(
262            !self.pending.contains_key(&id),
263            protocol("duplicate client request id"),
264        )?;
265        self.pending.insert(id, pending);
266        Ok(())
267    }
268
269    pub fn take_pending(&mut self, id: u64) -> Result<Pending<R>, Error> {
270        self.pending
271            .remove(&id)
272            .ok_or_else(|| protocol("unexpected or duplicate app-server response id"))
273    }
274
275    pub fn track_tool(
276        &mut self,
277        key: &str,
278        serial: u64,
279        call: impl Into<String>,
280        id: &Value,
281    ) -> Result<ToolToken, Error> {
282        let (id, rpc_key) = parse_rpc_id(id)?;
283        let token = (key.to_owned(), serial, call.into());
284        ensure(
285            !self.rpc_ids.contains(&rpc_key),
286            protocol("duplicate app-server request id"),
287        )?;
288        ensure(
289            !self.tools.contains_key(&token),
290            protocol("duplicate dynamic tool call id"),
291        )?;
292        self.rpc_ids.insert(rpc_key.clone());
293        self.tools
294            .insert(token.clone(), PendingTool { id, rpc_key });
295        Ok(token)
296    }
297
298    pub fn take_tool(&mut self, token: &ToolToken) -> Option<PendingTool> {
299        let pending = self.tools.remove(token)?;
300        self.rpc_ids.remove(&pending.rpc_key);
301        Some(pending)
302    }
303
304    pub fn take_turn_tools(&mut self, key: &str, serial: u64) -> Vec<(ToolToken, PendingTool)> {
305        self.tools
306            .keys()
307            .filter(|(owner, turn, _)| owner == key && *turn == serial)
308            .cloned()
309            .collect::<Vec<_>>()
310            .into_iter()
311            .filter_map(|token| self.take_tool(&token).map(|pending| (token, pending)))
312            .collect()
313    }
314
315    pub fn resolve_tool(&mut self, id: &Value) -> Result<Option<ToolToken>, Error> {
316        let (_, rpc_key) = parse_rpc_id(id)?;
317        let token = self
318            .tools
319            .iter()
320            .find_map(|(token, pending)| (pending.rpc_key == rpc_key).then(|| token.clone()));
321        if let Some(token) = &token {
322            self.take_tool(token);
323        }
324        Ok(token)
325    }
326}
327
328fn empty_conversation<S>() -> Conversation<S> {
329    Conversation {
330        thread: None,
331        active: None,
332        closing: false,
333    }
334}
335
336fn allocate(counter: &mut u64, exhausted: &'static str) -> Result<u64, Error> {
337    let id = *counter;
338    *counter = id.checked_add(1).ok_or_else(|| protocol(exhausted))?;
339    Ok(id)
340}
341
342fn ensure(valid: bool, error: Error) -> Result<(), Error> {
343    valid.then_some(()).ok_or(error)
344}
345
346fn busy(message: &'static str) -> Error {
347    Error::new(ErrorKind::Busy, message)
348}
349
350fn protocol(message: impl Into<String>) -> Error {
351    Error::new(ErrorKind::Protocol, message)
352}
353
354pub fn thread_start_params(config: &Config) -> Value {
355    let tools = config
356        .tools
357        .iter()
358        .map(|tool| {
359            json!({
360                "name": tool.name,
361                "description": tool.description,
362                "inputSchema": tool.input_schema
363            })
364        })
365        .collect::<Vec<_>>();
366    json!({
367        "model": config.model,
368        "cwd": config.working_directory,
369        "approvalPolicy": "never",
370        "sandbox": "read-only",
371        "baseInstructions": config.base_instructions,
372        "serviceName": "kcode-k1-codex-adapter",
373        "dynamicTools": tools
374    })
375}
376
377pub fn turn_start_params(thread: &str, input: impl Into<String>) -> Value {
378    json!({"threadId": thread, "input": [{"type": "text", "text": input.into()}]})
379}
380
381pub fn turn_steer_params(thread: &str, expected_turn: &str, input: impl Into<String>) -> Value {
382    json!({
383        "threadId": thread,
384        "expectedTurnId": expected_turn,
385        "input": [{"type": "text", "text": input.into()}]
386    })
387}
388
389pub fn parse_scope(params: Option<&Value>) -> Result<(&str, &str), Error> {
390    let params = params.ok_or_else(|| protocol("scoped message omitted params"))?;
391    let thread = params
392        .get("threadId")
393        .and_then(Value::as_str)
394        .ok_or_else(|| protocol("scoped message omitted threadId"))?;
395    let turn = params
396        .get("turnId")
397        .and_then(Value::as_str)
398        .or_else(|| params.pointer("/turn/id").and_then(Value::as_str))
399        .ok_or_else(|| protocol("scoped message omitted turn id"))?;
400    Ok((thread, turn))
401}
402
403pub fn parse_rpc_id(id: &Value) -> Result<(Value, String), Error> {
404    match id {
405        Value::String(value) => Ok((id.clone(), format!("s:{value}"))),
406        Value::Number(value) => Ok((id.clone(), format!("n:{value}"))),
407        _ => Err(protocol("server request id must be a string or number")),
408    }
409}
410
411pub fn is_model_reroute(method: &str) -> bool {
412    let method = method.to_ascii_lowercase();
413    method.contains("model") && method.contains("rerout")
414}
415
416#[cfg(test)]
417mod tests {
418    use super::*;
419
420    fn active() -> (State<(), &'static str>, u64) {
421        let mut state = State::default();
422        state.set_thread("key", "thread-a").unwrap();
423        let serial = state.begin_turn("key", ()).unwrap();
424        state.set_native_turn("key", serial, "turn-b").unwrap();
425        (state, serial)
426    }
427
428    #[test]
429    fn protocol_and_steering_are_exact() {
430        let config = Config {
431            executable: "codex".into(),
432            working_directory: "/work".into(),
433            model: "model".into(),
434            reasoning_effort: None,
435            base_instructions: String::new(),
436            tools: Vec::new(),
437        };
438        let start = thread_start_params(&config);
439        assert_eq!(
440            (start["approvalPolicy"].as_str(), start["sandbox"].as_str()),
441            (Some("never"), Some("read-only"))
442        );
443        assert_eq!(
444            turn_steer_params("thread-a", "turn-b", "next"),
445            json!({"threadId": "thread-a", "expectedTurnId": "turn-b", "input": [{"type": "text", "text": "next"}]})
446        );
447
448        let (mut state, serial) = active();
449        let (id, thread, turn) = state.begin_steer("key", serial, "first").unwrap();
450        assert_eq!((thread.as_str(), turn.as_str()), ("thread-a", "turn-b"));
451        let before = (state.next_id, state.pending.len());
452        for (key, turn) in [("key", serial), ("wrong", serial), ("key", serial + 1)] {
453            assert_eq!(
454                state.begin_steer(key, turn, "bad").unwrap_err().kind,
455                ErrorKind::Busy
456            );
457            assert_eq!((state.next_id, state.pending.len()), before);
458        }
459        state.set_thread("other", "thread-c").unwrap();
460        let other = state.begin_turn("other", ()).unwrap();
461        state.set_native_turn("other", other, "turn-d").unwrap();
462        assert!(state.begin_steer("other", other, "other").is_ok());
463        assert!(
464            matches!(state.take_pending(id).unwrap(), Pending::Steer { key, serial: got, reply: "first" } if key == "key" && got == serial)
465        );
466        assert!(state.begin_steer("key", serial, "again").is_ok());
467    }
468
469    #[test]
470    fn duplicate_pending_is_transactional() {
471        let mut state: State<(), ()> = State::default();
472        state.insert_pending(1, Pending::Interrupt).unwrap();
473        assert!(state.insert_pending(1, Pending::Interrupt).is_err());
474        assert!(matches!(state.take_pending(1).unwrap(), Pending::Interrupt));
475    }
476}