#![forbid(unsafe_code)]
#![doc = include_str!("../Documentation.md")]
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::{
AccessContext, AccessPolicy, DurableThread, PreparedMailboxFlush as Flush, ProfileId,
ToolCallId, TransitionError,
};
use kcode_k1_chat_thread_session_view::SessionView;
use kcode_k1_codex_adapter::{Adapter, ShimOutput};
use kcode_k1_codex_websearch::Runner as WebSearchRunner;
use std::{
sync::Arc,
time::{Duration, Instant},
};
use tokio::{sync::mpsc, task::JoinHandle};
type ActorResult = Result<bool, String>;
type InferenceResult = Result<ShimOutput<BoxValue>, String>;
type RestartAccess = (AccessContext, ProfileId, AccessPolicy);
type SearchResult = Result<String, String>;
type UnitResult = Result<(), String>;
const WEB_SEARCH_TOOL: &str = "WebSearch";
const WEB_SEARCH_TIMEOUT: Duration = Duration::from_secs(60 * 60);
const NO_ACTIVE_INFERENCE: &str = "no active K1 inference";
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)?;
Ok(spawn_actor(adapter, key, durable, 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)?;
Ok(spawn_actor(adapter, key, durable, web_search))
}
fn spawn_actor(
adapter: Adapter,
key: impl Into<String>,
durable: DurableThread,
web_search: WebSearchRunner,
) -> Handle {
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,
web_search,
search_epoch: 0,
searches: Vec::new(),
view: SessionView::new(event_receiver),
sender,
receiver,
usage: ModelUsageSession::default(),
inference: None,
mailbox_flush: None,
stage_reply: None,
job: None,
};
tokio::spawn(actor.run());
handle
}
#[derive(Clone, Copy, PartialEq)]
struct SearchKey {
epoch: u64,
id: ToolCallId,
}
struct SearchTask {
key: SearchKey,
task: JoinHandle<()>,
}
struct SearchDropGuard {
key: SearchKey,
sender: mpsc::UnboundedSender<Message>,
result: Option<SearchResult>,
}
impl Drop for SearchDropGuard {
fn drop(&mut self) {
let result = self
.result
.take()
.unwrap_or_else(|| Err("WebSearch failed: Codex execution failed".to_owned()));
let _ = self.sender.send(Message::WebSearchCompleted(
self.key.epoch,
self.key.id,
result,
));
}
}
struct Actor {
durable: DurableThread,
shim: Option<ActorShim>,
adapter: Adapter,
base_key: String,
active_key: String,
generation: u64,
web_search: WebSearchRunner,
search_epoch: u64,
searches: Vec<SearchTask>,
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) {
loop {
if let Err(error) = self.drive().await {
self.fatal(error);
break;
}
let searching = !self.searches.is_empty();
self.view.wake(&self.durable, searching);
let stop = tokio::select! {
message = self.receiver.recv() => match message {
Some(message) => match self.handle(message).await {
Ok(stop) => stop,
Err(error) => self.fatal(error),
},
None => true,
},
reply = self.view.receive_event_query() => {
if let Some(reply) = reply {
self.view.answer_event_query(
&self.durable,
searching,
reply,
);
}
false
}
usage = self.usage.receive(&mut self.durable) => usage.is_err(),
};
if stop {
break;
}
}
self.abort_searches();
let inference_active = self.inference.is_some();
let _ = self
.usage
.shutdown(&mut self.durable, inference_active)
.await;
if let Some(reply) = self.stage_reply.take() {
let _ = reply.send(Err("K1 actor is closed".to_owned()));
}
let tasks = [self.mailbox_flush.take(), self.inference.take()];
for task in tasks.into_iter().flatten() {
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) => {
Ok(reply_transition(
reply,
self.durable.accept_external_box(
box_type,
contents,
hidden_type,
hidden_contents,
),
))
}
Message::AcceptUser(context, profile_id, policy, contents, reply) => {
Ok(reply_transition(
reply,
self.durable
.accept_user(context, profile_id, policy, contents),
))
}
Message::Return(id, result, reply) => {
let result = self.durable.accept_return(id, result);
let result = result.map_err(|_| ActorError::Closed);
let stop = result.is_err();
Ok(answer(reply, result, stop))
}
Message::Stage(text, boxes, reply) => self.stage(text, boxes, reply),
Message::WebSearchCompleted(epoch, id, result) => {
self.search_done(SearchKey { epoch, id }, result)
}
Message::MailboxFlushCompleted(job, prepared, result) => {
self.flush_done(job, prepared, result)
}
Message::Inferred(job, shim, result) => self.finish(job, shim, result).await,
Message::Snapshot(reply) => {
let searching = !self.searches.is_empty();
let snapshot = self.view.snapshot(&self.durable, searching);
Ok(answer(reply, Ok(snapshot), false))
}
Message::Wait(reply) => {
let searching = !self.searches.is_empty();
self.view.wait(&self.durable, searching, reply);
Ok(false)
}
Message::Restart(context, profile_id, policy, reply) => {
let result = self.restart((context, profile_id, policy));
if let Err(error) = result {
return Ok(answer(reply, Err(error), false));
}
if self.usage.restart(&mut self.durable).await.is_err() {
Ok(answer(reply, Err(ActorError::Closed), true))
} else {
Ok(answer(reply, Ok(()), false))
}
}
Message::Abandon => Ok(true),
}
}
fn stage(&mut self, text: String, boxes: Vec<BoxValue>, reply: StageReply) -> ActorResult {
if self.stage_reply.is_some() || self.mailbox_flush.is_some() {
let _ = reply.send(Err("overlapping Codex stages or mailbox flush".to_owned()));
return Ok(true);
}
self.stage_reply = Some(reply);
let job = self.job.ok_or_else(|| NO_ACTIVE_INFERENCE.to_owned())?;
let calls = self.durable.prepare_stage(job, text, boxes)?;
for call in calls {
match call.disposition().clone() {
PreparedCallDisposition::ImmediateError(error) => {
self.durable.accept_tool_return_v2(
call.tool_call_id,
Err(error.message.clone()),
error.metadata_type.clone(),
error.metadata_contents.clone(),
)
}
PreparedCallDisposition::External => {
self.launch_external(call.tool_call_id, &call.name, &call.arguments)
}
}?;
}
self.continue_flush(true)
}
fn launch_external(&mut self, id: ToolCallId, name: &str, arguments: &str) -> UnitResult {
if !kcode_k1_ktool_docs::is_known_ktool(name) {
return self
.durable
.accept_tool_return(id, Err("unknown Ktool".to_owned()));
}
if name == WEB_SEARCH_TOOL {
return self
.launch_web_search(id, arguments)
.or_else(|error| self.durable.accept_tool_return(id, Err(error)));
}
let result = self.durable.launch_action(name, arguments);
self.durable.accept_tool_return(id, result)
}
fn launch_web_search(&mut self, id: ToolCallId, arguments: &str) -> UnitResult {
if self.searches.iter().any(|search| search.key.id == id) {
return Err("WebSearch failed: duplicate ToolCallId".to_owned());
}
let deadline = Instant::now() + WEB_SEARCH_TIMEOUT;
let request = kcode_k1_chat_websearch_request::parse(arguments, deadline)?;
let epoch = self.search_epoch.checked_add(1);
let epoch = epoch.ok_or_else(|| "WebSearch failed: task epoch exhausted".to_owned())?;
self.search_epoch = epoch;
let key = SearchKey { epoch, id };
let runner = self.web_search.clone();
let sender = self.sender.clone();
let task = tokio::spawn(async move {
let mut guard = SearchDropGuard {
key,
sender,
result: None,
};
guard.result = Some(runner.run(request).await);
});
self.searches.push(SearchTask { key, task });
Ok(())
}
fn search_done(&mut self, key: SearchKey, result: SearchResult) -> ActorResult {
let Some(index) = self.searches.iter().position(|search| search.key == key) else {
return Ok(false);
};
drop(self.searches.swap_remove(index));
self.durable.accept_tool_return(key.id, result)?;
if self.job.is_none() {
return Ok(false);
}
self.continue_flush(false)
}
fn abort_searches(&mut self) {
for search in self.searches.drain(..) {
search.task.abort();
}
}
fn start_mailbox_flush(&mut self) -> Result<bool, String> {
if self.mailbox_flush.is_some() {
return Ok(false);
}
let job = self.job.ok_or_else(|| NO_ACTIVE_INFERENCE.to_owned())?;
let prepared = self.durable.prepare_mailbox_flush(job)?;
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;
let result = result.map_err(|error| error.to_string());
let _ = sender.send(Message::MailboxFlushCompleted(job, prepared, result));
}));
Ok(false)
}
fn flush_done(&mut self, job: u64, prepared: Flush, result: UnitResult) -> ActorResult {
if self.mailbox_flush.take().is_none() || self.job != Some(job) {
return Err("stale Codex mailbox-flush completion".to_owned());
}
result?;
self.durable.commit_mailbox_flush(prepared)?;
self.continue_flush(true)
}
fn continue_flush(&mut self, acknowledge: bool) -> ActorResult {
let done = self.start_mailbox_flush()?;
Ok(done && acknowledge && self.finish_stage())
}
fn finish_stage(&mut self) -> bool {
self.stage_reply
.take()
.is_some_and(|reply| reply.send(Ok(())).is_err())
}
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) -> UnitResult {
if self.inference.is_some() || self.shim.is_none() {
return Ok(());
}
let Some((job, input)) = self.durable.begin_input()? else {
return Ok(());
};
if !self.usage.is_subscribed() {
let subscription = self
.usage
.subscribe(&self.adapter, self.active_key.clone())
.await;
subscription.map_err(|_| "model-usage subscription failed".to_owned())?;
}
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.take().is_none() || self.mailbox_flush.is_some() || self.job != Some(job)
{
return Err("stale Codex inference completion".to_owned());
}
self.job = None;
let key = self.active_key.clone();
let (error, resume, terminal_id) = match result {
Err(error) => {
self.abort_searches();
(Some(error), false, None)
}
Ok(output) => {
let completion = self.durable.complete_with_terminal_response(job, output)?;
let (resume, terminal_id) = completion;
(None, resume, Some(terminal_id))
}
};
let usage = self
.usage
.finish_inference(&mut self.durable, &key, terminal_id)
.await;
if usage.is_err() {
return Ok(true);
}
if let Some(error) = error {
if let Some(reply) = self.stage_reply.take() {
let _ = reply.send(Err(error.clone()));
}
self.durable.fail(job, error, true);
return Ok(false);
}
self.shim = Some(shim);
if !resume && self.searches.is_empty() {
self.durable.clear_authorization();
}
Ok(false)
}
fn restart(&mut self, (context, profile_id, policy): RestartAccess) -> Result<(), ActorError> {
let generation = self.generation.checked_add(1);
let generation = generation.ok_or(ActorError::NotRestartable)?;
let restart = self.durable.restart(context, profile_id, policy);
restart.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 reply_transition(reply: Reply<()>, result: Result<(), TransitionError>) -> bool {
let stop = matches!(result, Err(TransitionError::Internal(_)));
answer(reply, result.map_err(map_transition), stop)
}
fn answer<T>(reply: Reply<T>, result: Result<T, ActorError>, stop: bool) -> bool {
let _ = reply.send(result);
stop
}