use crate::message::{Message, ToolCall};
use crate::provider::{
ChatRequest, ChatResponse, FinishReason, Provider, ProviderError, StreamEvent, Usage,
};
use futures::stream::BoxStream;
use std::collections::VecDeque;
use std::sync::{Arc, Mutex};
const EXHAUSTED_MSG: &str = "fake provider: script exhausted";
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum FakeReply {
Text(String),
TextWithReasoning {
content: String,
reasoning: String,
},
ToolCalls {
content: String,
calls: Vec<ToolCall>,
},
WithUsage {
reply: Box<FakeReply>,
usage: Usage,
},
Error(ProviderError),
}
impl FakeReply {
pub fn text_with_usage(content: impl Into<String>, usage: Usage) -> Self {
Self::WithUsage {
reply: Box::new(Self::Text(content.into())),
usage,
}
}
}
#[derive(Debug, Default)]
pub struct FakeProvider {
replies: Mutex<VecDeque<FakeReply>>,
requests: Mutex<Vec<ChatRequest>>,
}
impl Clone for FakeProvider {
fn clone(&self) -> Self {
let replies = self
.replies
.lock()
.expect("FakeProvider internal lock poisoned")
.clone();
let requests = self
.requests
.lock()
.expect("FakeProvider internal lock poisoned")
.clone();
Self {
replies: Mutex::new(replies),
requests: Mutex::new(requests),
}
}
}
impl FakeProvider {
pub fn new(replies: impl IntoIterator<Item = FakeReply>) -> Self {
Self {
replies: Mutex::new(replies.into_iter().collect()),
requests: Mutex::new(Vec::new()),
}
}
pub fn push(&self, reply: FakeReply) {
self.replies
.lock()
.expect("FakeProvider internal lock poisoned")
.push_back(reply);
}
pub fn requests(&self) -> Vec<ChatRequest> {
self.requests
.lock()
.expect("FakeProvider internal lock poisoned")
.clone()
}
}
#[async_trait::async_trait]
impl Provider for FakeProvider {
async fn chat(&self, request: ChatRequest) -> Result<ChatResponse, ProviderError> {
self.requests
.lock()
.expect("FakeProvider internal lock poisoned")
.push(request);
match self
.replies
.lock()
.expect("FakeProvider internal lock poisoned")
.pop_front()
{
None => Err(ProviderError::Api {
status: 0,
message: EXHAUSTED_MSG.to_string(),
}),
Some(reply) => chat_response(reply),
}
}
async fn stream_chat(
&self,
request: ChatRequest,
) -> Result<BoxStream<'static, Result<StreamEvent, ProviderError>>, ProviderError> {
self.requests
.lock()
.expect("FakeProvider internal lock poisoned")
.push(request);
let (reply, usage_override) = match self
.replies
.lock()
.expect("FakeProvider internal lock poisoned")
.pop_front()
{
None => {
return Err(ProviderError::Api {
status: 0,
message: EXHAUSTED_MSG.to_string(),
});
}
Some(reply) => unwrap_usage(reply),
};
let mut events = stream_events(reply)?;
if let Some(usage) = usage_override {
if let Some(Ok(StreamEvent::Done {
usage: done_usage, ..
})) = events.last_mut()
{
*done_usage = Some(usage);
}
}
Ok(Box::pin(futures::stream::iter(events)))
}
}
fn chat_response(reply: FakeReply) -> Result<ChatResponse, ProviderError> {
match reply {
FakeReply::Error(e) => Err(e),
FakeReply::WithUsage { reply, usage } => {
let mut response = chat_response(*reply)?;
response.usage = usage;
Ok(response)
}
FakeReply::Text(content) => Ok(ChatResponse {
message: Message::assistant(content),
finish_reason: FinishReason::Stop,
usage: Usage::default(),
}),
FakeReply::TextWithReasoning { content, reasoning } => Ok(ChatResponse {
message: Message::assistant_with_reasoning(content, reasoning),
finish_reason: FinishReason::Stop,
usage: Usage::default(),
}),
FakeReply::ToolCalls { content, calls } => Ok(ChatResponse {
message: Message::Assistant {
content,
reasoning: None,
tool_calls: calls,
},
finish_reason: FinishReason::Stop,
usage: Usage::default(),
}),
}
}
fn unwrap_usage(reply: FakeReply) -> (FakeReply, Option<Usage>) {
match reply {
FakeReply::WithUsage { reply, usage } => {
let (inner, inner_usage) = unwrap_usage(*reply);
(inner, Some(usage).or(inner_usage))
}
other => (other, None),
}
}
fn stream_events(
reply: FakeReply,
) -> Result<Vec<Result<StreamEvent, ProviderError>>, ProviderError> {
match reply {
FakeReply::Error(e) => Err(e),
FakeReply::Text(content) => Ok(vec![
Ok(StreamEvent::Delta(content)),
Ok(StreamEvent::Done {
reason: FinishReason::Stop,
usage: Some(Usage::default()),
}),
]),
FakeReply::TextWithReasoning { content, reasoning } => Ok(vec![
Ok(StreamEvent::Delta(content)),
Ok(StreamEvent::Reasoning(reasoning)),
Ok(StreamEvent::Done {
reason: FinishReason::Stop,
usage: Some(Usage::default()),
}),
]),
FakeReply::ToolCalls { content, calls } => {
let mut events = Vec::new();
if !content.is_empty() {
events.push(Ok(StreamEvent::Delta(content)));
}
events.extend(calls.into_iter().map(|c| {
Ok(StreamEvent::ToolCall {
id: c.id,
name: c.name,
arguments: c.arguments,
})
}));
events.push(Ok(StreamEvent::Done {
reason: FinishReason::Stop,
usage: Some(Usage::default()),
}));
Ok(events)
}
FakeReply::WithUsage { .. } => Err(ProviderError::Api {
status: 0,
message: "internal error: WithUsage must be unwrapped before stream_events".into(),
}),
}
}
#[async_trait::async_trait]
impl Provider for Arc<FakeProvider> {
async fn chat(&self, request: ChatRequest) -> Result<ChatResponse, ProviderError> {
self.as_ref().chat(request).await
}
async fn stream_chat(
&self,
request: ChatRequest,
) -> Result<BoxStream<'static, Result<StreamEvent, ProviderError>>, ProviderError> {
self.as_ref().stream_chat(request).await
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::message::Message;
use futures::StreamExt;
#[tokio::test]
async fn chat_consumes_script_in_order() {
let fake = FakeProvider::new([FakeReply::Text("a".into()), FakeReply::Text("b".into())]);
let r1 = fake.chat(ChatRequest::default()).await.unwrap();
assert_eq!(r1.message, Message::assistant("a"));
assert_eq!(r1.finish_reason, FinishReason::Stop);
let r2 = fake.chat(ChatRequest::default()).await.unwrap();
assert_eq!(r2.message, Message::assistant("b"));
}
#[tokio::test]
async fn chat_returns_error_when_script_exhausted() {
let fake = FakeProvider::new([FakeReply::Text("only".into())]);
fake.chat(ChatRequest::default()).await.unwrap();
let err = fake.chat(ChatRequest::default()).await.unwrap_err();
assert!(matches!(err, ProviderError::Api { message: m, .. } if m.contains("exhausted")));
assert_eq!(fake.requests().len(), 2);
}
#[tokio::test]
async fn chat_with_text_and_reasoning() {
let fake = FakeProvider::new([FakeReply::TextWithReasoning {
content: "answer".into(),
reasoning: "think".into(),
}]);
let r = fake.chat(ChatRequest::default()).await.unwrap();
assert_eq!(
r.message,
Message::assistant_with_reasoning("answer", "think")
);
}
#[tokio::test]
async fn chat_with_tool_calls() {
let calls = vec![ToolCall {
id: "c1".into(),
name: "add".into(),
arguments: r#"{"a":1,"b":2}"#.into(),
}];
let fake = FakeProvider::new([FakeReply::ToolCalls {
content: String::new(),
calls: calls.clone(),
}]);
let r = fake.chat(ChatRequest::default()).await.unwrap();
match r.message {
Message::Assistant {
content,
reasoning,
tool_calls,
} => {
assert_eq!(content, "");
assert_eq!(reasoning, None);
assert_eq!(tool_calls, calls);
}
other => panic!("expected assistant, got {other:?}"),
}
}
#[tokio::test]
async fn chat_with_usage_override() {
let usage = Usage::new(7, 3);
let fake = FakeProvider::new([FakeReply::text_with_usage("hi", usage)]);
let r = fake.chat(ChatRequest::default()).await.unwrap();
assert_eq!(r.message, Message::assistant("hi"));
assert_eq!(r.usage, usage);
}
#[tokio::test]
async fn stream_with_usage_override() {
let usage = Usage::new(7, 3);
let fake = FakeProvider::new([FakeReply::text_with_usage("hi", usage)]);
let mut stream = fake.stream_chat(ChatRequest::default()).await.unwrap();
let events: Vec<StreamEvent> = stream.by_ref().map(|e| e.unwrap()).collect().await;
assert_eq!(
events.last(),
Some(&StreamEvent::Done {
reason: FinishReason::Stop,
usage: Some(usage),
})
);
}
#[tokio::test]
async fn stream_nested_usage_override_no_panic() {
let usage = Usage::new(7, 3);
let fake = FakeProvider::new([FakeReply::WithUsage {
reply: Box::new(FakeReply::WithUsage {
reply: Box::new(FakeReply::Text("hi".into())),
usage: Usage::new(1, 1),
}),
usage,
}]);
let mut stream = fake.stream_chat(ChatRequest::default()).await.unwrap();
let events: Vec<StreamEvent> = stream.by_ref().map(|e| e.unwrap()).collect().await;
assert_eq!(events[0], StreamEvent::Delta("hi".into()));
assert_eq!(
events.last(),
Some(&StreamEvent::Done {
reason: FinishReason::Stop,
usage: Some(usage),
})
);
}
#[tokio::test]
async fn error_reply_passthrough_and_script_continues() {
let fake = FakeProvider::new([
FakeReply::Error(ProviderError::Api {
status: 429,
message: "rate limited".into(),
}),
FakeReply::Text("ok".into()),
]);
let err = fake.chat(ChatRequest::default()).await.unwrap_err();
assert!(
matches!(err, ProviderError::Api { status: 429, message: m } if m == "rate limited")
);
let r = fake.chat(ChatRequest::default()).await.unwrap();
assert_eq!(r.message, Message::assistant("ok"));
}
#[tokio::test]
async fn stream_chat_text_reply() {
let fake = FakeProvider::new([FakeReply::TextWithReasoning {
content: "hi".into(),
reasoning: "think".into(),
}]);
let mut stream = fake.stream_chat(ChatRequest::default()).await.unwrap();
assert_eq!(
stream.next().await.unwrap().unwrap(),
StreamEvent::Delta("hi".into())
);
assert_eq!(
stream.next().await.unwrap().unwrap(),
StreamEvent::Reasoning("think".into())
);
assert_eq!(
stream.next().await.unwrap().unwrap(),
StreamEvent::Done {
reason: FinishReason::Stop,
usage: Some(Usage::default()),
}
);
assert!(stream.next().await.is_none());
}
#[tokio::test]
async fn stream_chat_tool_call_reply() {
let fake = FakeProvider::new([FakeReply::ToolCalls {
content: "thinking aloud".into(),
calls: vec![ToolCall {
id: "c1".into(),
name: "add".into(),
arguments: r#"{"a":1}"#.into(),
}],
}]);
let mut stream = fake.stream_chat(ChatRequest::default()).await.unwrap();
assert_eq!(
stream.next().await.unwrap().unwrap(),
StreamEvent::Delta("thinking aloud".into())
);
assert_eq!(
stream.next().await.unwrap().unwrap(),
StreamEvent::ToolCall {
id: "c1".into(),
name: "add".into(),
arguments: r#"{"a":1}"#.into(),
}
);
assert_eq!(
stream.next().await.unwrap().unwrap(),
StreamEvent::Done {
reason: FinishReason::Stop,
usage: Some(Usage::default()),
}
);
assert!(stream.next().await.is_none());
}
#[tokio::test]
async fn stream_chat_error_reply_returns_err() {
let fake = FakeProvider::new([FakeReply::Error(ProviderError::Api {
status: 400,
message: "boom".into(),
})]);
match fake.stream_chat(ChatRequest::default()).await {
Err(ProviderError::Api {
status: 400,
message: m,
}) => assert_eq!(m, "boom"),
Ok(_) => panic!("expected error"),
Err(other) => panic!("expected Api error, got {other:?}"),
}
}
#[tokio::test]
async fn stream_chat_exhausted_returns_err() {
let fake = FakeProvider::new([FakeReply::Text("only".into())]);
drop(fake.stream_chat(ChatRequest::default()).await.unwrap());
match fake.stream_chat(ChatRequest::default()).await {
Err(ProviderError::Api { message: m, .. }) => assert!(m.contains("exhausted")),
Ok(_) => panic!("expected error"),
Err(other) => panic!("expected Api error, got {other:?}"),
}
}
#[tokio::test]
async fn requests_records_messages_passed() {
let fake = FakeProvider::new([FakeReply::Text("hi".into())]);
let messages = vec![Message::system("sys"), Message::user("hello")];
fake.chat(ChatRequest {
messages: messages.clone(),
..Default::default()
})
.await
.unwrap();
let requests = fake.requests();
assert_eq!(requests.len(), 1);
assert_eq!(requests[0].messages, messages);
}
#[tokio::test]
async fn push_appends_to_script() {
let fake = FakeProvider::new([FakeReply::Text("first".into())]);
fake.push(FakeReply::Text("second".into()));
let r1 = fake.chat(ChatRequest::default()).await.unwrap();
assert_eq!(r1.message, Message::assistant("first"));
let r2 = fake.chat(ChatRequest::default()).await.unwrap();
assert_eq!(r2.message, Message::assistant("second"));
}
#[tokio::test]
async fn works_as_boxed_dyn_provider() {
let fake: Box<dyn Provider> = Box::new(FakeProvider::new([FakeReply::Text("hi".into())]));
let r = fake.chat(ChatRequest::default()).await.unwrap();
assert_eq!(r.message, Message::assistant("hi"));
}
}