kcode-k1-chat-thread-actions 0.1.0

Tool scheduling and update delivery for K1 chat threads
Documentation
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(),
        )
    })
}