use std::collections::VecDeque;
use std::sync::Mutex;
use anyhow::Result;
use async_trait::async_trait;
use super::{LlmResponse, LlmService, Message, ToolCall, ToolSchema};
pub struct MockLlm {
response: String,
sql: Option<String>,
}
impl MockLlm {
pub fn new(response: impl Into<String>) -> Self {
Self { response: response.into(), sql: None }
}
pub fn with_default_sql() -> Self {
let sql = "SELECT p.category, ROUND(SUM(oi.quantity * oi.unit_price), 2) AS revenue \
FROM order_items oi \
JOIN products p ON p.id = oi.product_id \
JOIN orders o ON o.id = oi.order_id \
WHERE o.status NOT IN ('cancelled','refunded') \
GROUP BY p.category ORDER BY revenue DESC";
Self {
response: format!("```sql\n{sql};\n```"),
sql: Some(sql.to_string()),
}
}
}
#[async_trait]
impl LlmService for MockLlm {
async fn submit_prompt(&self, _messages: Vec<Message>) -> Result<String> {
Ok(self.response.clone())
}
async fn chat(&self, messages: Vec<Message>, tools: &[ToolSchema]) -> Result<LlmResponse> {
let run_sql_available = tools.iter().any(|t| t.name == "run_sql");
let viz_available = tools.iter().any(|t| t.name == "visualize_data");
let has_tool_result = messages.iter().any(|m| m.role == "tool");
let viz_attempted = messages.iter().any(|m| m.role == "tool" && m.content.contains("chart"));
if run_sql_available && !has_tool_result {
if let Some(sql) = &self.sql {
return Ok(LlmResponse {
text: None,
tool_calls: vec![ToolCall {
id: "mock-run-sql".into(),
name: "run_sql".into(),
args: serde_json::json!({ "sql": sql }),
}],
});
}
}
if viz_available && has_tool_result && !viz_attempted {
if let Some(filename) = find_results_filename(&messages) {
return Ok(LlmResponse {
text: None,
tool_calls: vec![ToolCall {
id: "mock-visualize".into(),
name: "visualize_data".into(),
args: serde_json::json!({ "filename": filename, "title": "Results" }),
}],
});
}
}
if has_tool_result {
return Ok(LlmResponse::text("Here are the results from your database."));
}
Ok(LlmResponse::text(self.response.clone()))
}
}
fn find_results_filename(messages: &[Message]) -> Option<String> {
for m in messages.iter().rev() {
if m.role == "tool" {
for tok in m.content.split_whitespace() {
let t = tok.trim_matches(|c| c == '*' || c == '.' || c == ':');
if t.starts_with("query_results_") && t.ends_with(".json") {
return Some(t.to_string());
}
}
}
}
None
}
pub struct ScriptedMockLlm {
responses: Mutex<VecDeque<String>>,
last: Mutex<String>,
}
impl ScriptedMockLlm {
pub fn new(responses: Vec<String>) -> Self {
let last = responses.last().cloned().unwrap_or_default();
Self {
responses: Mutex::new(responses.into()),
last: Mutex::new(last),
}
}
}
#[async_trait]
impl LlmService for ScriptedMockLlm {
async fn submit_prompt(&self, _messages: Vec<Message>) -> Result<String> {
let mut queue = self.responses.lock().unwrap();
match queue.pop_front() {
Some(r) => {
*self.last.lock().unwrap() = r.clone();
Ok(r)
}
None => Ok(self.last.lock().unwrap().clone()),
}
}
}
pub struct ScriptedToolLlm {
responses: Mutex<VecDeque<LlmResponse>>,
last: Mutex<LlmResponse>,
}
impl ScriptedToolLlm {
pub fn new(responses: Vec<LlmResponse>) -> Self {
let last = responses.last().cloned().unwrap_or_default();
Self {
responses: Mutex::new(responses.into()),
last: Mutex::new(last),
}
}
fn next(&self) -> LlmResponse {
let mut queue = self.responses.lock().unwrap();
match queue.pop_front() {
Some(r) => {
*self.last.lock().unwrap() = r.clone();
r
}
None => self.last.lock().unwrap().clone(),
}
}
}
#[async_trait]
impl LlmService for ScriptedToolLlm {
async fn submit_prompt(&self, _messages: Vec<Message>) -> Result<String> {
Ok(self.next().text.unwrap_or_default())
}
async fn chat(&self, _messages: Vec<Message>, _tools: &[ToolSchema]) -> Result<LlmResponse> {
Ok(self.next())
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::llm::{ToolCall, ToolSchema};
use serde_json::json;
#[tokio::test]
async fn scripted_tool_llm_returns_tool_call_then_text() {
let llm = ScriptedToolLlm::new(vec![
LlmResponse {
text: None,
tool_calls: vec![ToolCall {
id: "call-1".into(),
name: "calculator".into(),
args: json!({ "a": 2, "b": 2 }),
}],
},
LlmResponse::text("The answer is 4."),
]);
let tools = vec![ToolSchema {
name: "calculator".into(),
description: "adds two numbers".into(),
parameters: json!({
"type": "object",
"properties": { "a": {"type":"number"}, "b": {"type":"number"} },
"required": ["a", "b"]
}),
}];
let r1 = llm.chat(vec![Message::user("what is 2+2?")], &tools).await.unwrap();
assert!(r1.is_tool_call());
assert_eq!(r1.tool_calls.len(), 1);
assert_eq!(r1.tool_calls[0].name, "calculator");
assert_eq!(r1.tool_calls[0].args["a"], 2);
assert_eq!(r1.tool_calls[0].args["b"], 2);
let r2 = llm.chat(vec![], &tools).await.unwrap();
assert!(!r2.is_tool_call());
assert_eq!(r2.text.as_deref(), Some("The answer is 4."));
}
}