Skip to main content

kcode_k1_chat_thread_session_actor/
lib.rs

1#![forbid(unsafe_code)]
2#![doc = include_str!("../Documentation.md")]
3
4use kcode_k1_access_kmap::K1AccessKmap;
5use kcode_k1_chat_model_usage_session::ModelUsageSession;
6use kcode_k1_chat_persistence::Session;
7use kcode_k1_chat_thread_actor_channel::{
8    ActorError, ActorShim, BoxValue, Handle, Message, ProviderInput, ProviderInputKind, Reply,
9    channel_with_events, new_shim,
10};
11use kcode_k1_chat_thread_durable_state::{
12    AccessContext, AccessPolicy, DurableThread, ProfileId, SetLaunchNodeKtool, TransitionError,
13};
14use kcode_k1_chat_thread_session_code_runtime::SessionCodeRuntime;
15use kcode_k1_chat_thread_session_stage_runtime::{EventContext, StageRuntime};
16use kcode_k1_chat_thread_session_view::SessionView;
17use kcode_k1_codex_adapter::{Adapter, ShimOutput};
18use kcode_k1_codex_websearch::Runner as WebSearchRunner;
19use kcode_k1_rust_code_ktool_service::RustCodeKtoolService;
20use kcode_k1_web_code_ktool_service::K1WebCodeKtoolService;
21use std::sync::Arc;
22use tokio::{sync::mpsc, task::JoinHandle};
23
24type ActorResult = Result<bool, String>;
25type InferenceResult = Result<ShimOutput<BoxValue>, String>;
26type UnitResult = Result<(), String>;
27
28pub fn open(
29    adapter: Adapter,
30    key: impl Into<String>,
31    session: Session,
32    kmap: Arc<K1AccessKmap>,
33    web_search: WebSearchRunner,
34) -> Result<Handle, String> {
35    let durable = DurableThread::recover(session, kmap)?;
36    spawn_actor(adapter, key, durable, None, None, web_search)
37}
38
39pub fn open_with_social(
40    adapter: Adapter,
41    key: impl Into<String>,
42    session: Session,
43    kmap: Arc<K1AccessKmap>,
44    social: kcode_k1_ktool_social::SocialKtools,
45    web_search: WebSearchRunner,
46) -> Result<Handle, String> {
47    let durable = DurableThread::recover_with_social(session, kmap, social)?;
48    spawn_actor(adapter, key, durable, None, None, web_search)
49}
50
51pub fn open_with_social_and_set_launch_node(
52    adapter: Adapter,
53    key: impl Into<String>,
54    session: Session,
55    kmap: Arc<K1AccessKmap>,
56    social: kcode_k1_ktool_social::SocialKtools,
57    set_launch_node: SetLaunchNodeKtool,
58    web_search: WebSearchRunner,
59) -> Result<Handle, String> {
60    let durable = DurableThread::recover_with_social_and_set_launch_node(
61        session,
62        kmap,
63        social,
64        set_launch_node,
65    )?;
66    spawn_actor(adapter, key, durable, None, None, web_search)
67}
68
69#[allow(clippy::too_many_arguments)]
70pub fn open_with_social_and_set_launch_node_and_rust_code(
71    adapter: Adapter,
72    key: impl Into<String>,
73    session: Session,
74    kmap: Arc<K1AccessKmap>,
75    social: kcode_k1_ktool_social::SocialKtools,
76    set_launch_node: SetLaunchNodeKtool,
77    rust_code: Arc<RustCodeKtoolService>,
78    web_search: WebSearchRunner,
79) -> Result<Handle, String> {
80    let durable = DurableThread::recover_with_social_and_set_launch_node(
81        session,
82        kmap,
83        social,
84        set_launch_node,
85    )?;
86    spawn_actor(adapter, key, durable, Some(rust_code), None, web_search)
87}
88
89#[allow(clippy::too_many_arguments)]
90pub fn open_with_social_and_set_launch_node_and_rust_code_and_web_code(
91    adapter: Adapter,
92    key: impl Into<String>,
93    session: Session,
94    kmap: Arc<K1AccessKmap>,
95    social: kcode_k1_ktool_social::SocialKtools,
96    set_launch_node: SetLaunchNodeKtool,
97    rust_code: Arc<RustCodeKtoolService>,
98    web_code: K1WebCodeKtoolService,
99    web_search: WebSearchRunner,
100) -> Result<Handle, String> {
101    let durable = DurableThread::recover_with_social_and_set_launch_node(
102        session,
103        kmap,
104        social,
105        set_launch_node,
106    )?;
107    spawn_actor(
108        adapter,
109        key,
110        durable,
111        Some(rust_code),
112        Some(web_code),
113        web_search,
114    )
115}
116
117fn spawn_actor(
118    adapter: Adapter,
119    key: impl Into<String>,
120    durable: DurableThread,
121    rust_code: Option<Arc<RustCodeKtoolService>>,
122    web_code: Option<K1WebCodeKtoolService>,
123    web_search: WebSearchRunner,
124) -> Result<Handle, String> {
125    let code = SessionCodeRuntime::recover(&durable, rust_code, web_code)?;
126    let stage = StageRuntime::new(code, web_search);
127    let (handle, sender, receiver, event_receiver) = channel_with_events();
128    let base_key = key.into();
129    let actor = Actor {
130        durable,
131        shim: Some(new_shim(adapter.clone(), base_key.clone(), sender.clone())),
132        adapter,
133        active_key: base_key.clone(),
134        base_key,
135        generation: 0,
136        stage,
137        view: SessionView::new(event_receiver),
138        sender,
139        receiver,
140        usage: ModelUsageSession::default(),
141        inference: None,
142        job: None,
143    };
144    tokio::spawn(actor.run());
145    Ok(handle)
146}
147
148struct Actor {
149    durable: DurableThread,
150    shim: Option<ActorShim>,
151    adapter: Adapter,
152    base_key: String,
153    active_key: String,
154    generation: u64,
155    stage: StageRuntime,
156    view: SessionView,
157    sender: mpsc::UnboundedSender<Message>,
158    receiver: mpsc::UnboundedReceiver<Message>,
159    usage: ModelUsageSession,
160    inference: Option<JoinHandle<()>>,
161    job: Option<u64>,
162}
163
164impl Actor {
165    async fn run(mut self) {
166        loop {
167            if let Err(error) = self.drive().await {
168                self.fatal(error);
169                break;
170            }
171            let searching = self.stage.search_active();
172            let can_receive_code = self.stage.can_receive_code();
173            self.view.wake(&self.durable, searching);
174            let stop = tokio::select! {
175                message = self.receiver.recv() => match message {
176                    Some(message) => match self.handle(message).await {
177                        Ok(stop) => stop,
178                        Err(error) => self.fatal(error),
179                    },
180                    None => true,
181                },
182                result = self.stage.receive_code(), if can_receive_code => match result {
183                    Ok(()) => match self.handle_code_completion() {
184                        Ok(stop) => stop,
185                        Err(error) => self.fatal(error),
186                    },
187                    Err(error) => self.fatal(error),
188                },
189                reply = self.view.receive_event_query() => {
190                    if let Some(reply) = reply {
191                        self.view.answer_event_query(&self.durable, searching, reply);
192                    }
193                    false
194                }
195                usage = self.usage.receive(&mut self.durable) => usage.is_err(),
196            };
197            if stop {
198                break;
199            }
200        }
201        self.stage.abort();
202        let inference_active = self.inference.is_some();
203        let _ = self
204            .usage
205            .shutdown(&mut self.durable, inference_active)
206            .await;
207        self.stage.shutdown();
208        if let Some(task) = self.inference.take() {
209            task.abort();
210        }
211        self.view.close();
212    }
213
214    async fn handle(&mut self, message: Message) -> ActorResult {
215        match message {
216            Message::Accept((box_type, contents, hidden_type, hidden_contents), reply) => {
217                Ok(reply_transition(
218                    reply,
219                    self.durable.accept_external_box(
220                        box_type,
221                        contents,
222                        hidden_type,
223                        hidden_contents,
224                    ),
225                ))
226            }
227            Message::AcceptUser(context, profile_id, policy, contents, reply) => {
228                let result = self.stage.accept_user(
229                    &mut self.durable,
230                    context,
231                    profile_id,
232                    policy,
233                    contents,
234                );
235                Ok(reply_transition(reply, result))
236            }
237            Message::Return(id, result, reply) => {
238                let result = self
239                    .durable
240                    .accept_return(id, result)
241                    .map_err(|_| ActorError::Closed);
242                let stop = result.is_err();
243                Ok(answer(reply, result, stop))
244            }
245            Message::Stage(text, boxes, reply) => {
246                let mut event = EventContext {
247                    durable: &mut self.durable,
248                    adapter: &self.adapter,
249                    active_key: &self.active_key,
250                    view: &mut self.view,
251                    sender: &self.sender,
252                    job: self.job,
253                };
254                self.stage.handle_stage(&mut event, text, boxes, reply)
255            }
256            Message::WebSearchCompleted(epoch, id, result) => {
257                let mut event = EventContext {
258                    durable: &mut self.durable,
259                    adapter: &self.adapter,
260                    active_key: &self.active_key,
261                    view: &mut self.view,
262                    sender: &self.sender,
263                    job: self.job,
264                };
265                self.stage
266                    .handle_web_search_completion(&mut event, epoch, id, result)
267            }
268            Message::MailboxFlushCompleted(job, prepared, result) => {
269                let mut event = EventContext {
270                    durable: &mut self.durable,
271                    adapter: &self.adapter,
272                    active_key: &self.active_key,
273                    view: &mut self.view,
274                    sender: &self.sender,
275                    job: self.job,
276                };
277                self.stage
278                    .handle_mailbox_flush_completion(&mut event, job, prepared, result)
279            }
280            Message::Inferred(job, shim, result) => self.finish(job, shim, result).await,
281            Message::Snapshot(reply) => {
282                let snapshot = self
283                    .view
284                    .snapshot(&self.durable, self.stage.search_active());
285                Ok(answer(reply, Ok(snapshot), false))
286            }
287            Message::Wait(reply) => {
288                self.view
289                    .wait(&self.durable, self.stage.search_active(), reply);
290                Ok(false)
291            }
292            Message::Restart(context, profile_id, policy, reply) => {
293                let result = self.restart(context, profile_id, policy);
294                if let Err(error) = result {
295                    return Ok(answer(reply, Err(error), false));
296                }
297                if self.usage.restart(&mut self.durable).await.is_err() {
298                    Ok(answer(reply, Err(ActorError::Closed), true))
299                } else {
300                    Ok(answer(reply, Ok(()), false))
301                }
302            }
303            Message::Abandon => Ok(true),
304        }
305    }
306
307    fn handle_code_completion(&mut self) -> ActorResult {
308        let mut event = EventContext {
309            durable: &mut self.durable,
310            adapter: &self.adapter,
311            active_key: &self.active_key,
312            view: &mut self.view,
313            sender: &self.sender,
314            job: self.job,
315        };
316        self.stage.handle_code_completion(&mut event)
317    }
318
319    fn fatal(&mut self, error: String) -> bool {
320        self.stage.fail(error)
321    }
322
323    async fn drive(&mut self) -> UnitResult {
324        if self.inference.is_some() || self.shim.is_none() {
325            return Ok(());
326        }
327        let Some((job, input)) = self.durable.begin_input()? else {
328            return Ok(());
329        };
330        if !self.usage.is_subscribed() {
331            self.usage
332                .subscribe(&self.adapter, self.active_key.clone())
333                .await
334                .map_err(|_| "model-usage subscription failed".to_owned())?;
335        }
336        self.view.set_model_input(ProviderInput {
337            kind: ProviderInputKind::Turn,
338            text: input.clone(),
339        });
340        let mut shim = self.shim.take().expect("shim was checked");
341        self.job = Some(job);
342        let sender = self.sender.clone();
343        self.inference = Some(tokio::spawn(async move {
344            let result = shim.infer(input).await.map_err(|error| error.to_string());
345            let _ = sender.send(Message::Inferred(job, shim, result));
346        }));
347        Ok(())
348    }
349
350    async fn finish(&mut self, job: u64, shim: ActorShim, result: InferenceResult) -> ActorResult {
351        if self.inference.take().is_none() || self.stage.mailbox_pending() || self.job != Some(job)
352        {
353            return Err("stale Codex inference completion".to_owned());
354        }
355        if self.stage.code_pending() {
356            return Err("Codex inference completed while Rust code work remains".to_owned());
357        }
358        self.job = None;
359        let key = self.active_key.clone();
360        let (error, resume, terminal_id) = match result {
361            Err(error) => {
362                self.stage.abort();
363                (Some(error), false, None)
364            }
365            Ok(output) => {
366                let (resume, terminal) =
367                    self.durable.complete_with_terminal_response(job, output)?;
368                (None, resume, Some(terminal))
369            }
370        };
371        if self
372            .usage
373            .finish_inference(&mut self.durable, &key, terminal_id)
374            .await
375            .is_err()
376        {
377            return Ok(true);
378        }
379        if let Some(error) = error {
380            self.stage.fail(error.clone());
381            self.durable.fail(job, error, true);
382            return Ok(false);
383        }
384        self.shim = Some(shim);
385        if !resume && !self.stage.search_active() {
386            self.stage.clear_authorization(&mut self.durable);
387        }
388        Ok(false)
389    }
390
391    fn restart(
392        &mut self,
393        context: AccessContext,
394        profile_id: ProfileId,
395        policy: AccessPolicy,
396    ) -> Result<(), ActorError> {
397        let generation = self
398            .generation
399            .checked_add(1)
400            .ok_or(ActorError::NotRestartable)?;
401        self.stage
402            .restart(&mut self.durable, context, profile_id, policy)?;
403        self.generation = generation;
404        self.active_key = format!("{}#restart-{generation}", self.base_key);
405        self.shim = Some(new_shim(
406            self.adapter.clone(),
407            self.active_key.clone(),
408            self.sender.clone(),
409        ));
410        Ok(())
411    }
412}
413
414fn map_transition(error: TransitionError) -> ActorError {
415    match error {
416        TransitionError::Unauthorized => ActorError::Unauthorized,
417        TransitionError::NotStalled => ActorError::NotStalled,
418        TransitionError::NotRestartable => ActorError::NotRestartable,
419        TransitionError::Internal(_) => ActorError::Closed,
420    }
421}
422
423fn reply_transition(reply: Reply<()>, result: Result<(), TransitionError>) -> bool {
424    let stop = matches!(result, Err(TransitionError::Internal(_)));
425    answer(reply, result.map_err(map_transition), stop)
426}
427
428fn answer<T>(reply: Reply<T>, result: Result<T, ActorError>, stop: bool) -> bool {
429    let _ = reply.send(result);
430    stop
431}