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, PreflightItem, ProviderInput,
9    ProviderInputKind, Reply, 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_preflight_runtime::PreflightRuntime;
15use kcode_k1_chat_thread_session_code_runtime::SessionCodeRuntime;
16use kcode_k1_chat_thread_session_inference_settlement::settle_stage_inference;
17use kcode_k1_chat_thread_session_stage_runtime::{EventContext, StageRuntime};
18use kcode_k1_chat_thread_session_view::SessionView;
19use kcode_k1_codex_adapter::{Adapter, ShimOutput};
20use kcode_k1_codex_websearch::Runner as WebSearchRunner;
21use kcode_k1_rust_code_ktool_service::RustCodeKtoolService;
22use kcode_k1_web_code_ktool_service::K1WebCodeKtoolService;
23use std::sync::Arc;
24use tokio::{sync::mpsc, task::JoinHandle};
25
26type ActorResult = Result<bool, String>;
27type InferenceResult = Result<ShimOutput<BoxValue>, String>;
28type UnitResult = Result<(), String>;
29
30pub fn open(
31    adapter: Adapter,
32    key: impl Into<String>,
33    session: Session,
34    kmap: Arc<K1AccessKmap>,
35    web_search: WebSearchRunner,
36) -> Result<Handle, String> {
37    let durable = DurableThread::recover(session, kmap)?;
38    spawn_actor(adapter, key, durable, None, None, web_search)
39}
40
41pub fn open_with_social(
42    adapter: Adapter,
43    key: impl Into<String>,
44    session: Session,
45    kmap: Arc<K1AccessKmap>,
46    social: kcode_k1_ktool_social::SocialKtools,
47    web_search: WebSearchRunner,
48) -> Result<Handle, String> {
49    let durable = DurableThread::recover_with_social(session, kmap, social)?;
50    spawn_actor(adapter, key, durable, None, None, web_search)
51}
52
53pub fn open_with_social_and_set_launch_node(
54    adapter: Adapter,
55    key: impl Into<String>,
56    session: Session,
57    kmap: Arc<K1AccessKmap>,
58    social: kcode_k1_ktool_social::SocialKtools,
59    set_launch_node: SetLaunchNodeKtool,
60    web_search: WebSearchRunner,
61) -> Result<Handle, String> {
62    let durable = DurableThread::recover_with_social_and_set_launch_node(
63        session,
64        kmap,
65        social,
66        set_launch_node,
67    )?;
68    spawn_actor(adapter, key, durable, None, None, web_search)
69}
70
71#[allow(clippy::too_many_arguments)]
72pub fn open_with_social_and_set_launch_node_and_rust_code(
73    adapter: Adapter,
74    key: impl Into<String>,
75    session: Session,
76    kmap: Arc<K1AccessKmap>,
77    social: kcode_k1_ktool_social::SocialKtools,
78    set_launch_node: SetLaunchNodeKtool,
79    rust_code: Arc<RustCodeKtoolService>,
80    web_search: WebSearchRunner,
81) -> Result<Handle, String> {
82    let durable = DurableThread::recover_with_social_and_set_launch_node(
83        session,
84        kmap,
85        social,
86        set_launch_node,
87    )?;
88    spawn_actor(adapter, key, durable, Some(rust_code), None, web_search)
89}
90
91#[allow(clippy::too_many_arguments)]
92pub fn open_with_social_and_set_launch_node_and_rust_code_and_web_code(
93    adapter: Adapter,
94    key: impl Into<String>,
95    session: Session,
96    kmap: Arc<K1AccessKmap>,
97    social: kcode_k1_ktool_social::SocialKtools,
98    set_launch_node: SetLaunchNodeKtool,
99    rust_code: Arc<RustCodeKtoolService>,
100    web_code: K1WebCodeKtoolService,
101    web_search: WebSearchRunner,
102) -> Result<Handle, String> {
103    let durable = DurableThread::recover_with_social_and_set_launch_node(
104        session,
105        kmap,
106        social,
107        set_launch_node,
108    )?;
109    spawn_actor(
110        adapter,
111        key,
112        durable,
113        Some(rust_code),
114        Some(web_code),
115        web_search,
116    )
117}
118
119fn spawn_actor(
120    adapter: Adapter,
121    key: impl Into<String>,
122    durable: DurableThread,
123    rust_code: Option<Arc<RustCodeKtoolService>>,
124    web_code: Option<K1WebCodeKtoolService>,
125    web_search: WebSearchRunner,
126) -> Result<Handle, String> {
127    let code = SessionCodeRuntime::recover(&durable, rust_code, web_code)?;
128    let stage = StageRuntime::new(code, web_search);
129    let preflight = PreflightRuntime::recover(
130        durable.preflight_calls(),
131        durable.boxes(),
132        &durable.events(),
133    )?;
134    let (handle, sender, receiver, event_receiver) = channel_with_events();
135    let base_key = key.into();
136    let actor = Actor {
137        durable,
138        shim: Some(new_shim(adapter.clone(), base_key.clone(), sender.clone())),
139        adapter,
140        active_key: base_key.clone(),
141        base_key,
142        generation: 0,
143        stage,
144        preflight,
145        view: SessionView::new(event_receiver),
146        sender,
147        receiver,
148        usage: ModelUsageSession::default(),
149        inference: None,
150        job: None,
151    };
152    tokio::spawn(actor.run());
153    Ok(handle)
154}
155
156struct Actor {
157    durable: DurableThread,
158    shim: Option<ActorShim>,
159    adapter: Adapter,
160    base_key: String,
161    active_key: String,
162    generation: u64,
163    stage: StageRuntime,
164    preflight: PreflightRuntime,
165    view: SessionView,
166    sender: mpsc::UnboundedSender<Message>,
167    receiver: mpsc::UnboundedReceiver<Message>,
168    usage: ModelUsageSession,
169    inference: Option<JoinHandle<()>>,
170    job: Option<u64>,
171}
172
173impl Actor {
174    async fn run(mut self) {
175        loop {
176            if let Err(error) = self.drive().await {
177                self.fatal(error);
178                break;
179            }
180            let active = self.work_active();
181            let can_receive_code = self.stage.can_receive_code();
182            self.view.wake(&self.durable, active);
183            let stop = tokio::select! {
184                message = self.receiver.recv() => match message {
185                    Some(message) => match self.handle(message).await {
186                        Ok(stop) => stop,
187                        Err(error) => self.fatal(error),
188                    },
189                    None => true,
190                },
191                result = self.stage.receive_code(), if can_receive_code => match result {
192                    Ok(()) => match self.handle_code_completion() {
193                        Ok(stop) => stop,
194                        Err(error) => self.fatal(error),
195                    },
196                    Err(error) => self.fatal(error),
197                },
198                reply = self.view.receive_event_query() => {
199                    if let Some(reply) = reply {
200                        self.view.answer_event_query(&self.durable, active, reply);
201                    }
202                    false
203                }
204                usage = self.usage.receive(&mut self.durable) => usage.is_err(),
205            };
206            if stop {
207                break;
208            }
209        }
210        self.stage.abort();
211        let inference_active = self.inference.is_some();
212        let _ = self
213            .usage
214            .shutdown(&mut self.durable, inference_active)
215            .await;
216        self.stage.shutdown();
217        if let Some(task) = self.inference.take() {
218            task.abort();
219        }
220        self.view.close();
221    }
222
223    async fn handle(&mut self, message: Message) -> ActorResult {
224        match message {
225            Message::Accept((box_type, contents, hidden_type, hidden_contents), reply) => {
226                Ok(reply_transition(
227                    reply,
228                    self.durable.accept_external_box(
229                        box_type,
230                        contents,
231                        hidden_type,
232                        hidden_contents,
233                    ),
234                ))
235            }
236            Message::PreparePreflight(context, profile_id, policy, items, reply) => {
237                let result = self.prepare_preflight(context, profile_id, policy, items);
238                Ok(reply_transition(reply, result))
239            }
240            Message::ResumePreflight(context, profile_id, policy, reply) => {
241                let result = self.resume_preflight(context, profile_id, policy);
242                Ok(reply_transition(reply, result))
243            }
244            Message::AcceptUser(context, profile_id, policy, contents, reply) => {
245                let result = self.stage.accept_user(
246                    &mut self.durable,
247                    context,
248                    profile_id,
249                    policy,
250                    contents,
251                );
252                Ok(reply_transition(reply, result))
253            }
254            Message::Return(id, result, reply) => {
255                let result = self
256                    .durable
257                    .accept_return(id, result)
258                    .map_err(|_| ActorError::Closed);
259                let stop = result.is_err();
260                Ok(answer(reply, result, stop))
261            }
262            Message::Stage(text, boxes, reply) => {
263                let mut event = event_context(
264                    &mut self.durable,
265                    &self.adapter,
266                    &self.active_key,
267                    &mut self.view,
268                    &self.sender,
269                    self.job,
270                );
271                self.stage.handle_stage(&mut event, text, boxes, reply)
272            }
273            Message::PreflightCompleted(id, result) => {
274                let outcome = {
275                    let mut event = event_context(
276                        &mut self.durable,
277                        &self.adapter,
278                        &self.active_key,
279                        &mut self.view,
280                        &self.sender,
281                        self.job,
282                    );
283                    self.stage
284                        .handle_preflight_completion(&mut event, id, result)
285                }?;
286                self.preflight.complete(id)?;
287                if self.job.is_none() && !self.preflight.work_active() {
288                    self.durable.clear_authorization();
289                }
290                Ok(outcome)
291            }
292            Message::WebSearchCompleted(epoch, id, result) => {
293                let mut event = event_context(
294                    &mut self.durable,
295                    &self.adapter,
296                    &self.active_key,
297                    &mut self.view,
298                    &self.sender,
299                    self.job,
300                );
301                self.stage
302                    .handle_web_search_completion(&mut event, epoch, id, result)
303            }
304            Message::MailboxFlushCompleted(job, prepared, result) => {
305                let mut event = event_context(
306                    &mut self.durable,
307                    &self.adapter,
308                    &self.active_key,
309                    &mut self.view,
310                    &self.sender,
311                    self.job,
312                );
313                self.stage
314                    .handle_mailbox_flush_completion(&mut event, job, prepared, result)
315            }
316            Message::Inferred(job, shim, result) => self.finish(job, shim, result).await,
317            Message::Snapshot(reply) => {
318                let snapshot = self.view.snapshot(&self.durable, self.work_active());
319                Ok(answer(reply, Ok(snapshot), false))
320            }
321            Message::Wait(reply) => {
322                self.view.wait(&self.durable, self.work_active(), reply);
323                Ok(false)
324            }
325            Message::Restart(context, profile_id, policy, reply) => {
326                let result = self.restart(context, profile_id, policy);
327                if let Err(error) = result {
328                    return Ok(answer(reply, Err(error), false));
329                }
330                if self.usage.restart(&mut self.durable).await.is_err() {
331                    Ok(answer(reply, Err(ActorError::Closed), true))
332                } else {
333                    Ok(answer(reply, Ok(()), false))
334                }
335            }
336            Message::Abandon => Ok(true),
337        }
338    }
339
340    fn prepare_preflight(
341        &mut self,
342        context: AccessContext,
343        profile_id: ProfileId,
344        policy: AccessPolicy,
345        items: Vec<PreflightItem>,
346    ) -> Result<(), TransitionError> {
347        self.durable
348            .prepare_preflight(context, profile_id, policy, items)?;
349        self.refresh_preflight()
350            .map_err(TransitionError::Internal)?;
351        self.launch_preflight();
352        if !self.preflight.work_active() {
353            self.durable.clear_authorization();
354        }
355        Ok(())
356    }
357
358    fn resume_preflight(
359        &mut self,
360        context: AccessContext,
361        profile_id: ProfileId,
362        policy: AccessPolicy,
363    ) -> Result<(), TransitionError> {
364        self.refresh_preflight()
365            .map_err(TransitionError::Internal)?;
366        if !self.preflight.work_active() {
367            return Ok(());
368        }
369        self.durable
370            .authorize_preflight(context, profile_id, policy)?;
371        self.launch_preflight();
372        Ok(())
373    }
374
375    fn refresh_preflight(&mut self) -> Result<(), String> {
376        self.preflight.refresh(
377            self.durable.preflight_calls(),
378            self.durable.boxes(),
379            &self.durable.events(),
380        )
381    }
382
383    fn launch_preflight(&mut self) {
384        let executor = self.durable.preflight_executor();
385        for call in self.preflight.take_unlaunched() {
386            let executor = executor.clone();
387            let sender = self.sender.clone();
388            tokio::task::spawn_blocking(move || {
389                let result = executor.launch(&call.name, &call.arguments);
390                let _ = sender.send(Message::PreflightCompleted(call.tool_call_id, result));
391            });
392        }
393    }
394
395    fn work_active(&self) -> bool {
396        self.stage.search_active() || self.preflight.work_active()
397    }
398
399    fn handle_code_completion(&mut self) -> ActorResult {
400        let mut event = event_context(
401            &mut self.durable,
402            &self.adapter,
403            &self.active_key,
404            &mut self.view,
405            &self.sender,
406            self.job,
407        );
408        self.stage.handle_code_completion(&mut event)
409    }
410
411    fn fatal(&mut self, error: String) -> bool {
412        self.stage.fail(error)
413    }
414
415    async fn drive(&mut self) -> UnitResult {
416        if self.inference.is_some()
417            || self.shim.is_none()
418            || !self.preflight.allows_inference(self.durable.boxes())
419        {
420            return Ok(());
421        }
422        let Some((job, input)) = self.durable.begin_input()? else {
423            return Ok(());
424        };
425        if !self.usage.is_subscribed() {
426            self.usage
427                .subscribe(&self.adapter, self.active_key.clone())
428                .await
429                .map_err(|_| "model-usage subscription failed".to_owned())?;
430        }
431        self.view.set_model_input(ProviderInput {
432            kind: ProviderInputKind::Turn,
433            text: input.clone(),
434        });
435        let mut shim = self.shim.take().expect("shim was checked");
436        self.job = Some(job);
437        let sender = self.sender.clone();
438        self.inference = Some(tokio::spawn(async move {
439            let result = shim.infer(input).await.map_err(|error| error.to_string());
440            let _ = sender.send(Message::Inferred(job, shim, result));
441        }));
442        Ok(())
443    }
444
445    async fn finish(&mut self, job: u64, shim: ActorShim, result: InferenceResult) -> ActorResult {
446        if self.inference.is_none() || self.stage.mailbox_pending() || self.job != Some(job) {
447            return Err("stale Codex inference completion".to_owned());
448        }
449        drop(self.inference.take());
450        self.job = None;
451        let active_key = self.active_key.clone();
452        let settlement = settle_stage_inference(
453            &mut self.durable,
454            &mut self.usage,
455            &mut self.stage,
456            &active_key,
457            job,
458            result,
459        )
460        .await?;
461        let should_stop = settlement.should_stop();
462        let restore_shim = settlement.restore_shim();
463        let stage_error = settlement.into_stage_error();
464        if restore_shim {
465            self.shim = Some(shim);
466        }
467        if let Some(error) = stage_error {
468            self.stage.fail(error);
469        }
470        Ok(should_stop)
471    }
472
473    fn restart(
474        &mut self,
475        context: AccessContext,
476        profile_id: ProfileId,
477        policy: AccessPolicy,
478    ) -> Result<(), ActorError> {
479        let generation = self
480            .generation
481            .checked_add(1)
482            .ok_or(ActorError::NotRestartable)?;
483        let active_key = format!("{}#restart-{generation}", self.base_key);
484        let shim = new_shim(
485            self.adapter.clone(),
486            active_key.clone(),
487            self.sender.clone(),
488        );
489        self.stage
490            .restart(&mut self.durable, context, profile_id, policy)?;
491        self.generation = generation;
492        self.active_key = active_key;
493        self.shim = Some(shim);
494        Ok(())
495    }
496}
497
498fn event_context<'a>(
499    durable: &'a mut DurableThread,
500    adapter: &'a Adapter,
501    active_key: &'a str,
502    view: &'a mut SessionView,
503    sender: &'a mpsc::UnboundedSender<Message>,
504    job: Option<u64>,
505) -> EventContext<'a> {
506    EventContext {
507        durable,
508        adapter,
509        active_key,
510        view,
511        sender,
512        job,
513    }
514}
515
516fn map_transition(error: TransitionError) -> ActorError {
517    match error {
518        TransitionError::Unauthorized => ActorError::Unauthorized,
519        TransitionError::NotStalled => ActorError::NotStalled,
520        TransitionError::NotRestartable => ActorError::NotRestartable,
521        TransitionError::Internal(_) => ActorError::Closed,
522    }
523}
524
525fn reply_transition(reply: Reply<()>, result: Result<(), TransitionError>) -> bool {
526    let stop = matches!(result, Err(TransitionError::Internal(_)));
527    answer(reply, result.map_err(map_transition), stop)
528}
529
530fn answer<T>(reply: Reply<T>, result: Result<T, ActorError>, stop: bool) -> bool {
531    let _ = reply.send(result);
532    stop
533}