pub use kcode_k1_chat_codex_codec::{BoxValue, Call};
pub use kcode_k1_chat_state::{ActionId, ChatBox};
use kcode_k1_chat_codex_codec::project;
use kcode_k1_chat_state::{ActorState, DispatchCall, DispatchOutcome};
use kcode_k1_codex_adapter::{ShimItem, ShimOutput};
#[derive(Clone, Debug, Eq, PartialEq)]
pub enum Arrival {
System(String),
User(String),
Attachment,
Return {
action_id: ActionId,
result: Result<String, String>,
},
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub struct Start {
pub job: u64,
pub boxes: Vec<BoxValue>,
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub struct Launch {
pub action_id: ActionId,
pub result: Result<String, String>,
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub enum Status {
Running,
Quiet,
Stalled { message: String, restartable: bool },
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum RestartError {
NotStalled,
NotRestartable,
}
#[derive(Clone)]
struct Launched {
action_id: ActionId,
call: Call,
}
struct Round {
job: u64,
ledger: Vec<Launched>,
arrivals: Vec<Arrival>,
}
enum Mode {
Idle,
Running(Round),
Stalled { message: String, restartable: bool },
}
pub struct ConversationState {
state: ActorState,
session: [u8; 12],
sequence: u64,
submitted: usize,
mode: Mode,
}
impl ConversationState {
pub fn new(session: [u8; 12]) -> Self {
Self {
state: ActorState::new(false),
session,
sequence: 0,
submitted: 0,
mode: Mode::Idle,
}
}
pub fn boxes(&self) -> &[ChatBox] {
self.state.boxes()
}
pub fn status(&self) -> Status {
match &self.mode {
Mode::Running(_) => Status::Running,
Mode::Idle if self.state.quiet() => Status::Quiet,
Mode::Idle => Status::Running,
Mode::Stalled {
message,
restartable,
} => Status::Stalled {
message: message.clone(),
restartable: *restartable,
},
}
}
pub fn accept(&mut self, arrival: Arrival) -> Result<(), String> {
if let Mode::Running(round) = &mut self.mode {
round.arrivals.push(arrival);
Ok(())
} else {
self.apply(arrival)
}
}
pub fn begin(&mut self) -> Result<Option<Start>, String> {
if !matches!(self.mode, Mode::Idle) {
return Ok(None);
}
let Some(start) = self.state.begin_inference().map_err(debug)? else {
return Ok(None);
};
let boxes = self.state.boxes();
let projected = boxes[self.submitted..].iter().map(project).collect();
self.submitted = boxes.len();
self.mode = Mode::Running(Round {
job: start.job,
ledger: Vec::new(),
arrivals: Vec::new(),
});
Ok(Some(Start {
job: start.job,
boxes: projected,
}))
}
pub fn launch(&mut self, job: u64, call: Call) -> Result<Launch, String> {
let Mode::Running(round) = &mut self.mode else {
return Err("no K1 inference is running".to_owned());
};
if round.job != job {
return Err("stale K1 inference launch".to_owned());
}
let sequence = self
.sequence
.checked_add(1)
.ok_or_else(|| "K1 ActionId space was exhausted".to_owned())?;
let action_id = ActionId::new(self.session, sequence);
self.state
.collect_provider_call(job, call.name.clone(), call.arguments.clone())
.map_err(debug)?;
round.ledger.push(Launched {
action_id,
call: call.clone(),
});
self.sequence = sequence;
let result = kcode_k1_chat_thread_actions::launch(&call.name, &call.arguments);
Ok(Launch { action_id, result })
}
pub fn complete(&mut self, job: u64, output: ShimOutput<BoxValue>) -> Result<(), String> {
let round = self.take_round(job)?;
let text = match validate(output, &round.ledger) {
Ok(text) => text,
Err(error) => {
self.preserve(round, error.clone(), false);
return Err(error);
}
};
if let Err(error) = self.commit(&round, Some(&text)) {
let message = format!("failed to commit Codex output: {error}");
let _ = self.state.halt(message.clone());
self.mode = Mode::Stalled {
message: message.clone(),
restartable: false,
};
return Err(message);
}
self.mode = Mode::Idle;
Ok(())
}
pub fn fail(&mut self, job: u64, message: String, restartable_before_launch: bool) {
if let Ok(round) = self.take_round(job) {
self.preserve(round, message, restartable_before_launch);
}
}
pub fn restart(&mut self) -> Result<(), RestartError> {
match &self.mode {
Mode::Stalled {
restartable: true, ..
} => {}
Mode::Stalled { .. } => return Err(RestartError::NotRestartable),
_ => return Err(RestartError::NotStalled),
}
self.state
.restart()
.map_err(|_| RestartError::NotRestartable)?;
self.submitted = 0;
self.mode = Mode::Idle;
Ok(())
}
fn take_round(&mut self, job: u64) -> Result<Round, String> {
let mode = std::mem::replace(&mut self.mode, Mode::Idle);
match mode {
Mode::Running(round) if round.job == job => Ok(round),
other => {
self.mode = other;
Err("stale K1 inference completion".to_owned())
}
}
}
fn commit(&mut self, round: &Round, text: Option<&str>) -> Result<(), String> {
if let Some(text) = text {
self.state
.append_kennedy_text(round.job, text)
.map_err(debug)?;
}
let action_ids = round.ledger.iter().map(|entry| entry.action_id).collect();
let dispatches = self
.state
.complete_provider_output(round.job, action_ids)
.map_err(debug)?;
if !dispatches_match(&round.ledger, &dispatches) {
return Err("committed calls did not match launch ledger".to_owned());
}
if !dispatches.is_empty() {
self.state
.complete_dispatch(vec![DispatchOutcome::Pending; dispatches.len()])
.map_err(debug)?;
}
self.apply_all(round.arrivals.clone())
}
fn preserve(&mut self, round: Round, message: String, restartable_before_launch: bool) {
let preserved = if round.ledger.is_empty() {
self.state
.stall_inference(round.job, message.clone())
.map_err(debug)
.and_then(|_| self.apply_all(round.arrivals))
.is_ok()
} else {
let preserved = self.commit(&round, None).is_ok();
let _ = self.state.halt(message.clone());
preserved
};
self.mode = Mode::Stalled {
message,
restartable: round.ledger.is_empty() && restartable_before_launch && preserved,
};
}
fn apply_all(&mut self, arrivals: Vec<Arrival>) -> Result<(), String> {
for arrival in arrivals {
self.apply(arrival)?;
}
Ok(())
}
fn apply(&mut self, arrival: Arrival) -> Result<(), String> {
match arrival {
Arrival::System(text) => self.state.accept_system(text).map_err(debug),
Arrival::User(text) => self.state.accept_user(text).map_err(debug),
Arrival::Attachment => self.state.accept_attachment().map_err(debug),
Arrival::Return { action_id, result } => self
.state
.accept_async_return(action_id, result)
.map_err(debug),
}
}
}
fn validate(output: ShimOutput<BoxValue>, ledger: &[Launched]) -> Result<String, String> {
let mut text = String::new();
let mut calls = Vec::new();
for item in output.items {
match item {
ShimItem::Text(value) => text.push_str(&value),
ShimItem::Box(BoxValue::Call(Ok(call))) => calls.push(call),
ShimItem::Box(_) => return Err("shim returned an invalid call box".to_owned()),
}
}
if calls
!= ledger
.iter()
.map(|entry| entry.call.clone())
.collect::<Vec<_>>()
{
return Err("shim output did not match the launch ledger".to_string());
}
Ok(text)
}
fn dispatches_match(ledger: &[Launched], dispatches: &[DispatchCall]) -> bool {
ledger.len() == dispatches.len()
&& ledger.iter().zip(dispatches).all(|(expected, actual)| {
expected.action_id == actual.action_id
&& expected.call.name == actual.name
&& expected.call.arguments == actual.arguments
})
}
fn debug(error: impl std::fmt::Debug) -> String {
format!("{error:?}")
}