#![forbid(unsafe_code)]
use kcode_k1_access_kmap::K1AccessKmap;
use kcode_k1_chat_persistence::Session;
use kcode_k1_chat_state::USER_MESSAGE_TYPE;
use kcode_k1_chat_thread_actor_channel::{
AcceptedBox, ActorError, ActorShim, BoxValue, Handle, Message, ProviderInput,
ProviderInputKind, Reply, Snapshot, StageReply, channel, new_shim,
};
use kcode_k1_chat_thread_durable_state::{
DurableThread, PreparedSteer, Status, ToolCallId, TransitionError,
};
use kcode_k1_chat_thread_web_search::{WebSearchAction, WebSearchToolEvent};
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) = channel();
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,
sender,
receiver,
waiters: Vec::new(),
inference: None,
steering: None,
searches: Vec::new(),
search_epoch: 0,
stage_reply: None,
job: None,
model_input: None,
web_search,
};
tokio::spawn(actor.run());
Ok(handle)
}
struct Actor {
durable: DurableThread,
shim: Option<ActorShim>,
adapter: Adapter,
base_key: String,
active_key: String,
generation: u64,
sender: mpsc::UnboundedSender<Message>,
receiver: mpsc::UnboundedReceiver<Message>,
waiters: Vec<Reply<Snapshot>>,
inference: Option<JoinHandle<()>>,
steering: Option<JoinHandle<()>>,
searches: Vec<(ToolCallId, JoinHandle<()>)>,
search_epoch: u64,
stage_reply: Option<StageReply>,
job: Option<u64>,
model_input: Option<ProviderInput>,
web_search: WebSearchAction,
}
impl Actor {
async fn run(mut self) {
'actor: loop {
if self.drive() {
break;
}
self.wake();
let Some(mut message) = self.receiver.recv().await else {
break;
};
loop {
if self.handle(message) {
break 'actor;
}
match self.receiver.try_recv() {
Ok(next) => message = next,
Err(_) => break,
}
}
}
self.abort_searches();
if let Some(reply) = self.stage_reply.take() {
let _ = reply.send(Err("K1 actor is closed".to_owned()));
}
for task in [self.steering.take(), self.inference.take()]
.into_iter()
.flatten()
{
task.abort();
}
for waiter in self.waiters {
let _ = waiter.send(Err(ActorError::Closed));
}
}
fn handle(&mut self, message: Message) -> bool {
match message {
Message::Accept(value, reply) => self.accept(value, reply),
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::Steered(job, prepared, result) => self.steered(job, prepared, result),
Message::Inferred(job, shim, result) => self.finish(job, shim, result),
Message::Snapshot(reply) => {
let _ = reply.send(Ok(self.snapshot()));
false
}
Message::Wait(reply) => {
if self.running() {
self.waiters.push(reply);
} else {
let _ = reply.send(Ok(self.snapshot()));
}
false
}
Message::Restart(context, profile_id, policy, reply) => {
let result = self.restart(context, profile_id, policy);
let _ = reply.send(result);
false
}
Message::Abandon => true,
}
}
fn accept(&mut self, value: AcceptedBox, reply: Reply<()>) -> bool {
let (box_type, contents, hidden_type, hidden_contents) = value;
if box_type == USER_MESSAGE_TYPE {
let _ = reply.send(Err(ActorError::Unauthorized));
return false;
}
let result = self
.durable
.accept_box(box_type, contents, hidden_type, hidden_contents);
let stop = result.is_err();
let _ = reply.send(result.map_err(|_| ActorError::Closed));
stop
}
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.steering.is_some() {
return reject_stage(reply, "overlapping Codex stages or steer".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 {
if call.name == self.web_search.tool_name() {
if self.search_live(&call.tool_call_id) {
return self.fatal("duplicate WebSearch ToolCallId".into());
}
self.launch_search(call.tool_call_id, call.arguments);
} 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_steer() {
Ok(true) => self.finish_stage(),
Ok(false) => false,
Err(error) => self.fatal(error),
}
}
fn launch_search(&mut self, id: ToolCallId, arguments: String) {
let mut search = self.web_search.launch(&arguments);
let sender = self.sender.clone();
let epoch = self.search_epoch;
let event_id = id;
let task = tokio::spawn(async move {
while let Some(event) = search.recv().await {
let terminal = matches!(&event, WebSearchToolEvent::Result { .. });
if sender
.send(Message::WebSearch(epoch, event_id, event))
.is_err()
{
return;
}
if terminal {
return;
}
}
let _ = sender.send(Message::WebSearchEnded(epoch, event_id));
});
self.searches.push((id, task));
}
fn web_search_event(&mut self, epoch: u64, id: ToolCallId, event: WebSearchToolEvent) -> bool {
if epoch != self.search_epoch {
return false;
}
match event {
WebSearchToolEvent::Message { contents } => {
if !self.search_live(&id) {
return self.fatal("WebSearch Message has no live task".into());
}
if let Err(error) = self.durable.accept_tool_message(id, contents) {
return self.fatal(error);
}
false
}
WebSearchToolEvent::Result {
result,
metadata_type,
metadata_contents,
} => {
let Some(task) = self.remove_search(&id) else {
return self.fatal("WebSearch Result has no live task".into());
};
task.abort();
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_steer() {
Ok(_) => false,
Err(error) => self.fatal(error),
}
}
}
}
fn web_search_ended(&mut self, epoch: u64, id: ToolCallId) -> bool {
if epoch != self.search_epoch {
return false;
}
let Some(task) = self.remove_search(&id) else {
return self.fatal("ended WebSearch has no live task".into());
};
task.abort();
self.fatal("WebSearch ended without a terminal Result".into())
}
fn start_steer(&mut self) -> Result<bool, String> {
let job = self
.job
.ok_or_else(|| "no active K1 inference".to_owned())?;
let prepared = self.durable.prepare_steer(job)?;
if self.steering.is_some() {
return Ok(false);
}
let Some(prepared) = prepared else {
return Ok(true);
};
let input = self.durable.prepared_input(&prepared)?;
self.model_input = Some(ProviderInput {
kind: ProviderInputKind::Steer,
text: input.clone(),
});
let adapter = self.adapter.clone();
let key = self.active_key.clone();
let sender = self.sender.clone();
self.steering = Some(tokio::spawn(async move {
let result = adapter
.steer(key, input)
.await
.map_err(|error| error.to_string());
let _ = sender.send(Message::Steered(job, prepared, result));
}));
Ok(false)
}
fn steered(&mut self, job: u64, prepared: PreparedSteer, result: Result<(), String>) -> bool {
if self.steering.take().is_none() || self.job != Some(job) {
return self.fatal("stale Codex steer completion".into());
}
if let Err(error) = result {
return self.fatal(error);
}
if let Err(error) = self.durable.commit_steer(prepared) {
return self.fatal(error);
}
if self.stage_reply.is_some() && self.finish_stage() {
return true;
}
match self.start_steer() {
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
}
fn search_live(&self, id: &ToolCallId) -> bool {
self.searches.iter().any(|(live, _)| live == id)
}
fn remove_search(&mut self, id: &ToolCallId) -> Option<JoinHandle<()>> {
let index = self.searches.iter().position(|(live, _)| live == id)?;
Some(self.searches.swap_remove(index).1)
}
fn abort_searches(&mut self) {
for (_, task) in self.searches.drain(..) {
task.abort();
}
}
fn cancel_searches(&mut self) -> Result<(), String> {
self.abort_searches();
self.search_epoch = self
.search_epoch
.checked_add(1)
.ok_or_else(|| "WebSearch cancellation epoch was exhausted".to_owned())?;
Ok(())
}
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),
};
self.model_input = Some(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
}
fn finish(
&mut self,
job: u64,
shim: ActorShim,
result: Result<ShimOutput<BoxValue>, String>,
) -> bool {
if self.inference.take().is_none() || self.steering.is_some() || self.job != Some(job) {
return self.fatal("stale Codex inference completion".into());
}
self.job = None;
match result {
Err(error) => {
let fence = self.cancel_searches();
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 = match self.durable.complete(job, output) {
Ok(resume) => resume,
Err(error) => return self.fatal(error),
};
self.shim = Some(shim);
if !resume && self.searches.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 running(&self) -> bool {
!self.searches.is_empty() || matches!(self.durable.status(), Status::Running)
}
fn snapshot(&self) -> Snapshot {
let status = if self.searches.is_empty() {
self.durable.status()
} else {
Status::Running
};
Snapshot {
boxes: self.durable.boxes().to_vec(),
status,
model_input: self.model_input.clone(),
}
}
fn wake(&mut self) {
if !self.running() {
let snapshot = self.snapshot();
for waiter in std::mem::take(&mut self.waiters) {
let _ = waiter.send(Ok(snapshot.clone()));
}
}
}
}
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
}