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)
);
}