#![forbid(unsafe_code)]
#![doc = include_str!("../Documentation.md")]
use kcode_k1_access_kmap::K1AccessKmap;
use kcode_k1_chat_model_usage_session::ModelUsageSession;
use kcode_k1_chat_persistence::Session;
use kcode_k1_chat_thread_actor_channel::{
ActorError, ActorShim, BoxValue, Handle, Message, PreflightItem, ProviderInput,
ProviderInputKind, Reply, channel_with_events, new_shim,
};
use kcode_k1_chat_thread_durable_state::{
AccessContext, AccessPolicy, ChatDiagnostic, DurableThread, ProfileId, SetLaunchNodeKtool,
TransitionError,
};
use kcode_k1_chat_thread_preflight_runtime::PreflightRuntime;
use kcode_k1_chat_thread_session_code_runtime::SessionCodeRuntime;
use kcode_k1_chat_thread_session_inference_settlement::{
settle_stage_failure, settle_stage_inference,
};
use kcode_k1_chat_thread_session_stage_runtime::{EventContext, StageRuntime};
use kcode_k1_chat_thread_session_view::SessionView;
use kcode_k1_codex_adapter::{Adapter, ShimOutput};
use kcode_k1_codex_websearch::Runner as WebSearchRunner;
use kcode_k1_rust_code_ktool_service::RustCodeKtoolService;
use kcode_k1_web_code_ktool_service::K1WebCodeKtoolService;
use std::sync::Arc;
use tokio::{sync::mpsc, task::JoinHandle};
type ActorResult = Result<bool, String>;
type InferenceResult = Result<ShimOutput<BoxValue>, String>;
type UnitResult = Result<(), String>;
const SAFE_CRITICAL_FAILURE: &str =
"This chat stopped because an internal integrity failure occurred.";
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
enum FailureSeverity {
OptionalDiagnostic,
RecoverableTurn,
CriticalIntegrity,
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
struct FailureClass {
severity: FailureSeverity,
diagnostic: ChatDiagnostic,
}
impl FailureClass {
fn optional(diagnostic: ChatDiagnostic) -> Self {
Self {
severity: FailureSeverity::OptionalDiagnostic,
diagnostic,
}
}
fn recoverable(diagnostic: ChatDiagnostic) -> Self {
Self {
severity: FailureSeverity::RecoverableTurn,
diagnostic,
}
}
fn critical() -> Self {
Self {
severity: FailureSeverity::CriticalIntegrity,
diagnostic: ChatDiagnostic::CriticalIntegrity,
}
}
}
pub fn open(
adapter: Adapter,
key: impl Into<String>,
session: Session,
kmap: Arc<K1AccessKmap>,
web_search: WebSearchRunner,
) -> Result<Handle, String> {
let durable = DurableThread::recover(session, kmap)?;
spawn_actor(adapter, key, durable, None, None, web_search)
}
pub fn open_with_social(
adapter: Adapter,
key: impl Into<String>,
session: Session,
kmap: Arc<K1AccessKmap>,
social: kcode_k1_ktool_social::SocialKtools,
web_search: WebSearchRunner,
) -> Result<Handle, String> {
let durable = DurableThread::recover_with_social(session, kmap, social)?;
spawn_actor(adapter, key, durable, None, None, web_search)
}
pub fn open_with_social_and_set_launch_node(
adapter: Adapter,
key: impl Into<String>,
session: Session,
kmap: Arc<K1AccessKmap>,
social: kcode_k1_ktool_social::SocialKtools,
set_launch_node: SetLaunchNodeKtool,
web_search: WebSearchRunner,
) -> Result<Handle, String> {
let durable = DurableThread::recover_with_social_and_set_launch_node(
session,
kmap,
social,
set_launch_node,
)?;
spawn_actor(adapter, key, durable, None, None, web_search)
}
#[allow(clippy::too_many_arguments)]
pub fn open_with_social_and_set_launch_node_and_rust_code(
adapter: Adapter,
key: impl Into<String>,
session: Session,
kmap: Arc<K1AccessKmap>,
social: kcode_k1_ktool_social::SocialKtools,
set_launch_node: SetLaunchNodeKtool,
rust_code: Arc<RustCodeKtoolService>,
web_search: WebSearchRunner,
) -> Result<Handle, String> {
let durable = DurableThread::recover_with_social_and_set_launch_node(
session,
kmap,
social,
set_launch_node,
)?;
spawn_actor(adapter, key, durable, Some(rust_code), None, web_search)
}
#[allow(clippy::too_many_arguments)]
pub fn open_with_social_and_set_launch_node_and_rust_code_and_web_code(
adapter: Adapter,
key: impl Into<String>,
session: Session,
kmap: Arc<K1AccessKmap>,
social: kcode_k1_ktool_social::SocialKtools,
set_launch_node: SetLaunchNodeKtool,
rust_code: Arc<RustCodeKtoolService>,
web_code: K1WebCodeKtoolService,
web_search: WebSearchRunner,
) -> Result<Handle, String> {
let durable = DurableThread::recover_with_social_and_set_launch_node(
session,
kmap,
social,
set_launch_node,
)?;
spawn_actor(
adapter,
key,
durable,
Some(rust_code),
Some(web_code),
web_search,
)
}
fn spawn_actor(
adapter: Adapter,
key: impl Into<String>,
durable: DurableThread,
rust_code: Option<Arc<RustCodeKtoolService>>,
web_code: Option<K1WebCodeKtoolService>,
web_search: WebSearchRunner,
) -> Result<Handle, String> {
let code = SessionCodeRuntime::recover(&durable, rust_code, web_code)?;
let stage = StageRuntime::new(code, web_search);
let preflight = PreflightRuntime::recover(
durable.preflight_calls(),
durable.boxes(),
&durable.events(),
)?;
let (handle, sender, receiver, event_receiver) = channel_with_events();
let base_key = key.into();
let actor = Actor {
durable,
shim: Some(new_shim(adapter.clone(), base_key.clone(), sender.clone())),
adapter,
active_key: base_key.clone(),
base_key,
generation: 0,
stage,
preflight,
view: SessionView::new(event_receiver),
sender,
receiver,
usage: ModelUsageSession::default(),
usage_enabled: true,
inference: None,
job: None,
};
tokio::spawn(actor.run());
Ok(handle)
}
struct Actor {
durable: DurableThread,
shim: Option<ActorShim>,
adapter: Adapter,
base_key: String,
active_key: String,
generation: u64,
stage: StageRuntime,
preflight: PreflightRuntime,
view: SessionView,
sender: mpsc::UnboundedSender<Message>,
receiver: mpsc::UnboundedReceiver<Message>,
usage: ModelUsageSession,
usage_enabled: bool,
inference: Option<JoinHandle<()>>,
job: Option<u64>,
}
impl Actor {
async fn run(mut self) {
loop {
if let Err(error) = self.drive().await {
self.critical(error);
break;
}
let active = self.work_active();
let can_receive_code = self.stage.can_receive_code();
self.view.wake(&self.durable, active);
let stop = tokio::select! {
message = self.receiver.recv() => match message {
Some(message) => match self.handle(message).await {
Ok(stop) => stop,
Err(error) => self.critical(error),
},
None => true,
},
result = self.stage.receive_code(), if can_receive_code => match result {
Ok(()) => match self.handle_code_completion() {
Ok(stop) => stop,
Err(error) => self.critical(error),
},
Err(error) => self.critical(error),
},
reply = self.view.receive_event_query() => {
if let Some(reply) = reply {
self.view.answer_event_query(&self.durable, active, reply);
}
false
}
usage = self.usage.receive(&mut self.durable), if self.usage_enabled => {
match usage {
Ok(_) => false,
Err(_) => {
self.disable_usage(FailureClass::optional(
ChatDiagnostic::ModelUsageReceive,
));
false
}
}
},
};
if stop {
break;
}
}
self.stage.abort();
let inference_active = self.inference.is_some();
if self.usage_enabled
&& self
.usage
.shutdown(&mut self.durable, inference_active)
.await
.is_err()
{
self.disable_usage(FailureClass::optional(ChatDiagnostic::ModelUsageShutdown));
}
self.stage.shutdown();
if let Some(task) = self.inference.take() {
task.abort();
}
self.view.close();
}
async fn handle(&mut self, message: Message) -> ActorResult {
match message {
Message::Accept((box_type, contents, hidden_type, hidden_contents), reply) => {
reply_transition(
reply,
self.durable.accept_external_box(
box_type,
contents,
hidden_type,
hidden_contents,
),
)
}
Message::PreparePreflight(context, profile_id, policy, items, reply) => {
let result = self.prepare_preflight(context, profile_id, policy, items);
reply_transition(reply, result)
}
Message::ResumePreflight(context, profile_id, policy, reply) => {
let result = self.resume_preflight(context, profile_id, policy);
reply_transition(reply, result)
}
Message::AcceptUser(context, profile_id, policy, contents, reply) => {
let result = self.stage.accept_user(
&mut self.durable,
context,
profile_id,
policy,
contents,
);
reply_transition(reply, result)
}
Message::Return(id, result, reply) => match self.durable.accept_return(id, result) {
Ok(()) => Ok(answer(reply, Ok(()), false)),
Err(error) => {
let _ = reply.send(Err(ActorError::Closed));
Err(error)
}
},
Message::Stage(text, boxes, reply) => {
let mut event = event_context(
&mut self.durable,
&self.adapter,
&self.active_key,
&mut self.view,
&self.sender,
self.job,
);
continue_after_stage(self.stage.handle_stage(&mut event, text, boxes, reply))
}
Message::PreflightCompleted(id, result) => {
let mut event = event_context(
&mut self.durable,
&self.adapter,
&self.active_key,
&mut self.view,
&self.sender,
self.job,
);
let _ = self
.stage
.handle_preflight_completion(&mut event, id, result)?;
self.preflight.complete(id)?;
if self.job.is_none() && !self.preflight.work_active() {
self.durable.clear_authorization();
}
Ok(false)
}
Message::WebSearchCompleted(epoch, id, result) => {
let mut event = event_context(
&mut self.durable,
&self.adapter,
&self.active_key,
&mut self.view,
&self.sender,
self.job,
);
continue_after_stage(
self.stage
.handle_web_search_completion(&mut event, epoch, id, result),
)
}
Message::MailboxFlushCompleted(job, prepared, result) => {
let transport_failed = result.is_err();
if transport_failed && (!self.stage.mailbox_pending() || self.job != Some(job)) {
return Err("stale failed Codex mailbox-flush completion".to_owned());
}
let outcome = {
let mut event = event_context(
&mut self.durable,
&self.adapter,
&self.active_key,
&mut self.view,
&self.sender,
self.job,
);
self.stage
.handle_mailbox_flush_completion(&mut event, job, prepared, result)
};
if !transport_failed {
return continue_after_stage(outcome);
}
if outcome.is_ok() {
return Err("failed Codex mailbox transport completed successfully".to_owned());
}
self.recover_turn(
job,
FailureClass::recoverable(ChatDiagnostic::MailboxTransport),
)
.await
}
Message::Inferred(job, shim, result) => self.finish(job, shim, result).await,
Message::Snapshot(reply) => {
let snapshot = self.view.snapshot(&self.durable, self.work_active());
Ok(answer(reply, Ok(snapshot), false))
}
Message::Wait(reply) => {
self.view.wait(&self.durable, self.work_active(), reply);
Ok(false)
}
Message::Restart(context, profile_id, policy, reply) => {
match self.restart(context, profile_id, policy) {
Err(ActorError::Closed) => {
let _ = reply.send(Err(ActorError::Closed));
Err("chat restart failed internally".to_owned())
}
Err(error) => Ok(answer(reply, Err(error), false)),
Ok(()) => {
if self.usage_enabled
&& self.usage.restart(&mut self.durable).await.is_err()
{
self.disable_usage(FailureClass::optional(
ChatDiagnostic::ModelUsageRestart,
));
}
Ok(answer(reply, Ok(()), false))
}
}
}
Message::Abandon => Ok(true),
}
}
fn prepare_preflight(
&mut self,
context: AccessContext,
profile_id: ProfileId,
policy: AccessPolicy,
items: Vec<PreflightItem>,
) -> Result<(), TransitionError> {
self.durable
.prepare_preflight(context, profile_id, policy, items)?;
self.refresh_preflight()
.map_err(TransitionError::Internal)?;
self.launch_preflight();
if !self.preflight.work_active() {
self.durable.clear_authorization();
}
Ok(())
}
fn resume_preflight(
&mut self,
context: AccessContext,
profile_id: ProfileId,
policy: AccessPolicy,
) -> Result<(), TransitionError> {
self.refresh_preflight()
.map_err(TransitionError::Internal)?;
if !self.preflight.work_active() {
return Ok(());
}
self.durable
.authorize_preflight(context, profile_id, policy)?;
self.launch_preflight();
Ok(())
}
fn refresh_preflight(&mut self) -> Result<(), String> {
self.preflight.refresh(
self.durable.preflight_calls(),
self.durable.boxes(),
&self.durable.events(),
)
}
fn launch_preflight(&mut self) {
let executor = self.durable.preflight_executor();
for call in self.preflight.take_unlaunched() {
let executor = executor.clone();
let sender = self.sender.clone();
tokio::task::spawn_blocking(move || {
let result = executor.launch(&call.name, &call.arguments);
let _ = sender.send(Message::PreflightCompleted(call.tool_call_id, result));
});
}
}
fn work_active(&self) -> bool {
self.stage.search_active() || self.preflight.work_active()
}
fn handle_code_completion(&mut self) -> ActorResult {
let mut event = event_context(
&mut self.durable,
&self.adapter,
&self.active_key,
&mut self.view,
&self.sender,
self.job,
);
continue_after_stage(self.stage.handle_code_completion(&mut event))
}
fn disable_usage(&mut self, failure: FailureClass) {
debug_assert_eq!(failure.severity, FailureSeverity::OptionalDiagnostic);
self.usage = ModelUsageSession::default();
self.usage_enabled = false;
let _ = self.durable.record_diagnostic(failure.diagnostic);
}
fn critical(&mut self, _error: String) -> bool {
let failure = FailureClass::critical();
debug_assert_eq!(failure.severity, FailureSeverity::CriticalIntegrity);
let _ = self.durable.record_diagnostic(failure.diagnostic);
let _ = self.durable.halt_critical(SAFE_CRITICAL_FAILURE.to_owned());
self.stage.fail(SAFE_CRITICAL_FAILURE.to_owned());
self.view.wake(&self.durable, false);
true
}
async fn drive(&mut self) -> UnitResult {
if self.inference.is_some()
|| self.shim.is_none()
|| !self.preflight.allows_inference(self.durable.boxes())
{
return Ok(());
}
let Some((job, input)) = self.durable.begin_input()? else {
return Ok(());
};
if self.usage_enabled
&& !self.usage.is_subscribed()
&& self
.usage
.subscribe(&self.adapter, self.active_key.clone())
.await
.is_err()
{
self.disable_usage(FailureClass::optional(ChatDiagnostic::ModelUsageSubscribe));
}
self.view.set_model_input(ProviderInput {
kind: ProviderInputKind::Turn,
text: input.clone(),
});
let mut shim = self.shim.take().expect("shim was checked");
self.job = Some(job);
let sender = self.sender.clone();
self.inference = Some(tokio::spawn(async move {
let result = shim.infer(input).await.map_err(|error| error.to_string());
let _ = sender.send(Message::Inferred(job, shim, result));
}));
Ok(())
}
async fn finish(&mut self, job: u64, shim: ActorShim, result: InferenceResult) -> ActorResult {
if self.inference.is_none() || self.stage.mailbox_pending() || self.job != Some(job) {
return Err("stale Codex inference completion".to_owned());
}
drop(self.inference.take());
self.job = None;
let active_key = self.active_key.clone();
let settlement = settle_stage_inference(
&mut self.durable,
&mut self.usage,
&mut self.stage,
&active_key,
job,
result,
)
.await?;
let should_stop = settlement.should_stop();
let recoverable = settlement.recoverable_failure();
let restore_shim = settlement.restore_shim();
let usage_diagnostic = settlement.usage_diagnostic();
let stage_error = settlement.into_stage_error();
if let Some(diagnostic) = usage_diagnostic {
self.disable_usage(FailureClass::optional(diagnostic));
}
if should_stop {
return Err("inference settlement requested a critical stop".to_owned());
}
if recoverable {
self.install_fresh_provider("recovery")?;
} else if restore_shim {
self.shim = Some(shim);
} else {
return Err("inference settlement lost the provider shim".to_owned());
}
if let Some(error) = stage_error {
self.stage.fail(error);
}
Ok(false)
}
async fn recover_turn(&mut self, job: u64, failure: FailureClass) -> ActorResult {
debug_assert_eq!(failure.severity, FailureSeverity::RecoverableTurn);
if self.inference.is_none() || self.stage.mailbox_pending() || self.job != Some(job) {
return Err("recoverable turn failure lost active inference".to_owned());
}
if let Some(task) = self.inference.take() {
task.abort();
}
self.job = None;
let active_key = self.active_key.clone();
let settlement = settle_stage_failure(
&mut self.durable,
&mut self.usage,
&mut self.stage,
&active_key,
job,
failure.diagnostic,
)
.await?;
let should_stop = settlement.should_stop();
let recoverable = settlement.recoverable_failure();
let usage_diagnostic = settlement.usage_diagnostic();
let stage_error = settlement.into_stage_error();
if let Some(diagnostic) = usage_diagnostic {
self.disable_usage(FailureClass::optional(diagnostic));
}
if should_stop || !recoverable {
return Err("turn failure settlement was not recoverable".to_owned());
}
self.install_fresh_provider("recovery")?;
if let Some(error) = stage_error {
self.stage.fail(error);
}
Ok(false)
}
fn install_fresh_provider(&mut self, label: &str) -> Result<(), String> {
let generation = self
.generation
.checked_add(1)
.ok_or_else(|| "provider generation space was exhausted".to_owned())?;
let active_key = format!("{}#{label}-{generation}", self.base_key);
self.shim = Some(new_shim(
self.adapter.clone(),
active_key.clone(),
self.sender.clone(),
));
self.generation = generation;
self.active_key = active_key;
Ok(())
}
fn restart(
&mut self,
context: AccessContext,
profile_id: ProfileId,
policy: AccessPolicy,
) -> Result<(), ActorError> {
let generation = self
.generation
.checked_add(1)
.ok_or(ActorError::NotRestartable)?;
let active_key = format!("{}#restart-{generation}", self.base_key);
let shim = new_shim(
self.adapter.clone(),
active_key.clone(),
self.sender.clone(),
);
self.stage
.restart(&mut self.durable, context, profile_id, policy)?;
self.generation = generation;
self.active_key = active_key;
self.shim = Some(shim);
Ok(())
}
}
fn event_context<'a>(
durable: &'a mut DurableThread,
adapter: &'a Adapter,
active_key: &'a str,
view: &'a mut SessionView,
sender: &'a mpsc::UnboundedSender<Message>,
job: Option<u64>,
) -> EventContext<'a> {
EventContext {
durable,
adapter,
active_key,
view,
sender,
job,
}
}
fn continue_after_stage(result: ActorResult) -> ActorResult {
result.map(|_| false)
}
fn map_transition(error: TransitionError) -> ActorError {
match error {
TransitionError::Unauthorized => ActorError::Unauthorized,
TransitionError::NotStalled => ActorError::NotStalled,
TransitionError::NotRestartable => ActorError::NotRestartable,
TransitionError::Internal(_) => ActorError::Closed,
}
}
fn reply_transition(reply: Reply<()>, result: Result<(), TransitionError>) -> ActorResult {
match result {
Ok(()) => Ok(answer(reply, Ok(()), false)),
Err(TransitionError::Internal(error)) => {
let _ = reply.send(Err(ActorError::Closed));
Err(error)
}
Err(error) => Ok(answer(reply, Err(map_transition(error)), false)),
}
}
fn answer<T>(reply: Reply<T>, result: Result<T, ActorError>, stop: bool) -> bool {
let _ = reply.send(result);
stop
}