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, ChatDiagnostic, DurableThread, ProfileId, SetLaunchNodeKtool,
13    TransitionError,
14};
15use kcode_k1_chat_thread_preflight_runtime::PreflightRuntime;
16use kcode_k1_chat_thread_session_code_runtime::SessionCodeRuntime;
17use kcode_k1_chat_thread_session_inference_settlement::{
18    settle_stage_failure, settle_stage_inference,
19};
20use kcode_k1_chat_thread_session_stage_runtime::{EventContext, StageRuntime};
21use kcode_k1_chat_thread_session_view::SessionView;
22use kcode_k1_codex_adapter::{Adapter, ShimOutput};
23use kcode_k1_codex_websearch::Runner as WebSearchRunner;
24use kcode_k1_rust_code_ktool_service::RustCodeKtoolService;
25use kcode_k1_web_code_ktool_service::K1WebCodeKtoolService;
26use std::sync::Arc;
27use tokio::{sync::mpsc, task::JoinHandle};
28
29type ActorResult = Result<bool, String>;
30type InferenceResult = Result<ShimOutput<BoxValue>, String>;
31type UnitResult = Result<(), String>;
32
33const SAFE_CRITICAL_FAILURE: &str =
34    "This chat stopped because an internal integrity failure occurred.";
35
36#[derive(Clone, Copy, Debug, Eq, PartialEq)]
37enum FailureSeverity {
38    OptionalDiagnostic,
39    RecoverableTurn,
40    CriticalIntegrity,
41}
42
43#[derive(Clone, Copy, Debug, Eq, PartialEq)]
44struct FailureClass {
45    severity: FailureSeverity,
46    diagnostic: ChatDiagnostic,
47}
48
49impl FailureClass {
50    fn optional(diagnostic: ChatDiagnostic) -> Self {
51        Self {
52            severity: FailureSeverity::OptionalDiagnostic,
53            diagnostic,
54        }
55    }
56
57    fn recoverable(diagnostic: ChatDiagnostic) -> Self {
58        Self {
59            severity: FailureSeverity::RecoverableTurn,
60            diagnostic,
61        }
62    }
63
64    fn critical() -> Self {
65        Self {
66            severity: FailureSeverity::CriticalIntegrity,
67            diagnostic: ChatDiagnostic::CriticalIntegrity,
68        }
69    }
70}
71
72pub fn open(
73    adapter: Adapter,
74    key: impl Into<String>,
75    session: Session,
76    kmap: Arc<K1AccessKmap>,
77    web_search: WebSearchRunner,
78) -> Result<Handle, String> {
79    let durable = DurableThread::recover(session, kmap)?;
80    spawn_actor(adapter, key, durable, None, None, web_search)
81}
82
83pub fn open_with_social(
84    adapter: Adapter,
85    key: impl Into<String>,
86    session: Session,
87    kmap: Arc<K1AccessKmap>,
88    social: kcode_k1_ktool_social::SocialKtools,
89    web_search: WebSearchRunner,
90) -> Result<Handle, String> {
91    let durable = DurableThread::recover_with_social(session, kmap, social)?;
92    spawn_actor(adapter, key, durable, None, None, web_search)
93}
94
95pub fn open_with_social_and_set_launch_node(
96    adapter: Adapter,
97    key: impl Into<String>,
98    session: Session,
99    kmap: Arc<K1AccessKmap>,
100    social: kcode_k1_ktool_social::SocialKtools,
101    set_launch_node: SetLaunchNodeKtool,
102    web_search: WebSearchRunner,
103) -> Result<Handle, String> {
104    let durable = DurableThread::recover_with_social_and_set_launch_node(
105        session,
106        kmap,
107        social,
108        set_launch_node,
109    )?;
110    spawn_actor(adapter, key, durable, None, None, web_search)
111}
112
113#[allow(clippy::too_many_arguments)]
114pub fn open_with_social_and_set_launch_node_and_rust_code(
115    adapter: Adapter,
116    key: impl Into<String>,
117    session: Session,
118    kmap: Arc<K1AccessKmap>,
119    social: kcode_k1_ktool_social::SocialKtools,
120    set_launch_node: SetLaunchNodeKtool,
121    rust_code: Arc<RustCodeKtoolService>,
122    web_search: WebSearchRunner,
123) -> Result<Handle, String> {
124    let durable = DurableThread::recover_with_social_and_set_launch_node(
125        session,
126        kmap,
127        social,
128        set_launch_node,
129    )?;
130    spawn_actor(adapter, key, durable, Some(rust_code), None, web_search)
131}
132
133#[allow(clippy::too_many_arguments)]
134pub fn open_with_social_and_set_launch_node_and_rust_code_and_web_code(
135    adapter: Adapter,
136    key: impl Into<String>,
137    session: Session,
138    kmap: Arc<K1AccessKmap>,
139    social: kcode_k1_ktool_social::SocialKtools,
140    set_launch_node: SetLaunchNodeKtool,
141    rust_code: Arc<RustCodeKtoolService>,
142    web_code: K1WebCodeKtoolService,
143    web_search: WebSearchRunner,
144) -> Result<Handle, String> {
145    let durable = DurableThread::recover_with_social_and_set_launch_node(
146        session,
147        kmap,
148        social,
149        set_launch_node,
150    )?;
151    spawn_actor(
152        adapter,
153        key,
154        durable,
155        Some(rust_code),
156        Some(web_code),
157        web_search,
158    )
159}
160
161fn spawn_actor(
162    adapter: Adapter,
163    key: impl Into<String>,
164    durable: DurableThread,
165    rust_code: Option<Arc<RustCodeKtoolService>>,
166    web_code: Option<K1WebCodeKtoolService>,
167    web_search: WebSearchRunner,
168) -> Result<Handle, String> {
169    let code = SessionCodeRuntime::recover(&durable, rust_code, web_code)?;
170    let stage = StageRuntime::new(code, web_search);
171    let preflight = PreflightRuntime::recover(
172        durable.preflight_calls(),
173        durable.boxes(),
174        &durable.events(),
175    )?;
176    let (handle, sender, receiver, event_receiver) = channel_with_events();
177    let base_key = key.into();
178    let actor = Actor {
179        durable,
180        shim: Some(new_shim(adapter.clone(), base_key.clone(), sender.clone())),
181        adapter,
182        active_key: base_key.clone(),
183        base_key,
184        generation: 0,
185        stage,
186        preflight,
187        view: SessionView::new(event_receiver),
188        sender,
189        receiver,
190        usage: ModelUsageSession::default(),
191        usage_enabled: true,
192        inference: None,
193        job: None,
194    };
195    tokio::spawn(actor.run());
196    Ok(handle)
197}
198
199struct Actor {
200    durable: DurableThread,
201    shim: Option<ActorShim>,
202    adapter: Adapter,
203    base_key: String,
204    active_key: String,
205    generation: u64,
206    stage: StageRuntime,
207    preflight: PreflightRuntime,
208    view: SessionView,
209    sender: mpsc::UnboundedSender<Message>,
210    receiver: mpsc::UnboundedReceiver<Message>,
211    usage: ModelUsageSession,
212    usage_enabled: bool,
213    inference: Option<JoinHandle<()>>,
214    job: Option<u64>,
215}
216
217impl Actor {
218    async fn run(mut self) {
219        loop {
220            if let Err(error) = self.drive().await {
221                self.critical(error);
222                break;
223            }
224            let active = self.work_active();
225            let can_receive_code = self.stage.can_receive_code();
226            self.view.wake(&self.durable, active);
227            let stop = tokio::select! {
228                message = self.receiver.recv() => match message {
229                    Some(message) => match self.handle(message).await {
230                        Ok(stop) => stop,
231                        Err(error) => self.critical(error),
232                    },
233                    None => true,
234                },
235                result = self.stage.receive_code(), if can_receive_code => match result {
236                    Ok(()) => match self.handle_code_completion() {
237                        Ok(stop) => stop,
238                        Err(error) => self.critical(error),
239                    },
240                    Err(error) => self.critical(error),
241                },
242                reply = self.view.receive_event_query() => {
243                    if let Some(reply) = reply {
244                        self.view.answer_event_query(&self.durable, active, reply);
245                    }
246                    false
247                }
248                usage = self.usage.receive(&mut self.durable), if self.usage_enabled => {
249                    match usage {
250                        Ok(_) => false,
251                        Err(_) => {
252                            self.disable_usage(FailureClass::optional(
253                                ChatDiagnostic::ModelUsageReceive,
254                            ));
255                            false
256                        }
257                    }
258                },
259            };
260            if stop {
261                break;
262            }
263        }
264        self.stage.abort();
265        let inference_active = self.inference.is_some();
266        if self.usage_enabled
267            && self
268                .usage
269                .shutdown(&mut self.durable, inference_active)
270                .await
271                .is_err()
272        {
273            self.disable_usage(FailureClass::optional(ChatDiagnostic::ModelUsageShutdown));
274        }
275        self.stage.shutdown();
276        if let Some(task) = self.inference.take() {
277            task.abort();
278        }
279        self.view.close();
280    }
281
282    async fn handle(&mut self, message: Message) -> ActorResult {
283        match message {
284            Message::Accept((box_type, contents, hidden_type, hidden_contents), reply) => {
285                reply_transition(
286                    reply,
287                    self.durable.accept_external_box(
288                        box_type,
289                        contents,
290                        hidden_type,
291                        hidden_contents,
292                    ),
293                )
294            }
295            Message::PreparePreflight(context, profile_id, policy, items, reply) => {
296                let result = self.prepare_preflight(context, profile_id, policy, items);
297                reply_transition(reply, result)
298            }
299            Message::ResumePreflight(context, profile_id, policy, reply) => {
300                let result = self.resume_preflight(context, profile_id, policy);
301                reply_transition(reply, result)
302            }
303            Message::AcceptUser(context, profile_id, policy, contents, reply) => {
304                let result = self.stage.accept_user(
305                    &mut self.durable,
306                    context,
307                    profile_id,
308                    policy,
309                    contents,
310                );
311                reply_transition(reply, result)
312            }
313            Message::Return(id, result, reply) => match self.durable.accept_return(id, result) {
314                Ok(()) => Ok(answer(reply, Ok(()), false)),
315                Err(error) => {
316                    let _ = reply.send(Err(ActorError::Closed));
317                    Err(error)
318                }
319            },
320            Message::Stage(text, boxes, reply) => {
321                let mut event = event_context(
322                    &mut self.durable,
323                    &self.adapter,
324                    &self.active_key,
325                    &mut self.view,
326                    &self.sender,
327                    self.job,
328                );
329                continue_after_stage(self.stage.handle_stage(&mut event, text, boxes, reply))
330            }
331            Message::PreflightCompleted(id, result) => {
332                let mut event = event_context(
333                    &mut self.durable,
334                    &self.adapter,
335                    &self.active_key,
336                    &mut self.view,
337                    &self.sender,
338                    self.job,
339                );
340                let _ = self
341                    .stage
342                    .handle_preflight_completion(&mut event, id, result)?;
343                self.preflight.complete(id)?;
344                if self.job.is_none() && !self.preflight.work_active() {
345                    self.durable.clear_authorization();
346                }
347                Ok(false)
348            }
349            Message::WebSearchCompleted(epoch, id, result) => {
350                let mut event = event_context(
351                    &mut self.durable,
352                    &self.adapter,
353                    &self.active_key,
354                    &mut self.view,
355                    &self.sender,
356                    self.job,
357                );
358                continue_after_stage(
359                    self.stage
360                        .handle_web_search_completion(&mut event, epoch, id, result),
361                )
362            }
363            Message::MailboxFlushCompleted(job, prepared, result) => {
364                let transport_failed = result.is_err();
365                if transport_failed && (!self.stage.mailbox_pending() || self.job != Some(job)) {
366                    return Err("stale failed Codex mailbox-flush completion".to_owned());
367                }
368                let outcome = {
369                    let mut event = event_context(
370                        &mut self.durable,
371                        &self.adapter,
372                        &self.active_key,
373                        &mut self.view,
374                        &self.sender,
375                        self.job,
376                    );
377                    self.stage
378                        .handle_mailbox_flush_completion(&mut event, job, prepared, result)
379                };
380                if !transport_failed {
381                    return continue_after_stage(outcome);
382                }
383                if outcome.is_ok() {
384                    return Err("failed Codex mailbox transport completed successfully".to_owned());
385                }
386                self.recover_turn(
387                    job,
388                    FailureClass::recoverable(ChatDiagnostic::MailboxTransport),
389                )
390                .await
391            }
392            Message::Inferred(job, shim, result) => self.finish(job, shim, result).await,
393            Message::Snapshot(reply) => {
394                let snapshot = self.view.snapshot(&self.durable, self.work_active());
395                Ok(answer(reply, Ok(snapshot), false))
396            }
397            Message::Wait(reply) => {
398                self.view.wait(&self.durable, self.work_active(), reply);
399                Ok(false)
400            }
401            Message::Restart(context, profile_id, policy, reply) => {
402                match self.restart(context, profile_id, policy) {
403                    Err(ActorError::Closed) => {
404                        let _ = reply.send(Err(ActorError::Closed));
405                        Err("chat restart failed internally".to_owned())
406                    }
407                    Err(error) => Ok(answer(reply, Err(error), false)),
408                    Ok(()) => {
409                        if self.usage_enabled
410                            && self.usage.restart(&mut self.durable).await.is_err()
411                        {
412                            self.disable_usage(FailureClass::optional(
413                                ChatDiagnostic::ModelUsageRestart,
414                            ));
415                        }
416                        Ok(answer(reply, Ok(()), false))
417                    }
418                }
419            }
420            Message::Abandon => Ok(true),
421        }
422    }
423
424    fn prepare_preflight(
425        &mut self,
426        context: AccessContext,
427        profile_id: ProfileId,
428        policy: AccessPolicy,
429        items: Vec<PreflightItem>,
430    ) -> Result<(), TransitionError> {
431        self.durable
432            .prepare_preflight(context, profile_id, policy, items)?;
433        self.refresh_preflight()
434            .map_err(TransitionError::Internal)?;
435        self.launch_preflight();
436        if !self.preflight.work_active() {
437            self.durable.clear_authorization();
438        }
439        Ok(())
440    }
441
442    fn resume_preflight(
443        &mut self,
444        context: AccessContext,
445        profile_id: ProfileId,
446        policy: AccessPolicy,
447    ) -> Result<(), TransitionError> {
448        self.refresh_preflight()
449            .map_err(TransitionError::Internal)?;
450        if !self.preflight.work_active() {
451            return Ok(());
452        }
453        self.durable
454            .authorize_preflight(context, profile_id, policy)?;
455        self.launch_preflight();
456        Ok(())
457    }
458
459    fn refresh_preflight(&mut self) -> Result<(), String> {
460        self.preflight.refresh(
461            self.durable.preflight_calls(),
462            self.durable.boxes(),
463            &self.durable.events(),
464        )
465    }
466
467    fn launch_preflight(&mut self) {
468        let executor = self.durable.preflight_executor();
469        for call in self.preflight.take_unlaunched() {
470            let executor = executor.clone();
471            let sender = self.sender.clone();
472            tokio::task::spawn_blocking(move || {
473                let result = executor.launch(&call.name, &call.arguments);
474                let _ = sender.send(Message::PreflightCompleted(call.tool_call_id, result));
475            });
476        }
477    }
478
479    fn work_active(&self) -> bool {
480        self.stage.search_active() || self.preflight.work_active()
481    }
482
483    fn handle_code_completion(&mut self) -> ActorResult {
484        let mut event = event_context(
485            &mut self.durable,
486            &self.adapter,
487            &self.active_key,
488            &mut self.view,
489            &self.sender,
490            self.job,
491        );
492        continue_after_stage(self.stage.handle_code_completion(&mut event))
493    }
494
495    fn disable_usage(&mut self, failure: FailureClass) {
496        debug_assert_eq!(failure.severity, FailureSeverity::OptionalDiagnostic);
497        self.usage = ModelUsageSession::default();
498        self.usage_enabled = false;
499        let _ = self.durable.record_diagnostic(failure.diagnostic);
500    }
501
502    fn critical(&mut self, _error: String) -> bool {
503        let failure = FailureClass::critical();
504        debug_assert_eq!(failure.severity, FailureSeverity::CriticalIntegrity);
505        let _ = self.durable.record_diagnostic(failure.diagnostic);
506        let _ = self.durable.halt_critical(SAFE_CRITICAL_FAILURE.to_owned());
507        self.stage.fail(SAFE_CRITICAL_FAILURE.to_owned());
508        self.view.wake(&self.durable, false);
509        true
510    }
511
512    async fn drive(&mut self) -> UnitResult {
513        if self.inference.is_some()
514            || self.shim.is_none()
515            || !self.preflight.allows_inference(self.durable.boxes())
516        {
517            return Ok(());
518        }
519        let Some((job, input)) = self.durable.begin_input()? else {
520            return Ok(());
521        };
522        if self.usage_enabled
523            && !self.usage.is_subscribed()
524            && self
525                .usage
526                .subscribe(&self.adapter, self.active_key.clone())
527                .await
528                .is_err()
529        {
530            self.disable_usage(FailureClass::optional(ChatDiagnostic::ModelUsageSubscribe));
531        }
532        self.view.set_model_input(ProviderInput {
533            kind: ProviderInputKind::Turn,
534            text: input.clone(),
535        });
536        let mut shim = self.shim.take().expect("shim was checked");
537        self.job = Some(job);
538        let sender = self.sender.clone();
539        self.inference = Some(tokio::spawn(async move {
540            let result = shim.infer(input).await.map_err(|error| error.to_string());
541            let _ = sender.send(Message::Inferred(job, shim, result));
542        }));
543        Ok(())
544    }
545
546    async fn finish(&mut self, job: u64, shim: ActorShim, result: InferenceResult) -> ActorResult {
547        if self.inference.is_none() || self.stage.mailbox_pending() || self.job != Some(job) {
548            return Err("stale Codex inference completion".to_owned());
549        }
550        drop(self.inference.take());
551        self.job = None;
552        let active_key = self.active_key.clone();
553        let settlement = settle_stage_inference(
554            &mut self.durable,
555            &mut self.usage,
556            &mut self.stage,
557            &active_key,
558            job,
559            result,
560        )
561        .await?;
562        let should_stop = settlement.should_stop();
563        let recoverable = settlement.recoverable_failure();
564        let restore_shim = settlement.restore_shim();
565        let usage_diagnostic = settlement.usage_diagnostic();
566        let stage_error = settlement.into_stage_error();
567        if let Some(diagnostic) = usage_diagnostic {
568            self.disable_usage(FailureClass::optional(diagnostic));
569        }
570        if should_stop {
571            return Err("inference settlement requested a critical stop".to_owned());
572        }
573        if recoverable {
574            self.install_fresh_provider("recovery")?;
575        } else if restore_shim {
576            self.shim = Some(shim);
577        } else {
578            return Err("inference settlement lost the provider shim".to_owned());
579        }
580        if let Some(error) = stage_error {
581            self.stage.fail(error);
582        }
583        Ok(false)
584    }
585
586    async fn recover_turn(&mut self, job: u64, failure: FailureClass) -> ActorResult {
587        debug_assert_eq!(failure.severity, FailureSeverity::RecoverableTurn);
588        if self.inference.is_none() || self.stage.mailbox_pending() || self.job != Some(job) {
589            return Err("recoverable turn failure lost active inference".to_owned());
590        }
591        if let Some(task) = self.inference.take() {
592            task.abort();
593        }
594        self.job = None;
595        let active_key = self.active_key.clone();
596        let settlement = settle_stage_failure(
597            &mut self.durable,
598            &mut self.usage,
599            &mut self.stage,
600            &active_key,
601            job,
602            failure.diagnostic,
603        )
604        .await?;
605        let should_stop = settlement.should_stop();
606        let recoverable = settlement.recoverable_failure();
607        let usage_diagnostic = settlement.usage_diagnostic();
608        let stage_error = settlement.into_stage_error();
609        if let Some(diagnostic) = usage_diagnostic {
610            self.disable_usage(FailureClass::optional(diagnostic));
611        }
612        if should_stop || !recoverable {
613            return Err("turn failure settlement was not recoverable".to_owned());
614        }
615        self.install_fresh_provider("recovery")?;
616        if let Some(error) = stage_error {
617            self.stage.fail(error);
618        }
619        Ok(false)
620    }
621
622    fn install_fresh_provider(&mut self, label: &str) -> Result<(), String> {
623        let generation = self
624            .generation
625            .checked_add(1)
626            .ok_or_else(|| "provider generation space was exhausted".to_owned())?;
627        let active_key = format!("{}#{label}-{generation}", self.base_key);
628        self.shim = Some(new_shim(
629            self.adapter.clone(),
630            active_key.clone(),
631            self.sender.clone(),
632        ));
633        self.generation = generation;
634        self.active_key = active_key;
635        Ok(())
636    }
637
638    fn restart(
639        &mut self,
640        context: AccessContext,
641        profile_id: ProfileId,
642        policy: AccessPolicy,
643    ) -> Result<(), ActorError> {
644        let generation = self
645            .generation
646            .checked_add(1)
647            .ok_or(ActorError::NotRestartable)?;
648        let active_key = format!("{}#restart-{generation}", self.base_key);
649        let shim = new_shim(
650            self.adapter.clone(),
651            active_key.clone(),
652            self.sender.clone(),
653        );
654        self.stage
655            .restart(&mut self.durable, context, profile_id, policy)?;
656        self.generation = generation;
657        self.active_key = active_key;
658        self.shim = Some(shim);
659        Ok(())
660    }
661}
662
663fn event_context<'a>(
664    durable: &'a mut DurableThread,
665    adapter: &'a Adapter,
666    active_key: &'a str,
667    view: &'a mut SessionView,
668    sender: &'a mpsc::UnboundedSender<Message>,
669    job: Option<u64>,
670) -> EventContext<'a> {
671    EventContext {
672        durable,
673        adapter,
674        active_key,
675        view,
676        sender,
677        job,
678    }
679}
680
681fn continue_after_stage(result: ActorResult) -> ActorResult {
682    result.map(|_| false)
683}
684
685fn map_transition(error: TransitionError) -> ActorError {
686    match error {
687        TransitionError::Unauthorized => ActorError::Unauthorized,
688        TransitionError::NotStalled => ActorError::NotStalled,
689        TransitionError::NotRestartable => ActorError::NotRestartable,
690        TransitionError::Internal(_) => ActorError::Closed,
691    }
692}
693
694fn reply_transition(reply: Reply<()>, result: Result<(), TransitionError>) -> ActorResult {
695    match result {
696        Ok(()) => Ok(answer(reply, Ok(()), false)),
697        Err(TransitionError::Internal(error)) => {
698            let _ = reply.send(Err(ActorError::Closed));
699            Err(error)
700        }
701        Err(error) => Ok(answer(reply, Err(map_transition(error)), false)),
702    }
703}
704
705fn answer<T>(reply: Reply<T>, result: Result<T, ActorError>, stop: bool) -> bool {
706    let _ = reply.send(result);
707    stop
708}