kcode-k1-chat-core 0.1.0

Core chat contracts, update delivery, and inference retry behavior
Documentation
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;