Skip to main content

kcode_k1_chat_thread_durable_state/
lib.rs

1#![forbid(unsafe_code)]
2
3use kcode_k1_access_kmap::K1AccessKmap;
4use kcode_k1_chat_persistence::Session;
5use kcode_k1_chat_state::USER_MESSAGE_TYPE;
6use kcode_k1_chat_thread_actions::ChatThreadActions;
7pub use kcode_k1_chat_thread_actions::{AccessContext, AccessPolicy, ProfileId};
8pub use kcode_k1_chat_thread_durable_turn::{
9    BoxValue, ChatBox, PreparedCall, PreparedSteer, Status, ToolCallId,
10};
11use kcode_k1_chat_thread_durable_turn::{DurableTurn, RestartError, ShimOutput};
12use std::sync::Arc;
13
14#[derive(Clone, Debug, Eq, PartialEq)]
15pub enum TransitionError {
16    Unauthorized,
17    NotStalled,
18    NotRestartable,
19    Internal(String),
20}
21
22pub struct DurableThread {
23    turn: DurableTurn,
24    actions: ChatThreadActions,
25    authorized: bool,
26}
27
28impl DurableThread {
29    pub fn recover(session: Session, kmap: Arc<K1AccessKmap>) -> Result<Self, String> {
30        Ok(Self {
31            turn: DurableTurn::recover(session)?,
32            actions: ChatThreadActions::new(kmap),
33            authorized: false,
34        })
35    }
36
37    pub fn boxes(&self) -> &[ChatBox] {
38        self.turn.boxes()
39    }
40
41    pub fn status(&self) -> Status {
42        self.turn.status()
43    }
44
45    pub fn accept_box(
46        &mut self,
47        box_type: String,
48        contents: String,
49        hidden_type: String,
50        hidden_contents: String,
51    ) -> Result<(), String> {
52        self.turn
53            .accept(box_type, contents, hidden_type, hidden_contents)
54    }
55
56    pub fn accept_user(
57        &mut self,
58        context: AccessContext,
59        profile_id: ProfileId,
60        policy: AccessPolicy,
61        contents: String,
62    ) -> Result<(), TransitionError> {
63        let installed = self.bind_authorization(context, profile_id, policy)?;
64        match self.turn.accept(
65            USER_MESSAGE_TYPE.into(),
66            contents,
67            String::new(),
68            String::new(),
69        ) {
70            Ok(()) => Ok(()),
71            Err(error) => {
72                if installed {
73                    self.clear_authorization();
74                }
75                Err(TransitionError::Internal(error))
76            }
77        }
78    }
79
80    pub fn accept_return(
81        &mut self,
82        id: ToolCallId,
83        result: Result<String, String>,
84    ) -> Result<(), String> {
85        self.turn.accept_tool_return(id, result)
86    }
87
88    pub fn prepare_stage(
89        &mut self,
90        job: u64,
91        text: String,
92        boxes: Vec<BoxValue>,
93    ) -> Result<Vec<PreparedCall>, String> {
94        self.turn.prepare_stage(job, text, boxes)
95    }
96
97    pub fn launch_action(&mut self, name: &str, arguments: &str) -> Result<String, String> {
98        self.actions.launch(name, arguments)
99    }
100
101    pub fn accept_tool_message(&mut self, id: ToolCallId, contents: String) -> Result<(), String> {
102        self.turn.accept_tool_message(id, contents)
103    }
104
105    pub fn accept_tool_return(
106        &mut self,
107        id: ToolCallId,
108        result: Result<String, String>,
109    ) -> Result<(), String> {
110        self.turn.accept_tool_return(id, result)
111    }
112
113    pub fn accept_tool_return_v2(
114        &mut self,
115        id: ToolCallId,
116        result: Result<String, String>,
117        metadata_type: String,
118        metadata_contents: String,
119    ) -> Result<(), String> {
120        self.turn
121            .accept_tool_return_v2(id, result, metadata_type, metadata_contents)
122    }
123
124    pub fn prepare_steer(&mut self, job: u64) -> Result<Option<PreparedSteer>, String> {
125        self.turn.prepare_steer(job)
126    }
127
128    pub fn prepared_input(&self, prepared: &PreparedSteer) -> Result<String, String> {
129        self.turn.validate_steer(prepared)?;
130        render_input(prepared.values())
131    }
132
133    pub fn commit_steer(&mut self, prepared: PreparedSteer) -> Result<(), String> {
134        self.turn.commit_steer(prepared)
135    }
136
137    pub fn begin_input(&mut self) -> Result<Option<(u64, String)>, String> {
138        let Some(start) = self.turn.begin()? else {
139            return Ok(None);
140        };
141        Ok(Some((start.job, render_input(&start.values)?)))
142    }
143
144    pub fn complete(&mut self, job: u64, output: ShimOutput<BoxValue>) -> Result<bool, String> {
145        self.turn.complete(job, output)
146    }
147
148    pub fn fail(&mut self, job: u64, error: String, restartable: bool) {
149        self.turn.fail(job, error, restartable);
150        self.clear_authorization();
151    }
152
153    pub fn restart(
154        &mut self,
155        context: AccessContext,
156        profile_id: ProfileId,
157        policy: AccessPolicy,
158    ) -> Result<(), TransitionError> {
159        let installed = self.bind_authorization(context, profile_id, policy)?;
160        if let Err(error) = self.turn.restart().map_err(|error| match error {
161            RestartError::NotStalled => TransitionError::NotStalled,
162            RestartError::ProviderActionAccepted => TransitionError::NotRestartable,
163        }) {
164            if installed {
165                self.clear_authorization();
166            }
167            return Err(error);
168        }
169        Ok(())
170    }
171
172    pub fn clear_authorization(&mut self) {
173        self.actions.clear_authorization();
174        self.authorized = false;
175    }
176
177    fn bind_authorization(
178        &mut self,
179        context: AccessContext,
180        profile_id: ProfileId,
181        policy: AccessPolicy,
182    ) -> Result<bool, TransitionError> {
183        let installed = !self.authorized;
184        if self
185            .actions
186            .bind_authorization(context, profile_id, policy)
187            .is_err()
188        {
189            if installed {
190                self.actions.clear_authorization();
191            }
192            return Err(TransitionError::Unauthorized);
193        }
194        self.authorized = true;
195        Ok(installed)
196    }
197}
198
199fn render_input(values: &[BoxValue]) -> Result<String, String> {
200    let mut output = String::new();
201    for value in values {
202        let BoxValue::History(section) = value else {
203            return Err("Codex provider input contains a non-history value".into());
204        };
205        if section.is_empty() {
206            continue;
207        }
208        if !output.is_empty() && !output.ends_with('\n') {
209            output.push('\n');
210        }
211        output.push_str(section);
212    }
213    Ok(output)
214}