1use std::{sync::Arc, time::Duration};
2
3use kcode_k1_chat_core::{SubmittedUpdate, UpdateSink};
4use tokio::{
5 task::JoinHandle,
6 time::{Instant, sleep_until},
7};
8
9pub use kcode_k1_chat_core::{
10 ActionId, ChatError, Runtime, ToolMode, ToolOutput, ToolRequest, ToolStart, Updates,
11};
12
13const FAST_BOUNDARY: Duration = Duration::from_secs(2);
14type ToolTask = JoinHandle<(ToolOutput, Instant)>;
15
16pub trait ActionSink: Send + Sync + 'static {
17 fn submit(&self, event: ActionEvent) -> Result<(), ChatError>;
18}
19
20pub enum ActionEvent {
21 Update {
22 action: ActionId,
23 identity: u64,
24 update: SubmittedUpdate,
25 },
26 ToolReplies {
27 entries: Vec<Immediate>,
28 finished: bool,
29 },
30 ToolDone {
31 action: ActionId,
32 output: ToolOutput,
33 started: Option<Instant>,
34 },
35}
36
37pub struct ToolCall {
38 pub index: usize,
39 pub action: ActionId,
40 pub request: ToolRequest,
41}
42
43pub struct Immediate {
44 pub index: usize,
45 pub action: ActionId,
46 pub kind: ImmediateKind,
47 pub complete: bool,
48}
49
50pub enum ImmediateKind {
51 Plain(String),
52 Tool {
53 output: ToolOutput,
54 started: Option<Instant>,
55 },
56 Worker(String),
57}
58
59impl ImmediateKind {
60 pub fn text(&self) -> &str {
61 match self {
62 Self::Plain(text) | Self::Worker(text) => text,
63 Self::Tool { output, .. } => &output.text,
64 }
65 }
66}
67
68impl Immediate {
69 pub fn plain(index: usize, action: ActionId, text: String, complete: bool) -> Self {
70 Self {
71 index,
72 action,
73 kind: ImmediateKind::Plain(text),
74 complete,
75 }
76 }
77
78 pub fn worker(index: usize, action: ActionId, text: String, complete: bool) -> Self {
79 Self {
80 index,
81 action,
82 kind: ImmediateKind::Worker(text),
83 complete,
84 }
85 }
86}
87
88#[derive(Clone)]
89struct BoundSink {
90 sink: Arc<dyn ActionSink>,
91}
92
93impl UpdateSink for BoundSink {
94 fn submit(
95 &self,
96 action: ActionId,
97 identity: u64,
98 update: SubmittedUpdate,
99 ) -> Result<(), ChatError> {
100 self.sink.submit(ActionEvent::Update {
101 action,
102 identity,
103 update,
104 })
105 }
106}
107
108pub fn start_tool_batch(
109 runtime: Arc<dyn Runtime>,
110 calls: Vec<ToolCall>,
111 mut immediate: Vec<Immediate>,
112 sink: Arc<dyn ActionSink>,
113) {
114 let deadline = Instant::now() + FAST_BOUNDARY;
115 let update_sink: Arc<dyn UpdateSink> = Arc::new(BoundSink { sink: sink.clone() });
116 let mut tools = Vec::with_capacity(calls.len());
117
118 for call in calls {
119 match runtime.start_tool(
120 call.request,
121 Updates::bind(call.action, update_sink.clone()),
122 ) {
123 Ok(start) => tools.push(ToolSpec {
124 index: call.index,
125 action: call.action,
126 start,
127 started: Instant::now(),
128 }),
129 Err(error) => immediate.push(Immediate {
130 index: call.index,
131 action: call.action,
132 kind: ImmediateKind::Tool {
133 output: ToolOutput {
134 text: error,
135 cost_cents: Default::default(),
136 },
137 started: None,
138 },
139 complete: true,
140 }),
141 }
142 }
143
144 launch(tools, immediate, deadline, sink);
145}
146
147pub fn redrive_tool(
148 runtime: Arc<dyn Runtime>,
149 action: ActionId,
150 request: ToolRequest,
151 sink: Arc<dyn ActionSink>,
152) {
153 let update_sink: Arc<dyn UpdateSink> = Arc::new(BoundSink { sink: sink.clone() });
154 match runtime.start_tool(request, Updates::bind(action, update_sink)) {
155 Ok(start) => {
156 let started = Instant::now();
157 tokio::spawn(async move {
158 let output = start.future.await;
159 let _ = sink.submit(ActionEvent::ToolDone {
160 action,
161 output,
162 started: Some(started),
163 });
164 });
165 }
166 Err(error) => {
167 let _ = sink.submit(ActionEvent::ToolDone {
168 action,
169 output: ToolOutput {
170 text: error,
171 cost_cents: Default::default(),
172 },
173 started: None,
174 });
175 }
176 }
177}
178
179struct ToolSpec {
180 index: usize,
181 action: ActionId,
182 start: ToolStart,
183 started: Instant,
184}
185
186struct Slot {
187 index: usize,
188 action: ActionId,
189 queued: String,
190 running: Option<(ToolMode, ToolTask, Instant)>,
191 immediate: Option<ImmediateKind>,
192}
193
194enum Terminal {
195 Ready(ToolOutput, Instant),
196 Task(ToolTask, Instant),
197}
198
199fn launch(
200 specs: Vec<ToolSpec>,
201 entries: Vec<Immediate>,
202 deadline: Instant,
203 sink: Arc<dyn ActionSink>,
204) {
205 tokio::spawn(async move {
206 let mut slots = entries
207 .into_iter()
208 .map(|entry| Slot {
209 index: entry.index,
210 action: entry.action,
211 queued: String::new(),
212 running: None,
213 immediate: Some(entry.kind),
214 })
215 .collect::<Vec<_>>();
216
217 for spec in specs {
218 let ToolStart {
219 mode,
220 queued,
221 future,
222 } = spec.start;
223 let task = tokio::spawn(async move { (future.await, Instant::now()) });
224 slots.push(Slot {
225 index: spec.index,
226 action: spec.action,
227 queued,
228 running: Some((mode, task, spec.started)),
229 immediate: None,
230 });
231 }
232 slots.sort_by_key(|slot| slot.index);
233
234 let count = slots.len();
235 for (position, slot) in slots.into_iter().enumerate() {
236 let finished = position + 1 == count;
237 let mut terminal = None;
238 let immediate = if let Some(kind) = slot.immediate {
239 Immediate {
240 index: slot.index,
241 action: slot.action,
242 complete: !matches!(kind, ImmediateKind::Plain(_)),
243 kind,
244 }
245 } else {
246 let (mode, mut task, started) = slot.running.expect("tool slot is populated");
247 let observed = match mode {
248 ToolMode::Queued => None,
249 ToolMode::Fast => tokio::select! {
250 biased;
251 result = &mut task => Some(join_tool(result)),
252 _ = sleep_until(deadline) => None,
253 },
254 };
255 match observed {
256 Some((output, completed_at)) if completed_at <= deadline => Immediate {
257 index: slot.index,
258 action: slot.action,
259 kind: ImmediateKind::Tool {
260 output,
261 started: Some(started),
262 },
263 complete: true,
264 },
265 Some((output, _)) => {
266 terminal = Some(Terminal::Ready(output, started));
267 Immediate::plain(slot.index, slot.action, slot.queued, false)
268 }
269 None => {
270 terminal = Some(Terminal::Task(task, started));
271 Immediate::plain(slot.index, slot.action, slot.queued, false)
272 }
273 }
274 };
275
276 if sink
277 .submit(ActionEvent::ToolReplies {
278 entries: vec![immediate],
279 finished,
280 })
281 .is_err()
282 {
283 return;
284 }
285 match terminal {
286 Some(Terminal::Ready(output, started)) => {
287 let _ = sink.submit(ActionEvent::ToolDone {
288 action: slot.action,
289 output,
290 started: Some(started),
291 });
292 }
293 Some(Terminal::Task(task, started)) => {
294 let sink = sink.clone();
295 let action = slot.action;
296 tokio::spawn(async move {
297 let (output, _) = join_tool(task.await);
298 let _ = sink.submit(ActionEvent::ToolDone {
299 action,
300 output,
301 started: Some(started),
302 });
303 });
304 }
305 None => {}
306 }
307 }
308 });
309}
310
311fn join_tool(
312 result: Result<(ToolOutput, Instant), tokio::task::JoinError>,
313) -> (ToolOutput, Instant) {
314 result.unwrap_or_else(|error| {
315 (
316 ToolOutput {
317 text: error.to_string(),
318 cost_cents: Default::default(),
319 },
320 Instant::now(),
321 )
322 })
323}