use std::collections::VecDeque;
use std::sync::Mutex;
use anyhow::Result;
use tokio_util::sync::CancellationToken;
use super::contract::{
ChatChunk, ChatRequest, ChatStream, EmbedRole, Embedder, EngineBackend, FinishReason,
};
pub struct MockBackend {
scripts: Mutex<VecDeque<Vec<ChatChunk>>>,
wait_for_cancel: bool,
cycle: bool,
delay_ms: u64,
}
impl MockBackend {
#[cfg(test)]
pub fn scripted(script: Vec<ChatChunk>) -> Self {
Self {
scripts: Mutex::new(VecDeque::from([script])),
wait_for_cancel: false,
cycle: false,
delay_ms: 0,
}
}
#[cfg(test)]
pub fn sequence(scripts: Vec<Vec<ChatChunk>>) -> Self {
Self {
scripts: Mutex::new(VecDeque::from(scripts)),
wait_for_cancel: false,
cycle: false,
delay_ms: 0,
}
}
#[cfg(test)]
pub fn cancellable(prefix: Vec<ChatChunk>) -> Self {
Self {
scripts: Mutex::new(VecDeque::from([prefix])),
wait_for_cancel: true,
cycle: false,
delay_ms: 0,
}
}
pub fn cycling(scripts: Vec<Vec<ChatChunk>>, delay_ms: u64) -> Self {
Self {
scripts: Mutex::new(VecDeque::from(scripts)),
wait_for_cancel: false,
cycle: true,
delay_ms,
}
}
}
#[async_trait::async_trait]
impl EngineBackend for MockBackend {
async fn chat_stream(
&self,
_req: ChatRequest,
cancel: CancellationToken,
) -> Result<ChatStream> {
let script = {
let mut q = self.scripts.lock().unwrap();
if self.cycle {
let s = q.pop_front().unwrap_or_default();
q.push_back(s.clone());
s
} else if q.len() > 1 {
q.pop_front().unwrap()
} else {
q.front().cloned().unwrap_or_default()
}
};
let wait = self.wait_for_cancel;
let delay = self.delay_ms;
let s = async_stream::stream! {
for chunk in script {
if cancel.is_cancelled() {
yield ChatChunk::Finished(FinishReason::Cancelled);
return;
}
if delay > 0 {
tokio::time::sleep(std::time::Duration::from_millis(delay)).await;
}
yield chunk;
}
if wait {
cancel.cancelled().await;
yield ChatChunk::Finished(FinishReason::Cancelled);
}
};
Ok(Box::pin(s))
}
}
pub struct MockEmbedder {
dim: usize,
}
impl MockEmbedder {
pub fn new(dim: usize) -> Self {
Self { dim }
}
}
#[async_trait::async_trait]
impl Embedder for MockEmbedder {
async fn embed(&self, texts: Vec<String>, _role: EmbedRole) -> Result<Vec<Vec<f32>>> {
Ok(texts.iter().map(|t| embed_text(t, self.dim)).collect())
}
}
fn embed_text(text: &str, dim: usize) -> Vec<f32> {
let mut v = vec![0.0_f32; dim];
for ch in text.to_lowercase().chars().filter(|c| c.is_alphanumeric()) {
v[(ch as usize) % dim] += 1.0;
}
let norm = v.iter().map(|x| x * x).sum::<f32>().sqrt();
if norm > 0.0 {
for x in &mut v {
*x /= norm;
}
}
v
}