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