Skip to main content

kcode_k1_chat_codex_state/
lib.rs

1#![forbid(unsafe_code)]
2
3mod recovery;
4
5pub use kcode_k1_chat_codex_codec::{BoxValue, Call};
6pub use kcode_k1_chat_state::{
7    AGENT_ATTACHMENT_TYPE, AGENT_MESSAGE_TYPE, AGENT_RESPONSE_TYPE, ActorState, BoxId, ChatBox,
8    ProviderGenerated, SYSTEM_MESSAGE_TYPE, TOOL_ATTACHMENT_TYPE, TOOL_CALL_TYPE,
9    TOOL_MESSAGE_TYPE, TOOL_RESULT_TYPE, ToolCallId, USER_ATTACHMENT_TYPE, USER_MESSAGE_TYPE,
10};
11pub use kcode_k1_codex_adapter::{ShimItem, ShimOutput};
12
13use std::sync::Arc;
14use std::sync::atomic::AtomicU8;
15
16use kcode_k1_chat_codex_codec::{open_agent_response, project};
17use kcode_k1_chat_state::{ProviderCall, StateError};
18use recovery::recovered_sequence;
19
20pub const MALFORMED_NATIVE_CALL_METADATA_TYPE: &str = "k1.malformed-native-call.v1";
21
22#[derive(Clone, Debug)]
23pub struct Start {
24    pub job: u64,
25    pub values: Vec<BoxValue>,
26    pub attempt: Arc<AtomicU8>,
27}
28
29#[derive(Clone, Debug, Eq, PartialEq)]
30pub struct PreparedCall {
31    pub tool_call_id: ToolCallId,
32    pub name: String,
33    pub arguments: String,
34    disposition: PreparedCallDisposition,
35}
36impl PreparedCall {
37    pub fn disposition(&self) -> &PreparedCallDisposition {
38        &self.disposition
39    }
40}
41
42#[derive(Clone, Debug, Eq, PartialEq)]
43pub enum PreparedCallDisposition {
44    External,
45    ImmediateError(Box<ImmediateToolError>),
46}
47
48#[derive(Clone, Debug, Eq, PartialEq)]
49pub struct ImmediateToolError {
50    pub message: String,
51    pub native_tool: String,
52    pub attempted_ktool: Option<String>,
53    pub validation_code: String,
54    pub path: String,
55    pub expected: String,
56    pub received: String,
57    pub native_arguments: String,
58    pub metadata_type: String,
59    pub metadata_contents: String,
60}
61
62#[derive(Clone, Debug)]
63pub struct PreparedMailboxFlush(Arc<Prepared>);
64#[derive(Debug)]
65struct Prepared {
66    token: u64,
67    values: Vec<BoxValue>,
68    job: u64,
69    external_count: usize,
70}
71impl PreparedMailboxFlush {
72    pub fn values(&self) -> &[BoxValue] {
73        &self.0.values
74    }
75}
76
77#[derive(Clone, Debug, Eq, PartialEq)]
78pub enum Status {
79    Running,
80    Quiet,
81    Stalled { message: String, restartable: bool },
82}
83#[derive(Clone, Copy, Debug, Eq, PartialEq)]
84pub enum RestartError {
85    NotStalled,
86    ProviderActionAccepted,
87}
88#[derive(Clone, Copy, Debug, Eq, PartialEq)]
89enum Phase {
90    ProviderActive,
91    ChatendBoundary,
92    PendingGeneration,
93}
94struct Round {
95    job: u64,
96    accepted_provider_action: bool,
97    phase: Phase,
98    mailbox_flush_needed: bool,
99    restartable: bool,
100}
101enum Mode {
102    Idle,
103    Running(Round),
104    Stalled(Status, bool),
105}
106
107pub struct ConversationState {
108    state: ActorState,
109    session: [u8; 12],
110    sequence: u64,
111    unsubmitted: Vec<BoxValue>,
112    queued_trigger: bool,
113    token: u64,
114    prepared: Option<Arc<Prepared>>,
115    mode: Mode,
116}
117
118impl ConversationState {
119    pub fn new(session: [u8; 12]) -> Self {
120        Self {
121            state: ActorState::new(false),
122            session,
123            sequence: 0,
124            unsubmitted: Vec::new(),
125            queued_trigger: false,
126            token: 0,
127            prepared: None,
128            mode: Mode::Idle,
129        }
130    }
131    pub fn recover(session: [u8; 12], boxes: Vec<ChatBox>, force: bool) -> Result<Self, String> {
132        let sequence = recovered_sequence(session, &boxes)?;
133        let state = ActorState::recover(boxes, force).map_err(debug)?;
134        Ok(Self {
135            unsubmitted: state.boxes().iter().map(project).collect(),
136            state,
137            sequence,
138            ..Self::new(session)
139        })
140    }
141    pub fn boxes(&self) -> &[ChatBox] {
142        self.state.boxes()
143    }
144    pub fn status(&self) -> Status {
145        match &self.mode {
146            Mode::Running(_) => Status::Running,
147            Mode::Idle if self.state.quiet() => Status::Quiet,
148            Mode::Idle => Status::Running,
149            Mode::Stalled(status, _) => status.clone(),
150        }
151    }
152
153    pub fn accept(
154        &mut self,
155        box_type: String,
156        contents: String,
157        hidden_type: String,
158        hidden_contents: String,
159    ) -> Result<(), String> {
160        self.accept_arrival(true, |state| {
161            state.accept_box(box_type, contents, hidden_type, hidden_contents)
162        })
163    }
164    pub fn accept_tool_message(&mut self, id: ToolCallId, message: String) -> Result<(), String> {
165        self.accept_arrival(false, |state| state.accept_tool_message(id, message))
166    }
167    pub fn accept_tool_return(
168        &mut self,
169        id: ToolCallId,
170        result: Result<String, String>,
171    ) -> Result<(), String> {
172        self.accept_arrival(true, |state| state.accept_async_return(id, result))
173    }
174    pub fn accept_tool_return_v2(
175        &mut self,
176        id: ToolCallId,
177        result: Result<String, String>,
178        metadata_type: String,
179        metadata_contents: String,
180    ) -> Result<(), String> {
181        self.accept_arrival(true, |state| {
182            state.accept_async_return_v2(id, result, metadata_type, metadata_contents)
183        })
184    }
185
186    pub fn accept_idle_context_box(
187        &mut self,
188        box_type: String,
189        contents: String,
190        hidden_type: String,
191        hidden_contents: String,
192    ) -> Result<(), String> {
193        self.accept_idle(|state| {
194            state
195                .accept_idle_context_box(box_type, contents, hidden_type, hidden_contents)
196                .map(|_| ())
197        })
198    }
199    pub fn accept_idle_context_tool_call(
200        &mut self,
201        name: String,
202        arguments: String,
203    ) -> Result<ToolCallId, String> {
204        let sequence = next_sequence(self.sequence)?;
205        let id = ToolCallId::new(self.session, sequence);
206        self.accept_idle(|state| {
207            state
208                .accept_idle_context_tool_call(ProviderCall {
209                    tool_call_id: id,
210                    name,
211                    arguments,
212                })
213                .map(|_| ())
214        })?;
215        self.sequence = sequence;
216        Ok(id)
217    }
218    pub fn accept_idle_context_tool_return(
219        &mut self,
220        id: ToolCallId,
221        result: Result<String, String>,
222    ) -> Result<(), String> {
223        self.accept_idle(|state| {
224            state
225                .accept_idle_context_tool_return(id, result)
226                .map(|_| ())
227        })
228    }
229
230    pub fn begin(&mut self) -> Result<Option<Start>, String> {
231        if !matches!(self.mode, Mode::Idle) {
232            return Ok(None);
233        }
234        let Some(start) = self.state.begin_inference().map_err(debug)? else {
235            return Ok(None);
236        };
237        let promised_id = self.promised_id()?;
238        let mut values = std::mem::take(&mut self.unsubmitted);
239        values.push(open_agent_response(promised_id));
240        let mailbox_flush_needed = std::mem::take(&mut self.queued_trigger);
241        self.mode = Mode::Running(Round {
242            job: start.job,
243            accepted_provider_action: false,
244            phase: Phase::ProviderActive,
245            mailbox_flush_needed,
246            restartable: true,
247        });
248        Ok(Some(Start {
249            job: start.job,
250            values,
251            attempt: start.attempt,
252        }))
253    }
254
255    pub fn prepare_stage(
256        &mut self,
257        job: u64,
258        text: String,
259        values: Vec<BoxValue>,
260    ) -> Result<Vec<PreparedCall>, String> {
261        if !matches!(&self.mode,Mode::Running(round) if round.job==job&&round.phase==Phase::ProviderActive)
262        {
263            return Err("stale Codex inference stage".into());
264        }
265        if self.prepared.is_some() {
266            return Err("a Codex mailbox flush remains uncommitted".into());
267        }
268        let accepted_provider_action = !values.is_empty();
269        let mut sequence = self.sequence;
270        let mut generated = Vec::with_capacity(values.len());
271        let mut prepared = Vec::new();
272        for value in values {
273            match value {
274                BoxValue::AgentMessage(Ok(contents)) => {
275                    generated.push(ProviderGenerated::AgentMessage { contents })
276                }
277                BoxValue::Call(Ok(call)) => {
278                    sequence = next_sequence(sequence)?;
279                    let id = ToolCallId::new(self.session, sequence);
280                    generated.push(ProviderGenerated::ToolCall(ProviderCall {
281                        tool_call_id: id,
282                        name: call.name.clone(),
283                        arguments: call.arguments.clone(),
284                    }));
285                    prepared.push(PreparedCall {
286                        tool_call_id: id,
287                        name: call.name,
288                        arguments: call.arguments,
289                        disposition: PreparedCallDisposition::External,
290                    });
291                }
292                BoxValue::MalformedNativeAction(action) => {
293                    sequence = next_sequence(sequence)?;
294                    let id = ToolCallId::new(self.session, sequence);
295                    let name = action
296                        .attempted_ktool()
297                        .unwrap_or_else(|| action.native_tool())
298                        .to_owned();
299                    let arguments = action.native_arguments_json();
300                    let error = ImmediateToolError {
301                        message: action.diagnostic(),
302                        native_tool: action.native_tool().to_owned(),
303                        attempted_ktool: action.attempted_ktool().map(str::to_owned),
304                        validation_code: action.validation_code().to_owned(),
305                        path: action.path().to_owned(),
306                        expected: action.expected().to_owned(),
307                        received: action.received().to_owned(),
308                        native_arguments: arguments.clone(),
309                        metadata_type: MALFORMED_NATIVE_CALL_METADATA_TYPE.to_owned(),
310                        metadata_contents: action.diagnostic_json(),
311                    };
312                    generated.push(ProviderGenerated::ToolCall(ProviderCall {
313                        tool_call_id: id,
314                        name: name.clone(),
315                        arguments: arguments.clone(),
316                    }));
317                    prepared.push(PreparedCall {
318                        tool_call_id: id,
319                        name,
320                        arguments,
321                        disposition: PreparedCallDisposition::ImmediateError(Box::new(error)),
322                    });
323                }
324                _ => return Err("stage contains a malformed provider action".into()),
325            }
326        }
327        let before = self.state.boxes().len();
328        self.state
329            .append_stage(job, text, generated)
330            .map_err(debug)?;
331        self.unsubmitted.extend(
332            self.state.boxes()[before..]
333                .iter()
334                .filter(|b| b.box_type() == TOOL_CALL_TYPE)
335                .map(project),
336        );
337        self.sequence = sequence;
338        if let Mode::Running(round) = &mut self.mode {
339            round.accepted_provider_action |= accepted_provider_action;
340            round.phase = Phase::ChatendBoundary;
341            round.mailbox_flush_needed = true;
342            round.restartable = false;
343        }
344        Ok(prepared)
345    }
346
347    pub fn mailbox_flush(&mut self, job: u64) -> Result<Vec<ChatBox>, String> {
348        if !matches!(&self.mode,Mode::Running(round) if round.job==job&&round.phase==Phase::ChatendBoundary)
349        {
350            return Err("stale Codex active-arrival mailbox flush".into());
351        }
352        let boxes = self.state.flush_active_arrivals(job).map_err(debug)?;
353        self.unsubmitted.extend(boxes.iter().map(project));
354        if let Mode::Running(round) = &mut self.mode {
355            round.phase = Phase::PendingGeneration;
356        }
357        Ok(boxes)
358    }
359    pub fn prepare_mailbox_flush(
360        &mut self,
361        job: u64,
362    ) -> Result<Option<PreparedMailboxFlush>, String> {
363        let (needed, phase) = match &self.mode {
364            Mode::Running(round) if round.job == job => (round.mailbox_flush_needed, round.phase),
365            _ => return Err("stale Codex inference mailbox flush".into()),
366        };
367        if let Some(value) = &self.prepared {
368            return Ok(Some(PreparedMailboxFlush(Arc::clone(value))));
369        }
370        if !needed {
371            return Ok(None);
372        }
373        if phase == Phase::ProviderActive {
374            return Ok(None);
375        }
376        if phase == Phase::ChatendBoundary {
377            self.mailbox_flush(job)?;
378        }
379        let token = self
380            .token
381            .checked_add(1)
382            .ok_or_else(|| "Codex mailbox-flush token space was exhausted".to_owned())?;
383        let external_count = self.unsubmitted.len();
384        let mut values = self.unsubmitted.clone();
385        values.push(open_agent_response(self.promised_id()?));
386        let prepared = Arc::new(Prepared {
387            token,
388            values,
389            job,
390            external_count,
391        });
392        self.token = token;
393        self.prepared = Some(Arc::clone(&prepared));
394        if let Mode::Running(round) = &mut self.mode {
395            round.mailbox_flush_needed = false;
396        }
397        Ok(Some(PreparedMailboxFlush(prepared)))
398    }
399    pub fn validate_mailbox_flush(&self, value: &PreparedMailboxFlush) -> Result<(), String> {
400        let p = &value.0;
401        let prefix =
402            self.unsubmitted.get(..p.external_count) == Some(&p.values[..p.external_count]);
403        let valid = matches!(&self.mode,Mode::Running(round) if round.job==p.job&&round.phase==Phase::PendingGeneration)
404            && self.token == p.token
405            && prefix
406            && self
407                .prepared
408                .as_ref()
409                .is_some_and(|current| Arc::ptr_eq(current, p));
410        if valid {
411            Ok(())
412        } else {
413            Err("stale or invalid Codex mailbox flush".into())
414        }
415    }
416    pub fn commit_mailbox_flush(&mut self, value: PreparedMailboxFlush) -> Result<(), String> {
417        self.validate_mailbox_flush(&value)?;
418        self.unsubmitted.drain(..value.0.external_count);
419        self.prepared = None;
420        if let Mode::Running(round) = &mut self.mode {
421            round.phase = Phase::ProviderActive;
422        }
423        Ok(())
424    }
425
426    pub fn complete(&mut self, job: u64, output: ShimOutput<BoxValue>) -> Result<(), String> {
427        let mut round = self.take_round(job)?;
428        if round.phase != Phase::ProviderActive {
429            let m = "Codex inference completed outside a provider generation".to_owned();
430            self.preserve(round, m.clone(), false);
431            return Err(m);
432        }
433        if self.prepared.is_some() {
434            let m = "Codex inference completed with an uncommitted mailbox flush".to_owned();
435            self.preserve(round, m.clone(), false);
436            return Err(m);
437        }
438        let mut text = String::new();
439        for item in output.items {
440            match item {
441                ShimItem::Text(value) => text.push_str(&value),
442                ShimItem::Box(_) => {
443                    round.accepted_provider_action = true;
444                    let m = "terminal Codex output contains a box".to_owned();
445                    self.preserve(round, m.clone(), false);
446                    return Err(m);
447                }
448            }
449        }
450        let before = self.state.boxes().len();
451        if let Err(error) = self.state.complete_inference(job, text).map_err(debug) {
452            self.preserve(round, error.clone(), false);
453            return Err(error);
454        }
455        self.unsubmitted
456            .extend(self.state.boxes()[before..].iter().skip(1).map(project));
457        self.queued_trigger = false;
458        self.mode = Mode::Idle;
459        Ok(())
460    }
461    pub fn fail(&mut self, job: u64, message: String, restartable: bool) {
462        if let Ok(round) = self.take_round(job) {
463            self.preserve(round, message, restartable)
464        }
465    }
466    pub fn restart(&mut self) -> Result<(), RestartError> {
467        match &self.mode {
468            Mode::Stalled(_, true) => return Err(RestartError::ProviderActionAccepted),
469            Mode::Stalled(
470                Status::Stalled {
471                    restartable: true, ..
472                },
473                false,
474            ) => {}
475            _ => return Err(RestartError::NotStalled),
476        }
477        self.state.restart().map_err(|_| RestartError::NotStalled)?;
478        self.unsubmitted = self.state.boxes().iter().map(project).collect();
479        self.queued_trigger = false;
480        self.prepared = None;
481        self.mode = Mode::Idle;
482        Ok(())
483    }
484
485    fn accept_idle<F>(&mut self, accept: F) -> Result<(), String>
486    where
487        F: FnOnce(&mut ActorState) -> Result<(), StateError>,
488    {
489        if !matches!(self.mode, Mode::Idle) || self.prepared.is_some() {
490            return Err("startup context requires an idle conversation".into());
491        }
492        let before = self.state.boxes().len();
493        accept(&mut self.state).map_err(debug)?;
494        self.unsubmitted
495            .extend(self.state.boxes()[before..].iter().map(project));
496        Ok(())
497    }
498    fn accept_arrival<F>(&mut self, triggering: bool, accept: F) -> Result<(), String>
499    where
500        F: FnOnce(&mut ActorState) -> Result<(), StateError>,
501    {
502        let before = self.state.boxes().len();
503        accept(&mut self.state).map_err(debug)?;
504        let appended = &self.state.boxes()[before..];
505        self.unsubmitted.extend(appended.iter().map(project));
506        match &mut self.mode {
507            Mode::Running(round) => round.mailbox_flush_needed |= triggering,
508            Mode::Idle if appended.is_empty() => self.queued_trigger |= triggering,
509            Mode::Idle => self.queued_trigger = false,
510            Mode::Stalled(_, _) => {}
511        }
512        Ok(())
513    }
514    fn promised_id(&self) -> Result<BoxId, String> {
515        let prior = self.state.boxes().last().map_or(0, |b| b.id().get());
516        prior
517            .checked_add(1)
518            .map(BoxId::new)
519            .ok_or_else(|| "BoxId space was exhausted".into())
520    }
521    fn take_round(&mut self, job: u64) -> Result<Round, String> {
522        match std::mem::replace(&mut self.mode, Mode::Idle) {
523            Mode::Running(round) if round.job == job => Ok(round),
524            other => {
525                self.mode = other;
526                Err("stale Codex inference completion".into())
527            }
528        }
529    }
530    fn preserve(&mut self, round: Round, message: String, restartable_before_launch: bool) {
531        self.prepared = None;
532        let restartable = restartable_before_launch
533            && round.restartable
534            && !round.accepted_provider_action
535            && self
536                .state
537                .stall_inference(round.job, message.clone())
538                .is_ok();
539        if !restartable {
540            let _ = self.state.halt(message.clone());
541        }
542        self.mode = Mode::Stalled(
543            Status::Stalled {
544                message,
545                restartable,
546            },
547            round.accepted_provider_action,
548        );
549    }
550}
551fn next_sequence(value: u64) -> Result<u64, String> {
552    value
553        .checked_add(1)
554        .ok_or_else(|| "ToolCallId space was exhausted".into())
555}
556fn debug(error: impl std::fmt::Debug) -> String {
557    format!("{error:?}")
558}
559
560#[cfg(test)]
561mod tests {
562    use super::*;
563    #[test]
564    fn idle_startup_context_is_ordered_nontriggering_and_correlated() {
565        let mut state = ConversationState::new([7; 12]);
566        state
567            .accept_idle_context_box(
568                SYSTEM_MESSAGE_TYPE.into(),
569                "prefix".into(),
570                String::new(),
571                String::new(),
572            )
573            .unwrap();
574        let id = state
575            .accept_idle_context_tool_call("KmapOpenNode".into(), "{}".into())
576            .unwrap();
577        state
578            .accept_idle_context_tool_return(id, Ok("loaded".into()))
579            .unwrap();
580        assert_eq!(state.status(), Status::Quiet);
581        assert!(state.begin().unwrap().is_none());
582        state
583            .accept(
584                USER_MESSAGE_TYPE.into(),
585                "kickoff".into(),
586                String::new(),
587                String::new(),
588            )
589            .unwrap();
590        assert!(state.begin().unwrap().is_some());
591        assert_eq!(
592            state
593                .boxes()
594                .iter()
595                .map(ChatBox::box_type)
596                .collect::<Vec<_>>(),
597            [
598                SYSTEM_MESSAGE_TYPE,
599                TOOL_CALL_TYPE,
600                TOOL_RESULT_TYPE,
601                USER_MESSAGE_TYPE
602            ]
603        );
604    }
605}