use std::collections::VecDeque;
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use async_trait::async_trait;
use parking_lot::Mutex;
use tokio::sync::broadcast;
use tracing::debug;
pub use crate::agent::MockAgentConfig;
use crate::backends::turn_engine::{
self, DispatchedResult, EmitCtx, EngineDeps, ResolvedCall, TurnProvider,
};
use crate::connections::{Connection, ConnectionStrategy, StepStream};
use crate::content::Content;
use crate::error::{Error, Result};
use crate::types::{Step, StepStatus, ToolResult, UsageMetadata};
const STEP_BROADCAST_CAPACITY: usize = 256;
#[derive(Debug, Clone)]
enum ScriptAction {
Text(String),
ToolCall {
name: String,
args: serde_json::Value,
},
}
#[derive(Debug, Clone, Default)]
pub struct ScriptedTurn {
actions: Vec<ScriptAction>,
usage: Option<UsageMetadata>,
}
impl ScriptedTurn {
pub fn new() -> Self {
Self::default()
}
pub fn text(mut self, text: impl Into<String>) -> Self {
self.actions.push(ScriptAction::Text(text.into()));
self
}
pub fn tool_call(mut self, name: impl Into<String>, args: serde_json::Value) -> Self {
self.actions.push(ScriptAction::ToolCall {
name: name.into(),
args,
});
self
}
pub fn with_usage(mut self, usage: UsageMetadata) -> Self {
self.usage = Some(usage);
self
}
}
#[derive(Default)]
pub struct MockConnectionBuilder {
turns: Vec<ScriptedTurn>,
conversation_id: Option<String>,
}
impl MockConnectionBuilder {
pub fn new() -> Self {
Self::default()
}
pub fn turn(mut self, f: impl FnOnce(ScriptedTurn) -> ScriptedTurn) -> Self {
self.turns.push(f(ScriptedTurn::new()));
self
}
pub fn push_turn(mut self, turn: ScriptedTurn) -> Self {
self.turns.push(turn);
self
}
pub fn turns(mut self, turns: Vec<ScriptedTurn>) -> Self {
self.turns = turns;
self
}
pub fn conversation_id(mut self, id: impl Into<String>) -> Self {
self.conversation_id = Some(id.into());
self
}
pub fn build(self) -> MockConnectionStrategy {
MockConnectionStrategy {
turns: Arc::new(self.turns),
conversation_id: self
.conversation_id
.unwrap_or_else(|| "mock-conversation".to_string()),
runners: MockRunners::default(),
}
}
}
pub type MockRunners = crate::backends::BackendRunners;
pub struct MockConnectionStrategy {
turns: Arc<Vec<ScriptedTurn>>,
conversation_id: String,
runners: MockRunners,
}
impl MockConnectionStrategy {
pub fn with_runners(mut self, runners: MockRunners) -> Self {
self.runners = runners;
self
}
}
#[cfg_attr(not(target_arch = "wasm32"), async_trait)]
#[cfg_attr(target_arch = "wasm32", async_trait(?Send))]
impl ConnectionStrategy for MockConnectionStrategy {
async fn connect(&self) -> Result<Arc<dyn Connection>> {
let (steps_tx, _) = broadcast::channel::<Step>(STEP_BROADCAST_CAPACITY);
let inner = Arc::new(MockInner {
turns: self.turns.clone(),
next_turn: AtomicUsize::new(0),
state: Arc::new(MockLoopState::new(steps_tx)),
conversation_id: self.conversation_id.clone().into(),
runners: self.runners.clone(),
});
Ok(Arc::new(MockConnection { inner }))
}
}
type MockMsg = String;
type MockLoopState = crate::backends::state::LoopState<MockMsg>;
#[derive(Clone)]
enum MockEvent {
Text(String),
Call { name: String, args: serde_json::Value },
Usage(UsageMetadata),
}
#[derive(Default)]
struct MockAccum {
calls: Vec<(String, serde_json::Value)>,
usage: Option<UsageMetadata>,
}
struct MockProvider;
impl TurnProvider for MockProvider {
type Message = MockMsg;
type Config = ();
type Request = ();
type Event = MockEvent;
type Accum = MockAccum;
fn build_request(_config: &(), _history: &[MockMsg]) {}
fn compaction_threshold(_config: &()) -> Option<u32> {
None
}
fn fold_event(
acc: &mut MockAccum,
ctx: &mut EmitCtx<'_, MockMsg>,
ev: MockEvent,
) -> Result<()> {
match ev {
MockEvent::Text(t) => ctx.push_text(&t),
MockEvent::Call { name, args } => acc.calls.push((name, args)),
MockEvent::Usage(u) => acc.usage = Some(u),
}
Ok(())
}
fn resolve_pending_calls(acc: &mut MockAccum) -> Vec<ResolvedCall> {
std::mem::take(&mut acc.calls)
.into_iter()
.map(|(name, args)| ResolvedCall {
id: None,
name,
args,
parse_error: None,
})
.collect()
}
fn round_usage(acc: &MockAccum) -> UsageMetadata {
acc.usage.clone().unwrap_or_default()
}
fn map_finish_reason(_acc: &MockAccum) -> (StepStatus, &'static str) {
(StepStatus::Done, "")
}
fn assemble_assistant_message(
_acc: MockAccum,
text: &str,
calls: &[ResolvedCall],
) -> Option<MockMsg> {
(!text.is_empty() || !calls.is_empty())
.then(|| format!("assistant:{text}:{}", calls.len()))
}
fn tool_result_messages(results: Vec<DispatchedResult>) -> Vec<MockMsg> {
results
.into_iter()
.map(|r| format!("tool:{}:{}", r.call.name, r.value))
.collect()
}
}
fn split_rounds(turn: ScriptedTurn) -> VecDeque<Vec<MockEvent>> {
let mut rounds: VecDeque<Vec<MockEvent>> = VecDeque::new();
let mut cur: Vec<MockEvent> = Vec::new();
let mut prev_was_call = false;
for action in turn.actions {
match action {
ScriptAction::Text(t) => {
if prev_was_call {
rounds.push_back(std::mem::take(&mut cur));
prev_was_call = false;
}
cur.push(MockEvent::Text(t));
}
ScriptAction::ToolCall { name, args } => {
cur.push(MockEvent::Call { name, args });
prev_was_call = true;
}
}
}
rounds.push_back(cur);
if let Some(u) = turn.usage {
if let Some(first) = rounds.front_mut() {
first.insert(0, MockEvent::Usage(u));
}
}
rounds
}
pub struct MockConnection {
inner: Arc<MockInner>,
}
struct MockInner {
turns: Arc<Vec<ScriptedTurn>>,
next_turn: AtomicUsize,
state: Arc<MockLoopState>,
conversation_id: Arc<str>,
runners: MockRunners,
}
impl MockConnection {
pub fn builder() -> MockConnectionBuilder {
MockConnectionBuilder::new()
}
}
impl MockInner {
async fn run_turn(&self, prompt: Content) {
let deps = EngineDeps::<MockProvider> {
config: (),
state: self.state.clone(),
tool_runner: self.runners.tool_runner.clone(),
hook_runner: self.runners.hook_runner.clone(),
session_ctx: self.runners.session_ctx.clone(),
};
let user = format!("user:{}", prompt.as_text().unwrap_or_default());
let rounds: Mutex<Option<VecDeque<Vec<MockEvent>>>> = Mutex::new(None);
let res = turn_engine::run_turn::<MockProvider, _, _, _, _, _>(
deps,
user,
prompt,
|_req| {
let evs = rounds
.lock()
.get_or_insert_with(|| {
let idx = self.next_turn.fetch_add(1, Ordering::Relaxed);
split_rounds(self.turns.get(idx).cloned().unwrap_or_default())
})
.pop_front()
.unwrap_or_default();
async move {
Ok(futures_util::stream::iter(
evs.into_iter().map(Ok::<_, Error>),
))
}
},
|| async {},
)
.await;
if let Err(e) = res {
debug!(error = %e, "mock turn ended with error");
}
}
}
#[cfg_attr(not(target_arch = "wasm32"), async_trait)]
#[cfg_attr(target_arch = "wasm32", async_trait(?Send))]
impl Connection for MockConnection {
fn is_idle(&self) -> bool {
self.inner.state.idle.load(Ordering::Acquire)
}
fn conversation_id(&self) -> &str {
&self.inner.conversation_id
}
async fn send(&self, content: Content) -> Result<()> {
let inner = self.inner.clone();
crate::runtime::spawn(async move {
inner.run_turn(content).await;
});
Ok(())
}
async fn send_trigger(&self, content: String) -> Result<()> {
self.send(Content::text(content)).await
}
async fn send_tool_results(&self, _results: Vec<ToolResult>) -> Result<()> {
Ok(())
}
fn subscribe_steps(&self) -> StepStream {
crate::backends::subscribe_step_stream(self.inner.state.steps.subscribe(), "mock")
}
async fn wait_for_idle(&self) -> Result<()> {
loop {
if self.is_idle() {
return Ok(());
}
self.inner.state.idle_notify.notified().await;
}
}
async fn shutdown(&self) -> Result<()> {
self.inner.state.idle.store(true, Ordering::Release);
self.inner.state.idle_notify.notify_waiters();
Ok(())
}
fn set_history_bytes(&self, _bytes: &[u8]) -> Result<()> {
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::agent::Agent;
use crate::policy;
use crate::tools::ClosureTool;
use parking_lot::Mutex;
use serde_json::json;
use std::sync::atomic::{AtomicBool, AtomicUsize};
#[tokio::test]
async fn scripted_tool_call_flow_runs_offline() {
let count = Arc::new(AtomicUsize::new(0));
let recorded: Arc<Mutex<Option<String>>> = Arc::new(Mutex::new(None));
let count_c = count.clone();
let recorded_c = recorded.clone();
let record_fact = ClosureTool::new(
"record_fact",
"Persist a fact",
json!({"type": "object", "properties": {"fact": {"type": "string"}}}),
move |args, _ctx| {
let count_c = count_c.clone();
let recorded_c = recorded_c.clone();
async move {
count_c.fetch_add(1, Ordering::SeqCst);
let fact = args["fact"].as_str().unwrap_or_default().to_string();
*recorded_c.lock() = Some(fact);
Ok(json!({"ok": true}))
}
},
);
let backend = MockConnection::builder()
.turn(|t| {
t.tool_call("record_fact", json!({"fact": "the sky is blue"}))
.text("logged")
})
.build();
let agent = Agent::start_mock(
MockAgentConfig::new(backend)
.with_tool(record_fact)
.with_policies(vec![policy::allow_all()]),
)
.await
.expect("mock agent starts");
let reply = agent
.chat("remember a fact")
.await
.expect("chat starts")
.text()
.await
.expect("turn completes");
assert_eq!(
count.load(Ordering::SeqCst),
1,
"the scripted tool must run exactly once",
);
assert_eq!(
recorded.lock().as_deref(),
Some("the sky is blue"),
"the tool received the scripted args",
);
assert_eq!(reply, "logged", "the scripted terminal text is returned");
agent.shutdown().await.expect("clean shutdown");
}
#[tokio::test]
async fn denied_tool_call_does_not_execute() {
let ran = Arc::new(AtomicBool::new(false));
let ran_c = ran.clone();
let tool = ClosureTool::new(
"danger",
"A blocked tool",
json!({"type": "object"}),
move |_args, _ctx| {
let ran_c = ran_c.clone();
async move {
ran_c.store(true, Ordering::SeqCst);
Ok(json!({"ok": true}))
}
},
);
let backend = MockConnection::builder()
.turn(|t| t.tool_call("danger", json!({})).text("attempted"))
.build();
let agent = Agent::start_mock(
MockAgentConfig::new(backend)
.with_tool(tool)
.with_policies(vec![policy::deny_all()]),
)
.await
.expect("mock agent starts");
let reply = agent.chat("go").await.unwrap().text().await.unwrap();
assert!(
!ran.load(Ordering::SeqCst),
"a denied tool must NOT execute its body",
);
assert_eq!(reply, "attempted", "the turn still completes");
agent.shutdown().await.unwrap();
}
#[tokio::test]
async fn scripted_tool_call_is_visible_on_the_stream() {
use futures_util::StreamExt;
let tool = ClosureTool::new(
"search",
"Search",
json!({"type": "object", "properties": {"q": {"type": "string"}}}),
|_args, _ctx| async move { Ok(json!({"hits": 0})) },
);
let backend = MockConnection::builder()
.turn(|t| t.tool_call("search", json!({"q": "rust"})).text("none found"))
.build();
let agent = Agent::start_mock(
MockAgentConfig::new(backend)
.with_tool(tool)
.with_policies(vec![policy::allow_all()]),
)
.await
.unwrap();
let resp = agent.chat("find rust").await.unwrap();
let mut calls = resp.tool_calls();
let first = calls
.next()
.await
.expect("a tool call is surfaced")
.expect("ok");
assert_eq!(first.name, "search");
assert_eq!(first.args, json!({"q": "rust"}));
agent.shutdown().await.unwrap();
}
#[tokio::test]
async fn turns_replay_in_order_with_usage() {
let backend = MockConnection::builder()
.turn(|t| {
t.text("first").with_usage(UsageMetadata {
total_token_count: Some(10),
..Default::default()
})
})
.turn(|t| {
t.text("second").with_usage(UsageMetadata {
total_token_count: Some(20),
..Default::default()
})
})
.build();
let agent = Agent::start_mock(MockAgentConfig::new(backend))
.await
.expect("mock agent starts");
let r1 = agent.chat("a").await.unwrap().text().await.unwrap();
assert_eq!(r1, "first");
let r2 = agent.chat("b").await.unwrap().text().await.unwrap();
assert_eq!(r2, "second");
assert_eq!(
agent.cumulative_usage().total_token_count,
Some(30),
"10 + 20, each turn counted once",
);
agent.shutdown().await.unwrap();
}
#[tokio::test]
async fn text_after_tool_call_rides_a_second_engine_round() {
let count = Arc::new(AtomicUsize::new(0));
let count_c = count.clone();
let tool = ClosureTool::new(
"ping",
"Ping",
json!({"type": "object"}),
move |_args, _ctx| {
let count_c = count_c.clone();
async move {
count_c.fetch_add(1, Ordering::SeqCst);
Ok(json!({"ok": true}))
}
},
);
let backend = MockConnection::builder()
.turn(|t| t.text("a").tool_call("ping", json!({})).text("b"))
.build();
let agent = Agent::start_mock(
MockAgentConfig::new(backend)
.with_tool(tool)
.with_policies(vec![policy::allow_all()]),
)
.await
.unwrap();
let reply = agent.chat("go").await.unwrap().text().await.unwrap();
assert_eq!(reply, "ab", "deltas from both rounds concatenate");
assert_eq!(count.load(Ordering::SeqCst), 1, "the tool ran exactly once");
agent.shutdown().await.unwrap();
}
#[tokio::test]
async fn scripted_finish_captures_summary_and_structured_output() {
use crate::types::StepType;
use futures_util::StreamExt;
let strategy = MockConnection::builder()
.turn(|t| {
t.text("working").tool_call(
crate::builtins::FINISH_TOOL_NAME,
json!({"summary": "all done", "output": {"x": 1}}),
)
})
.build();
let conn = strategy.connect().await.expect("connects");
let mut steps = conn.subscribe_steps();
conn.send(Content::text("go")).await.expect("send dispatches");
loop {
let step = steps
.next()
.await
.expect("steps flow")
.expect("no turn error");
if step.is_complete_response == Some(true) {
assert_eq!(step.kind, StepType::Finish);
assert_eq!(step.finish_summary.as_deref(), Some("all done"));
assert_eq!(step.structured_output, Some(json!({"x": 1})));
assert_eq!(step.content, "working");
break;
}
}
conn.shutdown().await.expect("clean shutdown");
}
}