kcode-k1-chat-core 0.1.0

Core chat contracts, update delivery, and inference retry behavior
Documentation
use std::{
    collections::VecDeque,
    sync::{
        Arc, Mutex,
        atomic::{AtomicU8, Ordering},
    },
    time::Duration,
};

use tokio::time::Instant;

use crate::{
    ChatError, Inference, LlmError, LlmFuture, LlmThread, SubmittedUpdate, UpdateSink, Updates,
    infer_with_retry,
};

type Trace = Arc<Mutex<Vec<(usize, String, u8)>>>;

struct RetryThread {
    outcomes: VecDeque<Result<Inference, LlmError>>,
    trace: Trace,
    attempt: Arc<AtomicU8>,
}

impl RetryThread {
    fn new(
        outcomes: Vec<Result<Inference, LlmError>>,
        trace: Trace,
        attempt: Arc<AtomicU8>,
    ) -> Self {
        Self {
            outcomes: outcomes.into(),
            trace,
            attempt,
        }
    }
}

impl LlmThread for RetryThread {
    fn infer<'a>(&'a mut self, delta: &'a str) -> LlmFuture<'a> {
        let identity = self as *const Self as usize;
        self.trace.lock().unwrap().push((
            identity,
            delta.to_owned(),
            self.attempt.load(Ordering::Relaxed),
        ));
        let outcome = self.outcomes.pop_front().unwrap();
        Box::pin(async move { outcome })
    }
}

fn inference(text: &str) -> Inference {
    Inference {
        text: text.to_owned(),
        calls: Vec::new(),
        continue_inference: false,
    }
}

#[tokio::test(start_paused = true)]
async fn transient_retries_keep_object_delta_timing_and_live_attempt() {
    let attempt = Arc::new(AtomicU8::new(0));
    let trace = Arc::new(Mutex::new(Vec::new()));
    let outcomes = (1..=5)
        .map(|number| Err(LlmError::Transient(format!("failure {number}"))))
        .collect();
    let start = Instant::now();
    let task = tokio::spawn(infer_with_retry(
        Box::new(RetryThread::new(outcomes, trace.clone(), attempt.clone())),
        "delta".to_owned(),
        attempt.clone(),
    ));

    tokio::task::yield_now().await;
    assert_eq!(attempt.load(Ordering::Relaxed), 1);
    let mut elapsed = 0;
    for (next, wait) in [(2, 10), (3, 20), (4, 40), (5, 80)] {
        tokio::time::advance(Duration::from_secs(wait - 1)).await;
        tokio::task::yield_now().await;
        assert_eq!(attempt.load(Ordering::Relaxed), next - 1);
        assert_eq!(start.elapsed(), Duration::from_secs(elapsed + wait - 1));
        tokio::time::advance(Duration::from_secs(1)).await;
        tokio::task::yield_now().await;
        elapsed += wait;
        assert_eq!(attempt.load(Ordering::Relaxed), next);
        assert_eq!(start.elapsed(), Duration::from_secs(elapsed));
    }

    let (_, result) = task.await.unwrap();
    assert_eq!(result, Err("failure 5".to_owned()));
    let trace = trace.lock().unwrap();
    assert_eq!(trace.len(), 5);
    let identity = trace[0].0;
    assert!(
        trace
            .iter()
            .all(|(seen, delta, _)| *seen == identity && delta == "delta")
    );
    assert_eq!(
        trace.iter().map(|(_, _, seen)| *seen).collect::<Vec<_>>(),
        vec![1, 2, 3, 4, 5]
    );
}

#[tokio::test(start_paused = true)]
async fn permanent_failure_has_no_wait_or_retry() {
    let attempt = Arc::new(AtomicU8::new(0));
    let trace = Arc::new(Mutex::new(Vec::new()));
    let outcomes = vec![
        Err(LlmError::Permanent("stop".to_owned())),
        Ok(inference("unused")),
    ];
    let start = Instant::now();
    let (_, result) = infer_with_retry(
        Box::new(RetryThread::new(outcomes, trace.clone(), attempt.clone())),
        "delta".to_owned(),
        attempt.clone(),
    )
    .await;

    assert_eq!(result, Err("stop".to_owned()));
    assert_eq!(start.elapsed(), Duration::ZERO);
    assert_eq!(attempt.load(Ordering::Relaxed), 1);
    assert_eq!(trace.lock().unwrap().len(), 1);
}

#[tokio::test]
async fn success_returns_the_object_and_output() {
    let attempt = Arc::new(AtomicU8::new(0));
    let trace = Arc::new(Mutex::new(Vec::new()));
    let outcomes = vec![Ok(inference("done")), Ok(inference("next"))];
    let (mut thread, result) = infer_with_retry(
        Box::new(RetryThread::new(outcomes, trace.clone(), attempt.clone())),
        "delta".to_owned(),
        attempt,
    )
    .await;

    assert_eq!(result, Ok(inference("done")));
    assert_eq!(thread.infer("after").await, Ok(inference("next")));
    let trace = trace.lock().unwrap();
    assert_eq!(trace.len(), 2);
    assert_eq!(trace[0].0, trace[1].0);
}

struct RecordingSink {
    submissions: Mutex<Vec<(u64, u64, SubmittedUpdate)>>,
    failure: Option<ChatError>,
}

impl RecordingSink {
    fn new(failure: Option<ChatError>) -> Self {
        Self {
            submissions: Mutex::new(Vec::new()),
            failure,
        }
    }
}

impl UpdateSink for RecordingSink {
    fn submit(&self, job: u64, identity: u64, update: SubmittedUpdate) -> Result<(), ChatError> {
        if let Some(error) = &self.failure {
            return Err(error.clone());
        }
        self.submissions
            .lock()
            .unwrap()
            .push((job, identity, update));
        Ok(())
    }
}

#[test]
fn prepared_clones_reuse_identity_and_distinct_preparations_increase() {
    let sink = Arc::new(RecordingSink::new(None));
    let updates = Updates::bind(42, sink.clone());
    let shared_updates = updates.clone();
    let activity = updates.activity("working".to_owned());
    let activity_clone = activity.clone();
    let append = shared_updates.append("answer".to_owned());

    activity.send().unwrap();
    activity_clone.send().unwrap();
    append.send().unwrap();

    assert_eq!(
        *sink.submissions.lock().unwrap(),
        vec![
            (42, 1, SubmittedUpdate::Activity("working".to_owned())),
            (42, 1, SubmittedUpdate::Activity("working".to_owned())),
            (42, 2, SubmittedUpdate::Append("answer".to_owned())),
        ]
    );
}

#[test]
fn sink_failure_passes_through_unchanged() {
    let sink = Arc::new(RecordingSink::new(Some(ChatError::Closed)));
    let updates = Updates::bind(7, sink);
    assert_eq!(
        updates.append("answer".to_owned()).send(),
        Err(ChatError::Closed)
    );
}