use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
use async_trait::async_trait;
use parking_lot::Mutex;
use wabot_feature_chat_bot::{
ChatAdapter, ChatAdapterError, ChatAdapterRequest, ChatAdapterResponse, ChatItem, ChatMessage,
FunctionCall, LanguageModelUsage,
};
pub enum ScriptedTurn {
Items(Vec<ChatItem>),
#[allow(clippy::type_complexity)]
Responder(Box<dyn Fn(&ChatAdapterRequest) -> Vec<ChatItem> + Send + Sync>),
}
#[derive(Debug, Clone)]
pub struct RecordedRequest {
pub system_prompt: String,
pub tool_names: Vec<String>,
pub models: Vec<String>,
pub prev_items: Vec<ChatItem>,
}
impl RecordedRequest {
fn of(request: &ChatAdapterRequest) -> Self {
Self {
system_prompt: request.system_prompt.clone(),
tool_names: request.tools.iter().map(|t| t.name.clone()).collect(),
models: request.models.iter().map(|m| m.model.clone()).collect(),
prev_items: request.prev_items.clone(),
}
}
}
pub struct MockChatAdapter {
queue: Mutex<std::collections::VecDeque<ScriptedTurn>>,
requests: Mutex<Vec<RecordedRequest>>,
call_ids: AtomicUsize,
fallback_reply: Mutex<Option<String>>,
}
impl Default for MockChatAdapter {
fn default() -> Self {
Self::new()
}
}
impl MockChatAdapter {
pub fn new() -> Self {
Self {
queue: Mutex::new(Default::default()),
requests: Mutex::new(Vec::new()),
call_ids: AtomicUsize::new(0),
fallback_reply: Mutex::new(None),
}
}
pub fn arc() -> Arc<Self> {
Arc::new(Self::new())
}
pub fn with_fallback_reply(self, text: impl Into<String>) -> Self {
*self.fallback_reply.lock() = Some(text.into());
self
}
pub fn reply(&self, text: impl Into<String>) -> &Self {
self.enqueue(ScriptedTurn::Items(vec![ChatItem::bot(ChatMessage::text(
text,
))]))
}
pub fn call_tool(&self, name: impl Into<String>, arguments: impl ToArguments) -> &Self {
let id = self.call_ids.fetch_add(1, Ordering::SeqCst) + 1;
self.enqueue(ScriptedTurn::Items(vec![ChatItem::call(FunctionCall {
id: format!("mock-call-{id}"),
name: name.into(),
arguments: Some(arguments.to_arguments()),
result: None,
signature: None,
})]))
}
pub fn enqueue(&self, turn: ScriptedTurn) -> &Self {
self.queue.lock().push_back(turn);
self
}
pub fn respond_with(
&self,
responder: impl Fn(&ChatAdapterRequest) -> Vec<ChatItem> + Send + Sync + 'static,
) -> &Self {
self.enqueue(ScriptedTurn::Responder(Box::new(responder)))
}
pub fn requests(&self) -> Vec<RecordedRequest> {
self.requests.lock().clone()
}
pub fn last_request(&self) -> Option<RecordedRequest> {
self.requests.lock().last().cloned()
}
pub fn call_count(&self) -> usize {
self.requests.lock().len()
}
pub fn pending(&self) -> usize {
self.queue.lock().len()
}
}
#[async_trait]
impl ChatAdapter for MockChatAdapter {
async fn next_items(
&self,
request: ChatAdapterRequest,
) -> Result<ChatAdapterResponse, ChatAdapterError> {
self.requests.lock().push(RecordedRequest::of(&request));
let turn = self.queue.lock().pop_front();
let next_items = match turn {
Some(ScriptedTurn::Items(items)) => items,
Some(ScriptedTurn::Responder(responder)) => responder(&request),
None => match self.fallback_reply.lock().clone() {
Some(text) => vec![ChatItem::bot(ChatMessage::text(text))],
None => {
return Err(ChatAdapterError::Other(
"MockChatAdapter: no scripted turn left. Queue one with reply() / \
call_tool() / enqueue(), or set with_fallback_reply(). Remember that \
after a call_tool() turn the chat loop calls the adapter again."
.into(),
))
}
},
};
Ok(ChatAdapterResponse {
next_items,
usage: LanguageModelUsage {
input_tokens: 1,
output_tokens: 1,
provider: Some("mock".into()),
model: request.models.first().map(|m| m.model.clone()),
..Default::default()
},
})
}
}
pub trait ToArguments {
fn to_arguments(self) -> String;
}
impl ToArguments for serde_json::Value {
fn to_arguments(self) -> String {
self.to_string()
}
}
impl ToArguments for &str {
fn to_arguments(self) -> String {
self.to_string()
}
}
impl ToArguments for String {
fn to_arguments(self) -> String {
self
}
}
pub struct NoArgs;
impl ToArguments for NoArgs {
fn to_arguments(self) -> String {
"{}".into()
}
}