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}