use std::{
future::Future,
pin::Pin,
sync::{
Arc,
atomic::{AtomicU8, AtomicU64, Ordering},
},
time::Duration,
};
pub type LlmFuture<'a> = Pin<Box<dyn Future<Output = Result<Inference, LlmError>> + Send + 'a>>;
pub type ToolFuture = Pin<Box<dyn Future<Output = String> + Send + 'static>>;
pub type CompactFuture = Pin<Box<dyn Future<Output = Result<String, String>> + Send + 'static>>;
pub trait Llm: Send + Sync {
fn start(&self) -> Box<dyn LlmThread>;
}
pub trait LlmThread: Send {
fn infer<'a>(&'a mut self, delta: &'a str) -> LlmFuture<'a>;
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum LlmError {
Transient(String),
Permanent(String),
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct Inference {
pub text: String,
pub calls: Vec<Call>,
pub continue_inference: bool,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum Call {
Tool(ToolRequest),
Worker(WorkerRequest),
Compact(CompactRequest),
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct ToolRequest {
pub name: String,
pub input: String,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct WorkerRequest {
pub llm: String,
pub prompt: String,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct CompactRequest {
pub instruction: String,
}
#[derive(Copy, Clone, Debug, PartialEq, Eq)]
pub enum ToolMode {
Fast,
Queued,
}
pub struct ToolStart {
pub mode: ToolMode,
pub queued: String,
pub future: ToolFuture,
}
pub struct WorkerStart {
pub llm: Arc<dyn Llm>,
pub queued: String,
}
pub trait Runtime: Send + Sync {
fn start_tool(&self, request: ToolRequest, updates: Updates) -> Result<ToolStart, String>;
fn start_worker(&self, request: &WorkerRequest) -> Result<WorkerStart, String>;
fn compact(&self, request: CompactRequest, primary: String) -> CompactFuture;
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum SubmittedUpdate {
Activity(String),
Append(String),
}
pub trait UpdateSink: Send + Sync + 'static {
fn submit(&self, job: u64, identity: u64, update: SubmittedUpdate) -> Result<(), ChatError>;
}
#[derive(Clone)]
pub struct Updates {
job: u64,
next_identity: Arc<AtomicU64>,
sink: Arc<dyn UpdateSink>,
}
impl Updates {
pub fn bind(job: u64, sink: Arc<dyn UpdateSink>) -> Self {
Self {
job,
next_identity: Arc::new(AtomicU64::new(1)),
sink,
}
}
pub fn activity(&self, activity: String) -> PreparedUpdate {
self.prepare(SubmittedUpdate::Activity(activity))
}
pub fn append(&self, text: String) -> PreparedUpdate {
self.prepare(SubmittedUpdate::Append(text))
}
fn prepare(&self, update: SubmittedUpdate) -> PreparedUpdate {
PreparedUpdate {
job: self.job,
identity: self.next_identity.fetch_add(1, Ordering::Relaxed),
update,
sink: self.sink.clone(),
}
}
}
#[derive(Clone)]
pub struct PreparedUpdate {
job: u64,
identity: u64,
update: SubmittedUpdate,
sink: Arc<dyn UpdateSink>,
}
impl PreparedUpdate {
pub fn send(&self) -> Result<(), ChatError> {
self.sink
.submit(self.job, self.identity, self.update.clone())
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct ChatView {
pub primary: String,
pub pending: String,
pub history: Vec<String>,
pub actions: Vec<PendingAction>,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum PendingAction {
Inference { attempt: u8 },
Tool { name: String },
Worker { llm: String },
Compaction,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum ChatEvent {
Text(String),
Activity(String),
Stalled(String),
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum ChatError {
Empty,
NotStalled,
Busy,
Closed,
}
pub async fn infer_with_retry(
mut thread: Box<dyn LlmThread>,
delta: String,
attempt: Arc<AtomicU8>,
) -> (Box<dyn LlmThread>, Result<Inference, String>) {
let mut attempt_number = 1;
loop {
attempt.store(attempt_number, Ordering::Relaxed);
match thread.infer(&delta).await {
Ok(inference) => return (thread, Ok(inference)),
Err(LlmError::Permanent(message)) => return (thread, Err(message)),
Err(LlmError::Transient(message)) => {
if attempt_number == 5 {
return (thread, Err(message));
}
tokio::time::sleep(Duration::from_secs(10 << (attempt_number - 1))).await;
attempt_number += 1;
}
}
}
}
#[cfg(test)]
mod tests;