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}