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}