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 steer_pending(&self, key: &str, serial: u64) -> bool {
156        self.pending.values().any(
157            |value| matches!(value, Pending::Steer { key: owner, serial: pending, .. } if owner == key && *pending == serial),
158        )
159    }
160
161    pub fn take_active(&mut self, key: &str, serial: u64) -> Option<Active<S>> {
162        let active = &mut self.conversations.get_mut(key)?.active;
163        (active.as_ref()?.serial == serial)
164            .then(|| active.take())
165            .flatten()
166    }
167
168    pub fn set_native_turn(
169        &mut self,
170        key: &str,
171        serial: u64,
172        turn: impl Into<String>,
173    ) -> Option<Vec<Value>> {
174        let active = self
175            .conversations
176            .get_mut(key)?
177            .active
178            .as_mut()
179            .filter(|value| value.serial == serial)?;
180        active.turn = Some(turn.into());
181        Some(std::mem::take(&mut active.early))
182    }
183
184    pub fn interrupt_target(&mut self, key: &str, serial: u64) -> Option<(String, String)> {
185        let conversation = self.conversations.get_mut(key)?;
186        let active = conversation.active.as_mut()?;
187        if active.serial != serial || active.interrupt_sent {
188            return None;
189        }
190        let target = (conversation.thread.clone()?, active.turn.clone()?);
191        active.interrupt_sent = true;
192        Some(target)
193    }
194
195    pub fn take_sinks(&mut self) -> Vec<S> {
196        self.conversations
197            .values_mut()
198            .filter_map(|value| value.active.take()?.sink)
199            .collect()
200    }
201
202    pub fn thread(&self, key: &str) -> Option<&str> {
203        self.conversations.get(key)?.thread.as_deref()
204    }
205
206    pub fn owner(&self, thread: &str) -> Option<&str> {
207        self.by_thread.get(thread).map(String::as_str)
208    }
209
210    pub fn set_thread(&mut self, key: &str, thread: impl Into<String>) -> Result<(), Error> {
211        let thread = thread.into();
212        ensure(
213            !self
214                .by_thread
215                .get(&thread)
216                .is_some_and(|owner| owner != key),
217            protocol("thread/start reused another conversation thread"),
218        )?;
219        let conversation = self
220            .conversations
221            .entry(key.to_owned())
222            .or_insert_with(empty_conversation);
223        if let Some(old) = conversation.thread.replace(thread.clone())
224            && old != thread
225        {
226            self.by_thread.remove(&old);
227        }
228        self.by_thread.insert(thread, key.to_owned());
229        Ok(())
230    }
231
232    pub fn begin_close(&mut self, key: &str) -> Result<Option<String>, Error> {
233        let Some(conversation) = self.conversations.get(key) else {
234            return Ok(None);
235        };
236        ensure(
237            conversation.active.is_none() && !conversation.closing,
238            busy("conversation is active or already closing"),
239        )?;
240        let Some(thread) = conversation.thread.clone() else {
241            self.conversations.remove(key);
242            return Ok(None);
243        };
244        self.conversations.get_mut(key).unwrap().closing = true;
245        Ok(Some(thread))
246    }
247
248    pub fn cancel_close(&mut self, key: &str) {
249        if let Some(conversation) = self.conversations.get_mut(key) {
250            conversation.closing = false;
251        }
252    }
253
254    pub fn finish_close(&mut self, key: &str, thread: &str) -> Result<(), Error> {
255        ensure(
256            self.conversations.get(key).is_some_and(|value| {
257                value.thread.as_deref() == Some(thread) && value.active.is_none() && value.closing
258            }),
259            protocol("thread/unsubscribe response did not match closing conversation"),
260        )?;
261        self.by_thread.remove(thread);
262        self.conversations.remove(key);
263        Ok(())
264    }
265
266    pub fn insert_pending(&mut self, id: u64, pending: Pending<R>) -> Result<(), Error> {
267        ensure(
268            !self.pending.contains_key(&id),
269            protocol("duplicate client request id"),
270        )?;
271        self.pending.insert(id, pending);
272        Ok(())
273    }
274
275    pub fn take_pending(&mut self, id: u64) -> Result<Pending<R>, Error> {
276        self.pending
277            .remove(&id)
278            .ok_or_else(|| protocol("unexpected or duplicate app-server response id"))
279    }
280
281    pub fn track_tool(
282        &mut self,
283        key: &str,
284        serial: u64,
285        call: impl Into<String>,
286        id: &Value,
287    ) -> Result<ToolToken, Error> {
288        let (id, rpc_key) = parse_rpc_id(id)?;
289        let token = (key.to_owned(), serial, call.into());
290        ensure(
291            !self.rpc_ids.contains(&rpc_key),
292            protocol("duplicate app-server request id"),
293        )?;
294        ensure(
295            !self.tools.contains_key(&token),
296            protocol("duplicate dynamic tool call id"),
297        )?;
298        self.rpc_ids.insert(rpc_key.clone());
299        self.tools
300            .insert(token.clone(), PendingTool { id, rpc_key });
301        Ok(token)
302    }
303
304    pub fn take_tool(&mut self, token: &ToolToken) -> Option<PendingTool> {
305        let pending = self.tools.remove(token)?;
306        self.rpc_ids.remove(&pending.rpc_key);
307        Some(pending)
308    }
309
310    pub fn take_turn_tools(&mut self, key: &str, serial: u64) -> Vec<(ToolToken, PendingTool)> {
311        self.tools
312            .keys()
313            .filter(|(owner, turn, _)| owner == key && *turn == serial)
314            .cloned()
315            .collect::<Vec<_>>()
316            .into_iter()
317            .filter_map(|token| self.take_tool(&token).map(|pending| (token, pending)))
318            .collect()
319    }
320
321    pub fn resolve_tool(&mut self, id: &Value) -> Result<Option<ToolToken>, Error> {
322        let (_, rpc_key) = parse_rpc_id(id)?;
323        let token = self
324            .tools
325            .iter()
326            .find_map(|(token, pending)| (pending.rpc_key == rpc_key).then(|| token.clone()));
327        if let Some(token) = &token {
328            self.take_tool(token);
329        }
330        Ok(token)
331    }
332}
333
334fn empty_conversation<S>() -> Conversation<S> {
335    Conversation {
336        thread: None,
337        active: None,
338        closing: false,
339    }
340}
341
342fn allocate(counter: &mut u64, exhausted: &'static str) -> Result<u64, Error> {
343    let id = *counter;
344    *counter = id.checked_add(1).ok_or_else(|| protocol(exhausted))?;
345    Ok(id)
346}
347
348fn ensure(valid: bool, error: Error) -> Result<(), Error> {
349    valid.then_some(()).ok_or(error)
350}
351
352fn busy(message: &'static str) -> Error {
353    Error::new(ErrorKind::Busy, message)
354}
355
356fn protocol(message: impl Into<String>) -> Error {
357    Error::new(ErrorKind::Protocol, message)
358}
359
360pub fn thread_start_params(config: &Config) -> Value {
361    let tools = config
362        .tools
363        .iter()
364        .map(|tool| {
365            json!({
366                "name": tool.name,
367                "description": tool.description,
368                "inputSchema": tool.input_schema
369            })
370        })
371        .collect::<Vec<_>>();
372    json!({
373        "model": config.model,
374        "cwd": config.working_directory,
375        "approvalPolicy": "never",
376        "sandbox": "read-only",
377        "baseInstructions": config.base_instructions,
378        "serviceName": "kcode-k1-codex-adapter",
379        "dynamicTools": tools
380    })
381}
382
383pub fn web_search_thread_start_params(config: &Config) -> Value {
384    json!({
385        "model": config.model,
386        "cwd": config.working_directory,
387        "approvalPolicy": "never",
388        "sandbox": "read-only",
389        "dynamicTools": [],
390        "ephemeral": true,
391        "config": {"web_search": "live"}
392    })
393}
394
395pub fn turn_start_params(thread: &str, input: impl Into<String>) -> Value {
396    json!({"threadId": thread, "input": [{"type": "text", "text": input.into()}]})
397}
398
399pub fn turn_steer_params(thread: &str, expected_turn: &str, input: impl Into<String>) -> Value {
400    json!({
401        "threadId": thread,
402        "expectedTurnId": expected_turn,
403        "input": [{"type": "text", "text": input.into()}]
404    })
405}
406
407pub fn parse_scope(params: Option<&Value>) -> Result<(&str, &str), Error> {
408    let params = params.ok_or_else(|| protocol("scoped message omitted params"))?;
409    let thread = params
410        .get("threadId")
411        .and_then(Value::as_str)
412        .ok_or_else(|| protocol("scoped message omitted threadId"))?;
413    let turn = params
414        .get("turnId")
415        .and_then(Value::as_str)
416        .or_else(|| params.pointer("/turn/id").and_then(Value::as_str))
417        .ok_or_else(|| protocol("scoped message omitted turn id"))?;
418    Ok((thread, turn))
419}
420
421pub fn parse_rpc_id(id: &Value) -> Result<(Value, String), Error> {
422    match id {
423        Value::String(value) => Ok((id.clone(), format!("s:{value}"))),
424        Value::Number(value) => Ok((id.clone(), format!("n:{value}"))),
425        _ => Err(protocol("server request id must be a string or number")),
426    }
427}
428
429pub fn is_model_reroute(method: &str) -> bool {
430    let method = method.to_ascii_lowercase();
431    method.contains("model") && method.contains("rerout")
432}
433
434#[cfg(test)]
435mod tests {
436    use super::*;
437
438    fn active() -> (State<(), &'static str>, u64) {
439        let mut state = State::default();
440        state.set_thread("key", "thread-a").unwrap();
441        let serial = state.begin_turn("key", ()).unwrap();
442        state.set_native_turn("key", serial, "turn-b").unwrap();
443        (state, serial)
444    }
445
446    #[test]
447    fn protocol_and_steering_are_exact() {
448        let config = Config {
449            executable: "codex".into(),
450            working_directory: "/work".into(),
451            model: "model".into(),
452            reasoning_effort: None,
453            base_instructions: "ordinary instructions".into(),
454            tools: vec![DynamicTool {
455                name: "ordinary_tool".into(),
456                description: "ordinary description".into(),
457                input_schema: json!({"type": "object"}),
458            }],
459        };
460        let expected_ordinary = json!({
461            "model": "model",
462            "cwd": "/work",
463            "approvalPolicy": "never",
464            "sandbox": "read-only",
465            "baseInstructions": "ordinary instructions",
466            "serviceName": "kcode-k1-codex-adapter",
467            "dynamicTools": [{
468                "name": "ordinary_tool",
469                "description": "ordinary description",
470                "inputSchema": {"type": "object"}
471            }]
472        });
473        assert_eq!(thread_start_params(&config), expected_ordinary);
474        let expected_search = json!({
475            "model": "model",
476            "cwd": "/work",
477            "approvalPolicy": "never",
478            "sandbox": "read-only",
479            "dynamicTools": [],
480            "ephemeral": true,
481            "config": {"web_search": "live"}
482        });
483        assert_eq!(web_search_thread_start_params(&config), expected_search);
484        assert_eq!(
485            turn_steer_params("thread-a", "turn-b", "next"),
486            json!({"threadId": "thread-a", "expectedTurnId": "turn-b", "input": [{"type": "text", "text": "next"}]})
487        );
488
489        let (mut state, serial) = active();
490        let (id, thread, turn) = state.begin_steer("key", serial, "first").unwrap();
491        assert_eq!((thread.as_str(), turn.as_str()), ("thread-a", "turn-b"));
492        let before = (state.next_id, state.pending.len());
493        for (key, turn) in [("key", serial), ("wrong", serial), ("key", serial + 1)] {
494            assert_eq!(
495                state.begin_steer(key, turn, "bad").unwrap_err().kind,
496                ErrorKind::Busy
497            );
498            assert_eq!((state.next_id, state.pending.len()), before);
499        }
500        state.set_thread("other", "thread-c").unwrap();
501        let other = state.begin_turn("other", ()).unwrap();
502        state.set_native_turn("other", other, "turn-d").unwrap();
503        assert!(state.begin_steer("other", other, "other").is_ok());
504        assert!(
505            matches!(state.take_pending(id).unwrap(), Pending::Steer { key, serial: got, reply: "first" } if key == "key" && got == serial)
506        );
507        assert!(state.begin_steer("key", serial, "again").is_ok());
508    }
509
510    #[test]
511    fn steer_pending_matches_and_clears() {
512        let (mut state, serial) = active();
513        let id = state.begin_steer("key", serial, "reply").unwrap().0;
514        assert!(state.steer_pending("key", serial));
515        assert!(!state.steer_pending("other", serial));
516        assert!(!state.steer_pending("key", serial + 1));
517        state.take_pending(id).unwrap();
518        assert!(!state.steer_pending("key", serial));
519    }
520
521    #[test]
522    fn duplicate_pending_is_transactional() {
523        let mut state: State<(), ()> = State::default();
524        state.insert_pending(1, Pending::Interrupt).unwrap();
525        assert!(state.insert_pending(1, Pending::Interrupt).is_err());
526        assert!(matches!(state.take_pending(1).unwrap(), Pending::Interrupt));
527    }
528}