Skip to main content

kcode_k1_chat_thread_actions/
lib.rs

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}