#![forbid(unsafe_code)]
use kcode_k1_access_kmap::K1AccessKmap;
use kcode_k1_chat_codex_state::PreparedCallDisposition;
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, ProviderInput, ProviderInputKind, Reply,
StageReply, channel_with_events, new_shim,
};
use kcode_k1_chat_thread_durable_state::{
DurableThread, PreparedMailboxFlush, ToolCallId, TransitionError,
};
use kcode_k1_chat_thread_session_view::SessionView;
use kcode_k1_chat_thread_web_search::{WebSearchAction, WebSearchToolEvent};
use kcode_k1_chat_thread_web_search_tasks::WebSearchTasks;
use kcode_k1_codex_adapter::{Adapter, ShimOutput};
use std::sync::Arc;
use tokio::{sync::mpsc, task::JoinHandle};
pub fn open(
adapter: Adapter,
key: impl Into<String>,
session: Session,
kmap: Arc<K1AccessKmap>,
web_search: WebSearchAction,
) -> Result<Handle, String> {
let durable = DurableThread::recover(session, kmap)?;
let (handle, sender, receiver, event_receiver) = channel_with_events();
let base_key = key.into();
let active_key = base_key.clone();
let actor = Actor {
durable,
shim: Some(new_shim(
adapter.clone(),
active_key.clone(),
sender.clone(),
)),
adapter,
base_key,
active_key,
generation: 0,
web_search: WebSearchTasks::new(web_search, sender.clone()),
view: SessionView::new(event_receiver),
sender,
receiver,
usage: ModelUsageSession::default(),
inference: None,
mailbox_flush: None,
stage_reply: 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,
web_search: WebSearchTasks,
view: SessionView,
sender: mpsc::UnboundedSender<Message>,
receiver: mpsc::UnboundedReceiver<Message>,
usage: ModelUsageSession,
inference: Option<JoinHandle<()>>,
mailbox_flush: Option<JoinHandle<()>>,
stage_reply: Option<StageReply>,
job: Option<u64>,
}
impl Actor {
async fn run(mut self) {
'actor: loop {
if self.drive().await {
break;
}
self.view.wake(&self.durable, !self.web_search.is_empty());
tokio::select! {
message = self.receiver.recv() => {
let Some(mut message) = message else { break; };
loop {
if self.handle(message).await {
break 'actor;
}
match self.receiver.try_recv() {
Ok(next) => message = next,
Err(_) => break,
}
}
}
reply = self.view.receive_event_query() => {
if let Some(reply) = reply {
self.view.answer_event_query(
&self.durable,
!self.web_search.is_empty(),
reply,
);
}
}
usage = self.usage.receive(&mut self.durable) => {
if usage.is_err() {
break;
}
}
}
}
let _ = self
.usage
.shutdown(&mut self.durable, self.inference.is_some())
.await;
self.web_search.abort();
if let Some(reply) = self.stage_reply.take() {
let _ = reply.send(Err("K1 actor is closed".to_owned()));
}
for task in [self.mailbox_flush.take(), self.inference.take()]
.into_iter()
.flatten()
{
task.abort();
}
self.view.close();
}
async fn handle(&mut self, message: Message) -> bool {
match message {
Message::Accept((box_type, contents, hidden_type, hidden_contents), reply) => {
let result = self.durable.accept_external_box(
box_type,
contents,
hidden_type,
hidden_contents,
);
let stop = matches!(result, Err(TransitionError::Internal(_)));
let _ = reply.send(result.map_err(map_transition));
stop
}
Message::AcceptUser(context, profile_id, policy, contents, reply) => {
self.accept_user(context, profile_id, policy, contents, reply)
}
Message::Return(id, result, reply) => self.accept_return(id, result, reply),
Message::Stage(text, boxes, reply) => self.stage(text, boxes, reply),
Message::WebSearch(epoch, id, event) => self.web_search_event(epoch, id, event),
Message::WebSearchEnded(epoch, id) => self.web_search_ended(epoch, id),
Message::MailboxFlushCompleted(job, prepared, result) => {
self.mailbox_flush_completed(job, prepared, result)
}
Message::Inferred(job, shim, result) => self.finish(job, shim, result).await,
Message::Snapshot(reply) => {
let snapshot = self
.view
.snapshot(&self.durable, !self.web_search.is_empty());
let _ = reply.send(Ok(snapshot));
false
}
Message::Wait(reply) => {
self.view
.wait(&self.durable, !self.web_search.is_empty(), reply);
false
}
Message::Restart(context, profile_id, policy, reply) => {
let result = self.restart(context, profile_id, policy);
if let Err(error) = result {
let _ = reply.send(Err(error));
false
} else if self.usage.restart(&mut self.durable).await.is_err() {
let _ = reply.send(Err(ActorError::Closed));
true
} else {
let _ = reply.send(Ok(()));
false
}
}
Message::Abandon => true,
}
}
fn accept_user(
&mut self,
context: kcode_k1_chat_thread_durable_state::AccessContext,
profile_id: kcode_k1_chat_thread_durable_state::ProfileId,
policy: kcode_k1_chat_thread_durable_state::AccessPolicy,
contents: String,
reply: Reply<()>,
) -> bool {
let result = self
.durable
.accept_user(context, profile_id, policy, contents);
let stop = matches!(result, Err(TransitionError::Internal(_)));
let _ = reply.send(result.map_err(map_transition));
stop
}
fn accept_return(
&mut self,
id: ToolCallId,
result: Result<String, String>,
reply: Reply<()>,
) -> bool {
let result = self.durable.accept_return(id, result);
let stop = result.is_err();
let _ = reply.send(result.map_err(|_| ActorError::Closed));
stop
}
fn stage(&mut self, text: String, boxes: Vec<BoxValue>, reply: StageReply) -> bool {
if self.stage_reply.is_some() || self.mailbox_flush.is_some() {
return reject_stage(reply, "overlapping Codex stages or mailbox flush".into());
}
let Some(job) = self.job else {
return reject_stage(reply, "no active K1 inference".into());
};
let calls = match self.durable.prepare_stage(job, text, boxes) {
Ok(calls) => calls,
Err(error) => return reject_stage(reply, error),
};
self.stage_reply = Some(reply);
for call in calls {
match call.disposition().clone() {
PreparedCallDisposition::ImmediateError(error) => {
if let Err(error) = self.durable.accept_tool_return_v2(
call.tool_call_id,
Err(error.message.clone()),
error.metadata_type.clone(),
error.metadata_contents.clone(),
) {
return self.fatal(error);
}
}
PreparedCallDisposition::External => {
if call.name == self.web_search.tool_name() {
if let Err(error) =
self.web_search.launch(call.tool_call_id, call.arguments)
{
return self.fatal(error);
}
} else {
let result = self.durable.launch_action(&call.name, &call.arguments);
if let Err(error) =
self.durable.accept_tool_return(call.tool_call_id, result)
{
return self.fatal(error);
}
}
}
}
}
match self.start_mailbox_flush() {
Ok(true) => self.finish_stage(),
Ok(false) => false,
Err(error) => self.fatal(error),
}
}
fn web_search_event(&mut self, epoch: u64, id: ToolCallId, event: WebSearchToolEvent) -> bool {
let event = match self.web_search.accept_event(epoch, id, event) {
Ok(Some(event)) => event,
Ok(None) => return false,
Err(error) => return self.fatal(error),
};
match event {
WebSearchToolEvent::Message { contents } => {
if let Err(error) = self.durable.accept_tool_message(id, contents) {
return self.fatal(error);
}
false
}
WebSearchToolEvent::Result {
result,
metadata_type,
metadata_contents,
} => {
if let Err(error) =
self.durable
.accept_tool_return_v2(id, result, metadata_type, metadata_contents)
{
return self.fatal(error);
}
if self.job.is_none() {
return false;
}
match self.start_mailbox_flush() {
Ok(_) => false,
Err(error) => self.fatal(error),
}
}
}
}
fn web_search_ended(&mut self, epoch: u64, id: ToolCallId) -> bool {
match self.web_search.accept_ended(epoch, id) {
Ok(()) => false,
Err(error) => self.fatal(error),
}
}
fn start_mailbox_flush(&mut self) -> Result<bool, String> {
let job = self
.job
.ok_or_else(|| "no active K1 inference".to_owned())?;
let prepared = self.durable.prepare_mailbox_flush(job)?;
if self.mailbox_flush.is_some() {
return Ok(false);
}
let Some(prepared) = prepared else {
return Ok(true);
};
let input = self.durable.prepared_input(&prepared)?;
self.view.set_model_input(ProviderInput {
kind: ProviderInputKind::MailboxFlush,
text: input.clone(),
});
let adapter = self.adapter.clone();
let key = self.active_key.clone();
let sender = self.sender.clone();
self.mailbox_flush = Some(tokio::spawn(async move {
let result = adapter
.steer(key, input)
.await
.map_err(|error| error.to_string());
let _ = sender.send(Message::MailboxFlushCompleted(job, prepared, result));
}));
Ok(false)
}
fn mailbox_flush_completed(
&mut self,
job: u64,
prepared: PreparedMailboxFlush,
result: Result<(), String>,
) -> bool {
if self.mailbox_flush.take().is_none() || self.job != Some(job) {
return self.fatal("stale Codex mailbox-flush completion".into());
}
if let Err(error) = result {
return self.fatal(error);
}
if let Err(error) = self.durable.commit_mailbox_flush(prepared) {
return self.fatal(error);
}
if self.stage_reply.is_some() && self.finish_stage() {
return true;
}
match self.start_mailbox_flush() {
Ok(_) => false,
Err(error) => self.fatal(error),
}
}
fn finish_stage(&mut self) -> bool {
match self.stage_reply.take() {
Some(reply) => reply.send(Ok(())).is_err(),
None => false,
}
}
fn fatal(&mut self, error: String) -> bool {
if let Some(reply) = self.stage_reply.take() {
let _ = reply.send(Err(error));
}
true
}
async fn drive(&mut self) -> bool {
if self.inference.is_some() || self.shim.is_none() {
return false;
}
let (job, input) = match self.durable.begin_input() {
Ok(Some(value)) => value,
Ok(None) => return false,
Err(error) => return self.fatal(error),
};
if !self.usage.is_subscribed()
&& self
.usage
.subscribe(&self.adapter, self.active_key.clone())
.await
.is_err()
{
return self.fatal("model-usage subscription failed".into());
}
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));
}));
false
}
async fn finish(
&mut self,
job: u64,
shim: ActorShim,
result: Result<ShimOutput<BoxValue>, String>,
) -> bool {
if self.inference.take().is_none() || self.mailbox_flush.is_some() || self.job != Some(job)
{
return self.fatal("stale Codex inference completion".into());
}
self.job = None;
let key = self.active_key.clone();
match result {
Err(error) => {
if self
.usage
.finish_inference(&mut self.durable, &key, None)
.await
.is_err()
{
return true;
}
let fence = self.web_search.cancel();
if let Some(reply) = self.stage_reply.take() {
let _ = reply.send(Err(error.clone()));
}
self.durable.fail(job, error, true);
match fence {
Ok(()) => false,
Err(error) => self.fatal(error),
}
}
Ok(output) => {
let (resume, terminal_id) =
match self.durable.complete_with_terminal_response(job, output) {
Ok(value) => value,
Err(error) => return self.fatal(error),
};
if self
.usage
.finish_inference(&mut self.durable, &key, Some(terminal_id))
.await
.is_err()
{
return true;
}
self.shim = Some(shim);
if !resume && self.web_search.is_empty() {
self.durable.clear_authorization();
}
false
}
}
}
fn restart(
&mut self,
context: kcode_k1_chat_thread_durable_state::AccessContext,
profile_id: kcode_k1_chat_thread_durable_state::ProfileId,
policy: kcode_k1_chat_thread_durable_state::AccessPolicy,
) -> Result<(), ActorError> {
let generation = self
.generation
.checked_add(1)
.ok_or(ActorError::NotRestartable)?;
self.durable
.restart(context, profile_id, policy)
.map_err(map_transition)?;
self.generation = generation;
self.active_key = format!("{}#restart-{generation}", self.base_key);
self.shim = Some(new_shim(
self.adapter.clone(),
self.active_key.clone(),
self.sender.clone(),
));
Ok(())
}
}
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 reject_stage(reply: StageReply, error: String) -> bool {
let _ = reply.send(Err(error));
true
}