use std::{sync::Arc, time::Duration};
use kcode_k1_chat_core::{SubmittedUpdate, UpdateSink};
use tokio::{
task::JoinHandle,
time::{Instant, sleep_until},
};
pub use kcode_k1_chat_core::{
ActionId, ChatError, Runtime, ToolMode, ToolOutput, ToolRequest, ToolStart, Updates,
};
const FAST_BOUNDARY: Duration = Duration::from_secs(2);
type ToolTask = JoinHandle<(ToolOutput, Instant)>;
pub trait ActionSink: Send + Sync + 'static {
fn submit(&self, event: ActionEvent) -> Result<(), ChatError>;
}
pub enum ActionEvent {
Update {
action: ActionId,
identity: u64,
update: SubmittedUpdate,
},
ToolReplies {
entries: Vec<Immediate>,
finished: bool,
},
ToolDone {
action: ActionId,
output: ToolOutput,
started: Option<Instant>,
},
}
pub struct ToolCall {
pub index: usize,
pub action: ActionId,
pub request: ToolRequest,
}
pub struct Immediate {
pub index: usize,
pub action: ActionId,
pub kind: ImmediateKind,
pub complete: bool,
}
pub enum ImmediateKind {
Plain(String),
Tool {
output: ToolOutput,
started: Option<Instant>,
},
Worker(String),
}
impl ImmediateKind {
pub fn text(&self) -> &str {
match self {
Self::Plain(text) | Self::Worker(text) => text,
Self::Tool { output, .. } => &output.text,
}
}
}
impl Immediate {
pub fn plain(index: usize, action: ActionId, text: String, complete: bool) -> Self {
Self {
index,
action,
kind: ImmediateKind::Plain(text),
complete,
}
}
pub fn worker(index: usize, action: ActionId, text: String, complete: bool) -> Self {
Self {
index,
action,
kind: ImmediateKind::Worker(text),
complete,
}
}
}
#[derive(Clone)]
struct BoundSink {
sink: Arc<dyn ActionSink>,
}
impl UpdateSink for BoundSink {
fn submit(
&self,
action: ActionId,
identity: u64,
update: SubmittedUpdate,
) -> Result<(), ChatError> {
self.sink.submit(ActionEvent::Update {
action,
identity,
update,
})
}
}
pub fn start_tool_batch(
runtime: Arc<dyn Runtime>,
calls: Vec<ToolCall>,
mut immediate: Vec<Immediate>,
sink: Arc<dyn ActionSink>,
) {
let deadline = Instant::now() + FAST_BOUNDARY;
let update_sink: Arc<dyn UpdateSink> = Arc::new(BoundSink { sink: sink.clone() });
let mut tools = Vec::with_capacity(calls.len());
for call in calls {
match runtime.start_tool(
call.request,
Updates::bind(call.action, update_sink.clone()),
) {
Ok(start) => tools.push(ToolSpec {
index: call.index,
action: call.action,
start,
started: Instant::now(),
}),
Err(error) => immediate.push(Immediate {
index: call.index,
action: call.action,
kind: ImmediateKind::Tool {
output: ToolOutput {
text: error,
cost_cents: Default::default(),
},
started: None,
},
complete: true,
}),
}
}
launch(tools, immediate, deadline, sink);
}
pub fn redrive_tool(
runtime: Arc<dyn Runtime>,
action: ActionId,
request: ToolRequest,
sink: Arc<dyn ActionSink>,
) {
let update_sink: Arc<dyn UpdateSink> = Arc::new(BoundSink { sink: sink.clone() });
match runtime.start_tool(request, Updates::bind(action, update_sink)) {
Ok(start) => {
let started = Instant::now();
tokio::spawn(async move {
let output = start.future.await;
let _ = sink.submit(ActionEvent::ToolDone {
action,
output,
started: Some(started),
});
});
}
Err(error) => {
let _ = sink.submit(ActionEvent::ToolDone {
action,
output: ToolOutput {
text: error,
cost_cents: Default::default(),
},
started: None,
});
}
}
}
struct ToolSpec {
index: usize,
action: ActionId,
start: ToolStart,
started: Instant,
}
struct Slot {
index: usize,
action: ActionId,
queued: String,
running: Option<(ToolMode, ToolTask, Instant)>,
immediate: Option<ImmediateKind>,
}
enum Terminal {
Ready(ToolOutput, Instant),
Task(ToolTask, Instant),
}
fn launch(
specs: Vec<ToolSpec>,
entries: Vec<Immediate>,
deadline: Instant,
sink: Arc<dyn ActionSink>,
) {
tokio::spawn(async move {
let mut slots = entries
.into_iter()
.map(|entry| Slot {
index: entry.index,
action: entry.action,
queued: String::new(),
running: None,
immediate: Some(entry.kind),
})
.collect::<Vec<_>>();
for spec in specs {
let ToolStart {
mode,
queued,
future,
} = spec.start;
let task = tokio::spawn(async move { (future.await, Instant::now()) });
slots.push(Slot {
index: spec.index,
action: spec.action,
queued,
running: Some((mode, task, spec.started)),
immediate: None,
});
}
slots.sort_by_key(|slot| slot.index);
let count = slots.len();
for (position, slot) in slots.into_iter().enumerate() {
let finished = position + 1 == count;
let mut terminal = None;
let immediate = if let Some(kind) = slot.immediate {
Immediate {
index: slot.index,
action: slot.action,
complete: !matches!(kind, ImmediateKind::Plain(_)),
kind,
}
} else {
let (mode, mut task, started) = slot.running.expect("tool slot is populated");
let observed = match mode {
ToolMode::Queued => None,
ToolMode::Fast => tokio::select! {
biased;
result = &mut task => Some(join_tool(result)),
_ = sleep_until(deadline) => None,
},
};
match observed {
Some((output, completed_at)) if completed_at <= deadline => Immediate {
index: slot.index,
action: slot.action,
kind: ImmediateKind::Tool {
output,
started: Some(started),
},
complete: true,
},
Some((output, _)) => {
terminal = Some(Terminal::Ready(output, started));
Immediate::plain(slot.index, slot.action, slot.queued, false)
}
None => {
terminal = Some(Terminal::Task(task, started));
Immediate::plain(slot.index, slot.action, slot.queued, false)
}
}
};
if sink
.submit(ActionEvent::ToolReplies {
entries: vec![immediate],
finished,
})
.is_err()
{
return;
}
match terminal {
Some(Terminal::Ready(output, started)) => {
let _ = sink.submit(ActionEvent::ToolDone {
action: slot.action,
output,
started: Some(started),
});
}
Some(Terminal::Task(task, started)) => {
let sink = sink.clone();
let action = slot.action;
tokio::spawn(async move {
let (output, _) = join_tool(task.await);
let _ = sink.submit(ActionEvent::ToolDone {
action,
output,
started: Some(started),
});
});
}
None => {}
}
}
});
}
fn join_tool(
result: Result<(ToolOutput, Instant), tokio::task::JoinError>,
) -> (ToolOutput, Instant) {
result.unwrap_or_else(|error| {
(
ToolOutput {
text: error.to_string(),
cost_cents: Default::default(),
},
Instant::now(),
)
})
}