use std::sync::{Arc, Mutex};
use async_trait::async_trait;
use serde_json::{Value, json};
#[cfg(test)]
use crate::model::ModelCall;
use crate::model::{
Completion, ModelError, ModelId, ModelProvider, ModelStreamEvent, ReasoningEffort, Request,
Usage,
};
#[derive(Debug, Clone, PartialEq)]
pub struct Ask {
pub model: ModelId,
pub prompt: Value,
pub max_output_tokens: u32,
pub reasoning_effort: Option<ReasoningEffort>,
pub schema: Option<Value>,
pub tools: Vec<crate::model::ToolDeclaration>,
pub exchanges: Vec<crate::model::ToolExchange>,
}
#[derive(Debug, Default)]
pub struct FakeProvider {
scripted: Mutex<std::collections::VecDeque<Result<Completion, ModelError>>>,
asked: Mutex<Vec<Ask>>,
streaming: Mutex<bool>,
}
impl FakeProvider {
#[must_use]
pub fn new() -> Arc<Self> {
Arc::new(Self::default())
}
pub fn will_answer(&self, completion: Completion) -> &Self {
self.scripted
.lock()
.expect("fake")
.push_back(Ok(completion));
self
}
pub fn will_fail(&self, error: ModelError) -> &Self {
self.scripted.lock().expect("fake").push_back(Err(error));
self
}
pub fn will_say(&self, text: impl Into<String>) -> &Self {
let text = text.into();
let usage = usage_for(&json!(&text));
self.will_answer(Completion {
tool_calls: Vec::new(),
text,
usage,
stop_reason: Some("end_turn".to_owned()),
truncated: false,
structured: None,
continuation: None,
})
}
pub fn will_call_tool(
&self,
id: impl Into<String>,
name: impl Into<String>,
arguments: serde_json::Value,
) -> &Self {
self.will_answer(Completion {
tool_calls: vec![crate::model::ToolCall {
id: id.into(),
name: name.into(),
arguments,
}],
text: String::new(),
usage: Usage::default(),
stop_reason: Some("tool_use".to_owned()),
truncated: false,
structured: None,
continuation: None,
})
}
pub fn streaming(&self) -> &Self {
*self.streaming.lock().expect("fake") = true;
self
}
#[must_use]
pub fn asked(&self) -> Vec<Ask> {
self.asked.lock().expect("fake").clone()
}
#[must_use]
pub fn calls(&self) -> usize {
self.asked.lock().expect("fake").len()
}
#[must_use]
pub fn script_exhausted(&self) -> bool {
self.scripted.lock().expect("fake").is_empty()
}
}
fn split_for_stream(text: &str) -> Vec<String> {
if text.is_empty() {
return Vec::new();
}
let mut chunks = Vec::new();
let mut current = String::new();
for ch in text.chars() {
current.push(ch);
if ch.is_whitespace() {
chunks.push(std::mem::take(&mut current));
}
}
if !current.is_empty() {
chunks.push(current);
}
chunks
}
fn usage_for(prompt: &Value) -> Usage {
let len = prompt.to_string().len() as u64;
Usage {
input_tokens: (len / 4).max(1),
output_tokens: (len / 8).max(1),
cache_write_tokens: 0,
cache_read_tokens: 0,
minor_units: 0,
}
}
fn echo(request: &Request<'_>) -> Completion {
let usage = usage_for(request.prompt);
match request.schema {
Some(schema) => {
let value = sample(schema);
Completion {
tool_calls: Vec::new(),
text: value.to_string(),
usage,
stop_reason: Some("end_turn".to_owned()),
truncated: false,
structured: Some(value),
continuation: None,
}
}
None => Completion {
tool_calls: Vec::new(),
text: format!("fake answer to {}", request.prompt),
usage,
stop_reason: Some("end_turn".to_owned()),
truncated: false,
structured: None,
continuation: None,
},
}
}
fn sample(schema: &Value) -> Value {
match schema.get("type").and_then(Value::as_str) {
Some("object") => {
let mut out = serde_json::Map::new();
if let Some(props) = schema.get("properties").and_then(Value::as_object) {
for (name, sub) in props {
out.insert(name.clone(), sample(sub));
}
}
Value::Object(out)
}
Some("array") => match schema.get("items") {
Some(items) => json!([sample(items)]),
None => json!([]),
},
Some("string") => json!("fake"),
Some("number" | "integer") => json!(0),
Some("boolean") => json!(false),
_ => Value::Null,
}
}
#[async_trait]
impl ModelProvider for FakeProvider {
async fn complete(&self, request: Request<'_>) -> Result<Completion, ModelError> {
self.asked.lock().expect("fake").push(Ask {
model: request.model.clone(),
prompt: request.prompt.clone(),
max_output_tokens: request.max_output_tokens,
reasoning_effort: request.reasoning_effort,
schema: request.schema.cloned(),
tools: request.tools.to_vec(),
exchanges: request.exchanges.to_vec(),
});
let scripted = self.scripted.lock().expect("fake").pop_front();
let answer = scripted.unwrap_or_else(|| Ok(echo(&request)));
if *self.streaming.lock().expect("fake")
&& let (Ok(completion), Some((observer, label))) = (&answer, request.stream)
{
for delta in split_for_stream(&completion.text) {
observer.event(crate::core::Tainted::with_label(
ModelStreamEvent::TextDelta(delta),
label.clone(),
));
}
observer.event(crate::core::Tainted::with_label(
ModelStreamEvent::Usage(completion.usage),
label.clone(),
));
}
answer
}
}
#[cfg(test)]
mod tests {
use super::*;
fn model() -> ModelId {
ModelId::new("fake", "m")
}
fn ask(schema: Option<&Value>) -> Request<'static> {
let prompt: &'static Value = Box::leak(Box::new(json!({"q": "what is the balance"})));
let model: &'static ModelId = Box::leak(Box::new(model()));
Request {
model,
prompt,
max_output_tokens: ModelCall::DEFAULT_MAX_OUTPUT_TOKENS,
reasoning_effort: None,
schema: schema.map(|s| &*Box::leak(Box::new(s.clone()))),
tools: &[],
exchanges: &[],
continuation: None,
stream: None,
}
}
#[tokio::test]
async fn the_same_question_gets_the_same_answer() {
let p = FakeProvider::new();
let a = p.complete(ask(None)).await.unwrap();
let b = p.complete(ask(None)).await.unwrap();
assert_eq!(a.text, b.text);
assert_eq!(a.usage, b.usage);
}
#[tokio::test]
async fn an_answer_is_never_free() {
let p = FakeProvider::new();
let c = p.complete(ask(None)).await.unwrap();
assert!(
c.usage.spend().tokens > 0,
"a fake reporting zero usage makes every ceiling test vacuous"
);
}
#[tokio::test]
async fn usage_grows_with_the_prompt() {
let p = FakeProvider::new();
let short: &'static Value = Box::leak(Box::new(json!("hi")));
let long: &'static Value = Box::leak(Box::new(json!("hi".repeat(500))));
let m = model();
let a = p
.complete(Request {
model: &m,
prompt: short,
max_output_tokens: ModelCall::DEFAULT_MAX_OUTPUT_TOKENS,
reasoning_effort: None,
schema: None,
tools: &[],
exchanges: &[],
continuation: None,
stream: None,
})
.await
.unwrap();
let b = p
.complete(Request {
model: &m,
prompt: long,
max_output_tokens: ModelCall::DEFAULT_MAX_OUTPUT_TOKENS,
reasoning_effort: None,
schema: None,
tools: &[],
exchanges: &[],
continuation: None,
stream: None,
})
.await
.unwrap();
assert!(b.usage.spend().tokens > a.usage.spend().tokens);
}
#[tokio::test]
async fn scripted_answers_come_back_in_order_then_the_default_takes_over() {
let p = FakeProvider::new();
p.will_say("first").will_say("second");
assert_eq!(p.complete(ask(None)).await.unwrap().text, "first");
assert_eq!(p.complete(ask(None)).await.unwrap().text, "second");
assert!(p.script_exhausted());
assert!(
p.complete(ask(None)).await.unwrap().text.contains("fake"),
"past the script, the default echo answers"
);
assert_eq!(p.calls(), 3);
}
#[tokio::test]
async fn a_scripted_failure_can_carry_usage() {
let p = FakeProvider::new();
p.will_fail(ModelError::Interrupted {
model: model(),
usage: Usage {
input_tokens: 100,
output_tokens: 300,
..Usage::default()
},
detail: "reset".to_owned(),
});
let e = p.complete(ask(None)).await.expect_err("scripted");
assert_eq!(e.usage().spend().tokens, 400);
}
#[tokio::test]
async fn a_schema_gets_json_shaped_like_it() {
let schema = json!({
"type": "object",
"properties": {
"verdict": {"type": "string"},
"score": {"type": "integer"},
"flags": {"type": "array", "items": {"type": "boolean"}},
},
});
let p = FakeProvider::new();
let c = p.complete(ask(Some(&schema))).await.unwrap();
let v = c.structured.expect("a schema was asked for");
assert_eq!(v["verdict"], json!("fake"));
assert_eq!(v["score"], json!(0));
assert_eq!(v["flags"], json!([false]));
assert_eq!(
c.text,
v.to_string(),
"`text` holds the raw string even when a schema was parsed"
);
}
#[tokio::test]
async fn no_schema_means_no_structured_value() {
let p = FakeProvider::new();
assert!(p.complete(ask(None)).await.unwrap().structured.is_none());
}
#[tokio::test]
async fn it_records_what_it_was_asked() {
let schema = json!({"type": "string"});
let p = FakeProvider::new();
p.complete(ask(None)).await.unwrap();
p.complete(ask(Some(&schema))).await.unwrap();
let asked = p.asked();
assert_eq!(asked.len(), 2);
assert_eq!(asked[0].model, model());
assert_eq!(asked[0].schema, None);
assert_eq!(asked[1].schema, Some(schema));
}
}