use std::sync::{
Arc, Mutex,
atomic::{AtomicU64, Ordering},
};
use futures::StreamExt;
use tracing::{Instrument, info_span, span::Id};
use super::{
completion::{Agent, PreparedCompletionRequest},
hook::{
AgentHook, CompletionCall, CompletionCallAction,
CompletionResponse as CompletionResponseEvent, HookContext, HookStack,
InvalidToolCallAction, ModelTurnAction, ModelTurnFinished, ObservationAction, RequestPatch,
ToolCall as ToolCallEvent, ToolCallAction, ToolResultAction, ToolResultEvent,
},
prompt_request::{
PromptResponse,
streaming::{
DriveItem, DriveStream, MultiTurnStreamItem, StreamingError, TurnSource, drive_agent,
drive_tool_calls, record_usage_on_span, streaming_error_into_prompt,
},
tool_result_output,
},
run::{
AgentRun, DEFAULT_OUTPUT_RETRIES, ModelTurn, ModelTurnOutcome, OutputMode, PendingToolCall,
},
};
use rig_core::{
memory::ConversationMemory,
message::{ToolCall, ToolChoice, UserContent},
};
use crate::{
completion::{CompletionError, CompletionModel, Document, Message, PromptError, Usage},
json_utils,
tool::{
ToolContext, ToolDispatch, ToolOutput, ToolResult,
server::{ToolRegistrySnapshot, ToolServerHandle},
},
};
use super::UNKNOWN_AGENT_NAME;
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
pub(crate) enum UnhandledInvalidToolCallPolicy {
#[default]
Fail,
IgnoreForExtractor,
}
macro_rules! build_chat_span {
($runner:expr, $effective_preamble:expr, $name:literal, $operation:literal) => {{
let system_instructions = $crate::core::telemetry::system_instructions_json(
$effective_preamble,
$runner.record_telemetry_content,
);
$crate::core::telemetry::completion_parent_span!(
target: "rig::agent_chat",
name: $name,
operation: $operation,
system_instructions: system_instructions.as_deref(),
gen_ai.agent.name = $runner.agent_name_or_default(),
)
}};
}
pub(crate) use build_chat_span;
pub(crate) fn observe_action(action: ObservationAction) -> Option<String> {
match action {
ObservationAction::Continue => None,
ObservationAction::Stop(reason) => Some(reason),
}
}
pub(crate) enum ModelTurnDecision {
Advance,
Retried,
Terminate(String),
}
pub(crate) fn resolve_model_turn_action(
run: &mut AgentRun,
action: ModelTurnAction,
) -> Result<ModelTurnDecision, PromptError> {
match action {
ModelTurnAction::Continue => Ok(ModelTurnDecision::Advance),
ModelTurnAction::Retry(request) => {
run.retry_model_turn(request)?;
Ok(ModelTurnDecision::Retried)
}
ModelTurnAction::Stop(reason) => Ok(ModelTurnDecision::Terminate(reason)),
}
}
pub(crate) enum ToolCallDecision {
Proceed,
ProceedWith(serde_json::Value),
Skip(String),
Terminate(String),
}
pub(crate) fn tool_call_decision(action: ToolCallAction) -> ToolCallDecision {
match action {
ToolCallAction::Run => ToolCallDecision::Proceed,
ToolCallAction::Rewrite(args) => ToolCallDecision::ProceedWith(args),
ToolCallAction::Skip(reason) => ToolCallDecision::Skip(reason),
ToolCallAction::Stop(reason) => ToolCallDecision::Terminate(reason),
}
}
pub(crate) enum ToolResultDecision {
Keep,
Replace(ToolOutput),
Terminate(String),
}
pub(crate) fn tool_result_decision(action: ToolResultAction) -> ToolResultDecision {
match action {
ToolResultAction::Keep => ToolResultDecision::Keep,
ToolResultAction::Rewrite(result) => ToolResultDecision::Replace(result),
ToolResultAction::Stop(reason) => ToolResultDecision::Terminate(reason),
}
}
pub(crate) enum CompletionCallDecision {
Proceed,
Patch(RequestPatch),
Terminate(String),
}
pub(crate) fn completion_call_decision(action: CompletionCallAction) -> CompletionCallDecision {
match action {
CompletionCallAction::Continue => CompletionCallDecision::Proceed,
CompletionCallAction::Patch(patch) => CompletionCallDecision::Patch(patch),
CompletionCallAction::Stop(reason) => CompletionCallDecision::Terminate(reason),
}
}
#[non_exhaustive]
pub struct AgentRunner<M>
where
M: CompletionModel,
{
pub(crate) prompt: Message,
pub(crate) chat_history: Option<Vec<Message>>,
pub(crate) max_turns: usize,
pub(crate) max_invalid_tool_call_retries: usize,
pub(crate) model: Arc<M>,
pub(crate) agent_name: Option<String>,
pub(crate) preamble: Option<String>,
pub(crate) static_context: Vec<Document>,
pub(crate) temperature: Option<f64>,
pub(crate) max_tokens: Option<u64>,
pub(crate) additional_params: Option<serde_json::Value>,
pub(crate) record_telemetry_content: bool,
pub(crate) tool_server_handle: ToolServerHandle,
pub(crate) tool_context: ToolContext,
pub(crate) tool_choice: Option<ToolChoice>,
pub(crate) output_schema: Option<schemars::Schema>,
pub(crate) output_mode: OutputMode,
pub(crate) output_tool_name: Option<String>,
pub(crate) output_tool_description: Option<String>,
pub(crate) augment_output_preamble: bool,
pub(crate) unhandled_invalid_tool_call_policy: UnhandledInvalidToolCallPolicy,
pub(crate) concurrency: usize,
pub(crate) memory: Option<Arc<dyn ConversationMemory>>,
pub(crate) conversation_id: Option<String>,
pub(crate) hooks: HookStack,
pub(crate) error_usage: Option<Arc<Mutex<Usage>>>,
}
impl<M> AgentRunner<M>
where
M: CompletionModel,
{
pub fn from_agent(agent: &Agent<M>, prompt: impl Into<Message>) -> Self {
Self {
prompt: prompt.into(),
chat_history: None,
max_turns: agent.default_max_turns.unwrap_or(1),
max_invalid_tool_call_retries: 0,
model: agent.model.clone(),
agent_name: agent.name.clone(),
preamble: agent.preamble.clone(),
static_context: agent.static_context.clone(),
temperature: agent.temperature,
max_tokens: agent.max_tokens,
additional_params: agent.additional_params.clone(),
record_telemetry_content: agent.record_telemetry_content,
tool_server_handle: agent.tool_server_handle.clone(),
tool_context: ToolContext::new(),
tool_choice: agent.tool_choice.clone(),
output_schema: agent.output_schema.clone(),
output_mode: agent.output_mode.clone(),
output_tool_name: None,
output_tool_description: None,
augment_output_preamble: true,
unhandled_invalid_tool_call_policy: UnhandledInvalidToolCallPolicy::Fail,
concurrency: 1,
memory: agent.memory.clone(),
conversation_id: agent.default_conversation_id.clone(),
hooks: agent.hooks.clone(),
error_usage: None,
}
}
pub fn add_hook<H>(mut self, hook: H) -> Self
where
H: AgentHook + 'static,
{
self.hooks.push(hook);
self
}
}
impl<M> AgentRunner<M>
where
M: CompletionModel,
{
pub fn max_turns(mut self, max_turns: usize) -> Self {
self.max_turns = max_turns;
self
}
pub fn tool_context(mut self, context: ToolContext) -> Self {
self.tool_context = context;
self
}
pub fn history<I, T>(mut self, history: I) -> Self
where
I: IntoIterator<Item = T>,
T: Into<Message>,
{
self.chat_history = Some(history.into_iter().map(Into::into).collect());
self
}
pub fn preamble(mut self, preamble: impl Into<String>) -> Self {
self.preamble = Some(preamble.into());
self
}
pub fn without_preamble(mut self) -> Self {
self.preamble = None;
self
}
pub fn document(mut self, document: Document) -> Self {
self.static_context.push(document);
self
}
pub fn documents(mut self, documents: impl IntoIterator<Item = Document>) -> Self {
self.static_context.extend(documents);
self
}
pub fn temperature(mut self, temperature: f64) -> Self {
self.temperature = Some(temperature);
self
}
pub fn without_temperature(mut self) -> Self {
self.temperature = None;
self
}
pub fn max_tokens(mut self, max_tokens: u64) -> Self {
self.max_tokens = Some(max_tokens);
self
}
pub fn without_max_tokens(mut self) -> Self {
self.max_tokens = None;
self
}
pub fn merge_additional_params(
mut self,
params: serde_json::Map<String, serde_json::Value>,
) -> Self {
let params = serde_json::Value::Object(params);
self.additional_params = Some(match self.additional_params.take() {
Some(baseline) if baseline.is_object() => crate::json_utils::merge(baseline, params),
_ => params,
});
self
}
pub fn replace_additional_params(mut self, params: serde_json::Value) -> Self {
self.additional_params = Some(params);
self
}
pub fn without_additional_params(mut self) -> Self {
self.additional_params = None;
self
}
pub fn tool_choice(mut self, tool_choice: ToolChoice) -> Self {
self.tool_choice = Some(tool_choice);
self
}
pub fn without_tool_choice(mut self) -> Self {
self.tool_choice = None;
self
}
pub(crate) fn output_tool(
mut self,
name: impl Into<String>,
description: impl Into<String>,
augment_preamble: bool,
) -> Self {
self.output_tool_name = Some(name.into());
self.output_tool_description = Some(description.into());
self.augment_output_preamble = augment_preamble;
self
}
pub(crate) fn ignore_unhandled_invalid_tool_calls(mut self) -> Self {
self.unhandled_invalid_tool_call_policy =
UnhandledInvalidToolCallPolicy::IgnoreForExtractor;
self
}
pub fn record_content_telemetry(mut self, enabled: bool) -> Self {
self.record_telemetry_content = enabled;
self
}
pub fn tool_concurrency(mut self, concurrency: usize) -> Self {
self.concurrency = concurrency.max(1);
self
}
pub fn conversation(mut self, id: impl Into<String>) -> Self {
self.conversation_id = Some(id.into());
self
}
pub fn without_memory(mut self) -> Self {
self.memory = None;
self.conversation_id = None;
self
}
pub fn max_invalid_tool_call_retries(mut self, retries: usize) -> Self {
self.max_invalid_tool_call_retries = retries;
self
}
pub(crate) fn agent_name_or_default(&self) -> &str {
self.agent_name.as_deref().unwrap_or(UNKNOWN_AGENT_NAME)
}
pub(crate) fn build_run(&self, history_override: Option<Vec<Message>>) -> AgentRun {
let run = build_agent_run(
self.prompt.clone(),
self.max_turns,
self.max_invalid_tool_call_retries,
self.output_schema.as_ref(),
history_override.or_else(|| self.chat_history.clone()),
self.tool_choice.clone(),
);
match &self.output_tool_name {
Some(name) => run.with_output_tool_name(name.clone()),
None => run,
}
}
}
pub(crate) fn build_agent_run(
prompt: Message,
max_turns: usize,
max_invalid_tool_call_retries: usize,
output_schema: Option<&schemars::Schema>,
history: Option<Vec<Message>>,
tool_choice: Option<ToolChoice>,
) -> AgentRun {
let mut run = AgentRun::new(prompt)
.max_turns(max_turns)
.max_invalid_tool_call_retries(max_invalid_tool_call_retries)
.with_output_validation(
output_schema.map(|schema| schema.as_value().clone()),
DEFAULT_OUTPUT_RETRIES,
);
if let Some(history) = history {
run = run.with_history(history);
}
if let Some(tool_choice) = tool_choice {
run = run.with_tool_choice(tool_choice);
}
run
}
pub(crate) fn acquire_agent_span(
agent_name: &str,
preamble: Option<&str>,
record_content: bool,
) -> (tracing::Span, bool) {
if tracing::Span::current().is_disabled() {
let system_instructions =
rig_core::telemetry::system_instructions_json(preamble, record_content);
let span = info_span!(
"invoke_agent",
gen_ai.operation.name = "invoke_agent",
gen_ai.agent.name = agent_name,
gen_ai.system_instructions = system_instructions.as_deref(),
gen_ai.prompt = tracing::field::Empty,
gen_ai.completion = tracing::field::Empty,
gen_ai.usage.input_tokens = tracing::field::Empty,
gen_ai.usage.output_tokens = tracing::field::Empty,
gen_ai.usage.cache_read.input_tokens = tracing::field::Empty,
gen_ai.usage.cache_creation.input_tokens = tracing::field::Empty,
gen_ai.usage.tool_use_prompt_tokens = tracing::field::Empty,
gen_ai.usage.reasoning_tokens = tracing::field::Empty,
);
(span, true)
} else {
(tracing::Span::current(), false)
}
}
pub(crate) enum CompletionCallOutcome {
Proceed(Option<RequestPatch>),
Terminate(String),
}
pub(crate) async fn resolve_completion_call(
hooks: &HookStack,
ctx: &HookContext,
prompt: &Message,
history: &[Message],
turn: usize,
) -> CompletionCallOutcome {
match completion_call_decision(
hooks
.on_completion_call(
ctx,
CompletionCall {
prompt,
history,
turn,
},
)
.await,
) {
CompletionCallDecision::Terminate(reason) => CompletionCallOutcome::Terminate(reason),
CompletionCallDecision::Patch(patch) => CompletionCallOutcome::Proceed(Some(patch)),
CompletionCallDecision::Proceed => CompletionCallOutcome::Proceed(None),
}
}
pub(crate) async fn append_run_messages(
memory_handle: Option<&(Arc<dyn ConversationMemory>, String)>,
messages: &[Message],
) {
if let Some((memory, id)) = memory_handle
&& let Err(err) = memory.append(id, messages.to_vec()).await
{
tracing::warn!(
error = %err,
conversation_id = %id,
"conversation memory append failed; surfacing final response anyway"
);
}
}
pub(crate) enum ToolExecution {
Executed(Box<ToolCall>),
Skipped,
}
pub(crate) struct ToolCallOutcome {
pub content: UserContent,
pub execution: ToolExecution,
}
pub(crate) async fn run_single_tool<M>(
runner: &AgentRunner<M>,
ctx: &HookContext,
tool_snapshot: &ToolRegistrySnapshot,
tool_call: &ToolCall,
internal_call_id: &str,
error_history: &[Message],
) -> Result<ToolCallOutcome, PromptError>
where
M: CompletionModel,
{
let hooks = &runner.hooks;
let tool_context = &runner.tool_context;
let record_content = runner.record_telemetry_content;
let tool_name = &tool_call.function.name;
let mut args = json_utils::serialize_json_value(&tool_call.function.arguments);
let tool_span = tracing::Span::current();
tool_span.record("gen_ai.tool.name", tool_name);
tool_span.record("gen_ai.tool.call.id", &tool_call.id);
if record_content {
tool_span.record("gen_ai.tool.call.arguments", &args);
}
let (action, salvaged_rewrite) = hooks
.resolve_tool_call(
ctx,
ToolCallEvent {
tool_name,
tool_call_id: tool_call.call_id.as_deref(),
internal_call_id,
args: &args,
},
)
.await;
if let Some(rewritten) = salvaged_rewrite.as_ref() {
args = json_utils::serialize_json_value(rewritten);
if record_content {
tool_span.record("gen_ai.tool.call.arguments", &args);
}
tracing::debug!(
tool_name = tool_name,
"tool-call arguments rewritten by a hook"
);
}
let mut skipped: Option<ToolResult> = None;
let effective_args: serde_json::Value = match tool_call_decision(action) {
ToolCallDecision::Terminate(reason) => {
return Err(PromptError::prompt_cancelled(
error_history.to_vec(),
reason,
));
}
ToolCallDecision::Skip(reason) => {
tracing::info!(tool_name = tool_name, reason = reason, "Tool call rejected");
skipped = Some(ToolResult::skipped(reason));
salvaged_rewrite.unwrap_or_else(|| tool_call.function.arguments.clone())
}
ToolCallDecision::ProceedWith(replacement) => {
args = json_utils::serialize_json_value(&replacement);
if record_content {
tool_span.record("gen_ai.tool.call.arguments", &args);
}
tracing::debug!(
tool_name = tool_name,
"tool-call arguments rewritten by a hook"
);
replacement
}
ToolCallDecision::Proceed => tool_call.function.arguments.clone(),
};
let (exec, execution, dispatch_context) = match skipped {
Some(exec) => (exec, ToolExecution::Skipped, tool_context.for_dispatch()),
None => {
let mut effective_tool_call = tool_call.clone();
effective_tool_call.function.arguments = effective_args;
let ToolDispatch {
result: exec,
context: dispatch_context,
} = tool_snapshot.dispatch(tool_name, &args, tool_context).await;
(
exec,
ToolExecution::Executed(Box::new(effective_tool_call)),
dispatch_context,
)
}
};
let result_decision = tool_result_decision(
hooks
.on_tool_result(
ctx,
ToolResultEvent {
tool_name,
tool_call_id: tool_call.call_id.as_deref(),
internal_call_id,
args: &args,
presentation: exec.output(),
raw_result: &exec,
tool_context: &dispatch_context,
},
)
.await,
);
record_tool_result(&tool_span, &exec);
match result_decision {
ToolResultDecision::Terminate(reason) => Err(PromptError::prompt_cancelled(
error_history.to_vec(),
reason,
)),
ToolResultDecision::Replace(replacement) => {
if record_content {
tool_span.record("gen_ai.tool.call.result", replacement.render());
}
Ok(ToolCallOutcome {
content: tool_result_output(
tool_call.id.clone(),
tool_call.call_id.clone(),
replacement,
),
execution,
})
}
ToolResultDecision::Keep => {
if record_content {
tool_span.record("gen_ai.tool.call.result", exec.output().render());
}
let content = tool_result_output(
tool_call.id.clone(),
tool_call.call_id.clone(),
exec.output().clone(),
);
Ok(ToolCallOutcome { content, execution })
}
}
}
fn record_tool_result(span: &tracing::Span, result: &ToolResult) {
span.record("gen_ai.tool.call.outcome", result.status_name());
if let Some(error) = result.error() {
span.record("gen_ai.tool.error.type", error.kind().as_str());
}
}
pub(crate) fn new_execute_tool_span() -> tracing::Span {
info_span!(
"execute_tool",
gen_ai.operation.name = "execute_tool",
gen_ai.tool.type = "function",
gen_ai.tool.name = tracing::field::Empty,
gen_ai.tool.call.id = tracing::field::Empty,
gen_ai.tool.call.arguments = tracing::field::Empty,
gen_ai.tool.call.result = tracing::field::Empty,
gen_ai.tool.call.outcome = tracing::field::Empty,
gen_ai.tool.error.type = tracing::field::Empty
)
}
pub(crate) struct UnaryTurnSource {
current_span_id: AtomicU64,
record_telemetry_content: bool,
}
impl UnaryTurnSource {
pub(crate) fn new(record_telemetry_content: bool) -> Self {
Self {
current_span_id: AtomicU64::new(0),
record_telemetry_content,
}
}
fn chain_span(&self, span: tracing::Span) -> tracing::Span {
let span = match self.current_span_id.load(Ordering::Relaxed) {
0 => span,
id => span.follows_from(Id::from_u64(id)).to_owned(),
};
if let Some(id) = span.id() {
self.current_span_id.store(id.into_u64(), Ordering::Relaxed);
}
span
}
}
impl<M> TurnSource<M> for UnaryTurnSource
where
M: CompletionModel,
{
type Raw = M::Response;
fn open_chat_span(
&self,
runner: &AgentRunner<M>,
effective_preamble: Option<&str>,
) -> tracing::Span {
let chat_span = build_chat_span!(runner, effective_preamble, "chat", "chat");
self.chain_span(chat_span)
}
fn run_model_turn<'a>(
&'a mut self,
runner: &'a AgentRunner<M>,
hook_ctx: &'a HookContext,
run: &'a mut AgentRun,
prepared: PreparedCompletionRequest<M>,
chat_span: tracing::Span,
_agent_span: &'a tracing::Span,
current_prompt: Message,
) -> DriveStream<'a, M::Response> {
Box::pin(async_stream::stream! {
let resp = match prepared.builder.send().instrument(chat_span.clone()).await {
Ok(resp) => resp,
Err(err) => {
yield Err(StreamingError::from(err));
return;
}
};
let mut outcome = match run.model_response(ModelTurn::new(
resp.message_id.clone(),
resp.choice.clone(),
resp.usage,
prepared.executable_tool_names,
prepared.allowed_tool_names,
)) {
Ok(outcome) => outcome,
Err(err) => {
yield Err(Box::new(err).into());
return;
}
};
loop {
match outcome {
ModelTurnOutcome::NeedsResolution(context) => {
let action = runner
.hooks
.on_invalid_tool_call(hook_ctx, &context)
.await;
let resolution = match action {
Some(action) => run.resolve_invalid_tool_call(action),
None
if runner.unhandled_invalid_tool_call_policy
== UnhandledInvalidToolCallPolicy::IgnoreForExtractor =>
{
run.ignore_invalid_tool_call()
}
None => run.resolve_invalid_tool_call(InvalidToolCallAction::fail()),
};
outcome = match resolution {
Ok(outcome) => outcome,
Err(err) => {
yield Err(Box::new(err).into());
return;
}
};
}
ModelTurnOutcome::TurnRetried => break,
ModelTurnOutcome::Continue {
response_hook_suppressed,
} => {
if !response_hook_suppressed {
if let Some(reason) = observe_action(
runner
.hooks
.on_completion_response(
hook_ctx,
CompletionResponseEvent {
prompt: ¤t_prompt,
content: &resp.choice,
usage: resp.usage,
message_id: resp.message_id.as_deref(),
},
)
.await,
) {
if runner.record_telemetry_content
&& let Some(choice) = run.accepted_turn_choice()
{
rig_core::telemetry::record_model_output(
&chat_span, &choice, true,
);
}
yield Err(StreamingError::Prompt(Box::new(run.cancel_error(reason))));
return;
}
let action = runner
.hooks
.on_model_turn_finished(
hook_ctx,
ModelTurnFinished {
turn: hook_ctx.turn(),
content: &resp.choice,
usage: resp.usage,
},
)
.await;
match resolve_model_turn_action(run, action) {
Ok(ModelTurnDecision::Advance) => {}
Ok(ModelTurnDecision::Retried) => break,
Ok(ModelTurnDecision::Terminate(reason)) => {
if runner.record_telemetry_content
&& let Some(choice) = run.accepted_turn_choice()
{
rig_core::telemetry::record_model_output(
&chat_span, &choice, true,
);
}
yield Err(StreamingError::Prompt(Box::new(
run.cancel_error(reason),
)));
return;
}
Err(err) => {
yield Err(StreamingError::Prompt(Box::new(err)));
return;
}
}
}
if runner.record_telemetry_content
&& let Some(choice) = run.accepted_turn_choice()
{
rig_core::telemetry::record_model_output(&chat_span, &choice, true);
}
break;
}
}
}
})
}
fn run_tool_calls<'a>(
&'a self,
runner: &'a AgentRunner<M>,
hook_ctx: &'a HookContext,
run: &'a mut AgentRun,
calls: Vec<PendingToolCall>,
tool_snapshot: Arc<ToolRegistrySnapshot>,
) -> DriveStream<'a, M::Response> {
drive_tool_calls(
runner,
hook_ctx,
run,
calls,
tool_snapshot,
|span| self.chain_span(span),
false,
)
}
fn record_run_level_telemetry(
&self,
agent_span: &tracing::Span,
response: &PromptResponse,
created_agent_span: bool,
) {
if created_agent_span {
if self.record_telemetry_content {
agent_span.record("gen_ai.completion", &response.output);
}
record_usage_on_span(agent_span, response.usage);
}
}
fn final_item(&self, _response: &PromptResponse) -> Option<MultiTurnStreamItem<M::Response>> {
None
}
}
impl<M> AgentRunner<M>
where
M: CompletionModel,
{
pub(crate) async fn run_with_error_usage(
mut self,
) -> (Result<PromptResponse, PromptError>, Usage) {
let usage = Arc::new(Mutex::new(Usage::new()));
self.error_usage = Some(usage.clone());
let result = self.run().await;
let observed = result.as_ref().map_or_else(
|_| *usage.lock().unwrap_or_else(|error| error.into_inner()),
|response| response.usage,
);
(result, observed)
}
pub async fn run(self) -> Result<PromptResponse, PromptError> {
let (agent_span, created_agent_span) = acquire_agent_span(
self.agent_name_or_default(),
self.preamble.as_deref(),
self.record_telemetry_content,
);
if self.record_telemetry_content
&& let Some(text) = self.prompt.rag_text()
{
agent_span.record("gen_ai.prompt", text);
}
let (history_override, memory_handle) = match &self.chat_history {
Some(_) => (None, None),
None => match (&self.memory, &self.conversation_id) {
(Some(memory), Some(id)) => {
let loaded = memory.load(id).await?;
(Some(loaded), Some((memory.clone(), id.clone())))
}
_ => (None, None),
},
};
let run = self.build_run(history_override);
let record_telemetry_content = self.record_telemetry_content;
let driver = drive_agent(
self,
UnaryTurnSource::new(record_telemetry_content),
run,
agent_span,
created_agent_span,
memory_handle,
false,
);
futures::pin_mut!(driver);
let mut response = None;
while let Some(item) = driver.next().await {
match item {
Ok(DriveItem::Done(done)) => response = Some(*done),
Ok(DriveItem::Item(_)) => {}
Err(err) => return Err(streaming_error_into_prompt(err)),
}
}
response.ok_or_else(|| {
PromptError::CompletionError(CompletionError::ResponseError(
"agent run ended without producing a final response".to_string(),
))
})
}
}
#[cfg(test)]
mod tests {
use std::sync::{
Arc, Mutex,
atomic::{AtomicUsize, Ordering},
};
use futures::StreamExt;
use serde_json::json;
use crate::{
agent::{AgentBuilder, AgentHook, HookContext, ToolResultAction, ToolResultEvent},
completion::{CompletionModel, Document},
test_utils::{MockCompletionModel, MockStreamEvent, MockTurn},
tool::{Tool, ToolContext, ToolErrorKind, ToolExecutionError},
};
use rig_core::message::ToolChoice;
struct MetadataFailingTool;
struct SnapshotValue {
value: usize,
clones: Arc<AtomicUsize>,
}
impl Clone for SnapshotValue {
fn clone(&self) -> Self {
self.clones.fetch_add(1, Ordering::SeqCst);
Self {
value: self.value,
clones: self.clones.clone(),
}
}
}
#[derive(Clone, Default)]
struct SnapshotMutatingTool(Arc<Mutex<Vec<usize>>>);
impl Tool for SnapshotMutatingTool {
const NAME: &'static str = "snapshot_mutator";
type Error = rig::tool::ToolExecutionError;
type Args = serde_json::Value;
type Output = String;
fn description(&self) -> String {
"Mutates its per-dispatch context snapshot".into()
}
fn parameters(&self) -> serde_json::Value {
json!({"type": "object", "properties": {}})
}
async fn call(
&self,
context: &mut ToolContext,
_args: Self::Args,
) -> Result<Self::Output, ToolExecutionError> {
let initial = context.require::<SnapshotValue>()?.value;
self.0.lock().expect("observed values").push(initial);
let updated = {
let value = context
.get_mut::<SnapshotValue>()
.expect("required snapshot value");
value.value += 1;
value.value
};
context.insert_result(updated);
Ok(updated.to_string())
}
}
#[derive(Clone, Default)]
struct SnapshotResults(Arc<Mutex<Vec<usize>>>);
impl AgentHook for SnapshotResults {
async fn on_tool_result(
&self,
_ctx: &HookContext,
event: ToolResultEvent<'_>,
) -> ToolResultAction {
self.0.lock().expect("result values").push(
*event
.tool_context
.require_result::<usize>()
.expect("per-dispatch result metadata"),
);
ToolResultAction::keep()
}
}
impl Tool for MetadataFailingTool {
const NAME: &'static str = "flaky_tool";
type Error = rig::tool::ToolExecutionError;
type Args = serde_json::Value;
type Output = String;
fn description(&self) -> String {
"Fails after attaching result metadata".into()
}
fn parameters(&self) -> serde_json::Value {
json!({"type": "object", "properties": {}})
}
async fn call(
&self,
context: &mut ToolContext,
_args: Self::Args,
) -> Result<Self::Output, ToolExecutionError> {
context.insert_result("shared-result-metadata".to_string());
Err(ToolExecutionError::timeout("raw timeout failure"))
}
}
#[derive(Clone, Default)]
struct Results(Arc<Mutex<Vec<(ToolErrorKind, String, String)>>>);
impl AgentHook for Results {
async fn on_tool_result(
&self,
_ctx: &HookContext,
event: ToolResultEvent<'_>,
) -> ToolResultAction {
if let Some(error) = event.raw_result.error() {
self.0.lock().expect("results").push((
error.kind(),
event.raw_result.output().render(),
event
.tool_context
.result::<String>()
.expect("tool result metadata")
.clone(),
));
}
ToolResultAction::rewrite("rewritten for model")
}
}
#[test]
fn agent_exposes_read_only_name_and_description() {
let named = AgentBuilder::new(MockCompletionModel::text("done"))
.name("researcher")
.description("Finds evidence")
.build();
assert_eq!(named.name(), Some("researcher"));
assert_eq!(named.description(), Some("Finds evidence"));
let unnamed = AgentBuilder::new(MockCompletionModel::text("done")).build();
assert_eq!(unnamed.name(), None);
assert_eq!(unnamed.description(), None);
}
#[tokio::test]
async fn runner_applies_per_run_request_overrides() {
let model = MockCompletionModel::text("done");
AgentBuilder::new(model.clone())
.preamble("baseline preamble")
.context("baseline document")
.temperature(0.1)
.max_tokens(10)
.additional_params(json!({"baseline": true}))
.build()
.runner("go")
.preamble("run preamble")
.document(Document {
id: "run-one".into(),
text: "first run document".into(),
additional_props: Default::default(),
})
.documents([Document {
id: "run-two".into(),
text: "second run document".into(),
additional_props: Default::default(),
}])
.temperature(0.7)
.max_tokens(42)
.replace_additional_params(json!({"override": true}))
.tool_choice(ToolChoice::None)
.run()
.await
.expect("runner request should succeed");
let requests = model.requests();
let request = requests.first().expect("one request");
assert!(request.chat_history.iter().any(
|message| matches!(message, crate::completion::Message::System { content } if content == "run preamble")
));
assert!(
request
.documents
.iter()
.any(|document| document.text == "baseline document")
);
assert!(
request
.documents
.iter()
.any(|document| document.id == "run-one")
);
assert!(
request
.documents
.iter()
.any(|document| document.id == "run-two")
);
assert_eq!(request.temperature, Some(0.7));
assert_eq!(request.max_tokens, Some(42));
assert_eq!(request.additional_params, Some(json!({"override": true})));
assert_eq!(request.tool_choice, Some(ToolChoice::None));
}
#[tokio::test]
async fn runner_can_merge_additional_params_into_the_baseline() {
let model = MockCompletionModel::text("done");
AgentBuilder::new(model.clone())
.additional_params(json!({"baseline": true, "winner": "baseline"}))
.build()
.runner("go")
.merge_additional_params(
json!({"override": true, "winner": "runner"})
.as_object()
.expect("object")
.clone(),
)
.run()
.await
.expect("runner request should succeed");
assert_eq!(
model
.requests()
.first()
.expect("one request")
.additional_params,
Some(json!({"baseline": true, "override": true, "winner": "runner"}))
);
}
#[tokio::test]
async fn runner_can_replace_additional_params_wholesale() {
let model = MockCompletionModel::text("done");
AgentBuilder::new(model.clone())
.additional_params(json!({"baseline": true}))
.build()
.runner("go")
.replace_additional_params(json!({"replacement": true}))
.run()
.await
.expect("runner request should succeed");
let requests = model.requests();
let request = requests.first().expect("one request");
assert_eq!(
request.additional_params,
Some(json!({"replacement": true}))
);
}
#[tokio::test]
async fn runner_can_clear_configured_request_defaults() {
let model = MockCompletionModel::text("done");
AgentBuilder::new(model.clone())
.preamble("baseline")
.temperature(0.1)
.max_tokens(10)
.additional_params(json!({"baseline": true}))
.tool_choice(ToolChoice::Required)
.build()
.runner("go")
.without_preamble()
.without_temperature()
.without_max_tokens()
.without_additional_params()
.without_tool_choice()
.run()
.await
.expect("runner request should succeed");
let requests = model.requests();
let request = requests.first().expect("one request");
assert!(
!request
.chat_history
.iter()
.any(|message| matches!(message, crate::completion::Message::System { .. }))
);
assert_eq!(request.temperature, None);
assert_eq!(request.max_tokens, None);
assert_eq!(request.additional_params, None);
assert_eq!(request.tool_choice, None);
}
#[tokio::test]
async fn direct_completion_model_requests_are_intentionally_hook_free() {
#[derive(Clone)]
struct CountCompletionCalls(Arc<AtomicUsize>);
impl AgentHook for CountCompletionCalls {
async fn on_completion_call(
&self,
_ctx: &HookContext,
_event: crate::agent::CompletionCallEvent<'_>,
) -> crate::agent::CompletionCallAction {
self.0.fetch_add(1, Ordering::SeqCst);
crate::agent::CompletionCallAction::Continue
}
}
let model = MockCompletionModel::text("raw response");
let calls = Arc::new(AtomicUsize::new(0));
let _agent = AgentBuilder::new(model.clone())
.add_hook(CountCompletionCalls(calls.clone()))
.build();
model
.completion_request("raw request")
.send()
.await
.expect("direct model request should succeed");
assert_eq!(calls.load(Ordering::SeqCst), 0);
assert_eq!(model.request_count(), 1);
}
#[tokio::test]
async fn blocking_and_streaming_preserve_raw_failure_while_rewriting_presentation() {
let blocking = Results::default();
let blocking_model = MockCompletionModel::from_turns([
MockTurn::tool_call("tc1", "flaky_tool", json!({})),
MockTurn::text("done"),
]);
AgentBuilder::new(blocking_model.clone())
.tool(MetadataFailingTool)
.add_hook(blocking.clone())
.build()
.runner("go")
.max_turns(3)
.run()
.await
.expect("blocking run");
let streaming = Results::default();
let streaming_model = MockCompletionModel::from_stream_turns([
vec![
MockStreamEvent::tool_call_name_delta("tc1", "ic1", "flaky_tool"),
MockStreamEvent::tool_call_arguments_delta("tc1", "ic1", "{}"),
MockStreamEvent::tool_call("tc1", "flaky_tool", json!({})),
MockStreamEvent::final_response_with_total_tokens(0),
],
vec![
MockStreamEvent::text("done"),
MockStreamEvent::final_response_with_total_tokens(0),
],
]);
let mut stream = AgentBuilder::new(streaming_model.clone())
.tool(MetadataFailingTool)
.add_hook(streaming.clone())
.build()
.runner("go")
.max_turns(3)
.stream()
.await;
while let Some(item) = stream.next().await {
item.expect("stream item");
}
assert_eq!(*blocking.0.lock().unwrap(), *streaming.0.lock().unwrap());
assert_eq!(
*blocking.0.lock().unwrap(),
vec![(
ToolErrorKind::Timeout,
"raw timeout failure".into(),
"shared-result-metadata".into()
)]
);
let blocking_history = serde_json::to_value(
&blocking_model
.requests()
.get(1)
.expect("second blocking request")
.chat_history,
)
.unwrap();
let streaming_history = serde_json::to_value(
&streaming_model
.requests()
.get(1)
.expect("second streaming request")
.chat_history,
)
.unwrap();
assert_eq!(blocking_history, streaming_history);
let history = blocking_history.to_string();
assert!(history.contains("rewritten for model"));
assert!(!history.contains("raw timeout failure"));
}
#[tokio::test]
async fn agent_dispatch_snapshot_clones_once_and_isolates_tool_mutations() {
let clones = Arc::new(AtomicUsize::new(0));
let mut context = ToolContext::new();
context.insert(SnapshotValue {
value: 0,
clones: clones.clone(),
});
let tool = SnapshotMutatingTool::default();
let results = SnapshotResults::default();
AgentBuilder::new(MockCompletionModel::from_turns([
MockTurn::tool_call("tc1", SnapshotMutatingTool::NAME, json!({})),
MockTurn::tool_call("tc2", SnapshotMutatingTool::NAME, json!({})),
MockTurn::text("done"),
]))
.tool(tool.clone())
.add_hook(results.clone())
.build()
.runner("go")
.tool_context(context)
.max_turns(4)
.run()
.await
.expect("agent run");
assert_eq!(*tool.0.lock().expect("observed values"), vec![0, 0]);
assert_eq!(*results.0.lock().expect("result values"), vec![1, 1]);
assert_eq!(
clones.load(Ordering::SeqCst),
2,
"each of the two agent dispatches should clone inbound context once"
);
}
}
#[cfg(test)]
#[allow(irrefutable_let_patterns, unreachable_patterns)]
mod migrated_tests {
use std::collections::HashMap;
use crate::agent::{
CompletionCallAction, CompletionCallEvent, HookStack, InvalidToolCallAction,
InvalidToolCallContext, ModelTurnAction, ModelTurnFinished, ObservationAction,
StreamResponseFinish, TextDelta, ToolCall, ToolCallAction, ToolCallDelta, ToolResultAction,
ToolResultEvent,
};
use std::sync::{
Arc, Mutex,
atomic::{AtomicBool, AtomicU32, AtomicU64, Ordering::SeqCst},
};
use futures::StreamExt;
use serde::Deserialize;
use serde_json::json;
use tokio::sync::{Barrier, Notify};
use crate::agent::AgentBuilder;
use crate::agent::hook::{AgentHook, HookContext, RequestPatch, StepEventKind};
use crate::agent::prompt_request::streaming::{MultiTurnStreamItem, StreamingError};
use crate::agent::run::OutputMode;
use crate::completion::{
CompletionError, CompletionModel, Message, Prompt, PromptError, Usage,
};
use crate::streaming::{StreamedAssistantContent, StreamedUserContent, StreamingPrompt};
use crate::test_utils::{
MockAddTool, MockBarrierTool, MockCompletionModel, MockOperationArgs, MockStreamEvent,
MockSubtractTool, MockToolError, MockTurn,
};
use crate::tool::{
Tool, ToolContext, ToolExecutionError, ToolSet,
server::{ToolServer, ToolServerHandle},
};
use rig_core::OneOrMany;
use rig_core::message::{
AssistantContent, ToolCall as MessageToolCall, ToolChoice, ToolFunction, UserContent,
};
use rig_core::vector_store::{
VectorSearchRequest, VectorStoreError, VectorStoreIndex, request::Filter,
};
use rig_core::wasm_compat::WasmCompatSend;
#[derive(Clone, Default)]
struct RecordingHook {
events: Arc<Mutex<Vec<StepEventKind>>>,
tool_results: Arc<Mutex<Vec<String>>>,
}
impl RecordingHook {
fn shared_events(&self) -> Vec<StepEventKind> {
self.events
.lock()
.expect("events lock")
.iter()
.copied()
.filter(|kind| {
matches!(
kind,
StepEventKind::CompletionCall
| StepEventKind::ToolCall
| StepEventKind::ToolResult
| StepEventKind::InvalidToolCall
)
})
.collect()
}
fn tool_results(&self) -> Vec<String> {
self.tool_results.lock().expect("results lock").clone()
}
fn count(&self, kind: StepEventKind) -> usize {
self.events
.lock()
.expect("events lock")
.iter()
.filter(|recorded| **recorded == kind)
.count()
}
}
impl RecordingHook {
fn record(&self, kind: StepEventKind) {
self.events.lock().expect("events lock").push(kind);
}
}
impl AgentHook for RecordingHook {
async fn on_completion_call(
&self,
_: &HookContext,
_: CompletionCallEvent<'_>,
) -> CompletionCallAction {
self.record(StepEventKind::CompletionCall);
CompletionCallAction::continue_run()
}
async fn on_completion_response(
&self,
_: &HookContext,
_: crate::agent::hook::CompletionResponse<'_>,
) -> ObservationAction {
self.record(StepEventKind::CompletionResponse);
ObservationAction::continue_run()
}
async fn on_model_turn_finished(
&self,
_: &HookContext,
_: ModelTurnFinished<'_>,
) -> ModelTurnAction {
self.record(StepEventKind::ModelTurnFinished);
ModelTurnAction::continue_run()
}
async fn on_invalid_tool_call(
&self,
_: &HookContext,
_: &InvalidToolCallContext,
) -> Option<InvalidToolCallAction> {
self.record(StepEventKind::InvalidToolCall);
None
}
async fn on_tool_call(&self, _: &HookContext, _: ToolCall<'_>) -> ToolCallAction {
self.record(StepEventKind::ToolCall);
ToolCallAction::run()
}
async fn on_tool_result(
&self,
_: &HookContext,
event: ToolResultEvent<'_>,
) -> ToolResultAction {
self.record(StepEventKind::ToolResult);
self.tool_results
.lock()
.expect("results lock")
.push(event.presentation.render());
ToolResultAction::keep()
}
async fn on_text_delta(&self, _: &HookContext, _: TextDelta<'_>) -> ObservationAction {
self.record(StepEventKind::TextDelta);
ObservationAction::continue_run()
}
async fn on_tool_call_delta(
&self,
_: &HookContext,
_: ToolCallDelta<'_>,
) -> ObservationAction {
self.record(StepEventKind::ToolCallDelta);
ObservationAction::continue_run()
}
async fn on_stream_response_finish(
&self,
_: &HookContext,
_: StreamResponseFinish<'_>,
) -> ObservationAction {
self.record(StepEventKind::StreamResponseFinish);
ObservationAction::continue_run()
}
}
#[derive(Clone, Debug, PartialEq)]
struct CanonicalResponseSnapshot {
prompt: Message,
content: OneOrMany<AssistantContent>,
usage: Usage,
message_id: Option<String>,
}
#[derive(Clone, Default)]
struct CanonicalResponseHook {
blocking: Arc<Mutex<Vec<CanonicalResponseSnapshot>>>,
streaming: Arc<Mutex<Vec<CanonicalResponseSnapshot>>>,
committed: Arc<Mutex<Vec<OneOrMany<AssistantContent>>>>,
}
impl AgentHook for CanonicalResponseHook {
async fn on_completion_response(
&self,
_ctx: &HookContext,
event: crate::agent::hook::CompletionResponse<'_>,
) -> ObservationAction {
self.blocking
.lock()
.expect("blocking snapshots")
.push(CanonicalResponseSnapshot {
prompt: event.prompt.clone(),
content: event.content.clone(),
usage: event.usage,
message_id: event.message_id.map(str::to_owned),
});
ObservationAction::continue_run()
}
async fn on_stream_response_finish(
&self,
_ctx: &HookContext,
event: StreamResponseFinish<'_>,
) -> ObservationAction {
self.streaming
.lock()
.expect("streaming snapshots")
.push(CanonicalResponseSnapshot {
prompt: event.prompt.clone(),
content: event.content.clone(),
usage: event.usage,
message_id: event.message_id.map(str::to_owned),
});
ObservationAction::continue_run()
}
async fn on_model_turn_finished(
&self,
_ctx: &HookContext,
event: ModelTurnFinished<'_>,
) -> ModelTurnAction {
self.committed
.lock()
.expect("committed snapshots")
.push(event.content.clone());
ModelTurnAction::continue_run()
}
}
#[derive(Clone, Default)]
struct FinishLifecycleHook {
snapshots: Arc<Mutex<Vec<CanonicalResponseSnapshot>>>,
model_turns: Arc<AtomicU32>,
stop: Arc<AtomicBool>,
}
impl FinishLifecycleHook {
fn stopping() -> Self {
let hook = Self::default();
hook.stop.store(true, SeqCst);
hook
}
}
impl AgentHook for FinishLifecycleHook {
async fn on_stream_response_finish(
&self,
_ctx: &HookContext,
event: StreamResponseFinish<'_>,
) -> ObservationAction {
self.snapshots
.lock()
.expect("finish snapshots")
.push(CanonicalResponseSnapshot {
prompt: event.prompt.clone(),
content: event.content.clone(),
usage: event.usage,
message_id: event.message_id.map(str::to_owned),
});
if self.stop.load(SeqCst) {
ObservationAction::stop("stop at stream EOF")
} else {
ObservationAction::continue_run()
}
}
async fn on_model_turn_finished(
&self,
_ctx: &HookContext,
_event: ModelTurnFinished<'_>,
) -> ModelTurnAction {
self.model_turns.fetch_add(1, SeqCst);
ModelTurnAction::continue_run()
}
}
fn canonical_usage() -> Usage {
Usage {
input_tokens: 11,
output_tokens: 7,
total_tokens: 18,
..Usage::new()
}
}
#[tokio::test]
async fn blocking_completion_response_hook_receives_canonical_fields() {
let hook = CanonicalResponseHook::default();
let prompt = Message::user("canonical prompt");
AgentBuilder::new(MockCompletionModel::new([MockTurn::text(
"canonical response",
)
.with_usage(canonical_usage())
.with_message_id("msg-canonical")]))
.add_hook(hook.clone())
.build()
.runner(prompt.clone())
.run()
.await
.expect("blocking response");
assert_eq!(
*hook.blocking.lock().expect("blocking snapshots"),
[CanonicalResponseSnapshot {
prompt,
content: OneOrMany::one(AssistantContent::text("canonical response")),
usage: canonical_usage(),
message_id: Some("msg-canonical".to_string()),
}]
);
}
#[tokio::test]
async fn streaming_response_finish_matches_blocking_canonical_fields() {
let prompt = Message::user("canonical prompt");
let blocking_hook = CanonicalResponseHook::default();
AgentBuilder::new(MockCompletionModel::new([MockTurn::text(
"canonical response",
)
.with_usage(canonical_usage())
.with_message_id("msg-canonical")]))
.add_hook(blocking_hook.clone())
.build()
.runner(prompt.clone())
.run()
.await
.expect("blocking response");
let streaming_hook = CanonicalResponseHook::default();
let mut stream = AgentBuilder::new(MockCompletionModel::from_stream_turns([[
MockStreamEvent::text("canonical response"),
MockStreamEvent::final_response(canonical_usage()),
MockStreamEvent::message_id("msg-canonical"),
]]))
.add_hook(streaming_hook.clone())
.build()
.runner(prompt)
.stream()
.await;
while let Some(item) = stream.next().await {
item.expect("stream item");
}
let blocking = blocking_hook
.blocking
.lock()
.expect("blocking snapshots")
.clone();
let streaming = streaming_hook
.streaming
.lock()
.expect("streaming snapshots")
.clone();
assert_eq!(streaming, blocking);
assert_eq!(streaming[0].usage, canonical_usage());
assert_eq!(streaming[0].message_id.as_deref(), Some("msg-canonical"));
}
#[tokio::test]
async fn streaming_response_finish_without_provider_message_id_reports_none() {
let hook = FinishLifecycleHook::default();
let mut stream = AgentBuilder::new(MockCompletionModel::from_stream_turns([[
MockStreamEvent::text("canonical response"),
MockStreamEvent::final_response(canonical_usage()),
]]))
.add_hook(hook.clone())
.build()
.runner("canonical prompt")
.stream()
.await;
while let Some(item) = stream.next().await {
item.expect("stream item");
}
let snapshots = hook.snapshots.lock().expect("finish snapshots");
assert_eq!(snapshots.len(), 1);
assert_eq!(snapshots[0].message_id, None);
}
#[tokio::test]
async fn streaming_response_finish_runs_before_buffered_final_is_exposed() {
let hook = FinishLifecycleHook::default();
let mut stream = AgentBuilder::new(MockCompletionModel::from_stream_turns([[
MockStreamEvent::text("canonical response"),
MockStreamEvent::final_response(canonical_usage()),
MockStreamEvent::message_id("msg-after-final"),
]]))
.add_hook(hook.clone())
.build()
.runner("canonical prompt")
.stream()
.await;
let mut provider_finals = 0;
while let Some(item) = stream.next().await {
if matches!(
item.expect("stream item"),
MultiTurnStreamItem::StreamAssistantItem(StreamedAssistantContent::Final(_))
) {
provider_finals += 1;
let snapshots = hook.snapshots.lock().expect("finish snapshots");
assert_eq!(snapshots.len(), 1, "hook must run before final exposure");
assert_eq!(snapshots[0].message_id.as_deref(), Some("msg-after-final"));
assert_eq!(
hook.model_turns.load(SeqCst),
1,
"the canonical turn hook must accept the turn before final exposure"
);
}
}
assert_eq!(provider_finals, 1);
assert_eq!(hook.snapshots.lock().expect("finish snapshots").len(), 1);
assert_eq!(hook.model_turns.load(SeqCst), 1);
}
#[tokio::test]
async fn streaming_response_finish_stop_suppresses_final_and_turn_commit() {
let hook = FinishLifecycleHook::stopping();
let prompt = Message::user("canonical prompt");
let mut stream = AgentBuilder::new(MockCompletionModel::from_stream_turns([[
MockStreamEvent::text("canonical response"),
MockStreamEvent::final_response(canonical_usage()),
MockStreamEvent::message_id("msg-after-final"),
]]))
.add_hook(hook.clone())
.build()
.runner(prompt.clone())
.stream()
.await;
let mut saw_provider_final = false;
let mut saw_run_final = false;
let mut error = None;
while let Some(item) = stream.next().await {
match item {
Ok(MultiTurnStreamItem::StreamAssistantItem(StreamedAssistantContent::Final(
_,
))) => saw_provider_final = true,
Ok(MultiTurnStreamItem::FinalResponse(_)) => saw_run_final = true,
Ok(_) => {}
Err(err) => error = Some(err),
}
}
assert!(!saw_provider_final, "the buffered final must remain hidden");
assert!(
!saw_run_final,
"the cancelled run must not produce a response"
);
assert_eq!(hook.snapshots.lock().expect("finish snapshots").len(), 1);
assert_eq!(hook.model_turns.load(SeqCst), 0);
assert!(matches!(
error,
Some(StreamingError::Prompt(error))
if matches!(
error.as_ref(),
PromptError::PromptCancelled { chat_history, reason }
if chat_history == &[prompt] && reason == "stop at stream EOF"
)
));
}
struct StopCompletedModelTurn;
impl AgentHook for StopCompletedModelTurn {
async fn on_model_turn_finished(
&self,
_ctx: &HookContext,
_event: ModelTurnFinished<'_>,
) -> ModelTurnAction {
ModelTurnAction::stop("stop completed model turn")
}
}
#[tokio::test]
async fn streaming_model_turn_stop_preserves_completed_provider_final() {
let prompt = Message::user("canonical prompt");
let mut stream = AgentBuilder::new(MockCompletionModel::from_stream_turns([[
MockStreamEvent::text("canonical response"),
MockStreamEvent::final_response(canonical_usage()),
]]))
.add_hook(StopCompletedModelTurn)
.build()
.runner(prompt.clone())
.stream()
.await;
let mut provider_finals = 0;
let mut saw_retry = false;
let mut saw_run_final = false;
let mut error = None;
while let Some(item) = stream.next().await {
match item {
Ok(MultiTurnStreamItem::StreamAssistantItem(StreamedAssistantContent::Final(
_,
))) => provider_finals += 1,
Ok(MultiTurnStreamItem::ModelTurnRetried { .. }) => saw_retry = true,
Ok(MultiTurnStreamItem::FinalResponse(_)) => saw_run_final = true,
Ok(_) => {}
Err(err) => error = Some(err),
}
}
assert_eq!(provider_finals, 1);
assert!(!saw_retry);
assert!(!saw_run_final);
assert!(matches!(
error,
Some(StreamingError::Prompt(error))
if matches!(
error.as_ref(),
PromptError::PromptCancelled { reason, .. }
if reason == "stop completed model turn"
)
));
}
#[tokio::test]
async fn provider_error_after_final_suppresses_finish_hook_and_buffered_final() {
let hook = FinishLifecycleHook::default();
let mut stream = AgentBuilder::new(MockCompletionModel::from_stream_turns([[
MockStreamEvent::text("canonical response"),
MockStreamEvent::final_response(canonical_usage()),
MockStreamEvent::error("post-final failure"),
]]))
.add_hook(hook.clone())
.build()
.runner("canonical prompt")
.stream()
.await;
let mut saw_provider_final = false;
let mut error = None;
while let Some(item) = stream.next().await {
match item {
Ok(MultiTurnStreamItem::StreamAssistantItem(StreamedAssistantContent::Final(
_,
))) => saw_provider_final = true,
Ok(_) => {}
Err(err) => error = Some(err),
}
}
assert!(!saw_provider_final, "the buffered final must remain hidden");
assert!(hook.snapshots.lock().expect("finish snapshots").is_empty());
assert_eq!(hook.model_turns.load(SeqCst), 0);
assert!(matches!(
error,
Some(StreamingError::Completion(CompletionError::ProviderError(message)))
if message == "post-final failure"
));
}
#[tokio::test]
async fn visible_assistant_items_after_final_are_rejected() {
let cases = [
("text", MockStreamEvent::text("late text")),
("reasoning", MockStreamEvent::reasoning("late reasoning")),
(
"reasoning delta",
MockStreamEvent::reasoning_delta(None::<String>, "late reasoning"),
),
(
"tool call",
MockStreamEvent::tool_call("late", "add", json!({"x": 1, "y": 2})),
),
(
"tool-call delta",
MockStreamEvent::tool_call_name_delta("late", "internal-late", "add"),
),
("unknown", MockStreamEvent::unknown(json!({"type": "late"}))),
];
for (case, visible_item) in cases {
let hook = FinishLifecycleHook::default();
let mut stream = AgentBuilder::new(MockCompletionModel::from_stream_turns([vec![
MockStreamEvent::text("canonical response"),
MockStreamEvent::final_response(canonical_usage()),
visible_item,
]]))
.add_hook(hook.clone())
.build()
.runner("canonical prompt")
.stream()
.await;
let mut saw_provider_final = false;
let mut error = None;
while let Some(item) = stream.next().await {
match item {
Ok(MultiTurnStreamItem::StreamAssistantItem(
StreamedAssistantContent::Final(_),
)) => saw_provider_final = true,
Ok(_) => {}
Err(err) => error = Some(err),
}
}
assert!(
!saw_provider_final,
"{case}: buffered final must remain hidden"
);
assert!(
hook.snapshots.lock().expect("finish snapshots").is_empty(),
"{case}: finish hook must not run"
);
assert_eq!(hook.model_turns.load(SeqCst), 0, "{case}");
assert!(
matches!(
error,
Some(StreamingError::Completion(CompletionError::ResponseError(ref message)))
if message.contains("visible assistant content after its final response")
),
"{case}: expected malformed-response error, got {error:?}"
);
}
}
#[tokio::test]
async fn visible_item_after_non_emittable_final_is_rejected() {
let hook = FinishLifecycleHook::default();
let mut stream = AgentBuilder::new(MockCompletionModel::from_stream_turns([[
MockStreamEvent::reasoning("think"),
MockStreamEvent::final_response(canonical_usage()),
MockStreamEvent::text("late text"),
]]))
.add_hook(hook.clone())
.build()
.runner("canonical prompt")
.stream()
.await;
let mut error = None;
while let Some(item) = stream.next().await {
if let Err(err) = item {
error = Some(err);
}
}
assert!(hook.snapshots.lock().expect("finish snapshots").is_empty());
assert_eq!(hook.model_turns.load(SeqCst), 0);
assert!(matches!(
error,
Some(StreamingError::Completion(CompletionError::ResponseError(message)))
if message.contains("visible assistant content after its final response")
));
}
#[tokio::test]
async fn streaming_response_finish_normalizes_interleaved_content() {
let hook = CanonicalResponseHook::default();
let mut stream = AgentBuilder::new(MockCompletionModel::from_stream_turns([
vec![
MockStreamEvent::reasoning("think"),
MockStreamEvent::tool_call("tc1", "add", json!({"x": 2, "y": 3})),
MockStreamEvent::text("answer"),
MockStreamEvent::final_response_with_total_tokens(0),
],
vec![
MockStreamEvent::text("done"),
MockStreamEvent::final_response_with_total_tokens(0),
],
]))
.tool(MockAddTool)
.add_hook(hook.clone())
.build()
.runner("go")
.max_turns(3)
.stream()
.await;
while let Some(item) = stream.next().await {
item.expect("stream item");
}
let snapshots = hook.streaming.lock().expect("streaming snapshots");
let committed = hook.committed.lock().expect("committed snapshots");
let kinds = snapshots[0]
.content
.iter()
.map(|content| match content {
AssistantContent::Reasoning(_) => "reasoning",
AssistantContent::Text(_) => "text",
AssistantContent::ToolCall(_) => "tool_call",
_ => "other",
})
.collect::<Vec<_>>();
assert_eq!(kinds, ["reasoning", "text", "tool_call"]);
assert_eq!(
snapshots[0].content, committed[0],
"finish hook and committed turn must share one canonical choice"
);
}
fn blocking_model() -> MockCompletionModel {
MockCompletionModel::from_turns([
MockTurn::tool_call("tc1", "add", json!({"x": 2, "y": 3})),
MockTurn::text("the answer is 5"),
])
}
fn streaming_model() -> MockCompletionModel {
MockCompletionModel::from_stream_turns([
vec![
MockStreamEvent::tool_call_name_delta("tc1", "ic1", "add"),
MockStreamEvent::tool_call_arguments_delta("tc1", "ic1", "{\"x\":2,\"y\":3}"),
MockStreamEvent::tool_call("tc1", "add", json!({"x": 2, "y": 3})),
MockStreamEvent::final_response_with_total_tokens(0),
],
vec![
MockStreamEvent::text("the answer is 5"),
MockStreamEvent::final_response_with_total_tokens(0),
],
])
}
#[tokio::test]
async fn from_agent_preserves_implicit_one_and_explicit_zero_budgets() {
let implicit_model = blocking_model();
let implicit_recorded = implicit_model.clone();
let implicit_agent = AgentBuilder::new(implicit_model).tool(MockAddTool).build();
let implicit_runner = super::AgentRunner::from_agent(&implicit_agent, "add 2 and 3");
assert_eq!(implicit_runner.max_turns, 1);
let implicit_err = implicit_runner
.run()
.await
.expect_err("implicit budget should reject the second model call");
assert!(matches!(
implicit_err,
PromptError::MaxTurnsError { max_turns: 1, .. }
));
assert_eq!(implicit_recorded.request_count(), 1);
let zero_model = MockCompletionModel::text("should not be requested");
let zero_recorded = zero_model.clone();
let zero_agent = AgentBuilder::new(zero_model).default_max_turns(0).build();
let zero_runner = super::AgentRunner::from_agent(&zero_agent, "do not call");
assert_eq!(zero_runner.max_turns, 0);
let zero_err = zero_runner
.run()
.await
.expect_err("explicit zero budget should reject the initial model call");
assert!(matches!(
zero_err,
PromptError::MaxTurnsError { max_turns: 0, .. }
));
assert_eq!(zero_recorded.request_count(), 0);
}
#[tokio::test]
async fn prompt_surfaces_reject_second_tool_roundtrip_request_at_budget_one() {
let blocking_model = blocking_model();
let blocking_recorded = blocking_model.clone();
let blocking_agent = AgentBuilder::new(blocking_model).tool(MockAddTool).build();
let blocking_err = blocking_agent
.prompt("add 2 and 3")
.max_turns(1)
.await
.expect_err("blocking prompt should reject request two");
assert!(matches!(
blocking_err,
PromptError::MaxTurnsError { max_turns: 1, .. }
));
assert_eq!(blocking_recorded.request_count(), 1);
let streaming_model = streaming_model();
let streaming_recorded = streaming_model.clone();
let streaming_agent = AgentBuilder::new(streaming_model).tool(MockAddTool).build();
let mut stream = streaming_agent
.stream_prompt("add 2 and 3")
.max_turns(1)
.await;
let mut streaming_err = None;
while let Some(item) = stream.next().await {
if let Err(err) = item {
streaming_err = Some(err);
break;
}
}
match streaming_err {
Some(StreamingError::Prompt(err)) => assert!(matches!(
*err,
PromptError::MaxTurnsError { max_turns: 1, .. }
)),
other => panic!("expected streaming max-turns error, got {other:?}"),
}
assert_eq!(streaming_recorded.request_count(), 1);
}
#[tokio::test]
async fn run_and_stream_behave_identically_for_a_tool_call() {
let blocking_hook = RecordingHook::default();
let blocking = AgentBuilder::new(blocking_model())
.tool(MockAddTool)
.build()
.runner("add 2 and 3")
.max_turns(2)
.add_hook(blocking_hook.clone())
.run()
.await
.expect("blocking run should succeed");
let streaming_hook = RecordingHook::default();
let mut stream = AgentBuilder::new(streaming_model())
.tool(MockAddTool)
.build()
.runner("add 2 and 3")
.max_turns(2)
.add_hook(streaming_hook.clone())
.stream()
.await;
let mut final_response = None;
while let Some(item) = stream.next().await {
if let Ok(MultiTurnStreamItem::FinalResponse(resp)) =
item.map_err(|err| panic!("stream item errored: {err}"))
{
final_response = Some(resp);
}
}
let final_response = final_response.expect("stream should yield a final response");
assert_eq!(blocking.output, "the answer is 5");
assert_eq!(final_response.output(), blocking.output);
assert_eq!(
blocking_hook.shared_events(),
streaming_hook.shared_events()
);
assert_eq!(
blocking_hook.shared_events(),
vec![
StepEventKind::CompletionCall,
StepEventKind::ToolCall,
StepEventKind::ToolResult,
StepEventKind::CompletionCall,
]
);
assert_eq!(blocking_hook.tool_results(), streaming_hook.tool_results());
assert_eq!(blocking_hook.tool_results(), vec!["5".to_string()]);
let blocking_messages = blocking.messages.expect("blocking messages");
let streaming_messages = final_response
.messages()
.expect("streaming history")
.to_vec();
assert_eq!(
serde_json::to_value(&blocking_messages).expect("serialize blocking"),
serde_json::to_value(&streaming_messages).expect("serialize streaming"),
);
}
mod structured_tool_results {
use std::sync::{Arc, Mutex};
use futures::StreamExt;
use serde_json::json;
use crate::agent::{
AgentBuilder, AgentHook, HookContext, HookStack, ToolCall, ToolCallAction,
ToolResultAction, ToolResultEvent,
};
use crate::test_utils::{
MockAddTool, MockCompletionModel, MockDeniedTool, MockFailingTool,
MockHandledFailureTool, MockMetadataTool, MockRequestId, MockStreamEvent, MockTurn,
};
use crate::tool::{ToolErrorKind, ToolResult};
#[derive(Clone, Default)]
struct OutcomeHook {
outcomes: Arc<Mutex<Vec<String>>>,
results: Arc<Mutex<Vec<String>>>,
}
impl OutcomeHook {
fn outcomes(&self) -> Vec<String> {
self.outcomes.lock().expect("outcomes").clone()
}
fn results(&self) -> Vec<String> {
self.results.lock().expect("results").clone()
}
}
fn outcome_label(result: &ToolResult) -> String {
if result.is_skipped() {
"skipped".to_string()
} else if result.is_refused() {
"denied".to_string()
} else if let Some(error) = result.error() {
format!("error:{}", error.kind().as_str())
} else {
"success".to_string()
}
}
impl AgentHook for OutcomeHook {
async fn on_tool_result(
&self,
_ctx: &HookContext,
event: ToolResultEvent<'_>,
) -> ToolResultAction {
if let ToolResultEvent {
presentation,
raw_result,
..
} = event
{
self.outcomes
.lock()
.expect("outcomes")
.push(outcome_label(raw_result));
self.results
.lock()
.expect("results")
.push(presentation.render());
}
ToolResultAction::keep()
}
}
fn model_one_tool_then_text(tool: &str) -> MockCompletionModel {
MockCompletionModel::from_turns([
MockTurn::tool_call("tc1", tool, json!({})),
MockTurn::text("done"),
])
}
fn stream_model_one_tool_then_text(tool: &str) -> MockCompletionModel {
MockCompletionModel::from_stream_turns([
vec![
MockStreamEvent::tool_call_name_delta("tc1", "ic1", tool),
MockStreamEvent::tool_call_arguments_delta("tc1", "ic1", "{}"),
MockStreamEvent::tool_call("tc1", tool, json!({})),
MockStreamEvent::final_response_with_total_tokens(0),
],
vec![
MockStreamEvent::text("done"),
MockStreamEvent::final_response_with_total_tokens(0),
],
])
}
#[tokio::test]
async fn timeout_failure_surfaces_structured_outcome() {
let hook = OutcomeHook::default();
AgentBuilder::new(model_one_tool_then_text("flaky_tool"))
.tool(MockFailingTool::new(ToolErrorKind::Timeout))
.add_hook(hook.clone())
.build()
.runner("go")
.max_turns(3)
.run()
.await
.expect("run should succeed; a tool timeout is model-visible feedback, not fatal");
assert_eq!(hook.outcomes(), vec!["error:timeout".to_string()]);
assert_eq!(hook.results(), vec!["mock tool call failed".to_string()]);
}
#[tokio::test]
async fn hook_terminates_after_repeated_timeouts() {
#[derive(Clone, Default)]
struct TimeoutCount(usize);
struct TimeoutTerminator;
impl AgentHook for TimeoutTerminator {
async fn on_tool_result(
&self,
ctx: &HookContext,
event: ToolResultEvent<'_>,
) -> ToolResultAction {
if let ToolResultEvent { raw_result, .. } = event
&& raw_result.is_error_kind(ToolErrorKind::Timeout)
{
let count = ctx.scratchpad().update(|c: &mut TimeoutCount| {
c.0 += 1;
c.0
});
if count >= 2 {
return ToolResultAction::stop("aborting after repeated tool timeouts");
}
}
ToolResultAction::keep()
}
}
let observer = OutcomeHook::default();
let err = AgentBuilder::new(MockCompletionModel::from_turns([
MockTurn::tool_call("tc1", "flaky_tool", json!({})),
MockTurn::tool_call("tc2", "flaky_tool", json!({})),
MockTurn::text("unreachable"),
]))
.tool(MockFailingTool::new(ToolErrorKind::Timeout))
.add_hook(observer.clone())
.add_hook(TimeoutTerminator)
.build()
.runner("go")
.max_turns(5)
.run()
.await
.expect_err("the run must terminate after two timeouts");
assert!(
err.to_string()
.contains("aborting after repeated tool timeouts"),
"unexpected error: {err}"
);
assert_eq!(
observer.outcomes(),
vec!["error:timeout".to_string(), "error:timeout".to_string()],
"both timeout outcomes must be observed before termination"
);
}
#[tokio::test]
async fn not_found_outcome_is_structured_and_non_fatal() {
let hook = OutcomeHook::default();
let status: Arc<Mutex<Option<u16>>> = Arc::new(Mutex::new(None));
struct StatusProbe(Arc<Mutex<Option<u16>>>);
impl AgentHook for StatusProbe {
async fn on_tool_result(
&self,
_ctx: &HookContext,
event: ToolResultEvent<'_>,
) -> ToolResultAction {
if let Some(error) = event.raw_result.error() {
*self.0.lock().expect("status") = error.http_status();
}
ToolResultAction::keep()
}
}
AgentBuilder::new(model_one_tool_then_text("flaky_tool"))
.tool(MockFailingTool::new(ToolErrorKind::NotFound))
.add_hook(hook.clone())
.add_hook(StatusProbe(status.clone()))
.build()
.runner("go")
.max_turns(3)
.run()
.await
.expect("a 404 must not terminate the run by default");
assert_eq!(hook.outcomes(), vec!["error:not_found".to_string()]);
assert_eq!(
*status.lock().expect("status"),
Some(404),
"the structured failure must carry the HTTP status"
);
}
#[tokio::test]
async fn handled_failure_delivers_model_output_and_error_outcome() {
let hook = OutcomeHook::default();
AgentBuilder::new(model_one_tool_then_text("lookup"))
.tool(MockHandledFailureTool)
.add_hook(hook.clone())
.build()
.runner("go")
.max_turns(3)
.run()
.await
.expect("a handled failure is not fatal");
assert_eq!(hook.outcomes(), vec!["error:not_found".to_string()]);
assert_eq!(
hook.results(),
vec!["no record found for id 42; try a different id".to_string()],
"the tool's model-visible output must survive alongside the error outcome"
);
}
#[tokio::test]
async fn flow_skip_produces_skipped_outcome() {
struct SkipHook;
impl AgentHook for SkipHook {
async fn on_tool_call(
&self,
_ctx: &HookContext,
event: ToolCall<'_>,
) -> ToolCallAction {
if let ToolCall { .. } = event {
ToolCallAction::skip("not executed (denied by policy); do not retry")
} else {
ToolCallAction::run()
}
}
}
let observer = OutcomeHook::default();
AgentBuilder::new(model_one_tool_then_text("flaky_tool"))
.tool(MockFailingTool::new(ToolErrorKind::Timeout))
.add_hook(SkipHook)
.add_hook(observer.clone())
.build()
.runner("go")
.max_turns(3)
.run()
.await
.expect("run should succeed after skipping the tool");
assert_eq!(observer.outcomes(), vec!["skipped".to_string()]);
assert_eq!(
observer.results(),
vec!["not executed (denied by policy); do not retry".to_string()]
);
}
#[tokio::test]
async fn tool_authored_denial_produces_denied_outcome() {
let hook = OutcomeHook::default();
AgentBuilder::new(model_one_tool_then_text("guarded"))
.tool(MockDeniedTool)
.add_hook(hook.clone())
.build()
.runner("go")
.max_turns(3)
.run()
.await
.expect("a tool-authored denial is not fatal");
assert_eq!(hook.outcomes(), vec!["denied".to_string()]);
assert_eq!(
hook.results(),
vec!["access to this resource is not permitted".to_string()],
"the model still receives the tool's denial message"
);
}
#[tokio::test]
async fn permission_denied_failure_is_not_a_tool_refusal() {
let hook = OutcomeHook::default();
AgentBuilder::new(model_one_tool_then_text("flaky_tool"))
.tool(MockFailingTool::new(ToolErrorKind::PermissionDenied))
.add_hook(hook.clone())
.build()
.runner("go")
.max_turns(3)
.run()
.await
.expect("a permission failure is model-visible feedback, not fatal");
assert_eq!(hook.outcomes(), vec!["error:permission_denied".to_string()]);
assert_eq!(hook.results(), vec!["mock tool call failed".to_string()]);
}
#[tokio::test]
async fn rewrite_args_then_skip_reports_rewritten_args() {
struct RewriteHook;
impl AgentHook for RewriteHook {
async fn on_tool_call(
&self,
_ctx: &HookContext,
event: ToolCall<'_>,
) -> ToolCallAction {
if let ToolCall { .. } = event {
ToolCallAction::rewrite(json!({ "x": 41, "y": 1 }))
} else {
ToolCallAction::run()
}
}
}
struct SkipHook;
impl AgentHook for SkipHook {
async fn on_tool_call(
&self,
_ctx: &HookContext,
event: ToolCall<'_>,
) -> ToolCallAction {
if let ToolCall { .. } = event {
ToolCallAction::skip("denied after rewrite")
} else {
ToolCallAction::run()
}
}
}
#[derive(Clone, Default)]
struct ArgsProbe {
args: Arc<Mutex<Option<String>>>,
outcome: Arc<Mutex<Option<String>>>,
}
impl AgentHook for ArgsProbe {
async fn on_tool_result(
&self,
_ctx: &HookContext,
event: ToolResultEvent<'_>,
) -> ToolResultAction {
if let ToolResultEvent {
args, raw_result, ..
} = event
{
*self.args.lock().expect("args") = Some(args.to_string());
*self.outcome.lock().expect("outcome") = Some(outcome_label(raw_result));
}
ToolResultAction::keep()
}
}
async fn run_surface(streaming: bool) -> (String, String) {
let probe = ArgsProbe::default();
if streaming {
let mut stream = AgentBuilder::new(stream_model_one_tool_then_text("add"))
.tool(MockAddTool)
.add_hook(RewriteHook)
.add_hook(SkipHook)
.add_hook(probe.clone())
.build()
.runner("go")
.max_turns(3)
.stream()
.await;
while let Some(item) = stream.next().await {
if let Err(err) = item {
panic!("stream item errored: {err}");
}
}
} else {
AgentBuilder::new(model_one_tool_then_text("add"))
.tool(MockAddTool)
.add_hook(RewriteHook)
.add_hook(SkipHook)
.add_hook(probe.clone())
.build()
.runner("go")
.max_turns(3)
.run()
.await
.expect("run should succeed after skipping the tool");
}
let args = probe.args.lock().expect("args").clone().expect("args seen");
let outcome = probe
.outcome
.lock()
.expect("outcome")
.clone()
.expect("outcome seen");
(args, outcome)
}
for streaming in [false, true] {
let (args, outcome) = run_surface(streaming).await;
assert_eq!(
outcome, "skipped",
"the skipped tool must produce a Skipped outcome (streaming={streaming})"
);
let parsed: serde_json::Value =
serde_json::from_str(&args).expect("ToolResult args are valid JSON");
assert_eq!(
parsed,
json!({ "x": 41, "y": 1 }),
"the skipped ToolResult must report the rewritten args, not the model's \
original {{}} (streaming={streaming}); got {args}"
);
}
}
#[tokio::test]
async fn nested_hook_stack_rewrite_then_skip_reports_rewritten_args() {
struct RewriteHook;
impl AgentHook for RewriteHook {
async fn on_tool_call(
&self,
_ctx: &HookContext,
event: ToolCall<'_>,
) -> ToolCallAction {
if let ToolCall { .. } = event {
ToolCallAction::rewrite(json!({ "x": 41, "y": 1 }))
} else {
ToolCallAction::run()
}
}
}
struct SkipHook;
impl AgentHook for SkipHook {
async fn on_tool_call(
&self,
_ctx: &HookContext,
event: ToolCall<'_>,
) -> ToolCallAction {
if let ToolCall { .. } = event {
ToolCallAction::skip("denied after nested rewrite")
} else {
ToolCallAction::run()
}
}
}
#[derive(Clone, Default)]
struct ArgsProbe {
args: Arc<Mutex<Option<String>>>,
outcome: Arc<Mutex<Option<String>>>,
}
impl AgentHook for ArgsProbe {
async fn on_tool_result(
&self,
_ctx: &HookContext,
event: ToolResultEvent<'_>,
) -> ToolResultAction {
if let ToolResultEvent {
args, raw_result, ..
} = event
{
*self.args.lock().expect("args") = Some(args.to_string());
*self.outcome.lock().expect("outcome") = Some(outcome_label(raw_result));
}
ToolResultAction::keep()
}
}
fn nested_stack() -> HookStack {
let mut nested = HookStack::new();
nested.push(RewriteHook);
nested.push(SkipHook);
nested
}
for streaming in [false, true] {
let probe = ArgsProbe::default();
if streaming {
let mut stream = AgentBuilder::new(stream_model_one_tool_then_text("add"))
.tool(MockAddTool)
.add_hook(nested_stack())
.add_hook(probe.clone())
.build()
.runner("go")
.max_turns(3)
.stream()
.await;
while let Some(item) = stream.next().await {
if let Err(err) = item {
panic!("stream item errored: {err}");
}
}
} else {
AgentBuilder::new(model_one_tool_then_text("add"))
.tool(MockAddTool)
.add_hook(nested_stack())
.add_hook(probe.clone())
.build()
.runner("go")
.max_turns(3)
.run()
.await
.expect("run should succeed after the nested stack skips the tool");
}
assert_eq!(
probe.outcome.lock().expect("outcome").clone(),
Some("skipped".to_string()),
"streaming={streaming}"
);
let args = probe.args.lock().expect("args").clone().expect("args seen");
let parsed: serde_json::Value =
serde_json::from_str(&args).expect("valid JSON args");
assert_eq!(
parsed,
json!({ "x": 41, "y": 1 }),
"the nested stack's rewrite must survive its skip and reach the ToolResult \
(streaming={streaming}); got {args}"
);
}
}
#[tokio::test]
async fn invalid_args_are_classified_as_invalid_args() {
let hook = OutcomeHook::default();
AgentBuilder::new(MockCompletionModel::from_turns([
MockTurn::tool_call("tc1", "add", json!({ "x": "not-a-number", "y": 1 })),
MockTurn::text("done"),
]))
.tool(MockAddTool)
.add_hook(hook.clone())
.build()
.runner("go")
.max_turns(3)
.run()
.await
.expect("an invalid-args failure is model-visible feedback, not fatal");
assert_eq!(hook.outcomes(), vec!["error:invalid_args".to_string()]);
}
#[tokio::test]
async fn success_result_metadata_reaches_hook_but_not_model() {
struct MetadataProbe {
seen: Arc<Mutex<Option<String>>>,
model_output: Arc<Mutex<Option<String>>>,
}
impl AgentHook for MetadataProbe {
async fn on_tool_result(
&self,
_ctx: &HookContext,
event: ToolResultEvent<'_>,
) -> ToolResultAction {
if let ToolResultEvent {
presentation,
tool_context,
..
} = event
{
*self.seen.lock().expect("seen") = tool_context
.result::<MockRequestId>()
.map(|id| id.0.clone());
*self.model_output.lock().expect("model_output") =
Some(presentation.render());
}
ToolResultAction::keep()
}
}
async fn run_surface(streaming: bool) -> (Option<String>, String) {
let seen: Arc<Mutex<Option<String>>> = Arc::new(Mutex::new(None));
let model_output: Arc<Mutex<Option<String>>> = Arc::new(Mutex::new(None));
let probe = MetadataProbe {
seen: seen.clone(),
model_output: model_output.clone(),
};
if streaming {
let mut stream =
AgentBuilder::new(stream_model_one_tool_then_text("with_meta"))
.tool(MockMetadataTool)
.add_hook(probe)
.build()
.runner("go")
.max_turns(3)
.stream()
.await;
while let Some(item) = stream.next().await {
if let Err(error) = item {
panic!("stream item errored: {error}");
}
}
} else {
AgentBuilder::new(model_one_tool_then_text("with_meta"))
.tool(MockMetadataTool)
.add_hook(probe)
.build()
.runner("go")
.max_turns(3)
.run()
.await
.expect("run should succeed");
}
let seen_value = seen.lock().expect("seen").clone();
let output = model_output
.lock()
.expect("model_output")
.clone()
.expect("output");
(seen_value, output)
}
for streaming in [false, true] {
let (seen, output) = run_surface(streaming).await;
assert_eq!(
seen,
Some("req-7".to_string()),
"the tool's result metadata must reach the hook (streaming={streaming})"
);
assert_eq!(output, "done");
assert!(
!output.contains("req-7"),
"result metadata must never leak into model output (streaming={streaming})"
);
}
}
#[tokio::test]
async fn rewrite_result_does_not_mask_the_structured_outcome() {
struct Redact;
impl AgentHook for Redact {
async fn on_tool_result(
&self,
_ctx: &HookContext,
event: ToolResultEvent<'_>,
) -> ToolResultAction {
if let ToolResultEvent { .. } = event {
ToolResultAction::rewrite("[REDACTED]")
} else {
ToolResultAction::keep()
}
}
}
let observer = OutcomeHook::default();
AgentBuilder::new(model_one_tool_then_text("flaky_tool"))
.tool(MockFailingTool::new(ToolErrorKind::NotFound))
.add_hook(Redact)
.add_hook(observer.clone())
.build()
.runner("go")
.max_turns(3)
.run()
.await
.expect("run should succeed");
assert_eq!(observer.outcomes(), vec!["error:not_found".to_string()]);
assert_eq!(observer.results(), vec!["[REDACTED]".to_string()]);
}
#[tokio::test]
async fn streaming_and_blocking_outcomes_match() {
let blocking = OutcomeHook::default();
AgentBuilder::new(model_one_tool_then_text("flaky_tool"))
.tool(MockFailingTool::new(ToolErrorKind::Timeout))
.add_hook(blocking.clone())
.build()
.runner("go")
.max_turns(3)
.run()
.await
.expect("blocking run should succeed");
let streaming = OutcomeHook::default();
let mut stream = AgentBuilder::new(stream_model_one_tool_then_text("flaky_tool"))
.tool(MockFailingTool::new(ToolErrorKind::Timeout))
.add_hook(streaming.clone())
.build()
.runner("go")
.max_turns(3)
.stream()
.await;
while let Some(item) = stream.next().await {
if let Err(err) = item {
panic!("stream item errored: {err}");
}
}
assert_eq!(blocking.outcomes(), vec!["error:timeout".to_string()]);
assert_eq!(blocking.outcomes(), streaming.outcomes());
assert_eq!(blocking.results(), streaming.results());
}
#[tokio::test]
async fn concurrent_tools_preserve_order_and_both_outcomes() {
use rig_core::message::{
AssistantContent, ToolCall as MessageToolCall, ToolFunction, UserContent,
};
let turn = MockTurn::from_contents([
AssistantContent::ToolCall(MessageToolCall::new(
"tc_add".to_string(),
ToolFunction::new("add".to_string(), json!({ "x": 2, "y": 3 })),
)),
AssistantContent::ToolCall(MessageToolCall::new(
"tc_flaky".to_string(),
ToolFunction::new("flaky_tool".to_string(), json!({})),
)),
])
.expect("two tool calls");
let observer = OutcomeHook::default();
let response = AgentBuilder::new(MockCompletionModel::from_turns([
turn,
MockTurn::text("done"),
]))
.tool(MockAddTool)
.tool(MockFailingTool::new(ToolErrorKind::Timeout))
.add_hook(observer.clone())
.build()
.runner("go")
.max_turns(3)
.tool_concurrency(2)
.run()
.await
.expect("run should succeed");
let mut outcomes = observer.outcomes();
outcomes.sort();
assert_eq!(
outcomes,
vec!["error:timeout".to_string(), "success".to_string()]
);
let messages = response.messages.expect("messages");
let tool_result_ids: Vec<String> = messages
.iter()
.flat_map(|message| match message {
crate::completion::Message::User { content } => content
.iter()
.filter_map(|c| match c {
UserContent::ToolResult(result) => Some(result.id.clone()),
_ => None,
})
.collect::<Vec<_>>(),
_ => Vec::new(),
})
.collect();
assert_eq!(
tool_result_ids,
vec!["tc_add".to_string(), "tc_flaky".to_string()],
"tool results must be persisted in call order"
);
}
}
mod span_safety_net {
use std::collections::{HashMap, HashSet};
use std::sync::{Arc, Mutex};
use futures::StreamExt;
use tracing::Instrument;
use tracing::field::{Field, Visit};
use tracing::span::{Attributes, Record};
use tracing::{Id, Subscriber};
use tracing_subscriber::layer::{Context, SubscriberExt};
use tracing_subscriber::{Layer, Registry, registry::LookupSpan};
use crate::agent::{
AgentBuilder, HookContext, MultiTurnStreamItem, ToolResultAction, ToolResultEvent,
};
use crate::completion::{
CompletionError, CompletionModel, CompletionRequest, CompletionResponse, Prompt,
PromptError, Usage,
};
use crate::streaming::StreamedAssistantContent;
use crate::streaming::StreamingCompletionResponse;
use crate::test_utils::{
MockAddTool, MockCompletionModel, MockResponse, MockStreamEvent, MockTurn,
};
use crate::tool::{ToolContext, ToolExecutionError};
use rig_core::telemetry::{CompletionOperation, CompletionSpanBuilder};
use super::{BoundedResponseRetry, StopCompletedModelTurn, TestRetryMode};
#[derive(Clone)]
struct CapturedSpan {
id: u64,
name: String,
target: String,
field_names: HashSet<String>,
u64_fields: HashMap<String, u64>,
string_fields: HashMap<String, Vec<String>>,
}
#[derive(Clone, Default)]
struct Captured {
spans: Arc<Mutex<Vec<CapturedSpan>>>,
follows: Arc<Mutex<Vec<(u64, u64)>>>,
}
impl Captured {
fn insert(&self, id: &Id, name: &str, target: &str) {
self.spans.lock().expect("spans").push(CapturedSpan {
id: id.into_u64(),
name: name.to_string(),
target: target.to_string(),
field_names: HashSet::new(),
u64_fields: HashMap::new(),
string_fields: HashMap::new(),
});
}
fn record(
&self,
id: &Id,
names: HashSet<String>,
u64s: HashMap<String, u64>,
strings: HashMap<String, String>,
) {
let id = id.into_u64();
if let Ok(mut spans) = self.spans.lock()
&& let Some(span) = spans.iter_mut().find(|s| s.id == id)
{
span.field_names.extend(names);
span.u64_fields.extend(u64s);
for (name, value) in strings {
span.string_fields.entry(name).or_default().push(value);
}
}
}
fn follows_from(&self, span: &Id, follows: &Id) {
self.follows
.lock()
.expect("follows")
.push((span.into_u64(), follows.into_u64()));
}
fn clear(&self) {
self.spans.lock().expect("spans").clear();
self.follows.lock().expect("follows").clear();
}
fn snapshot(&self) -> Vec<CapturedSpan> {
self.spans.lock().expect("spans").clone()
}
fn follows_edges(&self) -> Vec<(u64, u64)> {
self.follows.lock().expect("follows").clone()
}
}
struct CaptureLayer {
captured: Captured,
}
impl<S> Layer<S> for CaptureLayer
where
S: Subscriber + for<'l> LookupSpan<'l>,
{
fn on_new_span(&self, attrs: &Attributes<'_>, id: &Id, _ctx: Context<'_, S>) {
self.captured
.insert(id, attrs.metadata().name(), attrs.metadata().target());
}
fn on_record(&self, span: &Id, values: &Record<'_>, _ctx: Context<'_, S>) {
let mut visitor = FieldVisitor::default();
values.record(&mut visitor);
self.captured
.record(span, visitor.names, visitor.u64s, visitor.strings);
}
fn on_follows_from(&self, span: &Id, follows: &Id, _ctx: Context<'_, S>) {
self.captured.follows_from(span, follows);
}
}
#[derive(Default)]
struct FieldVisitor {
names: HashSet<String>,
u64s: HashMap<String, u64>,
strings: HashMap<String, String>,
}
impl Visit for FieldVisitor {
fn record_u64(&mut self, field: &Field, value: u64) {
self.names.insert(field.name().to_string());
self.u64s.insert(field.name().to_string(), value);
}
fn record_str(&mut self, field: &Field, value: &str) {
self.names.insert(field.name().to_string());
self.strings
.insert(field.name().to_string(), value.to_string());
}
fn record_debug(&mut self, field: &Field, value: &dyn std::fmt::Debug) {
self.names.insert(field.name().to_string());
self.strings
.insert(field.name().to_string(), format!("{value:?}"));
}
}
fn usage(input: u64, output: u64) -> Usage {
Usage {
input_tokens: input,
output_tokens: output,
..Usage::new()
}
}
fn tool_then_text_model() -> MockCompletionModel {
MockCompletionModel::from_turns([
MockTurn::tool_call("tc1", "add", serde_json::json!({"x": 2, "y": 3}))
.with_usage(usage(7, 11)),
MockTurn::text("the answer is 5").with_usage(usage(13, 17)),
])
}
#[derive(Clone)]
struct CompletionTelemetryModel {
inner: MockCompletionModel,
}
impl CompletionModel for CompletionTelemetryModel {
type Response = MockResponse;
type StreamingResponse = MockResponse;
type Client = ();
fn make(_client: &Self::Client, _model: impl Into<String>) -> Self {
Self {
inner: MockCompletionModel::default(),
}
}
async fn completion(
&self,
request: CompletionRequest,
) -> Result<CompletionResponse<Self::Response>, CompletionError> {
let span = CompletionSpanBuilder::new(
"fixture-provider",
"fixture-model",
CompletionOperation::Chat,
)
.build();
self.inner.completion(request).instrument(span).await
}
async fn stream(
&self,
request: CompletionRequest,
) -> Result<StreamingCompletionResponse<Self::StreamingResponse>, CompletionError>
{
let span = CompletionSpanBuilder::new(
"fixture-provider",
"fixture-model",
CompletionOperation::ChatStreaming,
)
.build();
self.inner.stream(request).instrument(span).await
}
}
async fn warm_blocking_callsites() {
let agent = AgentBuilder::new(tool_then_text_model())
.record_content_telemetry(true)
.tool(MockAddTool)
.build();
let _ = agent.runner("add 2 and 3").max_turns(3).run().await;
}
async fn run_blocking_response_retry_with_content_telemetry() {
AgentBuilder::new(MockCompletionModel::from_turns([
MockTurn::text("rejected"),
MockTurn::text("accepted"),
]))
.record_content_telemetry(true)
.add_hook(BoundedResponseRetry::new(
"rejected",
1,
TestRetryMode::Repeat,
))
.build()
.runner("question")
.max_turns(2)
.run()
.await
.expect("blocking retry should succeed");
}
async fn run_streaming_response_retry_with_content_telemetry() {
let mut stream = AgentBuilder::new(MockCompletionModel::from_stream_turns([
[
MockStreamEvent::text("rejected"),
MockStreamEvent::final_response_with_default_usage(),
],
[
MockStreamEvent::text("accepted"),
MockStreamEvent::final_response_with_default_usage(),
],
]))
.record_content_telemetry(true)
.add_hook(BoundedResponseRetry::new(
"rejected",
1,
TestRetryMode::Repeat,
))
.build()
.runner("question")
.max_turns(2)
.stream()
.await;
let mut saw_final = false;
while let Some(item) = stream.next().await {
if let MultiTurnStreamItem::FinalResponse(response) =
item.expect("streaming retry item")
{
saw_final = true;
assert_eq!(response.output, "accepted");
}
}
assert!(saw_final, "streaming retry should produce a final response");
}
async fn run_blocking_model_turn_stop_with_content_telemetry() {
let error = AgentBuilder::new(MockCompletionModel::from_turns([MockTurn::text(
"stopped blocking response",
)]))
.record_content_telemetry(true)
.add_hook(StopCompletedModelTurn)
.build()
.runner("question")
.run()
.await
.expect_err("blocking model-turn stop should cancel the run");
assert!(matches!(
error,
PromptError::PromptCancelled { reason, .. }
if reason == "stop completed model turn"
));
}
async fn run_streaming_model_turn_stop_with_content_telemetry() {
let mut stream = AgentBuilder::new(MockCompletionModel::from_stream_turns([[
MockStreamEvent::text("stopped streaming response"),
MockStreamEvent::final_response_with_default_usage(),
]]))
.record_content_telemetry(true)
.add_hook(StopCompletedModelTurn)
.build()
.runner("question")
.stream()
.await;
let mut provider_finals = 0;
let mut agent_finals = 0;
let mut retries = 0;
let mut errors = 0;
while let Some(item) = stream.next().await {
match item {
Ok(MultiTurnStreamItem::StreamAssistantItem(
StreamedAssistantContent::Final(_),
)) => provider_finals += 1,
Ok(MultiTurnStreamItem::FinalResponse(_)) => agent_finals += 1,
Ok(MultiTurnStreamItem::ModelTurnRetried { .. }) => retries += 1,
Ok(_) => {}
Err(error) => {
errors += 1;
assert!(matches!(
error,
super::StreamingError::Prompt(error)
if matches!(
error.as_ref(),
PromptError::PromptCancelled { reason, .. }
if reason == "stop completed model turn"
)
));
}
}
}
assert_eq!(provider_finals, 1);
assert_eq!(agent_finals, 0);
assert_eq!(retries, 0);
assert_eq!(errors, 1);
}
#[test]
fn chat_span_declares_the_full_completion_parent_contract() {
use rig_core::telemetry::{
COMPLETION_PARENT_MARKER_FIELD, COMPLETION_PARENT_REQUIRED_FIELDS,
};
let _isolation = crate::test_utils::scoped_tracing_subscriber_guard_blocking();
tracing::subscriber::with_default(Registry::default(), || {
let agent = AgentBuilder::new(MockCompletionModel::text("done"))
.name("contract-agent")
.build();
let runner = agent.runner("hello");
let span = build_chat_span!(runner, None, "chat", "chat");
let Some(metadata) = span.metadata() else {
panic!("chat span was disabled");
};
let declared: HashSet<&str> =
metadata.fields().iter().map(|field| field.name()).collect();
let expected: HashSet<&str> = COMPLETION_PARENT_REQUIRED_FIELDS
.iter()
.copied()
.chain([COMPLETION_PARENT_MARKER_FIELD, "gen_ai.agent.name"])
.collect();
assert_eq!(declared, expected);
assert_eq!(metadata.fields().len(), expected.len());
});
}
#[tokio::test]
async fn response_retry_records_only_accepted_content_on_both_surfaces() {
let _isolation = crate::test_utils::scoped_tracing_subscriber_guard().await;
let captured = Captured::default();
let subscriber = Registry::default().with(CaptureLayer {
captured: captured.clone(),
});
let _default = tracing::subscriber::set_default(subscriber);
run_blocking_response_retry_with_content_telemetry().await;
run_streaming_response_retry_with_content_telemetry().await;
tracing::callsite::rebuild_interest_cache();
captured.clear();
run_blocking_response_retry_with_content_telemetry().await;
let blocking = captured.snapshot();
let blocking_chats = blocking
.iter()
.filter(|span| span.name == "chat")
.collect::<Vec<_>>();
assert_eq!(blocking_chats.len(), 2);
assert!(
blocking_chats
.iter()
.all(|span| span.target == "rig::agent_chat")
);
assert!(
!blocking_chats[0]
.field_names
.contains("gen_ai.output.messages"),
"rejected blocking content must not be recorded as model output"
);
assert!(
blocking_chats[1]
.field_names
.contains("gen_ai.output.messages"),
"accepted blocking content must be recorded as model output"
);
let blocking_output = blocking_chats[1]
.string_fields
.get("gen_ai.output.messages")
.expect("accepted blocking output value");
assert!(
blocking_output
.iter()
.any(|value| value.contains("accepted"))
);
assert!(
blocking_output
.iter()
.all(|value| !value.contains("rejected"))
);
let blocking_completion = blocking
.iter()
.find(|span| span.name == "invoke_agent")
.and_then(|span| span.string_fields.get("gen_ai.completion"))
.expect("accepted blocking run-level completion");
assert_eq!(blocking_completion, &["accepted"]);
captured.clear();
run_streaming_response_retry_with_content_telemetry().await;
let streaming = captured.snapshot();
let streaming_chats = streaming
.iter()
.filter(|span| span.name == "chat_streaming")
.collect::<Vec<_>>();
assert_eq!(streaming_chats.len(), 2);
assert!(
streaming_chats
.iter()
.all(|span| span.target == "rig::agent_chat")
);
assert!(
!streaming_chats[0]
.field_names
.contains("gen_ai.output.messages"),
"rejected streaming content must not be recorded as model output"
);
assert!(
streaming_chats[1]
.field_names
.contains("gen_ai.output.messages"),
"accepted streaming content must be recorded as model output"
);
let streaming_output = streaming_chats[1]
.string_fields
.get("gen_ai.output.messages")
.expect("accepted streaming output value");
assert!(
streaming_output
.iter()
.any(|value| value.contains("accepted"))
);
assert!(
streaming_output
.iter()
.all(|value| !value.contains("rejected"))
);
let streaming_completion = streaming
.iter()
.find(|span| span.name == "invoke_agent")
.and_then(|span| span.string_fields.get("gen_ai.completion"))
.expect("accepted streaming run-level completion");
assert_eq!(streaming_completion, &["accepted"]);
}
#[tokio::test]
async fn model_turn_stop_preserves_completed_content_telemetry_on_both_surfaces() {
let _isolation = crate::test_utils::scoped_tracing_subscriber_guard().await;
let captured = Captured::default();
let subscriber = Registry::default().with(CaptureLayer {
captured: captured.clone(),
});
let _default = tracing::subscriber::set_default(subscriber);
run_blocking_model_turn_stop_with_content_telemetry().await;
run_streaming_model_turn_stop_with_content_telemetry().await;
tracing::callsite::rebuild_interest_cache();
captured.clear();
run_blocking_model_turn_stop_with_content_telemetry().await;
let blocking = captured.snapshot();
let blocking_output = blocking
.iter()
.find(|span| span.name == "chat")
.and_then(|span| span.string_fields.get("gen_ai.output.messages"))
.expect("stopped blocking turn should retain output telemetry");
assert!(
blocking_output
.iter()
.any(|value| value.contains("stopped blocking response"))
);
captured.clear();
run_streaming_model_turn_stop_with_content_telemetry().await;
let streaming = captured.snapshot();
let streaming_output = streaming
.iter()
.find(|span| span.name == "chat_streaming")
.and_then(|span| span.string_fields.get("gen_ai.output.messages"))
.expect("stopped streaming turn should retain output telemetry");
assert!(
streaming_output
.iter()
.any(|value| value.contains("stopped streaming response"))
);
let streaming_completion = streaming
.iter()
.find(|span| span.name == "invoke_agent")
.and_then(|span| span.string_fields.get("gen_ai.completion"))
.expect("stopped streaming turn should retain run-level completion telemetry");
assert_eq!(streaming_completion, &["stopped streaming response"]);
}
#[tokio::test]
async fn run_records_usage_and_chains_chat_spans_on_a_created_agent_span() {
let _isolation = crate::test_utils::scoped_tracing_subscriber_guard().await;
let captured = Captured::default();
let subscriber = Registry::default().with(CaptureLayer {
captured: captured.clone(),
});
let _default = tracing::subscriber::set_default(subscriber);
warm_blocking_callsites().await;
tracing::callsite::rebuild_interest_cache();
captured.clear();
let agent = AgentBuilder::new(tool_then_text_model())
.record_content_telemetry(true)
.tool(MockAddTool)
.build();
let response = agent
.runner("add 2 and 3")
.max_turns(3)
.run()
.await
.expect("blocking run should succeed");
assert_eq!(response.output, "the answer is 5");
let spans = captured.snapshot();
let chat_spans: Vec<&CapturedSpan> =
spans.iter().filter(|s| s.name == "chat").collect();
assert_eq!(chat_spans.len(), 2, "two model turns -> two chat spans");
assert!(
spans.iter().all(|s| s.name != "chat_streaming"),
"blocking driver must not emit chat_streaming spans"
);
let agent_span = spans
.iter()
.find(|s| s.name == "invoke_agent")
.expect("blocking run should create an invoke_agent span");
assert_eq!(
agent_span.u64_fields.get("gen_ai.usage.input_tokens"),
Some(&(7 + 13)),
);
assert_eq!(
agent_span.u64_fields.get("gen_ai.usage.output_tokens"),
Some(&(11 + 17)),
);
assert!(
agent_span.field_names.contains("gen_ai.completion"),
"the created agent span records the final completion text"
);
let tool_span = spans
.iter()
.find(|s| s.name == "execute_tool")
.expect("tool turn should emit an execute_tool span");
let edges = captured.follows_edges();
assert!(
edges.contains(&(tool_span.id, chat_spans[0].id)),
"execute_tool should follow_from the first chat span; edges={edges:?}"
);
assert!(
edges.contains(&(chat_spans[1].id, tool_span.id)),
"the second chat span should follow_from execute_tool; edges={edges:?}"
);
}
#[tokio::test]
async fn classic_completion_parent_is_enriched_without_duplicate_provider_span() {
let _isolation = crate::test_utils::scoped_tracing_subscriber_guard().await;
let captured = Captured::default();
let subscriber = Registry::default().with(CaptureLayer {
captured: captured.clone(),
});
let _default = tracing::subscriber::set_default(subscriber);
let warm = AgentBuilder::new(CompletionTelemetryModel {
inner: MockCompletionModel::text("warm"),
})
.build();
let _ = warm.prompt("warm").await;
tracing::callsite::rebuild_interest_cache();
captured.clear();
let agent = AgentBuilder::new(CompletionTelemetryModel {
inner: MockCompletionModel::text("done"),
})
.build();
let response = agent.prompt("hello").await.expect("prompt should succeed");
assert_eq!(response, "done");
let spans = captured.snapshot();
let chat_spans = spans
.iter()
.filter(|span| span.name == "chat")
.collect::<Vec<_>>();
assert_eq!(chat_spans.len(), 1, "provider telemetry must reuse chat");
assert_eq!(chat_spans[0].target, "rig::agent_chat");
assert!(
spans.iter().all(|span| span.target != "rig::completions"),
"an adopted classic completion parent must not gain a provider child"
);
assert_eq!(
chat_spans[0]
.string_fields
.get("gen_ai.provider.name")
.and_then(|values| values.first())
.map(String::as_str),
Some("fixture-provider")
);
assert_eq!(
chat_spans[0]
.string_fields
.get("gen_ai.request.model")
.and_then(|values| values.first())
.map(String::as_str),
Some("fixture-model")
);
}
#[tokio::test]
async fn run_does_not_record_usage_onto_a_caller_supplied_outer_span() {
let _isolation = crate::test_utils::scoped_tracing_subscriber_guard().await;
let captured = Captured::default();
let subscriber = Registry::default().with(CaptureLayer {
captured: captured.clone(),
});
let _default = tracing::subscriber::set_default(subscriber);
warm_blocking_callsites().await;
tracing::callsite::rebuild_interest_cache();
captured.clear();
let outer = tracing::info_span!(
"outer",
gen_ai.completion = tracing::field::Empty,
gen_ai.usage.input_tokens = tracing::field::Empty,
gen_ai.usage.output_tokens = tracing::field::Empty,
);
async {
let agent = AgentBuilder::new(tool_then_text_model())
.tool(MockAddTool)
.build();
agent
.runner("add 2 and 3")
.max_turns(3)
.run()
.await
.expect("blocking run should succeed");
}
.instrument(outer)
.await;
let spans = captured.snapshot();
assert!(
spans.iter().all(|s| s.name != "invoke_agent"),
"an ambient outer span should be adopted, not wrapped in invoke_agent"
);
let outer_span = spans
.iter()
.find(|s| s.name == "outer")
.expect("outer span should be captured");
assert!(
outer_span
.field_names
.iter()
.all(|name| !name.starts_with("gen_ai.usage.")),
"run-level usage must not be recorded onto a caller-supplied outer span"
);
assert!(
!outer_span.field_names.contains("gen_ai.completion"),
"run-level completion must not be recorded onto a caller-supplied outer span"
);
}
struct RawOutputTool;
impl crate::tool::Tool for RawOutputTool {
const NAME: &'static str = "raw_output";
type Error = rig::tool::ToolExecutionError;
type Args = serde_json::Value;
type Output = String;
fn description(&self) -> String {
"returns a raw output marker".to_string()
}
fn parameters(&self) -> serde_json::Value {
serde_json::json!({ "type": "object", "properties": {} })
}
async fn call(
&self,
_context: &mut ToolContext,
_args: Self::Args,
) -> Result<Self::Output, ToolExecutionError> {
Ok("RAW_EXECUTION_OUTPUT_42".to_string())
}
}
struct RedactResultHook;
impl crate::agent::AgentHook for RedactResultHook {
async fn on_tool_result(
&self,
_ctx: &HookContext,
event: ToolResultEvent<'_>,
) -> ToolResultAction {
if let crate::agent::ToolResultEvent { .. } = event {
crate::agent::ToolResultAction::rewrite("[REDACTED]")
} else {
crate::agent::ToolResultAction::keep()
}
}
}
struct StopOnResultHook;
impl crate::agent::AgentHook for StopOnResultHook {
async fn on_tool_result(
&self,
_ctx: &HookContext,
_event: ToolResultEvent<'_>,
) -> ToolResultAction {
ToolResultAction::stop("stop after raw result")
}
}
#[derive(Default)]
struct ResultValueVisitor {
values: Vec<String>,
}
impl Visit for ResultValueVisitor {
fn record_str(&mut self, field: &Field, value: &str) {
if field.name() == "gen_ai.tool.call.result" {
self.values.push(value.to_string());
}
}
fn record_debug(&mut self, field: &Field, value: &dyn std::fmt::Debug) {
if field.name() == "gen_ai.tool.call.result" {
self.values.push(format!("{value:?}"));
}
}
}
struct ResultValueLayer {
values: Arc<Mutex<Vec<String>>>,
}
impl<S> Layer<S> for ResultValueLayer
where
S: Subscriber + for<'l> LookupSpan<'l>,
{
fn on_record(&self, _span: &Id, values: &Record<'_>, _ctx: Context<'_, S>) {
let mut visitor = ResultValueVisitor::default();
values.record(&mut visitor);
if !visitor.values.is_empty() {
self.values.lock().expect("values").extend(visitor.values);
}
}
}
#[tokio::test]
async fn tool_result_rewrite_redacts_span_output() {
let _isolation = crate::test_utils::scoped_tracing_subscriber_guard().await;
let values: Arc<Mutex<Vec<String>>> = Arc::new(Mutex::new(Vec::new()));
let subscriber = Registry::default().with(ResultValueLayer {
values: values.clone(),
});
let _default = tracing::subscriber::set_default(subscriber);
warm_blocking_callsites().await;
tracing::callsite::rebuild_interest_cache();
values.lock().expect("values").clear();
let model = MockCompletionModel::from_turns([
MockTurn::tool_call("tc1", "raw_output", serde_json::json!({})),
MockTurn::text("ok"),
]);
let response = AgentBuilder::new(model)
.record_content_telemetry(true)
.tool(RawOutputTool)
.add_hook(RedactResultHook)
.build()
.runner("go")
.max_turns(3)
.run()
.await
.expect("run should succeed");
assert_eq!(response.output, "ok");
let captured = values.lock().expect("values").clone();
assert!(
captured.iter().any(|v| v.contains("[REDACTED]")),
"the rewritten presentation must reach telemetry; captured: {captured:?}"
);
assert!(
!captured
.iter()
.any(|v| v.contains("RAW_EXECUTION_OUTPUT_42")),
"the raw tool output must not leak through telemetry; captured: {captured:?}"
);
}
#[tokio::test]
async fn tool_result_stop_omits_span_output() {
let _isolation = crate::test_utils::scoped_tracing_subscriber_guard().await;
let values: Arc<Mutex<Vec<String>>> = Arc::new(Mutex::new(Vec::new()));
let subscriber = Registry::default().with(ResultValueLayer {
values: values.clone(),
});
let _default = tracing::subscriber::set_default(subscriber);
warm_blocking_callsites().await;
tracing::callsite::rebuild_interest_cache();
values.lock().expect("values").clear();
let result = AgentBuilder::new(MockCompletionModel::from_turns([MockTurn::tool_call(
"tc1",
"raw_output",
serde_json::json!({}),
)]))
.tool(RawOutputTool)
.add_hook(StopOnResultHook)
.build()
.runner("go")
.max_turns(2)
.run()
.await;
assert!(result.is_err(), "the result hook should stop the run");
let captured = values.lock().expect("values").clone();
assert!(
!captured
.iter()
.any(|value| value.contains("RAW_EXECUTION_OUTPUT_42")),
"a Stop must not leak raw execution telemetry; captured: {captured:?}"
);
}
}
fn tool_call_content(id: &str, args: serde_json::Value) -> AssistantContent {
AssistantContent::ToolCall(MessageToolCall::new(
id.to_string(),
ToolFunction::new("add".to_string(), args),
))
}
fn tool_result_text_in_history(messages: &[Message], expected: &str) -> bool {
messages.iter().any(|message| {
matches!(
message,
Message::User { content }
if content.iter().any(|item| matches!(
item,
UserContent::ToolResult(result)
if result.content.iter().any(|c| matches!(
c,
rig_core::message::ToolResultContent::Text(text)
if text.text == expected
))
))
)
})
}
fn tool_result_json_in_history(messages: &[Message], expected: &serde_json::Value) -> bool {
messages.iter().any(|message| {
matches!(
message,
Message::User { content }
if content.iter().any(|item| matches!(
item,
UserContent::ToolResult(result)
if result.content.iter().any(|content| matches!(
content,
rig_core::message::ToolResultContent::Json { value }
if value == expected
))
))
)
})
}
#[tokio::test]
async fn run_and_stream_same_message_history_for_parallel_tool_calls() {
let blocking_model = MockCompletionModel::from_turns([
MockTurn::from_contents([
tool_call_content("tc1", json!({"x": 2, "y": 3})),
tool_call_content("tc2", json!({"x": 10, "y": 20})),
])
.expect("two tool calls is a valid turn"),
MockTurn::text("done"),
]);
let blocking = AgentBuilder::new(blocking_model)
.tool(MockAddTool)
.build()
.runner("add two pairs")
.max_turns(3)
.tool_concurrency(4)
.run()
.await
.expect("blocking run should succeed");
let streaming_model = MockCompletionModel::from_stream_turns([
vec![
MockStreamEvent::tool_call("tc1", "add", json!({"x": 2, "y": 3})),
MockStreamEvent::tool_call("tc2", "add", json!({"x": 10, "y": 20})),
MockStreamEvent::final_response_with_total_tokens(0),
],
vec![
MockStreamEvent::text("done"),
MockStreamEvent::final_response_with_total_tokens(0),
],
]);
let mut stream = AgentBuilder::new(streaming_model)
.tool(MockAddTool)
.build()
.runner("add two pairs")
.max_turns(3)
.stream()
.await;
let mut final_response = None;
while let Some(item) = stream.next().await {
if let Ok(MultiTurnStreamItem::FinalResponse(resp)) =
item.map_err(|err| panic!("stream item errored: {err}"))
{
final_response = Some(resp);
}
}
let final_response = final_response.expect("stream should yield a final response");
let blocking_messages = blocking.messages.expect("blocking messages");
let streaming_messages = final_response
.messages()
.expect("streaming history")
.to_vec();
assert_eq!(
serde_json::to_value(&blocking_messages).expect("serialize blocking"),
serde_json::to_value(&streaming_messages).expect("serialize streaming"),
);
}
#[derive(Clone)]
struct OutOfOrderTool {
gate: Arc<tokio::sync::Notify>,
order: Arc<AtomicU32>,
}
impl Tool for OutOfOrderTool {
const NAME: &'static str = "add";
type Error = MockToolError;
type Args = MockOperationArgs;
type Output = i32;
fn description(&self) -> String {
MockAddTool.description()
}
fn parameters(&self) -> serde_json::Value {
MockAddTool.parameters()
}
async fn call(
&self,
_context: &mut ToolContext,
_args: Self::Args,
) -> Result<Self::Output, Self::Error> {
let nth = self.order.fetch_add(1, SeqCst);
if nth == 0 {
self.gate.notified().await;
} else {
self.gate.notify_one();
}
Ok(nth as i32)
}
}
#[tokio::test]
async fn run_preserves_tool_call_order_under_out_of_order_completion() {
let model = MockCompletionModel::from_turns([
MockTurn::from_contents([
tool_call_content("tc1", json!({"x": 1, "y": 0})),
tool_call_content("tc2", json!({"x": 2, "y": 0})),
])
.expect("two tool calls is a valid turn"),
MockTurn::text("done"),
]);
let response = AgentBuilder::new(model)
.tool(OutOfOrderTool {
gate: Arc::new(tokio::sync::Notify::new()),
order: Arc::new(AtomicU32::new(0)),
})
.build()
.runner("go")
.max_turns(3)
.tool_concurrency(4)
.run()
.await
.expect("run should succeed");
let messages = response.messages.expect("messages");
let result_ids: Vec<String> = messages
.iter()
.flat_map(|message| match message {
Message::User { content } => content
.iter()
.filter_map(|item| match item {
UserContent::ToolResult(result) => Some(result.id.clone()),
_ => None,
})
.collect::<Vec<_>>(),
_ => Vec::new(),
})
.collect();
assert_eq!(result_ids, vec!["tc1".to_string(), "tc2".to_string()]);
}
async fn drive_to_final_response<R: Send + 'static>(
mut stream: crate::agent::prompt_request::streaming::StreamingResult<R>,
) -> crate::agent::prompt_request::PromptResponse {
let mut final_response = None;
while let Some(item) = stream.next().await {
if let MultiTurnStreamItem::FinalResponse(resp) =
item.unwrap_or_else(|err| panic!("stream item errored: {err}"))
{
final_response = Some(resp);
}
}
final_response.expect("stream should yield a final response")
}
fn tool_result_ids(messages: &[Message]) -> Vec<String> {
messages
.iter()
.flat_map(|message| match message {
Message::User { content } => content
.iter()
.filter_map(|item| match item {
UserContent::ToolResult(result) => Some(result.id.clone()),
_ => None,
})
.collect::<Vec<_>>(),
_ => Vec::new(),
})
.collect()
}
#[tokio::test]
async fn stream_and_run_same_message_history_for_parallel_tool_calls_under_concurrency() {
let blocking_model = MockCompletionModel::from_turns([
MockTurn::from_contents([
tool_call_content("tc1", json!({"x": 2, "y": 3})),
tool_call_content("tc2", json!({"x": 10, "y": 20})),
])
.expect("two tool calls is a valid turn"),
MockTurn::text("done"),
]);
let blocking = AgentBuilder::new(blocking_model)
.tool(MockAddTool)
.build()
.runner("add two pairs")
.max_turns(3)
.tool_concurrency(4)
.run()
.await
.expect("blocking run should succeed");
let streaming_model = MockCompletionModel::from_stream_turns([
vec![
MockStreamEvent::tool_call("tc1", "add", json!({"x": 2, "y": 3})),
MockStreamEvent::tool_call("tc2", "add", json!({"x": 10, "y": 20})),
MockStreamEvent::final_response_with_total_tokens(0),
],
vec![
MockStreamEvent::text("done"),
MockStreamEvent::final_response_with_total_tokens(0),
],
]);
let stream = AgentBuilder::new(streaming_model)
.tool(MockAddTool)
.build()
.runner("add two pairs")
.max_turns(3)
.tool_concurrency(4)
.stream()
.await;
let final_response = drive_to_final_response(stream).await;
let blocking_messages = blocking.messages.expect("blocking messages");
let streaming_messages = final_response
.messages()
.expect("streaming history")
.to_vec();
assert_eq!(
serde_json::to_value(&blocking_messages).expect("serialize blocking"),
serde_json::to_value(&streaming_messages).expect("serialize streaming"),
);
}
#[tokio::test]
async fn stream_preserves_history_order_under_out_of_order_completion() {
let model = MockCompletionModel::from_stream_turns([
vec![
MockStreamEvent::tool_call("tc1", "add", json!({"x": 1, "y": 0})),
MockStreamEvent::tool_call("tc2", "add", json!({"x": 2, "y": 0})),
MockStreamEvent::final_response_with_total_tokens(0),
],
vec![
MockStreamEvent::text("done"),
MockStreamEvent::final_response_with_total_tokens(0),
],
]);
let stream = AgentBuilder::new(model)
.tool(OutOfOrderTool {
gate: Arc::new(tokio::sync::Notify::new()),
order: Arc::new(AtomicU32::new(0)),
})
.build()
.runner("go")
.max_turns(3)
.tool_concurrency(4)
.stream()
.await;
let final_response = tokio::time::timeout(
std::time::Duration::from_secs(5),
drive_to_final_response(stream),
)
.await
.expect("streamed tools must run concurrently, not deadlock on the first call");
let messages = final_response.messages().expect("history").to_vec();
assert_eq!(
tool_result_ids(&messages),
vec!["tc1".to_string(), "tc2".to_string()]
);
}
#[tokio::test]
async fn stream_emits_tool_results_in_call_order_after_batch_settles_under_concurrency() {
let model = MockCompletionModel::from_stream_turns([
vec![
MockStreamEvent::tool_call("tc1", "add", json!({"x": 1, "y": 0})),
MockStreamEvent::tool_call("tc2", "add", json!({"x": 2, "y": 0})),
MockStreamEvent::final_response_with_total_tokens(0),
],
vec![
MockStreamEvent::text("done"),
MockStreamEvent::final_response_with_total_tokens(0),
],
]);
let mut stream = AgentBuilder::new(model)
.tool(OutOfOrderTool {
gate: Arc::new(tokio::sync::Notify::new()),
order: Arc::new(AtomicU32::new(0)),
})
.build()
.runner("go")
.max_turns(3)
.tool_concurrency(4)
.stream()
.await;
let mut streamed_result_ids = Vec::new();
let mut final_response = None;
tokio::time::timeout(std::time::Duration::from_secs(5), async {
while let Some(item) = stream.next().await {
match item.unwrap_or_else(|err| panic!("stream item errored: {err}")) {
MultiTurnStreamItem::StreamUserItem(StreamedUserContent::ToolResult {
tool_result,
..
}) => streamed_result_ids.push(tool_result.id),
MultiTurnStreamItem::FinalResponse(resp) => final_response = Some(resp),
_ => {}
}
}
})
.await
.expect("streamed tools must run concurrently, not deadlock on the first call");
assert_eq!(
streamed_result_ids,
vec!["tc1".to_string(), "tc2".to_string()]
);
let final_response = final_response.expect("stream should yield a final response");
assert_eq!(
tool_result_ids(final_response.messages().expect("history")),
vec!["tc1".to_string(), "tc2".to_string()]
);
}
#[tokio::test]
async fn stream_executes_tools_concurrently_under_concurrency() {
let barrier = Arc::new(tokio::sync::Barrier::new(2));
let model = MockCompletionModel::from_stream_turns([
vec![
MockStreamEvent::tool_call("b1", "barrier_tool", json!({})),
MockStreamEvent::tool_call("b2", "barrier_tool", json!({})),
MockStreamEvent::final_response_with_total_tokens(0),
],
vec![
MockStreamEvent::text("done"),
MockStreamEvent::final_response_with_total_tokens(0),
],
]);
let stream = AgentBuilder::new(model)
.tool(MockBarrierTool::new(barrier))
.build()
.runner("hit the barrier twice")
.max_turns(3)
.tool_concurrency(2)
.stream()
.await;
tokio::time::timeout(
std::time::Duration::from_secs(5),
drive_to_final_response(stream),
)
.await
.expect("streamed tools must run concurrently, not deadlock at the barrier");
}
#[tokio::test]
async fn stream_emits_model_tool_calls_then_atomic_execution_items() {
async fn markers(concurrency: usize) -> Vec<&'static str> {
let model = MockCompletionModel::from_stream_turns([
vec![
MockStreamEvent::tool_call("tc1", "add", json!({"x": 1, "y": 1})),
MockStreamEvent::tool_call("tc2", "add", json!({"x": 2, "y": 2})),
MockStreamEvent::final_response_with_total_tokens(0),
],
vec![
MockStreamEvent::text("done"),
MockStreamEvent::final_response_with_total_tokens(0),
],
]);
let mut stream = AgentBuilder::new(model)
.tool(MockAddTool)
.build()
.runner("add two pairs")
.max_turns(3)
.tool_concurrency(concurrency)
.stream()
.await;
let mut markers = Vec::new();
while let Some(item) = stream.next().await {
match item.unwrap_or_else(|err| panic!("stream item errored: {err}")) {
MultiTurnStreamItem::StreamAssistantItem(
StreamedAssistantContent::ToolCall { .. },
) => markers.push("model-call"),
MultiTurnStreamItem::ToolExecutionCommitted { .. } => {
markers.push("exec-commit")
}
MultiTurnStreamItem::StreamUserItem(StreamedUserContent::ToolResult {
..
}) => markers.push("result"),
_ => {}
}
}
markers
}
let expected = vec![
"model-call",
"model-call",
"exec-commit",
"result",
"exec-commit",
"result",
];
assert_eq!(markers(1).await, expected);
assert_eq!(markers(4).await, expected);
}
struct TerminateAfterSiblingStartedHook {
sibling_started: Arc<tokio::sync::Notify>,
}
impl AgentHook for TerminateAfterSiblingStartedHook {
async fn on_tool_result(
&self,
_ctx: &HookContext,
event: ToolResultEvent<'_>,
) -> ToolResultAction {
if let ToolResultEvent { args, .. } = event
&& serde_json::from_str::<serde_json::Value>(args)
.ok()
.and_then(|v| v.get("x").and_then(serde_json::Value::as_i64))
== Some(1)
{
self.sibling_started.notified().await;
return ToolResultAction::stop("stop after a tool result");
}
ToolResultAction::keep()
}
}
#[derive(Clone)]
struct DrainProbeTool {
started: Arc<AtomicU32>,
completed: Arc<AtomicU32>,
slow_started: Arc<tokio::sync::Notify>,
}
impl Tool for DrainProbeTool {
const NAME: &'static str = "add";
type Error = MockToolError;
type Args = serde_json::Value;
type Output = i32;
fn description(&self) -> String {
MockAddTool.description()
}
fn parameters(&self) -> serde_json::Value {
MockAddTool.parameters()
}
async fn call(
&self,
_context: &mut ToolContext,
args: Self::Args,
) -> Result<Self::Output, Self::Error> {
self.started.fetch_add(1, SeqCst);
if args.get("x").and_then(serde_json::Value::as_i64) == Some(2) {
self.slow_started.notify_one();
for _ in 0..8 {
tokio::task::yield_now().await;
}
}
self.completed.fetch_add(1, SeqCst);
Ok(0)
}
}
#[tokio::test]
async fn stream_concurrent_tool_result_terminate_drains_in_flight_siblings() {
let started = Arc::new(AtomicU32::new(0));
let completed = Arc::new(AtomicU32::new(0));
let slow_started = Arc::new(tokio::sync::Notify::new());
let model = MockCompletionModel::from_stream_turns([
vec![
MockStreamEvent::tool_call("tc1", "add", json!({"x": 1, "y": 1})),
MockStreamEvent::tool_call("tc2", "add", json!({"x": 2, "y": 2})),
MockStreamEvent::final_response_with_total_tokens(0),
],
vec![
MockStreamEvent::text("done"),
MockStreamEvent::final_response_with_total_tokens(0),
],
]);
let mut stream = AgentBuilder::new(model)
.tool(DrainProbeTool {
started: started.clone(),
completed: completed.clone(),
slow_started: slow_started.clone(),
})
.build()
.runner("add two pairs")
.max_turns(3)
.tool_concurrency(2)
.add_hook(TerminateAfterSiblingStartedHook {
sibling_started: slow_started,
})
.stream()
.await;
let (saw_error, saw_final_response) =
tokio::time::timeout(std::time::Duration::from_secs(5), async move {
let mut saw_error = false;
let mut saw_final_response = false;
while let Some(item) = stream.next().await {
match item {
Ok(MultiTurnStreamItem::FinalResponse(_)) => saw_final_response = true,
Ok(_) => {}
Err(StreamingError::Prompt(_)) => saw_error = true,
Err(other) => panic!("unexpected streaming error: {other}"),
}
}
(saw_error, saw_final_response)
})
.await
.expect("draining the concurrent tools must not hang");
assert!(
saw_error,
"a terminate hook on the concurrent path must surface a StreamingError::Prompt"
);
assert!(
!saw_final_response,
"a terminated run must not yield a final response"
);
assert_eq!(
started.load(SeqCst),
2,
"both tools started (both in flight)"
);
assert_eq!(
completed.load(SeqCst),
2,
"the in-flight sibling must be drained to completion, not cancelled"
);
}
struct OrderedTerminateHook {
gate: Arc<tokio::sync::Notify>,
}
impl AgentHook for OrderedTerminateHook {
async fn on_tool_call(&self, _ctx: &HookContext, event: ToolCall<'_>) -> ToolCallAction {
if let ToolCall { args, .. } = event {
let x = serde_json::from_str::<serde_json::Value>(args)
.ok()
.and_then(|v| v.get("x").and_then(serde_json::Value::as_i64));
match x {
Some(2) => {
self.gate.notify_one();
return ToolCallAction::stop("terminated-by-tc2".to_string());
}
Some(1) => {
self.gate.notified().await;
return ToolCallAction::stop("terminated-by-tc1".to_string());
}
_ => {}
}
}
ToolCallAction::run()
}
}
fn two_terminating_tools_blocking_model() -> MockCompletionModel {
MockCompletionModel::from_turns([
MockTurn::from_contents([
tool_call_content("tc1", json!({"x": 1, "y": 1})),
tool_call_content("tc2", json!({"x": 2, "y": 2})),
])
.expect("two tool calls is non-empty"),
MockTurn::text("unreachable"),
])
}
fn two_terminating_tools_streaming_model() -> MockCompletionModel {
MockCompletionModel::from_stream_turns([
vec![
MockStreamEvent::tool_call("tc1", "add", json!({"x": 1, "y": 1})),
MockStreamEvent::tool_call("tc2", "add", json!({"x": 2, "y": 2})),
MockStreamEvent::final_response_with_total_tokens(0),
],
vec![
MockStreamEvent::text("unreachable"),
MockStreamEvent::final_response_with_total_tokens(0),
],
])
}
#[tokio::test]
async fn concurrent_simultaneous_tool_terminations_pick_call_order_on_both_drivers() {
let run_err = tokio::time::timeout(
std::time::Duration::from_secs(5),
AgentBuilder::new(two_terminating_tools_blocking_model())
.tool(MockAddTool)
.build()
.runner("go")
.max_turns(3)
.tool_concurrency(2)
.add_hook(OrderedTerminateHook {
gate: Arc::new(tokio::sync::Notify::new()),
})
.run(),
)
.await
.expect("blocking run must not hang")
.expect_err("the run must terminate");
let mut stream = AgentBuilder::new(two_terminating_tools_streaming_model())
.tool(MockAddTool)
.build()
.runner("go")
.max_turns(3)
.tool_concurrency(2)
.add_hook(OrderedTerminateHook {
gate: Arc::new(tokio::sync::Notify::new()),
})
.stream()
.await;
let stream_err = tokio::time::timeout(std::time::Duration::from_secs(5), async move {
while let Some(item) = stream.next().await {
if let Err(err) = item {
return Some(err);
}
}
None
})
.await
.expect("streamed run must not hang")
.expect("the stream must surface a terminate error");
let run_msg = run_err.to_string();
let stream_msg = stream_err.to_string();
assert!(
run_msg.contains("terminated-by-tc1"),
"blocking run should surface the first-called tool's reason, got: {run_msg}"
);
assert!(
stream_msg.contains("terminated-by-tc1"),
"stream should surface the first-called tool's reason, got: {stream_msg}"
);
assert!(
!run_msg.contains("terminated-by-tc2") && !stream_msg.contains("terminated-by-tc2"),
"neither driver should surface the later-completing tool's reason"
);
}
struct TerminateOnFirstToolHook;
impl AgentHook for TerminateOnFirstToolHook {
async fn on_tool_call(&self, _ctx: &HookContext, event: ToolCall<'_>) -> ToolCallAction {
if let ToolCall { args, .. } = event
&& serde_json::from_str::<serde_json::Value>(args)
.ok()
.and_then(|v| v.get("x").and_then(serde_json::Value::as_i64))
== Some(1)
{
return ToolCallAction::stop("stop".to_string());
}
ToolCallAction::run()
}
}
#[tokio::test]
async fn default_concurrency_terminate_skips_remaining_tools_on_both_drivers() {
let blocking_calls = Arc::new(AtomicU32::new(0));
AgentBuilder::new(two_terminating_tools_blocking_model())
.tool(CountingAddTool {
calls: blocking_calls.clone(),
})
.build()
.runner("go")
.max_turns(3)
.add_hook(TerminateOnFirstToolHook)
.run()
.await
.expect_err("the run terminates");
assert_eq!(
blocking_calls.load(SeqCst),
0,
"fail-fast: blocking run() must not start the second tool after the first terminates"
);
let streaming_calls = Arc::new(AtomicU32::new(0));
let mut stream = AgentBuilder::new(two_terminating_tools_streaming_model())
.tool(CountingAddTool {
calls: streaming_calls.clone(),
})
.build()
.runner("go")
.max_turns(3)
.add_hook(TerminateOnFirstToolHook)
.stream()
.await;
let mut saw_error = false;
while let Some(item) = stream.next().await {
if let Err(err) = item {
saw_error = true;
assert!(
err.to_string().contains("stop"),
"stream() should surface the terminate reason, got: {err}"
);
break;
}
}
assert!(saw_error, "stream() must surface the terminate error");
assert_eq!(
streaming_calls.load(SeqCst),
0,
"fail-fast: stream() must not start the second tool after the first terminates"
);
}
#[derive(Clone)]
struct RecordingArgsTool {
called: Arc<Mutex<Vec<i64>>>,
sibling_started: Arc<tokio::sync::Notify>,
}
impl Tool for RecordingArgsTool {
const NAME: &'static str = "add";
type Error = MockToolError;
type Args = serde_json::Value;
type Output = i32;
fn description(&self) -> String {
MockAddTool.description()
}
fn parameters(&self) -> serde_json::Value {
MockAddTool.parameters()
}
async fn call(
&self,
_context: &mut ToolContext,
args: Self::Args,
) -> Result<Self::Output, Self::Error> {
let x = args.get("x").and_then(serde_json::Value::as_i64);
if let Some(x) = x {
self.called.lock().expect("called").push(x);
}
if x == Some(1) {
self.sibling_started.notify_one();
for _ in 0..8 {
tokio::task::yield_now().await;
}
}
Ok(0)
}
}
fn three_tools_first_terminates_streaming_model() -> MockCompletionModel {
MockCompletionModel::from_stream_turns([
vec![
MockStreamEvent::tool_call("tc0", "add", json!({"x": 0, "y": 0})),
MockStreamEvent::tool_call("tc1", "add", json!({"x": 1, "y": 1})),
MockStreamEvent::tool_call("tc2", "add", json!({"x": 2, "y": 2})),
MockStreamEvent::final_response_with_total_tokens(0),
],
vec![
MockStreamEvent::text("unreachable"),
MockStreamEvent::final_response_with_total_tokens(0),
],
])
}
struct TerminateOnArgZeroAfterSiblingHook {
sibling_started: Arc<tokio::sync::Notify>,
}
impl AgentHook for TerminateOnArgZeroAfterSiblingHook {
async fn on_tool_call(&self, _ctx: &HookContext, event: ToolCall<'_>) -> ToolCallAction {
if let ToolCall { args, .. } = event
&& serde_json::from_str::<serde_json::Value>(args)
.ok()
.and_then(|v| v.get("x").and_then(serde_json::Value::as_i64))
== Some(0)
{
self.sibling_started.notified().await;
return ToolCallAction::stop("stop");
}
ToolCallAction::run()
}
}
#[tokio::test]
async fn concurrent_terminate_drops_beyond_window_sibling_but_drains_in_flight() {
let called = Arc::new(Mutex::new(Vec::new()));
let sibling_started = Arc::new(tokio::sync::Notify::new());
let mut stream = AgentBuilder::new(three_tools_first_terminates_streaming_model())
.tool(RecordingArgsTool {
called: called.clone(),
sibling_started: sibling_started.clone(),
})
.build()
.runner("go")
.max_turns(3)
.tool_concurrency(2)
.add_hook(TerminateOnArgZeroAfterSiblingHook { sibling_started })
.stream()
.await;
let (saw_error, saw_final) =
tokio::time::timeout(std::time::Duration::from_secs(5), async move {
let mut saw_error = false;
let mut saw_final = false;
while let Some(item) = stream.next().await {
match item {
Ok(MultiTurnStreamItem::FinalResponse(_)) => saw_final = true,
Ok(_) => {}
Err(_) => saw_error = true,
}
}
(saw_error, saw_final)
})
.await
.expect("the concurrent tool drive must not hang");
assert!(saw_error, "the terminated run must surface an error");
assert!(
!saw_final,
"a terminated run must not yield a final response"
);
let called = called.lock().expect("called").clone();
assert!(
called.contains(&1),
"the in-flight sibling (x==1) must be drained to completion; called args: {called:?}"
);
assert!(
!called.contains(&2),
"the not-yet-started sibling beyond the concurrency window (x==2) must be \
dropped, not executed; called args: {called:?}"
);
assert!(
!called.contains(&0),
"the terminator's own body never runs (its ToolCall hook terminated); \
called args: {called:?}"
);
}
#[derive(Clone)]
struct SignalOnRunTool {
a_ran: Arc<AtomicU32>,
a_done: Arc<tokio::sync::Notify>,
}
impl Tool for SignalOnRunTool {
const NAME: &'static str = "add";
type Error = MockToolError;
type Args = serde_json::Value;
type Output = i32;
fn description(&self) -> String {
MockAddTool.description()
}
fn parameters(&self) -> serde_json::Value {
MockAddTool.parameters()
}
async fn call(
&self,
_context: &mut ToolContext,
args: Self::Args,
) -> Result<Self::Output, Self::Error> {
if args.get("x").and_then(serde_json::Value::as_i64) == Some(1) {
self.a_ran.fetch_add(1, SeqCst);
self.a_done.notify_one();
}
Ok(0)
}
}
struct TerminateAfterSiblingDoneHook {
a_done: Arc<tokio::sync::Notify>,
}
impl AgentHook for TerminateAfterSiblingDoneHook {
async fn on_tool_call(&self, _ctx: &HookContext, event: ToolCall<'_>) -> ToolCallAction {
if let ToolCall { args, .. } = event
&& serde_json::from_str::<serde_json::Value>(args)
.ok()
.and_then(|v| v.get("x").and_then(serde_json::Value::as_i64))
== Some(2)
{
self.a_done.notified().await;
return ToolCallAction::stop("stop");
}
ToolCallAction::run()
}
}
#[tokio::test]
async fn concurrent_termination_surfaces_no_execution_items() {
let a_ran = Arc::new(AtomicU32::new(0));
let a_done = Arc::new(tokio::sync::Notify::new());
let model = MockCompletionModel::from_stream_turns([
vec![
MockStreamEvent::tool_call("tc1", "add", json!({"x": 1, "y": 1})),
MockStreamEvent::tool_call("tc2", "add", json!({"x": 2, "y": 2})),
MockStreamEvent::final_response_with_total_tokens(0),
],
vec![
MockStreamEvent::text("unreachable"),
MockStreamEvent::final_response_with_total_tokens(0),
],
]);
let mut stream = AgentBuilder::new(model)
.tool(SignalOnRunTool {
a_ran: a_ran.clone(),
a_done: a_done.clone(),
})
.build()
.runner("go")
.max_turns(3)
.tool_concurrency(2)
.add_hook(TerminateAfterSiblingDoneHook {
a_done: a_done.clone(),
})
.stream()
.await;
let (exec_commits, results, saw_error, saw_final) =
tokio::time::timeout(std::time::Duration::from_secs(5), async move {
let (mut exec_commits, mut results, mut saw_error, mut saw_final) =
(0, 0, false, false);
while let Some(item) = stream.next().await {
match item {
Ok(MultiTurnStreamItem::ToolExecutionCommitted { .. }) => exec_commits += 1,
Ok(MultiTurnStreamItem::StreamUserItem(
StreamedUserContent::ToolResult { .. },
)) => results += 1,
Ok(MultiTurnStreamItem::FinalResponse(_)) => saw_final = true,
Ok(_) => {}
Err(_) => saw_error = true,
}
}
(exec_commits, results, saw_error, saw_final)
})
.await
.expect("the concurrent tool drive must not hang");
assert!(saw_error, "the terminated run must surface an error");
assert!(
!saw_final,
"a terminated run must not yield a final response"
);
assert_eq!(
exec_commits, 0,
"a terminated batch surfaces no ToolExecutionCommitted events"
);
assert_eq!(
results, 0,
"a terminated batch surfaces no successful ToolResult"
);
assert_eq!(
a_ran.load(SeqCst),
1,
"the fast sibling did run (its side effect happened), but its result was suppressed"
);
}
#[tokio::test]
async fn stream_tool_execution_committed_carries_effective_rewritten_args() {
let model = MockCompletionModel::from_stream_turns([
vec![
MockStreamEvent::tool_call("tc1", "add", json!({"x": 2, "y": 3})),
MockStreamEvent::final_response_with_total_tokens(0),
],
vec![
MockStreamEvent::text("done"),
MockStreamEvent::final_response_with_total_tokens(0),
],
]);
let mut stream = AgentBuilder::new(model)
.tool(MockAddTool)
.add_hook(RewriteToolArgsHook(json!({"x": 2, "y": 40})))
.build()
.runner("go")
.max_turns(3)
.stream()
.await;
let mut model_args = None;
let mut exec_args = None;
while let Some(item) = stream.next().await {
match item.unwrap_or_else(|err| panic!("stream item errored: {err}")) {
MultiTurnStreamItem::StreamAssistantItem(StreamedAssistantContent::ToolCall {
tool_call,
..
}) => model_args = Some(tool_call.function.arguments),
MultiTurnStreamItem::ToolExecutionCommitted { tool_call, .. } => {
exec_args = Some(tool_call.function.arguments)
}
_ => {}
}
}
assert_eq!(
model_args,
Some(json!({"x": 2, "y": 3})),
"the model tool-call event carries the model's original arguments"
);
assert_eq!(
exec_args,
Some(json!({"x": 2, "y": 40})),
"the execution-commit event carries the hook-rewritten (effective) arguments"
);
}
#[tokio::test]
async fn stream_hook_skip_surfaces_result_without_execution_commit() {
struct SkipHook;
impl AgentHook for SkipHook {
async fn on_tool_call(
&self,
_ctx: &HookContext,
event: ToolCall<'_>,
) -> ToolCallAction {
if let ToolCall { .. } = event {
ToolCallAction::skip("blocked by policy")
} else {
ToolCallAction::run()
}
}
}
let calls = Arc::new(AtomicU32::new(0));
let model = MockCompletionModel::from_stream_turns([
vec![
MockStreamEvent::tool_call("tc1", "add", json!({"x": 1, "y": 2})),
MockStreamEvent::final_response_with_total_tokens(0),
],
vec![
MockStreamEvent::text("done"),
MockStreamEvent::final_response_with_total_tokens(0),
],
]);
let stream = AgentBuilder::new(model)
.tool(CountingAddTool {
calls: calls.clone(),
})
.add_hook(SkipHook)
.build()
.runner("go")
.max_turns(3)
.stream()
.await;
let mut exec_commits = 0;
let mut results = 0;
let mut final_response = None;
let mut stream = stream;
while let Some(item) = stream.next().await {
match item.unwrap_or_else(|err| panic!("stream item errored: {err}")) {
MultiTurnStreamItem::ToolExecutionCommitted { .. } => exec_commits += 1,
MultiTurnStreamItem::StreamUserItem(StreamedUserContent::ToolResult { .. }) => {
results += 1
}
MultiTurnStreamItem::FinalResponse(resp) => final_response = Some(resp),
_ => {}
}
}
assert_eq!(calls.load(SeqCst), 0, "a skipped tool's body never runs");
assert_eq!(
exec_commits, 0,
"a hook-skipped tool produces no execution-commit"
);
assert_eq!(
results, 1,
"the skip result is still surfaced to the consumer"
);
let final_response = final_response.expect("stream should yield a final response");
let history = final_response.messages().expect("history");
assert!(
history.iter().any(|m| serde_json::to_string(m)
.map(|s| s.contains("blocked by policy"))
.unwrap_or(false)),
"the skip result is committed to history"
);
}
#[tokio::test]
async fn required_with_empty_active_tools_errors_locally_without_provider_call() {
struct EmptyActiveToolsHook;
impl AgentHook for EmptyActiveToolsHook {
async fn on_completion_call(
&self,
_ctx: &HookContext,
event: CompletionCallEvent<'_>,
) -> CompletionCallAction {
if let CompletionCallEvent { .. } = event {
CompletionCallAction::patch(
RequestPatch::new().active_tools(Vec::<String>::new()),
)
} else {
CompletionCallAction::continue_run()
}
}
}
let model = MockCompletionModel::from_turns([MockTurn::text("unreachable")]);
let probe = model.clone();
let err = AgentBuilder::new(model)
.tool(MockAddTool)
.tool_choice(ToolChoice::Required)
.add_hook(EmptyActiveToolsHook)
.build()
.runner("go")
.run()
.await
.expect_err("Required with an empty active_tools filter must fail locally");
assert!(
probe.requests().is_empty(),
"the request must fail locally, with no provider round-trip"
);
let msg = err.to_string();
assert!(
msg.contains("Required"),
"error should mention Required: {msg}"
);
assert!(
msg.contains("active_tools"),
"error should name active_tools: {msg}"
);
}
#[tokio::test]
async fn specific_naming_filtered_out_tool_errors_locally_without_provider_call() {
struct FilterToAddHook;
impl AgentHook for FilterToAddHook {
async fn on_completion_call(
&self,
_ctx: &HookContext,
event: CompletionCallEvent<'_>,
) -> CompletionCallAction {
if let CompletionCallEvent { .. } = event {
CompletionCallAction::patch(RequestPatch::new().active_tools(["add"]))
} else {
CompletionCallAction::continue_run()
}
}
}
let model = MockCompletionModel::from_turns([MockTurn::text("unreachable")]);
let probe = model.clone();
let err = AgentBuilder::new(model)
.tool(MockAddTool)
.tool(MockSubtractTool)
.tool_choice(ToolChoice::Specific {
function_names: vec!["subtract".to_string()],
})
.add_hook(FilterToAddHook)
.build()
.runner("go")
.run()
.await
.expect_err("Specific naming a filtered-out tool must fail locally");
assert!(
probe.requests().is_empty(),
"the request must fail locally, with no provider round-trip"
);
let msg = err.to_string();
assert!(
msg.contains("subtract"),
"error should name the missing tool: {msg}"
);
assert!(
msg.contains("active_tools"),
"error should name active_tools: {msg}"
);
}
#[tokio::test]
async fn concurrent_tool_execution_stays_within_the_configured_bound() {
#[derive(Clone)]
struct ConcurrencyProbe {
barrier: Arc<Barrier>,
active: Arc<AtomicU32>,
max_active: Arc<AtomicU32>,
}
impl Tool for ConcurrencyProbe {
const NAME: &'static str = "add";
type Error = MockToolError;
type Args = serde_json::Value;
type Output = String;
fn description(&self) -> String {
"concurrency probe".to_string()
}
fn parameters(&self) -> serde_json::Value {
json!({"type": "object", "properties": {}})
}
async fn call(
&self,
_context: &mut ToolContext,
_args: Self::Args,
) -> Result<Self::Output, Self::Error> {
let now = self.active.fetch_add(1, SeqCst) + 1;
self.max_active.fetch_max(now, SeqCst);
self.barrier.wait().await;
self.active.fetch_sub(1, SeqCst);
Ok("ok".to_string())
}
}
let cap = 2usize;
let probe = ConcurrencyProbe {
barrier: Arc::new(Barrier::new(cap)),
active: Arc::new(AtomicU32::new(0)),
max_active: Arc::new(AtomicU32::new(0)),
};
let max_active = probe.max_active.clone();
let model = MockCompletionModel::from_turns([
MockTurn::from_contents([
tool_call_content("c1", json!({})),
tool_call_content("c2", json!({})),
tool_call_content("c3", json!({})),
tool_call_content("c4", json!({})),
])
.expect("four tool calls is a valid turn"),
MockTurn::text("done"),
]);
let _ = AgentBuilder::new(model)
.tool(probe)
.build()
.runner("probe concurrency")
.max_turns(3)
.tool_concurrency(cap)
.run()
.await
.expect("run should succeed");
let observed = max_active.load(SeqCst);
assert!(
observed > 1,
"tools actually ran concurrently (lower bound): max_active={observed}"
);
assert!(
observed <= cap as u32,
"in-flight never exceeded the configured bound {cap} (upper bound): max_active={observed}"
);
}
#[tokio::test]
async fn tool_concurrency_zero_is_clamped_and_does_not_hang() {
let model = MockCompletionModel::from_turns([
MockTurn::tool_call("tc1", "add", json!({"x": 1, "y": 2})),
MockTurn::text("done"),
]);
let run = AgentBuilder::new(model)
.tool(MockAddTool)
.build()
.runner("add")
.max_turns(3)
.tool_concurrency(0)
.run();
let response = tokio::time::timeout(std::time::Duration::from_secs(5), run)
.await
.expect("tool_concurrency(0) must clamp to 1, not hang on buffer_unordered(0)")
.expect("run should succeed");
assert_eq!(response.output, "done");
}
#[derive(Clone)]
struct CountingAddTool {
calls: Arc<AtomicU32>,
}
impl Tool for CountingAddTool {
const NAME: &'static str = "add";
type Error = MockToolError;
type Args = MockOperationArgs;
type Output = i32;
fn description(&self) -> String {
MockAddTool.description()
}
fn parameters(&self) -> serde_json::Value {
MockAddTool.parameters()
}
async fn call(
&self,
_context: &mut ToolContext,
args: Self::Args,
) -> Result<Self::Output, Self::Error> {
self.calls.fetch_add(1, SeqCst);
MockAddTool.call(_context, args).await
}
}
#[derive(Clone, Default)]
struct ToolOnlyHook {
text_delta_calls: Arc<AtomicU32>,
other_calls: Arc<AtomicU32>,
}
impl AgentHook for ToolOnlyHook {
async fn on_text_delta(&self, _: &HookContext, _: TextDelta<'_>) -> ObservationAction {
self.text_delta_calls.fetch_add(1, SeqCst);
ObservationAction::continue_run()
}
async fn on_completion_call(
&self,
_: &HookContext,
_: CompletionCallEvent<'_>,
) -> CompletionCallAction {
self.other_calls.fetch_add(1, SeqCst);
CompletionCallAction::continue_run()
}
fn observes(&self, kind: StepEventKind) -> bool {
kind != StepEventKind::TextDelta
}
}
#[tokio::test]
async fn observes_gates_text_delta_dispatch() {
let model = MockCompletionModel::from_stream_turns([vec![
MockStreamEvent::text("hel"),
MockStreamEvent::text("lo"),
MockStreamEvent::final_response_with_total_tokens(0),
]]);
let hook = ToolOnlyHook::default();
let mut stream = AgentBuilder::new(model)
.build()
.runner("hi")
.add_hook(hook.clone())
.stream()
.await;
while stream.next().await.is_some() {}
assert_eq!(
hook.text_delta_calls.load(SeqCst),
0,
"a hook that does not observe TextDelta must not be dispatched for it"
);
assert!(
hook.other_calls.load(SeqCst) > 0,
"the hook should still receive the events it observes"
);
}
struct TerminateOn(StepEventKind);
impl AgentHook for TerminateOn {
async fn on_completion_call(
&self,
_: &HookContext,
_: CompletionCallEvent<'_>,
) -> CompletionCallAction {
if self.0 == StepEventKind::CompletionCall {
CompletionCallAction::stop("stop here")
} else {
CompletionCallAction::continue_run()
}
}
async fn on_tool_call(&self, _: &HookContext, _: ToolCall<'_>) -> ToolCallAction {
if self.0 == StepEventKind::ToolCall {
ToolCallAction::stop("stop here")
} else {
ToolCallAction::run()
}
}
async fn on_tool_result(
&self,
_: &HookContext,
_: ToolResultEvent<'_>,
) -> ToolResultAction {
if self.0 == StepEventKind::ToolResult {
ToolResultAction::stop("stop here")
} else {
ToolResultAction::keep()
}
}
}
#[tokio::test]
async fn run_terminates_from_each_shared_event() {
for kind in [
StepEventKind::CompletionCall,
StepEventKind::ToolCall,
StepEventKind::ToolResult,
] {
let err = AgentBuilder::new(blocking_model())
.tool(MockAddTool)
.build()
.runner("add 2 and 3")
.max_turns(3)
.add_hook(TerminateOn(kind))
.run()
.await
.expect_err(&format!("terminate at {kind:?} must cancel the run"));
assert!(
matches!(err, PromptError::PromptCancelled { .. }),
"terminate at {kind:?} should cancel the run, got {err:?}"
);
}
}
#[tokio::test]
async fn stream_terminates_from_each_shared_event() {
for kind in [
StepEventKind::CompletionCall,
StepEventKind::ToolCall,
StepEventKind::ToolResult,
] {
let mut stream = AgentBuilder::new(streaming_model())
.tool(MockAddTool)
.build()
.runner("add 2 and 3")
.max_turns(3)
.add_hook(TerminateOn(kind))
.stream()
.await;
let mut saw_error = false;
let mut saw_final = false;
while let Some(item) = stream.next().await {
match item {
Ok(MultiTurnStreamItem::FinalResponse(_)) => saw_final = true,
Err(_) => saw_error = true,
_ => {}
}
}
assert!(saw_error, "terminate at {kind:?} must yield a stream error");
assert!(
!saw_final,
"terminate at {kind:?} must not also produce a final response"
);
}
}
#[tokio::test]
async fn multi_hook_stack_parity_across_run_and_stream() {
let a_block = RecordingHook::default();
let b_block = RecordingHook::default();
let blocking = AgentBuilder::new(blocking_model())
.tool(MockAddTool)
.build()
.runner("add 2 and 3")
.max_turns(3)
.add_hook(a_block.clone())
.add_hook(b_block.clone())
.run()
.await
.expect("blocking run should succeed");
let a_stream = RecordingHook::default();
let b_stream = RecordingHook::default();
let mut stream = AgentBuilder::new(streaming_model())
.tool(MockAddTool)
.build()
.runner("add 2 and 3")
.max_turns(3)
.add_hook(a_stream.clone())
.add_hook(b_stream.clone())
.stream()
.await;
while stream.next().await.is_some() {}
assert_eq!(a_block.shared_events(), b_block.shared_events());
assert_eq!(a_stream.shared_events(), b_stream.shared_events());
assert_eq!(a_block.shared_events(), a_stream.shared_events());
assert_eq!(
a_block.shared_events(),
vec![
StepEventKind::CompletionCall,
StepEventKind::ToolCall,
StepEventKind::ToolResult,
StepEventKind::CompletionCall,
]
);
assert_eq!(blocking.output, "the answer is 5");
}
struct RepairInvalidToHook(&'static str);
impl AgentHook for RepairInvalidToHook {
async fn on_invalid_tool_call(
&self,
_ctx: &HookContext,
event: &InvalidToolCallContext,
) -> Option<InvalidToolCallAction> {
Some(if let _ = event {
InvalidToolCallAction::repair(self.0)
} else {
InvalidToolCallAction::fail()
})
}
}
#[derive(Clone)]
struct CaptureAndRepairInvalidHook {
replacement: &'static str,
args: Arc<Mutex<Vec<Option<String>>>>,
}
impl AgentHook for CaptureAndRepairInvalidHook {
async fn on_invalid_tool_call(
&self,
_ctx: &HookContext,
event: &InvalidToolCallContext,
) -> Option<InvalidToolCallAction> {
self.args
.lock()
.expect("invalid args")
.push(event.args.clone());
Some(InvalidToolCallAction::repair(self.replacement))
}
}
#[tokio::test]
async fn invalid_tool_call_repair_parity_across_run_and_stream() {
let blocking_model = MockCompletionModel::from_turns([
MockTurn::tool_call("tc1", "default_api", json!({"x": 2, "y": 3})),
MockTurn::text("the answer is 5"),
]);
let blocking_hook = RecordingHook::default();
let blocking = AgentBuilder::new(blocking_model)
.tool(MockAddTool)
.build()
.runner("add 2 and 3")
.max_turns(3)
.add_hook(blocking_hook.clone())
.add_hook(RepairInvalidToHook("add"))
.run()
.await
.expect("blocking run should recover via repair");
let streaming_model = MockCompletionModel::from_stream_turns([
vec![
MockStreamEvent::tool_call("tc1", "default_api", json!({"x": 2, "y": 3})),
MockStreamEvent::final_response_with_total_tokens(0),
],
vec![
MockStreamEvent::text("the answer is 5"),
MockStreamEvent::final_response_with_total_tokens(0),
],
]);
let streaming_hook = RecordingHook::default();
let mut stream = AgentBuilder::new(streaming_model)
.tool(MockAddTool)
.build()
.runner("add 2 and 3")
.max_turns(3)
.add_hook(streaming_hook.clone())
.add_hook(RepairInvalidToHook("add"))
.stream()
.await;
let mut final_response = None;
while let Some(item) = stream.next().await {
if let Ok(MultiTurnStreamItem::FinalResponse(resp)) =
item.map_err(|err| panic!("stream item errored: {err}"))
{
final_response = Some(resp);
}
}
let final_response =
final_response.expect("stream should recover and yield a final response");
assert_eq!(blocking.output, "the answer is 5");
assert_eq!(final_response.output(), blocking.output);
assert_eq!(
blocking_hook.shared_events(),
streaming_hook.shared_events()
);
assert!(
blocking_hook
.shared_events()
.contains(&StepEventKind::InvalidToolCall),
"the hook must observe the invalid tool call"
);
assert_eq!(blocking_hook.tool_results(), streaming_hook.tool_results());
assert_eq!(blocking_hook.tool_results(), vec!["5".to_string()]);
let blocking_messages = blocking.messages.expect("blocking messages");
let streaming_messages = final_response
.messages()
.expect("streaming history")
.to_vec();
assert_eq!(
serde_json::to_value(&blocking_messages).expect("serialize blocking"),
serde_json::to_value(&streaming_messages).expect("serialize streaming"),
);
}
#[tokio::test]
async fn invalid_tool_call_scalar_args_are_canonical_across_run_and_complete_stream() {
let blocking_args = Arc::new(Mutex::new(Vec::new()));
let blocking_hook = RecordingHook::default();
let blocking = AgentBuilder::new(MockCompletionModel::from_turns([
MockTurn::tool_call("tc1", "unknown_echo", json!("payload")),
MockTurn::text("done"),
]))
.tool(EchoStringArgs)
.build()
.runner("echo a string")
.max_turns(3)
.add_hook(blocking_hook.clone())
.add_hook(CaptureAndRepairInvalidHook {
replacement: EchoStringArgs::NAME,
args: blocking_args.clone(),
})
.run()
.await
.expect("blocking scalar repair should succeed");
let streaming_args = Arc::new(Mutex::new(Vec::new()));
let streaming_hook = RecordingHook::default();
let mut stream = AgentBuilder::new(MockCompletionModel::from_stream_turns([
vec![
MockStreamEvent::tool_call("tc1", "unknown_echo", json!("payload")),
MockStreamEvent::final_response_with_total_tokens(0),
],
vec![
MockStreamEvent::text("done"),
MockStreamEvent::final_response_with_total_tokens(0),
],
]))
.tool(EchoStringArgs)
.build()
.runner("echo a string")
.max_turns(3)
.add_hook(streaming_hook.clone())
.add_hook(CaptureAndRepairInvalidHook {
replacement: EchoStringArgs::NAME,
args: streaming_args.clone(),
})
.stream()
.await;
let mut final_response = None;
while let Some(item) = stream.next().await {
if let MultiTurnStreamItem::FinalResponse(response) =
item.expect("streaming scalar repair should succeed")
{
final_response = Some(response);
}
}
let final_response = final_response.expect("stream should yield a final response");
let canonical_args = vec![Some(serde_json::to_string("payload").unwrap())];
assert_eq!(*blocking_args.lock().unwrap(), canonical_args);
assert_eq!(*streaming_args.lock().unwrap(), canonical_args);
assert_eq!(blocking_hook.tool_results(), vec!["payload"]);
assert_eq!(streaming_hook.tool_results(), vec!["payload"]);
assert_eq!(blocking.output, "done");
assert_eq!(final_response.output(), "done");
assert_eq!(
serde_json::to_value(blocking.messages.expect("blocking history")).unwrap(),
serde_json::to_value(final_response.messages().expect("streaming history")).unwrap()
);
}
#[derive(Clone)]
struct ScriptedToolCall {
id: &'static str,
name: &'static str,
args: serde_json::Value,
}
#[derive(Clone)]
enum ScriptedTurn {
Text(&'static str),
ToolCalls(Vec<ScriptedToolCall>),
}
#[derive(Clone, Copy)]
enum StreamShape {
Complete,
Chunked,
}
impl ScriptedTurn {
fn as_blocking_turn(&self) -> MockTurn {
match self {
ScriptedTurn::Text(text) => MockTurn::text(*text),
ScriptedTurn::ToolCalls(calls) => {
MockTurn::from_contents(calls.iter().map(|call| {
AssistantContent::ToolCall(MessageToolCall::new(
call.id.to_string(),
ToolFunction::new(call.name.to_string(), call.args.clone()),
))
}))
.expect("a scripted tool-call turn has at least one call")
}
}
}
fn as_stream_events(&self, shape: StreamShape) -> Vec<MockStreamEvent> {
let mut events = Vec::new();
match self {
ScriptedTurn::Text(text) => events.push(MockStreamEvent::text(*text)),
ScriptedTurn::ToolCalls(calls) => {
for call in calls {
if let StreamShape::Chunked = shape {
let internal = format!("ic-{}", call.id);
let args = serde_json::to_string(&call.args)
.expect("scripted args serialize to json");
events.push(MockStreamEvent::tool_call_name_delta(
call.id, &internal, call.name,
));
events.push(MockStreamEvent::tool_call_arguments_delta(
call.id, &internal, &args,
));
}
events.push(MockStreamEvent::tool_call(
call.id,
call.name,
call.args.clone(),
));
}
}
}
events.push(MockStreamEvent::final_response_with_total_tokens(0));
events
}
}
struct ParityOutcome {
output: String,
messages: Vec<Message>,
shared_events: Vec<StepEventKind>,
tool_results: Vec<String>,
}
async fn run_blocking_scenario(prompt: &'static str, turns: &[ScriptedTurn]) -> ParityOutcome {
let model =
MockCompletionModel::from_turns(turns.iter().map(ScriptedTurn::as_blocking_turn));
let hook = RecordingHook::default();
let response = AgentBuilder::new(model)
.tool(MockAddTool)
.build()
.runner(prompt)
.max_turns(8)
.add_hook(hook.clone())
.run()
.await
.expect("blocking scenario should succeed");
ParityOutcome {
output: response.output,
messages: response.messages.expect("blocking messages"),
shared_events: hook.shared_events(),
tool_results: hook.tool_results(),
}
}
async fn run_streaming_scenario(
prompt: &'static str,
turns: &[ScriptedTurn],
shape: StreamShape,
) -> ParityOutcome {
let model = MockCompletionModel::from_stream_turns(
turns.iter().map(|turn| turn.as_stream_events(shape)),
);
let hook = RecordingHook::default();
let mut stream = AgentBuilder::new(model)
.tool(MockAddTool)
.build()
.runner(prompt)
.max_turns(8)
.add_hook(hook.clone())
.stream()
.await;
let mut final_response = None;
while let Some(item) = stream.next().await {
if let Ok(MultiTurnStreamItem::FinalResponse(resp)) =
item.map_err(|err| panic!("stream item errored: {err}"))
{
final_response = Some(resp);
}
}
let final_response =
final_response.expect("streaming scenario should yield a final response");
ParityOutcome {
output: final_response.output().to_string(),
messages: final_response
.messages()
.expect("streaming history")
.to_vec(),
shared_events: hook.shared_events(),
tool_results: hook.tool_results(),
}
}
fn assert_outcomes_match(blocking: &ParityOutcome, streaming: &ParityOutcome, label: &str) {
assert_eq!(
blocking.output, streaming.output,
"{label}: final output diverged"
);
assert_eq!(
blocking.shared_events, streaming.shared_events,
"{label}: hook event sequence diverged"
);
assert_eq!(
blocking.tool_results, streaming.tool_results,
"{label}: tool-result content diverged"
);
assert_eq!(
serde_json::to_value(&blocking.messages).expect("serialize blocking"),
serde_json::to_value(&streaming.messages).expect("serialize streaming"),
"{label}: message history diverged"
);
}
async fn assert_run_stream_parity(prompt: &'static str, turns: &[ScriptedTurn]) {
let blocking = run_blocking_scenario(prompt, turns).await;
for (shape, label) in [
(StreamShape::Complete, "complete-stream"),
(StreamShape::Chunked, "chunked-stream"),
] {
let streaming = run_streaming_scenario(prompt, turns, shape).await;
assert_outcomes_match(&blocking, &streaming, label);
}
}
fn add_call(id: &'static str, x: i64, y: i64) -> ScriptedToolCall {
ScriptedToolCall {
id,
name: "add",
args: json!({ "x": x, "y": y }),
}
}
#[tokio::test]
async fn parity_text_only_run() {
assert_run_stream_parity("just say hi", &[ScriptedTurn::Text("hi there")]).await;
}
#[tokio::test]
async fn parity_single_tool_then_text() {
assert_run_stream_parity(
"add 2 and 3",
&[
ScriptedTurn::ToolCalls(vec![add_call("tc1", 2, 3)]),
ScriptedTurn::Text("the answer is 5"),
],
)
.await;
}
#[tokio::test]
async fn parity_multiple_tools_in_one_turn() {
assert_run_stream_parity(
"add two pairs",
&[
ScriptedTurn::ToolCalls(vec![add_call("tc1", 2, 3), add_call("tc2", 10, 20)]),
ScriptedTurn::Text("done"),
],
)
.await;
}
#[tokio::test]
async fn parity_multi_turn_sequential_tools() {
assert_run_stream_parity(
"chain two additions",
&[
ScriptedTurn::ToolCalls(vec![add_call("tc1", 1, 1)]),
ScriptedTurn::ToolCalls(vec![add_call("tc2", 2, 2)]),
ScriptedTurn::Text("chained"),
],
)
.await;
}
struct SkipInvalidHook(&'static str);
impl AgentHook for SkipInvalidHook {
async fn on_invalid_tool_call(
&self,
_ctx: &HookContext,
event: &InvalidToolCallContext,
) -> Option<InvalidToolCallAction> {
Some(if let _ = event {
InvalidToolCallAction::skip(self.0)
} else {
InvalidToolCallAction::fail()
})
}
}
#[tokio::test]
async fn invalid_tool_call_skip_parity_across_run_and_stream() {
let blocking_model = MockCompletionModel::from_turns([
MockTurn::tool_call("tc1", "default_api", json!({"x": 2, "y": 3})),
MockTurn::text("acknowledged"),
]);
let blocking_hook = RecordingHook::default();
let blocking = AgentBuilder::new(blocking_model)
.tool(MockAddTool)
.build()
.runner("do the thing")
.max_turns(3)
.add_hook(blocking_hook.clone())
.add_hook(SkipInvalidHook("tool not permitted"))
.run()
.await
.expect("blocking run should recover via skip");
let streaming_model = MockCompletionModel::from_stream_turns([
vec![
MockStreamEvent::tool_call("tc1", "default_api", json!({"x": 2, "y": 3})),
MockStreamEvent::final_response_with_total_tokens(0),
],
vec![
MockStreamEvent::text("acknowledged"),
MockStreamEvent::final_response_with_total_tokens(0),
],
]);
let streaming_hook = RecordingHook::default();
let mut stream = AgentBuilder::new(streaming_model)
.tool(MockAddTool)
.build()
.runner("do the thing")
.max_turns(3)
.add_hook(streaming_hook.clone())
.add_hook(SkipInvalidHook("tool not permitted"))
.stream()
.await;
let mut final_response = None;
while let Some(item) = stream.next().await {
if let Ok(MultiTurnStreamItem::FinalResponse(resp)) =
item.map_err(|err| panic!("stream item errored: {err}"))
{
final_response = Some(resp);
}
}
let final_response =
final_response.expect("stream should recover and yield a final response");
assert_eq!(blocking.output, "acknowledged");
assert_eq!(final_response.output(), blocking.output);
assert_eq!(
blocking_hook.shared_events(),
streaming_hook.shared_events()
);
assert!(
blocking_hook
.shared_events()
.contains(&StepEventKind::InvalidToolCall),
"the hook must observe the invalid tool call"
);
let blocking_messages = blocking.messages.expect("blocking messages");
let streaming_messages = final_response
.messages()
.expect("streaming history")
.to_vec();
assert_eq!(
serde_json::to_value(&blocking_messages).expect("serialize blocking"),
serde_json::to_value(&streaming_messages).expect("serialize streaming"),
);
assert!(
tool_result_text_in_history(&blocking_messages, "tool not permitted"),
"the verbatim invalid-tool skip reason must be the tool result content"
);
}
#[tokio::test]
async fn recovered_turn_suppresses_response_finish_hook_on_both_drivers() {
let blocking_model = MockCompletionModel::from_turns([
MockTurn::from_contents([
AssistantContent::text("let me compute that"),
AssistantContent::ToolCall(MessageToolCall::new(
"tc1".to_string(),
ToolFunction::new("default_api".to_string(), json!({"x": 2, "y": 3})),
)),
])
.expect("a text + tool-call turn is valid"),
MockTurn::text("the answer is 5"),
]);
let blocking_hook = RecordingHook::default();
let blocking = AgentBuilder::new(blocking_model)
.tool(MockAddTool)
.build()
.runner("compute")
.max_turns(3)
.add_hook(blocking_hook.clone())
.add_hook(RepairInvalidToHook("add"))
.run()
.await
.expect("blocking run should recover via repair");
let streaming_model = MockCompletionModel::from_stream_turns([
vec![
MockStreamEvent::text("let me compute that"),
MockStreamEvent::tool_call("tc1", "default_api", json!({"x": 2, "y": 3})),
MockStreamEvent::final_response_with_total_tokens(0),
],
vec![
MockStreamEvent::text("the answer is 5"),
MockStreamEvent::final_response_with_total_tokens(0),
],
]);
let streaming_hook = RecordingHook::default();
let mut stream = AgentBuilder::new(streaming_model)
.tool(MockAddTool)
.build()
.runner("compute")
.max_turns(3)
.add_hook(streaming_hook.clone())
.add_hook(RepairInvalidToHook("add"))
.stream()
.await;
while stream.next().await.is_some() {}
assert_eq!(blocking.output, "the answer is 5");
assert_eq!(
blocking_hook.count(StepEventKind::CompletionResponse),
1,
"the recovered turn must not fire CompletionResponse"
);
assert_eq!(
streaming_hook.count(StepEventKind::StreamResponseFinish),
1,
"the recovered turn must not fire StreamResponseFinish"
);
assert_eq!(
blocking_hook.count(StepEventKind::CompletionResponse),
streaming_hook.count(StepEventKind::StreamResponseFinish),
);
assert_eq!(
blocking_hook.count(StepEventKind::ModelTurnFinished),
1,
"the recovered turn must not fire ModelTurnFinished"
);
assert_eq!(
streaming_hook.count(StepEventKind::ModelTurnFinished),
1,
"the recovered turn must not fire ModelTurnFinished on the streaming surface either"
);
assert_eq!(
blocking_hook.count(StepEventKind::ModelTurnFinished),
streaming_hook.count(StepEventKind::ModelTurnFinished),
);
}
#[tokio::test]
async fn runner_add_hook_appends_to_agent_default_hooks() {
let agent_hook = RecordingHook::default();
let runner_hook = RecordingHook::default();
AgentBuilder::new(blocking_model())
.tool(MockAddTool)
.add_hook(agent_hook.clone())
.build()
.runner("add 2 and 3")
.max_turns(3)
.add_hook(runner_hook.clone())
.run()
.await
.expect("run should succeed");
assert!(
agent_hook.count(StepEventKind::CompletionCall) >= 1,
"the agent-default hook must still observe the run after a runner-level add_hook"
);
assert!(
runner_hook.count(StepEventKind::CompletionCall) >= 1,
"the runner-level hook must also observe the run"
);
assert_eq!(
agent_hook.count(StepEventKind::CompletionCall),
runner_hook.count(StepEventKind::CompletionCall),
"add_hook appends (both hooks observe every turn); it does not replace"
);
}
struct SkipToolCallHook(&'static str);
impl AgentHook for SkipToolCallHook {
async fn on_tool_call(&self, _ctx: &HookContext, event: ToolCall<'_>) -> ToolCallAction {
if let ToolCall { .. } = event {
ToolCallAction::skip(self.0)
} else {
ToolCallAction::run()
}
}
}
#[tokio::test]
async fn valid_tool_call_skip_parity_across_run_and_stream() {
let turns = [
ScriptedTurn::ToolCalls(vec![add_call("tc1", 2, 3)]),
ScriptedTurn::Text("acknowledged"),
];
let blocking_model =
MockCompletionModel::from_turns(turns.iter().map(ScriptedTurn::as_blocking_turn));
let blocking_hook = RecordingHook::default();
let blocking = AgentBuilder::new(blocking_model)
.tool(MockAddTool)
.build()
.runner("add 2 and 3")
.max_turns(3)
.add_hook(blocking_hook.clone())
.add_hook(SkipToolCallHook("skipped by policy"))
.run()
.await
.expect("blocking run should succeed with a skipped tool call");
let streaming_model = MockCompletionModel::from_stream_turns(
turns
.iter()
.map(|turn| turn.as_stream_events(StreamShape::Complete)),
);
let streaming_hook = RecordingHook::default();
let mut stream = AgentBuilder::new(streaming_model)
.tool(MockAddTool)
.build()
.runner("add 2 and 3")
.max_turns(3)
.add_hook(streaming_hook.clone())
.add_hook(SkipToolCallHook("skipped by policy"))
.stream()
.await;
let mut final_response = None;
while let Some(item) = stream.next().await {
if let Ok(MultiTurnStreamItem::FinalResponse(resp)) =
item.map_err(|err| panic!("stream item errored: {err}"))
{
final_response = Some(resp);
}
}
let final_response = final_response.expect("stream should yield a final response");
assert_eq!(blocking.output, "acknowledged");
assert_eq!(final_response.output(), blocking.output);
assert_eq!(
blocking_hook.shared_events(),
streaming_hook.shared_events()
);
assert_eq!(blocking_hook.tool_results(), streaming_hook.tool_results());
assert_eq!(
blocking_hook.tool_results(),
vec!["skipped by policy".to_string()],
"a skipped tool fires a ToolResult hook with the verbatim skip reason"
);
let blocking_messages = blocking.messages.expect("blocking messages");
let streaming_messages = final_response
.messages()
.expect("streaming history")
.to_vec();
assert_eq!(
serde_json::to_value(&blocking_messages).expect("serialize blocking"),
serde_json::to_value(&streaming_messages).expect("serialize streaming"),
);
assert!(
tool_result_text_in_history(&blocking_messages, "skipped by policy"),
"the verbatim skip reason must be the tool result content in the history"
);
}
struct RewriteToolArgsHook(serde_json::Value);
impl AgentHook for RewriteToolArgsHook {
async fn on_tool_call(&self, _ctx: &HookContext, event: ToolCall<'_>) -> ToolCallAction {
if let ToolCall { .. } = event {
ToolCallAction::rewrite(self.0.clone())
} else {
ToolCallAction::run()
}
}
}
struct EchoStringArgs;
impl Tool for EchoStringArgs {
const NAME: &'static str = "echo_string_args";
type Error = rig::tool::ToolExecutionError;
type Args = String;
type Output = String;
fn description(&self) -> String {
"Echo a JSON string argument".to_string()
}
fn parameters(&self) -> serde_json::Value {
json!({"type": "string"})
}
async fn call(
&self,
_context: &mut ToolContext,
args: Self::Args,
) -> Result<Self::Output, ToolExecutionError> {
Ok(args)
}
}
#[derive(serde::Deserialize)]
struct FirstGenerationArgs {
old: String,
}
struct FirstGenerationTool(Arc<AtomicU32>);
impl Tool for FirstGenerationTool {
const NAME: &'static str = "generation_pinned";
type Error = rig::tool::ToolExecutionError;
type Args = FirstGenerationArgs;
type Output = String;
fn description(&self) -> String {
"first generation schema".to_string()
}
fn parameters(&self) -> serde_json::Value {
json!({
"type": "object",
"properties": {"old": {"type": "string"}},
"required": ["old"]
})
}
async fn call(
&self,
_context: &mut ToolContext,
args: Self::Args,
) -> Result<Self::Output, ToolExecutionError> {
self.0.fetch_add(1, SeqCst);
Ok(format!("first:{}", args.old))
}
}
#[derive(serde::Deserialize)]
struct SecondGenerationArgs {
new: String,
}
struct SecondGenerationTool(Arc<AtomicU32>);
impl Tool for SecondGenerationTool {
const NAME: &'static str = FirstGenerationTool::NAME;
type Error = rig::tool::ToolExecutionError;
type Args = SecondGenerationArgs;
type Output = String;
fn description(&self) -> String {
"second generation schema".to_string()
}
fn parameters(&self) -> serde_json::Value {
json!({
"type": "object",
"properties": {"new": {"type": "string"}},
"required": ["new"]
})
}
async fn call(
&self,
_context: &mut ToolContext,
args: Self::Args,
) -> Result<Self::Output, ToolExecutionError> {
self.0.fetch_add(1, SeqCst);
Ok(format!("second:{}", args.new))
}
}
#[derive(Clone)]
struct PausingCompletionModel {
inner: MockCompletionModel,
request_started: Arc<Notify>,
release_response: Arc<Notify>,
requests: Arc<AtomicU32>,
}
impl PausingCompletionModel {
fn new(inner: MockCompletionModel) -> (Self, Arc<Notify>, Arc<Notify>) {
let request_started = Arc::new(Notify::new());
let release_response = Arc::new(Notify::new());
(
Self {
inner,
request_started: request_started.clone(),
release_response: release_response.clone(),
requests: Arc::new(AtomicU32::new(0)),
},
request_started,
release_response,
)
}
async fn inspect_and_pause(&self, request: &crate::completion::CompletionRequest) {
let request_index = self.requests.fetch_add(1, SeqCst);
let definition = request
.tools
.iter()
.find(|definition| definition.name == FirstGenerationTool::NAME)
.expect("generation tool must be advertised");
if request_index == 0 {
assert_eq!(definition.description, "first generation schema");
self.request_started.notify_one();
self.release_response.notified().await;
} else {
assert_eq!(definition.description, "second generation schema");
}
}
}
impl CompletionModel for PausingCompletionModel {
type Response = crate::test_utils::MockResponse;
type StreamingResponse = crate::test_utils::MockResponse;
type Client = ();
fn make(_: &Self::Client, _: impl Into<String>) -> Self {
Self::new(MockCompletionModel::default()).0
}
async fn completion(
&self,
request: crate::completion::CompletionRequest,
) -> Result<
crate::completion::CompletionResponse<Self::Response>,
crate::completion::CompletionError,
> {
self.inspect_and_pause(&request).await;
self.inner.completion(request).await
}
async fn stream(
&self,
request: crate::completion::CompletionRequest,
) -> Result<
crate::streaming::StreamingCompletionResponse<Self::StreamingResponse>,
crate::completion::CompletionError,
> {
self.inspect_and_pause(&request).await;
self.inner.stream(request).await
}
}
#[test]
fn one_hook_instance_attaches_to_distinct_completion_models() {
#[derive(Clone)]
struct ProviderIndependentHook;
impl AgentHook for ProviderIndependentHook {}
let hook = ProviderIndependentHook;
let _mock_agent = AgentBuilder::new(MockCompletionModel::default())
.add_hook(hook.clone())
.build();
let (other_model, _, _) = PausingCompletionModel::new(MockCompletionModel::default());
let _other_agent = AgentBuilder::new(other_model).add_hook(hook).build();
}
#[test]
fn rewrite_args_resolves_to_proceed_with_for_tool_call() {
let args = json!({"x": 1, "y": 2});
match super::tool_call_decision(ToolCallAction::rewrite(args.clone())) {
super::ToolCallDecision::ProceedWith(replacement) => assert_eq!(replacement, args),
_ => panic!("ToolCallAction::Rewrite should resolve to ProceedWith"),
}
assert_eq!(
ToolCallAction::try_rewrite(&json!({"x": 1, "y": 2})).expect("serializes"),
ToolCallAction::rewrite(json!({"x": 1, "y": 2})),
);
}
#[tokio::test]
async fn valid_tool_call_rewrite_args_parity_across_run_and_stream() {
let turns = [
ScriptedTurn::ToolCalls(vec![add_call("tc1", 2, 3)]),
ScriptedTurn::Text("acknowledged"),
];
let replacement = json!({"x": 2, "y": 40});
let blocking_model =
MockCompletionModel::from_turns(turns.iter().map(ScriptedTurn::as_blocking_turn));
let blocking_hook = RecordingHook::default();
let blocking = AgentBuilder::new(blocking_model)
.tool(MockAddTool)
.build()
.runner("add 2 and 3")
.max_turns(3)
.add_hook(blocking_hook.clone())
.add_hook(RewriteToolArgsHook(replacement.clone()))
.run()
.await
.expect("blocking run should succeed with rewritten tool arguments");
let streaming_model = MockCompletionModel::from_stream_turns(
turns
.iter()
.map(|turn| turn.as_stream_events(StreamShape::Complete)),
);
let streaming_hook = RecordingHook::default();
let mut stream = AgentBuilder::new(streaming_model)
.tool(MockAddTool)
.build()
.runner("add 2 and 3")
.max_turns(3)
.add_hook(streaming_hook.clone())
.add_hook(RewriteToolArgsHook(replacement))
.stream()
.await;
let mut final_response = None;
while let Some(item) = stream.next().await {
if let Ok(MultiTurnStreamItem::FinalResponse(resp)) =
item.map_err(|err| panic!("stream item errored: {err}"))
{
final_response = Some(resp);
}
}
let final_response = final_response.expect("stream should yield a final response");
assert_eq!(blocking_hook.tool_results(), vec!["42".to_string()]);
assert_eq!(blocking.output, "acknowledged");
assert_eq!(final_response.output(), blocking.output);
assert_eq!(
blocking_hook.shared_events(),
streaming_hook.shared_events()
);
assert_eq!(blocking_hook.tool_results(), streaming_hook.tool_results());
}
#[tokio::test]
async fn string_tool_call_without_rewrite_is_canonical_across_run_and_stream() {
let turns = [
ScriptedTurn::ToolCalls(vec![ScriptedToolCall {
id: "tc-string",
name: EchoStringArgs::NAME,
args: json!("original"),
}]),
ScriptedTurn::Text("done"),
];
let blocking_hook = RecordingHook::default();
let blocking = AgentBuilder::new(MockCompletionModel::from_turns(
turns.iter().map(ScriptedTurn::as_blocking_turn),
))
.tool(EchoStringArgs)
.build()
.runner("echo a string")
.max_turns(3)
.add_hook(blocking_hook.clone())
.run()
.await
.expect("blocking string call should execute");
let streaming_hook = RecordingHook::default();
let mut stream = AgentBuilder::new(MockCompletionModel::from_stream_turns(
turns
.iter()
.map(|turn| turn.as_stream_events(StreamShape::Complete)),
))
.tool(EchoStringArgs)
.build()
.runner("echo a string")
.max_turns(3)
.add_hook(streaming_hook.clone())
.stream()
.await;
let mut final_output = None;
while let Some(item) = stream.next().await {
if let MultiTurnStreamItem::FinalResponse(response) =
item.expect("streaming string call should execute")
{
final_output = Some(response.output().to_string());
}
}
assert_eq!(blocking.output, "done");
assert_eq!(final_output.as_deref(), Some("done"));
assert_eq!(blocking_hook.tool_results(), vec!["original"]);
assert_eq!(streaming_hook.tool_results(), vec!["original"]);
}
#[tokio::test]
async fn string_tool_call_rewrite_is_canonical_json_across_run_and_stream() {
let turns = [
ScriptedTurn::ToolCalls(vec![ScriptedToolCall {
id: "tc-string",
name: EchoStringArgs::NAME,
args: json!("original"),
}]),
ScriptedTurn::Text("done"),
];
let replacement = json!("sanitized");
let blocking_hook = RecordingHook::default();
let blocking = AgentBuilder::new(MockCompletionModel::from_turns(
turns.iter().map(ScriptedTurn::as_blocking_turn),
))
.tool(EchoStringArgs)
.build()
.runner("echo a string")
.max_turns(3)
.add_hook(blocking_hook.clone())
.add_hook(RewriteToolArgsHook(replacement.clone()))
.run()
.await
.expect("blocking string rewrite should execute");
let streaming_hook = RecordingHook::default();
let mut stream = AgentBuilder::new(MockCompletionModel::from_stream_turns(
turns
.iter()
.map(|turn| turn.as_stream_events(StreamShape::Complete)),
))
.tool(EchoStringArgs)
.build()
.runner("echo a string")
.max_turns(3)
.add_hook(streaming_hook.clone())
.add_hook(RewriteToolArgsHook(replacement))
.stream()
.await;
let mut final_output = None;
while let Some(item) = stream.next().await {
if let MultiTurnStreamItem::FinalResponse(response) =
item.expect("streaming string rewrite should execute")
{
final_output = Some(response.output().to_string());
}
}
assert_eq!(blocking.output, "done");
assert_eq!(final_output.as_deref(), Some("done"));
assert_eq!(blocking_hook.tool_results(), vec!["sanitized"]);
assert_eq!(streaming_hook.tool_results(), vec!["sanitized"]);
}
#[tokio::test]
async fn blocking_turn_dispatches_the_registry_generation_it_advertised() {
let first_calls = Arc::new(AtomicU32::new(0));
let second_calls = Arc::new(AtomicU32::new(0));
let handle: ToolServerHandle = ToolServer::new()
.tool(FirstGenerationTool(first_calls.clone()))
.run();
let inner = MockCompletionModel::from_turns([
MockTurn::tool_call(
"tc-generation",
FirstGenerationTool::NAME,
json!({"old": "payload"}),
),
MockTurn::text("done"),
]);
let (model, request_started, release_response) = PausingCompletionModel::new(inner);
let runner = AgentBuilder::new(model)
.tool_server_handle(handle.clone())
.build()
.runner("use the generation tool")
.max_turns(3);
let run = runner.run();
let replace = async {
request_started.notified().await;
handle
.add_tool(SecondGenerationTool(second_calls.clone()))
.await;
release_response.notify_one();
};
let (response, ()) = tokio::time::timeout(std::time::Duration::from_secs(2), async {
tokio::join!(run, replace)
})
.await
.expect("in-flight blocking replacement must not hang");
let response = response.expect("blocking run should use its pinned tool generation");
assert_eq!(response.output, "done");
assert_eq!(first_calls.load(SeqCst), 1);
assert_eq!(second_calls.load(SeqCst), 0);
}
#[tokio::test]
async fn streaming_turn_dispatches_the_registry_generation_it_advertised() {
let first_calls = Arc::new(AtomicU32::new(0));
let second_calls = Arc::new(AtomicU32::new(0));
let handle: ToolServerHandle = ToolServer::new()
.tool(FirstGenerationTool(first_calls.clone()))
.run();
let turns = [
ScriptedTurn::ToolCalls(vec![ScriptedToolCall {
id: "tc-generation",
name: FirstGenerationTool::NAME,
args: json!({"old": "payload"}),
}]),
ScriptedTurn::Text("done"),
];
let inner = MockCompletionModel::from_stream_turns(
turns
.iter()
.map(|turn| turn.as_stream_events(StreamShape::Complete)),
);
let (model, request_started, release_response) = PausingCompletionModel::new(inner);
let runner = AgentBuilder::new(model)
.tool_server_handle(handle.clone())
.build()
.runner("use the generation tool")
.max_turns(3);
let drive = async {
let mut stream = runner.stream().await;
let mut final_output = None;
while let Some(item) = stream.next().await {
if let MultiTurnStreamItem::FinalResponse(response) =
item.expect("streaming run should use its pinned tool generation")
{
final_output = Some(response.output().to_string());
}
}
final_output
};
let replace = async {
request_started.notified().await;
handle
.add_tool(SecondGenerationTool(second_calls.clone()))
.await;
release_response.notify_one();
};
let (final_output, ()) = tokio::time::timeout(std::time::Duration::from_secs(2), async {
tokio::join!(drive, replace)
})
.await
.expect("in-flight streaming replacement must not hang");
assert_eq!(final_output.as_deref(), Some("done"));
assert_eq!(first_calls.load(SeqCst), 1);
assert_eq!(second_calls.load(SeqCst), 0);
}
struct RewriteToolResultHook(&'static str);
impl AgentHook for RewriteToolResultHook {
async fn on_tool_result(
&self,
_ctx: &HookContext,
event: ToolResultEvent<'_>,
) -> ToolResultAction {
if let ToolResultEvent { .. } = event {
ToolResultAction::rewrite(self.0)
} else {
ToolResultAction::keep()
}
}
}
#[test]
fn rewrite_result_resolves_to_replace_for_tool_result() {
match super::tool_result_decision(ToolResultAction::rewrite("redacted")) {
super::ToolResultDecision::Replace(result) => {
assert_eq!(result.as_text(), Some("redacted"))
}
_ => panic!("ToolResultAction::Rewrite should resolve to Replace"),
}
}
#[tokio::test]
async fn valid_tool_result_rewrite_parity_across_run_and_stream() {
let turns = [
ScriptedTurn::ToolCalls(vec![add_call("tc1", 2, 3)]),
ScriptedTurn::Text("acknowledged"),
];
let blocking_model =
MockCompletionModel::from_turns(turns.iter().map(ScriptedTurn::as_blocking_turn));
let blocking_hook = RecordingHook::default();
let blocking = AgentBuilder::new(blocking_model)
.tool(MockAddTool)
.build()
.runner("add 2 and 3")
.max_turns(3)
.add_hook(blocking_hook.clone())
.add_hook(RewriteToolResultHook("redacted-result"))
.run()
.await
.expect("blocking run should succeed with a rewritten tool result");
let streaming_model = MockCompletionModel::from_stream_turns(
turns
.iter()
.map(|turn| turn.as_stream_events(StreamShape::Complete)),
);
let streaming_hook = RecordingHook::default();
let mut stream = AgentBuilder::new(streaming_model)
.tool(MockAddTool)
.build()
.runner("add 2 and 3")
.max_turns(3)
.add_hook(streaming_hook.clone())
.add_hook(RewriteToolResultHook("redacted-result"))
.stream()
.await;
let mut final_response = None;
while let Some(item) = stream.next().await {
if let Ok(MultiTurnStreamItem::FinalResponse(resp)) =
item.map_err(|err| panic!("stream item errored: {err}"))
{
final_response = Some(resp);
}
}
let final_response = final_response.expect("stream should yield a final response");
assert_eq!(blocking.output, "acknowledged");
assert_eq!(final_response.output(), blocking.output);
assert_eq!(blocking_hook.tool_results(), vec!["5".to_string()]);
assert_eq!(blocking_hook.tool_results(), streaming_hook.tool_results());
let blocking_messages = blocking.messages.expect("blocking messages");
let streaming_messages = final_response
.messages()
.expect("streaming history")
.to_vec();
assert_eq!(
serde_json::to_value(&blocking_messages).expect("serialize blocking"),
serde_json::to_value(&streaming_messages).expect("serialize streaming"),
);
assert!(
tool_result_text_in_history(&blocking_messages, "redacted-result"),
"the model-visible tool result must be the hook's replacement"
);
assert!(
!tool_result_text_in_history(&blocking_messages, "5"),
"the tool's original output must not reach the model after a rewrite"
);
}
#[tokio::test]
async fn rewrite_result_is_delivered_verbatim_not_reparsed() {
const IMAGE_JSON: &str = r#"{"type":"image","data":"abc","mimeType":"image/png"}"#;
let turns = [
ScriptedTurn::ToolCalls(vec![add_call("tc1", 2, 3)]),
ScriptedTurn::Text("done"),
];
let model =
MockCompletionModel::from_turns(turns.iter().map(ScriptedTurn::as_blocking_turn));
let result = AgentBuilder::new(model)
.tool(MockAddTool)
.build()
.runner("add 2 and 3")
.max_turns(3)
.add_hook(RewriteToolResultHook(IMAGE_JSON))
.run()
.await
.expect("run should succeed with a JSON-shaped rewritten result");
let messages = result.messages.expect("messages");
assert!(
tool_result_text_in_history(&messages, IMAGE_JSON),
"the JSON-shaped replacement must reach history verbatim as text, not be \
re-parsed into a structured/image content block"
);
}
struct PatchRequestHook;
impl AgentHook for PatchRequestHook {
async fn on_completion_call(
&self,
_ctx: &HookContext,
event: CompletionCallEvent<'_>,
) -> CompletionCallAction {
if let CompletionCallEvent { .. } = event {
CompletionCallAction::patch(
RequestPatch::new()
.preamble(OVERRIDE_PREAMBLE)
.temperature(0.25)
.max_tokens(OVERRIDE_MAX_TOKENS)
.tool_choice(ToolChoice::Required)
.active_tools(["add"])
.additional_params(json!({"injected": true})),
)
} else {
CompletionCallAction::continue_run()
}
}
}
const OVERRIDE_PREAMBLE: &str = "overridden: critical-step instructions";
const OVERRIDE_MAX_TOKENS: u64 = 512;
#[test]
fn patch_request_resolves_to_patch_for_completion_call() {
let patch = RequestPatch::new()
.temperature(0.25)
.tool_choice(ToolChoice::Required);
match super::completion_call_decision(CompletionCallAction::patch(patch.clone())) {
super::CompletionCallDecision::Patch(got) => assert_eq!(got, patch),
_ => panic!("PatchRequest should resolve to Patch for a completion call"),
}
}
#[tokio::test]
async fn patch_request_parity_across_run_and_stream() {
fn assert_request(req: &crate::completion::CompletionRequest) {
assert_eq!(
req.temperature,
Some(0.25),
"override temperature wins over the agent's 0.9"
);
assert_eq!(
req.max_tokens,
Some(OVERRIDE_MAX_TOKENS),
"override max_tokens wins over the agent's 64"
);
let system = req.chat_history.iter().find_map(|m| match m {
Message::System { content } => Some(content.as_str()),
_ => None,
});
assert_eq!(
system,
Some(OVERRIDE_PREAMBLE),
"override preamble wins over the agent's baseline and is the leading system message"
);
assert!(matches!(req.tool_choice, Some(ToolChoice::Required)));
let tool_names: Vec<&str> = req.tools.iter().map(|t| t.name.as_str()).collect();
assert_eq!(
tool_names,
["add"],
"active_tools narrows the advertised set to `add` (drops `subtract`)"
);
let params = req.additional_params.as_ref().expect("additional_params");
assert_eq!(params.get("runner").and_then(|v| v.as_str()), Some("keep"));
assert_eq!(params.get("injected").and_then(|v| v.as_bool()), Some(true));
assert!(params.get("baseline").is_none());
}
let blocking_model = MockCompletionModel::from_turns([MockTurn::text("done")]);
let blocking_probe = blocking_model.clone();
let blocking = AgentBuilder::new(blocking_model)
.tool(MockAddTool)
.tool(MockSubtractTool)
.preamble("baseline preamble")
.temperature(0.9)
.max_tokens(64)
.additional_params(json!({"baseline": "keep"}))
.add_hook(PatchRequestHook)
.build()
.runner("go")
.replace_additional_params(json!({"runner": "keep", "injected": false}))
.max_turns(2)
.run()
.await
.expect("blocking run should succeed");
assert_eq!(blocking.output, "done");
let blocking_requests = blocking_probe.requests();
assert_eq!(blocking_requests.len(), 1);
assert_request(&blocking_requests[0]);
let streaming_model = MockCompletionModel::from_stream_turns([
ScriptedTurn::Text("done").as_stream_events(StreamShape::Complete)
]);
let streaming_probe = streaming_model.clone();
let mut stream = AgentBuilder::new(streaming_model)
.tool(MockAddTool)
.tool(MockSubtractTool)
.preamble("baseline preamble")
.temperature(0.9)
.max_tokens(64)
.additional_params(json!({"baseline": "keep"}))
.add_hook(PatchRequestHook)
.build()
.runner("go")
.replace_additional_params(json!({"runner": "keep", "injected": false}))
.max_turns(2)
.stream()
.await;
while let Some(item) = stream.next().await {
let _ = item.map_err(|err| panic!("stream item errored: {err}"));
}
let streaming_requests = streaming_probe.requests();
assert_eq!(streaming_requests.len(), 1);
assert_request(&streaming_requests[0]);
}
fn hook_doc(id: &str, text: &str) -> crate::completion::Document {
crate::completion::Document {
id: id.to_string(),
text: text.to_string(),
additional_props: Default::default(),
}
}
struct ExtraContextHook {
id: &'static str,
text: &'static str,
}
impl AgentHook for ExtraContextHook {
async fn on_completion_call(
&self,
_ctx: &HookContext,
event: CompletionCallEvent<'_>,
) -> CompletionCallAction {
if let CompletionCallEvent { .. } = event {
CompletionCallAction::patch(
RequestPatch::new().context(hook_doc(self.id, self.text)),
)
} else {
CompletionCallAction::continue_run()
}
}
}
struct ExtraContextTurnOneHook;
impl AgentHook for ExtraContextTurnOneHook {
async fn on_completion_call(
&self,
_ctx: &HookContext,
event: CompletionCallEvent<'_>,
) -> CompletionCallAction {
if let CompletionCallEvent { turn, .. } = event
&& turn == 1
{
return CompletionCallAction::patch(
RequestPatch::new().context(hook_doc("turn-one", "only turn 1")),
);
}
CompletionCallAction::continue_run()
}
}
#[derive(Clone)]
struct RecordingContextIndex {
id: &'static str,
queries: Arc<Mutex<Vec<(String, u64)>>>,
}
impl VectorStoreIndex for RecordingContextIndex {
type Filter = Filter<serde_json::Value>;
async fn top_n<T: for<'a> Deserialize<'a> + WasmCompatSend>(
&self,
req: VectorSearchRequest,
) -> Result<Vec<(f64, String, T)>, VectorStoreError> {
self.queries
.lock()
.expect("context query recorder lock")
.push((req.query().to_string(), req.samples()));
let value = serde_json::from_value(json!({ "source": self.id }))?;
Ok(vec![(1.0, self.id.to_string(), value)])
}
async fn top_n_ids(
&self,
_req: VectorSearchRequest,
) -> Result<Vec<(f64, String)>, VectorStoreError> {
Ok(vec![(1.0, self.id.to_string())])
}
}
struct FailingContextIndex;
impl VectorStoreIndex for FailingContextIndex {
type Filter = Filter<serde_json::Value>;
async fn top_n<T: for<'a> Deserialize<'a> + WasmCompatSend>(
&self,
_req: VectorSearchRequest,
) -> Result<Vec<(f64, String, T)>, VectorStoreError> {
Err(VectorStoreError::BuilderError(
"context index unavailable".to_string(),
))
}
async fn top_n_ids(
&self,
_req: VectorSearchRequest,
) -> Result<Vec<(f64, String)>, VectorStoreError> {
Err(VectorStoreError::BuilderError(
"context index unavailable".to_string(),
))
}
}
struct QueryRecordingToolIndex {
queries: Arc<Mutex<Vec<String>>>,
}
impl VectorStoreIndex for QueryRecordingToolIndex {
type Filter = Filter<serde_json::Value>;
async fn top_n<T: for<'a> Deserialize<'a> + WasmCompatSend>(
&self,
_req: VectorSearchRequest,
) -> Result<Vec<(f64, String, T)>, VectorStoreError> {
Ok(Vec::new())
}
async fn top_n_ids(
&self,
req: VectorSearchRequest,
) -> Result<Vec<(f64, String)>, VectorStoreError> {
self.queries
.lock()
.expect("query recorder lock")
.push(req.query().to_string());
Ok(vec![(1.0, MockAddTool::NAME.to_string())])
}
}
fn one_text_stream_turn(text: &'static str) -> Vec<MockStreamEvent> {
vec![
MockStreamEvent::text(text),
MockStreamEvent::final_response_with_total_tokens(0),
]
}
#[tokio::test]
async fn extra_context_appears_after_static_context_on_both_surfaces() {
fn assert_docs(req: &crate::completion::CompletionRequest) {
let ids: Vec<&str> = req.documents.iter().map(|d| d.id.as_str()).collect();
let static_pos = ids
.iter()
.position(|id| id.starts_with("static_doc"))
.expect("static context document present");
let extra_pos = ids
.iter()
.position(|id| *id == "hook-doc")
.expect("hook extra_context document present");
assert!(
static_pos < extra_pos,
"static context precedes hook extras: {ids:?}"
);
assert!(
req.documents.iter().any(|d| d.text == "injected"),
"the hook document's text is present"
);
}
let blocking_model = MockCompletionModel::from_turns([MockTurn::text("done")]);
let blocking_probe = blocking_model.clone();
AgentBuilder::new(blocking_model)
.context("static context text")
.add_hook(ExtraContextHook {
id: "hook-doc",
text: "injected",
})
.build()
.runner("go")
.run()
.await
.expect("blocking run should succeed");
assert_docs(blocking_probe.requests().first().expect("one request"));
let streaming_model =
MockCompletionModel::from_stream_turns([one_text_stream_turn("done")]);
let streaming_probe = streaming_model.clone();
let mut stream = AgentBuilder::new(streaming_model)
.context("static context text")
.add_hook(ExtraContextHook {
id: "hook-doc",
text: "injected",
})
.build()
.runner("go")
.stream()
.await;
while let Some(item) = stream.next().await {
let _ = item.map_err(|err| panic!("stream item errored: {err}"));
}
assert_docs(streaming_probe.requests().first().expect("one request"));
}
#[tokio::test]
async fn multiple_hooks_extra_context_append_in_registration_order() {
let model = MockCompletionModel::from_turns([MockTurn::text("done")]);
let probe = model.clone();
AgentBuilder::new(model)
.add_hook(ExtraContextHook {
id: "first",
text: "1",
})
.add_hook(ExtraContextHook {
id: "second",
text: "2",
})
.build()
.runner("go")
.run()
.await
.expect("run should succeed");
let requests = probe.requests();
let req = requests.first().expect("one request");
let ids: Vec<&str> = req.documents.iter().map(|d| d.id.as_str()).collect();
assert_eq!(
ids,
vec!["first", "second"],
"hook extras append in registration order"
);
}
#[tokio::test]
async fn dynamic_context_preserves_query_selection_formatting_and_order_on_both_surfaces() {
fn assert_documents(request: &crate::completion::CompletionRequest) {
let documents = request
.documents
.iter()
.map(|document| (document.id.as_str(), document.text.as_str()))
.collect::<Vec<_>>();
assert_eq!(
documents,
vec![
("static_doc_0", "static context"),
("blocking", "{\n \"source\": \"blocking\"\n}"),
]
);
}
let blocking_queries = Arc::new(Mutex::new(Vec::new()));
let blocking_model = MockCompletionModel::from_turns([MockTurn::text("done")]);
let blocking_probe = blocking_model.clone();
AgentBuilder::new(blocking_model)
.context("static context")
.dynamic_context(
2,
RecordingContextIndex {
id: "blocking",
queries: blocking_queries.clone(),
},
)
.build()
.runner("current blocking query")
.history(vec![Message::user("ignored history query")])
.run()
.await
.expect("blocking dynamic-context run should succeed");
assert_eq!(
*blocking_queries.lock().expect("blocking queries"),
vec![("current blocking query".to_string(), 2)]
);
assert_documents(blocking_probe.requests().first().expect("one request"));
let streaming_queries = Arc::new(Mutex::new(Vec::new()));
let streaming_model =
MockCompletionModel::from_stream_turns([one_text_stream_turn("done")]);
let streaming_probe = streaming_model.clone();
let mut stream = AgentBuilder::new(streaming_model)
.dynamic_context(
3,
RecordingContextIndex {
id: "streaming",
queries: streaming_queries.clone(),
},
)
.build()
.runner(Message::User {
content: OneOrMany::one(UserContent::image_url(
"https://example.com/prompt.png",
None,
None,
)),
})
.history(vec![
Message::user("older history query"),
Message::user("latest history query"),
])
.stream()
.await;
while let Some(item) = stream.next().await {
item.expect("streaming dynamic-context run should succeed");
}
assert_eq!(
*streaming_queries.lock().expect("streaming queries"),
vec![("latest history query".to_string(), 3)]
);
let streaming_requests = streaming_probe.requests();
let request = streaming_requests.first().expect("one request");
assert_eq!(request.documents.len(), 1);
assert_eq!(request.documents[0].id, "streaming");
assert_eq!(
request.documents[0].text,
"{\n \"source\": \"streaming\"\n}"
);
}
#[tokio::test]
async fn dynamic_context_and_application_hooks_follow_registration_order() {
let queries = Arc::new(Mutex::new(Vec::new()));
let model = MockCompletionModel::from_turns([MockTurn::text("done")]);
let probe = model.clone();
AgentBuilder::new(model)
.context("static")
.add_hook(ExtraContextHook {
id: "before",
text: "before dynamic context",
})
.dynamic_context(
1,
RecordingContextIndex {
id: "first",
queries: queries.clone(),
},
)
.add_hook(ExtraContextHook {
id: "between",
text: "between dynamic contexts",
})
.dynamic_context(
2,
RecordingContextIndex {
id: "second",
queries: queries.clone(),
},
)
.add_hook(ExtraContextHook {
id: "after",
text: "after dynamic context",
})
.build()
.runner("query")
.run()
.await
.expect("run should succeed");
assert_eq!(
probe.requests()[0]
.documents
.iter()
.map(|document| document.id.as_str())
.collect::<Vec<_>>(),
vec![
"static_doc_0",
"before",
"first",
"between",
"second",
"after",
]
);
assert_eq!(
*queries.lock().expect("context queries"),
vec![("query".to_string(), 1), ("query".to_string(), 2)]
);
let skipped_queries = Arc::new(Mutex::new(Vec::new()));
let error = AgentBuilder::new(MockCompletionModel::from_turns([MockTurn::text("unused")]))
.add_hook(TerminateOn(StepEventKind::CompletionCall))
.dynamic_context(
1,
RecordingContextIndex {
id: "skipped",
queries: skipped_queries.clone(),
},
)
.build()
.runner("query")
.run()
.await
.expect_err("an earlier stop hook should terminate before retrieval");
assert!(matches!(error, PromptError::PromptCancelled { .. }));
assert!(skipped_queries.lock().expect("skipped queries").is_empty());
}
#[tokio::test]
async fn dynamic_context_retrieval_failure_stops_before_provider_io_on_both_surfaces() {
let blocking_model = MockCompletionModel::from_turns([MockTurn::text("unused")]);
let blocking_probe = blocking_model.clone();
let error = AgentBuilder::new(blocking_model)
.dynamic_context(1, FailingContextIndex)
.build()
.runner("retrieve this")
.run()
.await
.expect_err("failed retrieval should stop the run");
assert!(matches!(
error,
PromptError::PromptCancelled { reason, .. }
if reason.contains("context index unavailable")
));
assert_eq!(blocking_probe.request_count(), 0);
let streaming_model =
MockCompletionModel::from_stream_turns([one_text_stream_turn("unused")]);
let streaming_probe = streaming_model.clone();
let mut stream = AgentBuilder::new(streaming_model)
.dynamic_context(1, FailingContextIndex)
.build()
.runner("retrieve this")
.stream()
.await;
let error = stream
.next()
.await
.expect("stream should report retrieval failure")
.expect_err("failed retrieval should stop the stream");
assert!(matches!(
error,
StreamingError::Prompt(prompt_error)
if matches!(
prompt_error.as_ref(),
PromptError::PromptCancelled { reason, .. }
if reason.contains("context index unavailable")
)
));
assert_eq!(streaming_probe.request_count(), 0);
}
#[tokio::test]
async fn retrieved_tool_query_selection_is_unchanged_on_both_surfaces() {
let queries = Arc::new(Mutex::new(Vec::new()));
AgentBuilder::new(MockCompletionModel::from_turns([MockTurn::text("done")]))
.retrieved_tools(
1,
QueryRecordingToolIndex {
queries: queries.clone(),
},
ToolSet::from_tools(vec![MockAddTool]),
)
.build()
.runner("blocking retrieval query")
.history(vec![Message::user("blocking history query")])
.run()
.await
.expect("blocking run should succeed");
AgentBuilder::new(MockCompletionModel::from_turns([MockTurn::text("done")]))
.retrieved_tools(
1,
QueryRecordingToolIndex {
queries: queries.clone(),
},
ToolSet::from_tools(vec![MockAddTool]),
)
.build()
.runner(Message::User {
content: OneOrMany::one(UserContent::image_url(
"https://example.com/blocking.png",
None,
None,
)),
})
.history(vec![
Message::user("older blocking history query"),
Message::user("latest blocking history query"),
])
.run()
.await
.expect("blocking history fallback should succeed");
let mut stream = AgentBuilder::new(MockCompletionModel::from_stream_turns([
one_text_stream_turn("done"),
]))
.retrieved_tools(
1,
QueryRecordingToolIndex {
queries: queries.clone(),
},
ToolSet::from_tools(vec![MockAddTool]),
)
.build()
.runner("streaming retrieval query")
.history(vec![Message::user("streaming history query")])
.stream()
.await;
while let Some(item) = stream.next().await {
item.expect("stream item should succeed");
}
let mut stream = AgentBuilder::new(MockCompletionModel::from_stream_turns([
one_text_stream_turn("done"),
]))
.retrieved_tools(
1,
QueryRecordingToolIndex {
queries: queries.clone(),
},
ToolSet::from_tools(vec![MockAddTool]),
)
.build()
.runner(Message::User {
content: OneOrMany::one(UserContent::image_url(
"https://example.com/streaming.png",
None,
None,
)),
})
.history(vec![
Message::user("older streaming history query"),
Message::user("latest streaming history query"),
])
.stream()
.await;
while let Some(item) = stream.next().await {
item.expect("stream item should succeed");
}
assert_eq!(
*queries.lock().expect("query recorder lock"),
vec![
"blocking retrieval query",
"latest blocking history query",
"streaming retrieval query",
"latest streaming history query",
]
);
}
#[tokio::test]
async fn extra_context_is_per_turn_non_sticky() {
fn assert_turns(requests: &[crate::completion::CompletionRequest]) {
assert_eq!(requests.len(), 2, "two model turns");
let turn1 = requests.first().expect("turn 1");
let turn2 = requests.get(1).expect("turn 2");
assert!(
turn1.documents.iter().any(|d| d.id == "turn-one"),
"turn 1 carries the injected document"
);
assert!(
turn2.documents.iter().all(|d| d.id != "turn-one"),
"turn 2 does not inherit turn 1's per-turn document"
);
}
let blocking_probe = blocking_model();
let probe = blocking_probe.clone();
AgentBuilder::new(blocking_probe)
.tool(MockAddTool)
.add_hook(ExtraContextTurnOneHook)
.build()
.runner("add 2 and 3")
.max_turns(3)
.run()
.await
.expect("blocking run should succeed");
assert_turns(&probe.requests());
let streaming = streaming_model();
let stream_probe = streaming.clone();
let mut stream = AgentBuilder::new(streaming)
.tool(MockAddTool)
.add_hook(ExtraContextTurnOneHook)
.build()
.runner("add 2 and 3")
.max_turns(3)
.stream()
.await;
while let Some(item) = stream.next().await {
let _ = item.map_err(|err| panic!("stream item errored: {err}"));
}
assert_turns(&stream_probe.requests());
}
#[tokio::test]
async fn history_patch_changes_sent_messages_not_transcript_on_both_surfaces() {
const SENTINEL: &str = "COMPACTED-HISTORY-SENTINEL";
struct HistoryOverrideHook;
impl AgentHook for HistoryOverrideHook {
async fn on_completion_call(
&self,
_ctx: &HookContext,
event: CompletionCallEvent<'_>,
) -> CompletionCallAction {
if let CompletionCallEvent { .. } = event {
CompletionCallAction::patch(
RequestPatch::new().history([Message::user(SENTINEL)]),
)
} else {
CompletionCallAction::continue_run()
}
}
}
fn request_has_sentinel(req: &crate::completion::CompletionRequest) -> bool {
req.chat_history.iter().any(|m| match m {
Message::User { content } => content
.iter()
.any(|c| matches!(c, UserContent::Text(text) if text.text.contains(SENTINEL))),
_ => false,
})
}
fn messages_have_sentinel(messages: &[Message]) -> bool {
messages.iter().any(|m| match m {
Message::User { content } => content
.iter()
.any(|c| matches!(c, UserContent::Text(text) if text.text.contains(SENTINEL))),
_ => false,
})
}
let blocking_model = MockCompletionModel::from_turns([MockTurn::text("done")]);
let blocking_probe = blocking_model.clone();
let blocking = AgentBuilder::new(blocking_model)
.add_hook(HistoryOverrideHook)
.build()
.runner("real prompt")
.run()
.await
.expect("blocking run should succeed");
assert!(
request_has_sentinel(blocking_probe.requests().first().expect("one request")),
"the overridden history reaches the provider"
);
assert!(
!messages_have_sentinel(blocking.messages.as_deref().unwrap_or_default()),
"the persisted transcript is untouched by the per-turn history override"
);
let streaming_model =
MockCompletionModel::from_stream_turns([one_text_stream_turn("done")]);
let streaming_probe = streaming_model.clone();
let stream = AgentBuilder::new(streaming_model)
.add_hook(HistoryOverrideHook)
.build()
.runner("real prompt")
.stream()
.await;
let final_response = drive_to_final_response(stream).await;
assert!(
request_has_sentinel(streaming_probe.requests().first().expect("one request")),
"the overridden history reaches the provider on the streaming surface too"
);
assert!(
!messages_have_sentinel(final_response.messages().expect("history")),
"the persisted transcript is untouched by the per-turn history override on \
the streaming surface too"
);
}
#[tokio::test]
async fn model_turn_finished_fires_once_per_accepted_turn_including_tool_only() {
let blocking_hook = RecordingHook::default();
AgentBuilder::new(blocking_model())
.tool(MockAddTool)
.add_hook(blocking_hook.clone())
.build()
.runner("add 2 and 3")
.max_turns(3)
.run()
.await
.expect("blocking run should succeed");
assert_eq!(
blocking_hook.count(StepEventKind::ModelTurnFinished),
2,
"one ModelTurnFinished per accepted turn (tool turn + text turn)"
);
let streaming_hook = RecordingHook::default();
let mut stream = AgentBuilder::new(streaming_model())
.tool(MockAddTool)
.add_hook(streaming_hook.clone())
.build()
.runner("add 2 and 3")
.max_turns(3)
.stream()
.await;
while let Some(item) = stream.next().await {
let _ = item.map_err(|err| panic!("stream item errored: {err}"));
}
assert_eq!(
streaming_hook.count(StepEventKind::ModelTurnFinished),
2,
"ModelTurnFinished fires once per turn on the streaming surface too"
);
assert_eq!(
streaming_hook.count(StepEventKind::StreamResponseFinish),
1,
"the tool-only turn fires no StreamResponseFinish"
);
}
#[tokio::test]
async fn reasoning_only_turn_does_not_gain_stream_response_finish() {
let hook = RecordingHook::default();
let mut stream = AgentBuilder::new(MockCompletionModel::from_stream_turns([[
MockStreamEvent::reasoning("think"),
MockStreamEvent::final_response_with_total_tokens(0),
]]))
.add_hook(hook.clone())
.build()
.runner("reason")
.stream()
.await;
while let Some(item) = stream.next().await {
item.expect("reasoning-only stream item");
}
assert_eq!(
hook.count(StepEventKind::StreamResponseFinish),
0,
"reasoning-only turns must not fire StreamResponseFinish"
);
assert_eq!(
hook.count(StepEventKind::ModelTurnFinished),
1,
"the accepted reasoning-only turn still fires ModelTurnFinished"
);
}
#[derive(Clone, Default)]
struct CaptureFirstTurnContent {
kinds: Arc<Mutex<Option<Vec<&'static str>>>>,
}
impl AgentHook for CaptureFirstTurnContent {
async fn on_model_turn_finished(
&self,
_ctx: &HookContext,
event: ModelTurnFinished<'_>,
) -> ModelTurnAction {
if let ModelTurnFinished { turn, content, .. } = event
&& turn == 1
{
let kinds = content
.iter()
.map(|c| match c {
AssistantContent::Reasoning(_) => "reasoning",
AssistantContent::Text(_) => "text",
AssistantContent::ToolCall(_) => "tool_call",
_ => "other",
})
.collect();
*self.kinds.lock().expect("kinds") = Some(kinds);
}
ModelTurnAction::continue_run()
}
}
#[tokio::test]
async fn streaming_model_turn_finished_carries_canonical_committed_content() {
let model = MockCompletionModel::from_stream_turns([
vec![
MockStreamEvent::reasoning("think"),
MockStreamEvent::tool_call("tc1", "add", json!({"x": 2, "y": 3})),
MockStreamEvent::text("answer"),
MockStreamEvent::final_response_with_total_tokens(0),
],
vec![
MockStreamEvent::text("done"),
MockStreamEvent::final_response_with_total_tokens(0),
],
]);
let hook = CaptureFirstTurnContent::default();
let stream = AgentBuilder::new(model)
.tool(MockAddTool)
.add_hook(hook.clone())
.build()
.runner("go")
.max_turns(3)
.stream()
.await;
let _ = drive_to_final_response(stream).await;
assert_eq!(
hook.kinds.lock().expect("kinds").clone(),
Some(vec!["reasoning", "text", "tool_call"]),
"ModelTurnFinished carries the canonical reasoning->text->tool ordering \
from StreamedTurn::finish, not the raw stream.choice emission order"
);
}
#[tokio::test]
async fn chained_rewrites_compose_across_hooks() {
struct SetArg {
key: &'static str,
value: i64,
}
impl AgentHook for SetArg {
async fn on_tool_call(
&self,
_ctx: &HookContext,
event: ToolCall<'_>,
) -> ToolCallAction {
if let ToolCall { args, .. } = event {
let mut parsed: serde_json::Value =
serde_json::from_str(args).unwrap_or_else(|_| json!({}));
parsed[self.key] = json!(self.value);
ToolCallAction::rewrite(parsed)
} else {
ToolCallAction::run()
}
}
}
struct WrapResult(&'static str);
impl AgentHook for WrapResult {
async fn on_tool_result(
&self,
_ctx: &HookContext,
event: ToolResultEvent<'_>,
) -> ToolResultAction {
if let ToolResultEvent { presentation, .. } = event {
ToolResultAction::rewrite(format!("{}({})", self.0, presentation.render()))
} else {
ToolResultAction::keep()
}
}
}
let recorder = RecordingHook::default();
let blocking = AgentBuilder::new(blocking_model())
.tool(MockAddTool)
.add_hook(SetArg {
key: "y",
value: 40,
})
.add_hook(SetArg {
key: "x",
value: 100,
})
.add_hook(WrapResult("A"))
.add_hook(WrapResult("B"))
.add_hook(recorder.clone())
.build()
.runner("add 2 and 3")
.max_turns(3)
.run()
.await
.expect("blocking run should succeed");
assert_eq!(blocking.output, "the answer is 5");
assert_eq!(
recorder.tool_results(),
vec!["B(A(140))".to_string()],
"arg rewrites compose (100+40=140) and result rewrites nest B(A(...))"
);
let stream_recorder = RecordingHook::default();
let mut stream = AgentBuilder::new(streaming_model())
.tool(MockAddTool)
.add_hook(SetArg {
key: "y",
value: 40,
})
.add_hook(SetArg {
key: "x",
value: 100,
})
.add_hook(WrapResult("A"))
.add_hook(WrapResult("B"))
.add_hook(stream_recorder.clone())
.build()
.runner("add 2 and 3")
.max_turns(3)
.stream()
.await;
while let Some(item) = stream.next().await {
let _ = item.map_err(|err| panic!("stream item errored: {err}"));
}
assert_eq!(
stream_recorder.tool_results(),
vec!["B(A(140))".to_string()],
"chained rewrites compose identically on the streaming surface"
);
}
#[derive(serde::Deserialize, schemars::JsonSchema)]
#[allow(dead_code)]
struct Answer {
answer: String,
}
struct FinalResultTool;
impl Tool for FinalResultTool {
const NAME: &'static str = "final_result";
type Error = MockToolError;
type Args = serde_json::Value;
type Output = String;
fn description(&self) -> String {
"A real tool sharing the default output-tool name".to_string()
}
fn parameters(&self) -> serde_json::Value {
json!({ "type": "object", "properties": {} })
}
async fn call(
&self,
_context: &mut ToolContext,
_args: Self::Args,
) -> Result<Self::Output, Self::Error> {
Ok("real final_result output".to_string())
}
}
#[derive(Default)]
struct LateFinalResultIndex {
searches: AtomicU32,
}
impl VectorStoreIndex for LateFinalResultIndex {
type Filter = Filter<serde_json::Value>;
async fn top_n<T: for<'a> Deserialize<'a> + WasmCompatSend>(
&self,
_req: VectorSearchRequest,
) -> Result<Vec<(f64, String, T)>, VectorStoreError> {
Ok(Vec::new())
}
async fn top_n_ids(
&self,
_req: VectorSearchRequest,
) -> Result<Vec<(f64, String)>, VectorStoreError> {
if self.searches.fetch_add(1, SeqCst) == 0 {
Ok(Vec::new())
} else {
Ok(vec![(1.0, "final_result".to_string())])
}
}
}
#[derive(Clone)]
struct RegisterLateFinalResultTool {
handle: ToolServerHandle,
second_turn_patch: Option<RequestPatch>,
}
impl AgentHook for RegisterLateFinalResultTool {
async fn on_model_turn_finished(
&self,
ctx: &HookContext,
_event: ModelTurnFinished<'_>,
) -> ModelTurnAction {
if ctx.turn() == 1 {
self.handle.add_tool(FinalResultTool).await;
}
ModelTurnAction::continue_run()
}
async fn on_completion_call(
&self,
ctx: &HookContext,
_event: CompletionCallEvent<'_>,
) -> CompletionCallAction {
if ctx.turn() == 2
&& let Some(patch) = &self.second_turn_patch
{
return CompletionCallAction::patch(patch.clone());
}
CompletionCallAction::continue_run()
}
}
fn assert_structured_output_collision_error(message: &str) {
assert!(
message.contains("final_result"),
"error should name the conflicting tool: {message}"
);
assert!(
message.contains("structured-output") && message.contains("reserved"),
"error should explain the structured-output reservation: {message}"
);
assert!(
message.contains("rename or remove"),
"error should provide an actionable resolution: {message}"
);
}
#[tokio::test]
async fn initial_output_tool_collision_uses_a_unique_synthetic_name() {
let model = MockCompletionModel::from_turns([
MockTurn::tool_call("real", "final_result", json!({})),
MockTurn::tool_call("output", "final_result_1", json!({ "answer": "done" })),
]);
let probe = model.clone();
let response = AgentBuilder::new(model)
.tool(FinalResultTool)
.output_schema::<Answer>()
.output_mode(OutputMode::Tool)
.build()
.runner("go")
.max_turns(2)
.run()
.await
.expect("the real tool should dispatch before the unique output tool finalizes");
assert!(response.output.contains("done"));
let requests = probe.requests();
assert_eq!(
requests.len(),
2,
"real-tool dispatch must continue to a second model turn"
);
let tool_names = requests[0]
.tools
.iter()
.map(|tool| tool.name.as_str())
.collect::<Vec<_>>();
assert_eq!(tool_names.len(), 2);
for expected in ["final_result", "final_result_1"] {
assert_eq!(
tool_names.iter().filter(|name| **name == expected).count(),
1,
"the first request should advertise `{expected}` exactly once: {tool_names:?}"
);
}
assert!(
requests[1].chat_history.iter().any(|message| matches!(
message,
Message::User { content }
if content.iter().any(|item| matches!(
item,
UserContent::ToolResult(result)
if result.id == "real"
&& result.content.iter().any(|content| matches!(
content,
rig_core::message::ToolResultContent::Text(text)
if text.text == "real final_result output"
))
))
)),
"the real `final_result` call must execute normally and its result must reach the follow-up request"
);
}
#[tokio::test]
async fn late_output_tool_collision_fails_before_blocking_provider_for_all_choices() {
let cases = [
("inherited", None),
(
"required",
Some(RequestPatch::new().tool_choice(ToolChoice::Required)),
),
(
"none",
Some(RequestPatch::new().tool_choice(ToolChoice::None)),
),
(
"specific",
Some(RequestPatch::new().tool_choice(ToolChoice::Specific {
function_names: vec!["final_result".to_string()],
})),
),
];
for (case, second_turn_patch) in cases {
let handle = ToolServer::new().tool(MockAddTool).run();
let model = MockCompletionModel::from_turns([
MockTurn::tool_call("add-1", "add", json!({ "x": 1, "y": 2 })),
MockTurn::tool_call(
"shadowed",
"final_result",
json!({ "answer": "wrongly finalized" }),
),
]);
let probe = model.clone();
let err = AgentBuilder::new(model)
.tool_server_handle(handle.clone())
.output_schema::<Answer>()
.output_mode(OutputMode::Tool)
.add_hook(RegisterLateFinalResultTool {
handle,
second_turn_patch,
})
.build()
.runner("go")
.max_turns(3)
.run()
.await
.unwrap_err();
assert!(
matches!(
&err,
PromptError::CompletionError(CompletionError::RequestError(_))
),
"{case}: expected a local completion request error, got {err:?}"
);
assert_eq!(
probe.request_count(),
1,
"{case}: the colliding second request must not reach the provider"
);
assert_structured_output_collision_error(&err.to_string());
}
}
#[tokio::test]
async fn late_output_tool_collision_fails_before_streaming_provider() {
let handle = ToolServer::new().tool(MockAddTool).run();
let model = MockCompletionModel::from_stream_turns([
vec![
MockStreamEvent::tool_call("add-1", "add", json!({ "x": 1, "y": 2 })),
MockStreamEvent::final_response_with_total_tokens(0),
],
vec![
MockStreamEvent::tool_call(
"shadowed",
"final_result",
json!({ "answer": "wrongly finalized" }),
),
MockStreamEvent::final_response_with_total_tokens(0),
],
]);
let probe = model.clone();
let mut stream = AgentBuilder::new(model)
.tool_server_handle(handle.clone())
.output_schema::<Answer>()
.output_mode(OutputMode::Tool)
.add_hook(RegisterLateFinalResultTool {
handle,
second_turn_patch: None,
})
.build()
.runner("go")
.max_turns(3)
.stream()
.await;
let mut collisions = Vec::new();
let mut saw_final_response = false;
while let Some(item) = stream.next().await {
match item {
Err(err) => collisions.push(err),
Ok(MultiTurnStreamItem::FinalResponse(_)) => saw_final_response = true,
Ok(_) => {}
}
}
assert_eq!(
collisions.len(),
1,
"the stream should terminate with exactly one collision error"
);
assert!(
!saw_final_response,
"a collision error must terminate the stream without a final response"
);
let err = collisions.pop().expect("one collision error was asserted");
assert!(
matches!(
&err,
StreamingError::Completion(CompletionError::RequestError(_))
),
"expected a local streaming completion request error, got {err:?}"
);
assert_eq!(
probe.request_count(),
1,
"the colliding second stream must not reach the provider"
);
assert_structured_output_collision_error(&err.to_string());
}
#[tokio::test]
async fn late_output_tool_collision_is_checked_after_active_tools_filtering() {
let handle = ToolServer::new().tool(MockAddTool).run();
let model = MockCompletionModel::from_turns([
MockTurn::tool_call("add-1", "add", json!({ "x": 1, "y": 2 })),
MockTurn::tool_call("add-2", "add", json!({ "x": 3, "y": 4 })),
MockTurn::tool_call(
"shadowed",
"final_result",
json!({ "answer": "wrongly finalized" }),
),
]);
let probe = model.clone();
let err = AgentBuilder::new(model)
.tool_server_handle(handle.clone())
.output_schema::<Answer>()
.output_mode(OutputMode::Tool)
.add_hook(RegisterLateFinalResultTool {
handle,
second_turn_patch: Some(RequestPatch::new().active_tools(["add"])),
})
.build()
.runner("go")
.max_turns(4)
.run()
.await
.expect_err("the exposed third-turn collision should fail locally");
assert_eq!(
probe.request_count(),
2,
"the filtered second turn may run, but the exposed third turn may not"
);
let requests = probe.requests();
let second_turn_names = requests[1]
.tools
.iter()
.map(|tool| tool.name.as_str())
.collect::<Vec<_>>();
assert_eq!(second_turn_names.len(), 2);
for expected in ["add", "final_result"] {
assert_eq!(
second_turn_names
.iter()
.filter(|name| **name == expected)
.count(),
1,
"the second request should advertise `{expected}` exactly once: \
{second_turn_names:?}"
);
}
assert_structured_output_collision_error(&err.to_string());
}
#[tokio::test]
async fn retrieved_output_tool_collision_fails_before_provider_request() {
let mut retrieved_tools = ToolSet::default();
retrieved_tools.add_tool(FinalResultTool);
let handle = ToolServer::new()
.tool(MockAddTool)
.retrieved_tools(1, LateFinalResultIndex::default(), retrieved_tools)
.run();
let model = MockCompletionModel::from_turns([
MockTurn::tool_call("add-1", "add", json!({ "x": 1, "y": 2 })),
MockTurn::tool_call(
"shadowed",
"final_result",
json!({ "answer": "wrongly finalized" }),
),
]);
let probe = model.clone();
let err = AgentBuilder::new(model)
.tool_server_handle(handle)
.output_schema::<Answer>()
.output_mode(OutputMode::Tool)
.build()
.runner("go")
.max_turns(3)
.run()
.await
.expect_err("the retrieved second-turn collision should fail locally");
assert!(matches!(
&err,
PromptError::CompletionError(CompletionError::RequestError(_))
));
assert_eq!(
probe.request_count(),
1,
"the colliding retrieved tool must prevent the second provider request"
);
assert_structured_output_collision_error(&err.to_string());
}
struct ActiveToolsAddOnly;
impl AgentHook for ActiveToolsAddOnly {
async fn on_completion_call(
&self,
_ctx: &HookContext,
event: CompletionCallEvent<'_>,
) -> CompletionCallAction {
if let CompletionCallEvent { .. } = event {
CompletionCallAction::patch(RequestPatch::new().active_tools(["add"]))
} else {
CompletionCallAction::continue_run()
}
}
}
#[tokio::test]
async fn active_tools_filter_does_not_let_output_tool_collide_with_a_filtered_real_tool() {
let model = MockCompletionModel::from_turns([MockTurn::tool_call(
"out1",
"final_result_1",
json!({ "answer": "done" }),
)]);
let probe = model.clone();
let response = AgentBuilder::new(model)
.tool(MockAddTool)
.tool(FinalResultTool)
.output_schema::<Answer>()
.output_mode(OutputMode::Tool)
.add_hook(ActiveToolsAddOnly)
.build()
.runner("go")
.max_turns(2)
.run()
.await
.expect("run should finalize via the picked output tool `final_result_1`");
assert!(
response.output.contains("done"),
"the intercepted output-tool call should produce the structured result, \
got {:?}",
response.output
);
let requests = probe.requests();
assert!(
!requests.is_empty(),
"the first model request should be captured"
);
let tool_names: Vec<&str> = requests[0].tools.iter().map(|t| t.name.as_str()).collect();
assert!(
tool_names.contains(&"add"),
"active_tools keeps `add` advertised, saw {tool_names:?}"
);
assert!(
tool_names.contains(&"final_result_1"),
"the synthetic output tool must avoid the filtered real `final_result` name, \
saw {tool_names:?}"
);
assert!(
!tool_names.contains(&"final_result"),
"the real `final_result` is filtered out and the output tool must not reuse \
its name, saw {tool_names:?}"
);
}
#[derive(Clone, Default)]
struct CaptureOutputToolInModelTurn {
saw_output_tool_call: Arc<Mutex<bool>>,
}
impl AgentHook for CaptureOutputToolInModelTurn {
async fn on_model_turn_finished(
&self,
_ctx: &HookContext,
event: ModelTurnFinished<'_>,
) -> ModelTurnAction {
if let ModelTurnFinished { content, .. } = event
&& content.iter().any(|c| {
matches!(c, AssistantContent::ToolCall(tc) if tc.function.name == "final_result")
})
{
*self.saw_output_tool_call.lock().expect("lock") = true;
}
ModelTurnAction::continue_run()
}
}
#[tokio::test]
async fn model_turn_finished_content_carries_output_tool_call_in_tool_mode() {
let hook = CaptureOutputToolInModelTurn::default();
let response = AgentBuilder::new(MockCompletionModel::from_turns([MockTurn::tool_call(
"out1",
"final_result",
json!({ "answer": "done" }),
)]))
.output_schema::<Answer>()
.output_mode(OutputMode::Tool)
.add_hook(hook.clone())
.build()
.runner("go")
.max_turns(2)
.run()
.await
.expect("run should finalize via the output tool");
assert!(
*hook.saw_output_tool_call.lock().expect("lock"),
"ModelTurnFinished.content must carry the model-emitted output-tool call (blocking)"
);
assert!(
response.output.contains("done"),
"the run finalizes with the structured output, not the raw tool call: {:?}",
response.output
);
let s_hook = CaptureOutputToolInModelTurn::default();
let mut stream = AgentBuilder::new(MockCompletionModel::from_stream_turns([vec![
MockStreamEvent::tool_call("out1", "final_result", json!({ "answer": "done" })),
MockStreamEvent::final_response_with_total_tokens(0),
]]))
.output_schema::<Answer>()
.output_mode(OutputMode::Tool)
.add_hook(s_hook.clone())
.build()
.runner("go")
.max_turns(2)
.stream()
.await;
while stream.next().await.is_some() {}
assert!(
*s_hook.saw_output_tool_call.lock().expect("lock"),
"ModelTurnFinished.content must carry the model-emitted output-tool call (streaming)"
);
}
#[tokio::test]
async fn output_tool_finalization_emits_no_complete_tool_call_stream_item() {
let mut stream = AgentBuilder::new(MockCompletionModel::from_stream_turns([vec![
MockStreamEvent::tool_call("out1", "final_result", json!({ "answer": "done" })),
MockStreamEvent::final_response_with_total_tokens(0),
]]))
.output_schema::<Answer>()
.output_mode(OutputMode::Tool)
.build()
.runner("go")
.max_turns(2)
.stream()
.await;
let mut saw_complete_output_tool_call = false;
let mut final_has_output = false;
while let Some(item) = stream.next().await {
match item.expect("stream item") {
MultiTurnStreamItem::StreamAssistantItem(StreamedAssistantContent::ToolCall {
tool_call,
..
}) if tool_call.function.name == "final_result" => {
saw_complete_output_tool_call = true;
}
MultiTurnStreamItem::FinalResponse(res) => {
final_has_output = res.output().contains("done");
}
_ => {}
}
}
assert!(
!saw_complete_output_tool_call,
"the output-tool call finalizes the run, so no complete \
StreamAssistantItem::ToolCall item must be emitted for it"
);
assert!(
final_has_output,
"the structured output must be surfaced via the FinalResponse"
);
}
enum Decision {
Approve,
Deny(&'static str),
Edit(serde_json::Value),
Abort(&'static str),
}
#[derive(Clone)]
struct HumanApprovalHook {
decisions: Arc<Mutex<std::collections::VecDeque<Decision>>>,
reviewed: Arc<Mutex<Vec<String>>>,
}
impl HumanApprovalHook {
fn new(decisions: impl IntoIterator<Item = Decision>) -> Self {
Self {
decisions: Arc::new(Mutex::new(decisions.into_iter().collect())),
reviewed: Arc::new(Mutex::new(Vec::new())),
}
}
fn reviewed(&self) -> Vec<String> {
self.reviewed.lock().unwrap().clone()
}
}
impl AgentHook for HumanApprovalHook {
async fn on_tool_call(&self, _ctx: &HookContext, event: ToolCall<'_>) -> ToolCallAction {
let ToolCall {
tool_name, args, ..
} = event
else {
return ToolCallAction::run();
};
self.reviewed
.lock()
.unwrap()
.push(format!("{tool_name}({args})"));
let decision = self.decisions.lock().unwrap().pop_front();
match decision {
Some(Decision::Approve) => ToolCallAction::run(),
Some(Decision::Deny(reason)) => ToolCallAction::skip(reason),
Some(Decision::Edit(args)) => ToolCallAction::rewrite(args),
Some(Decision::Abort(reason)) => ToolCallAction::stop(reason),
None => ToolCallAction::skip("denied: no scripted decision (fail-closed)"),
}
}
}
#[tokio::test]
async fn human_in_the_loop_approve_deny_edit_parity_across_run_and_stream() {
let turns = [
ScriptedTurn::ToolCalls(vec![
add_call("tc1", 2, 3), add_call("tc2", 10, 20), add_call("tc3", 1, 1), ]),
ScriptedTurn::Text("done"),
];
let denial = "denied by reviewer: amount too large";
let decisions = || {
vec![
Decision::Approve,
Decision::Deny(denial),
Decision::Edit(json!({"x": 1, "y": 100})),
]
};
let blocking_model =
MockCompletionModel::from_turns(turns.iter().map(ScriptedTurn::as_blocking_turn));
let blocking_recorder = RecordingHook::default();
let blocking_approver = HumanApprovalHook::new(decisions());
let blocking = AgentBuilder::new(blocking_model)
.tool(MockAddTool)
.build()
.runner("carry out the plan")
.max_turns(3)
.add_hook(blocking_recorder.clone())
.add_hook(blocking_approver.clone())
.run()
.await
.expect("blocking HITL run should succeed");
let streaming_model = MockCompletionModel::from_stream_turns(
turns
.iter()
.map(|turn| turn.as_stream_events(StreamShape::Complete)),
);
let streaming_recorder = RecordingHook::default();
let streaming_approver = HumanApprovalHook::new(decisions());
let mut stream = AgentBuilder::new(streaming_model)
.tool(MockAddTool)
.build()
.runner("carry out the plan")
.max_turns(3)
.add_hook(streaming_recorder.clone())
.add_hook(streaming_approver.clone())
.stream()
.await;
let mut final_response = None;
while let Some(item) = stream.next().await {
if let Ok(MultiTurnStreamItem::FinalResponse(resp)) =
item.map_err(|err| panic!("stream item errored: {err}"))
{
final_response = Some(resp);
}
}
let final_response = final_response.expect("stream should yield a final response");
assert_eq!(
blocking_recorder.tool_results(),
vec![
"5".to_string(),
"denied by reviewer: amount too large".to_string(),
"101".to_string()
]
);
assert_eq!(
blocking_recorder.tool_results(),
streaming_recorder.tool_results()
);
assert!(
!blocking_recorder.tool_results().contains(&"30".to_string()),
"the denied call must not have executed"
);
let reviewed = blocking_approver.reviewed();
assert_eq!(reviewed.len(), 3);
assert_eq!(reviewed, streaming_approver.reviewed());
assert!(
reviewed[0].contains('2') && reviewed[0].contains('3'),
"first reviewed call should be add(2, 3): {reviewed:?}"
);
assert!(
reviewed[1].contains("10") && reviewed[1].contains("20"),
"the denied (second) call should be add(10, 20): {reviewed:?}"
);
assert_eq!(blocking.output, "done");
assert_eq!(final_response.output(), blocking.output);
assert_eq!(
blocking_recorder.shared_events(),
streaming_recorder.shared_events()
);
let blocking_messages = blocking.messages.expect("blocking messages");
let streaming_messages = final_response
.messages()
.expect("streaming history")
.to_vec();
assert_eq!(
serde_json::to_value(&blocking_messages).expect("serialize blocking"),
serde_json::to_value(&streaming_messages).expect("serialize streaming"),
);
assert!(
tool_result_text_in_history(&blocking_messages, denial),
"the denial reason must be the denied call's tool result in the history"
);
assert!(
tool_result_json_in_history(&blocking_messages, &json!(101)),
"the edited call must have executed with the rewritten arguments"
);
}
#[tokio::test]
async fn human_in_the_loop_abort_terminates_the_run() {
let turns = [
ScriptedTurn::ToolCalls(vec![add_call("tc1", 2, 3)]),
ScriptedTurn::Text("unreachable"),
];
const ABORT_REASON: &str = "aborted by the human reviewer";
let blocking_model =
MockCompletionModel::from_turns(turns.iter().map(ScriptedTurn::as_blocking_turn));
let err = AgentBuilder::new(blocking_model)
.tool(MockAddTool)
.build()
.runner("do the sensitive thing")
.max_turns(3)
.add_hook(HumanApprovalHook::new([Decision::Abort(ABORT_REASON)]))
.run()
.await
.expect_err("an aborted tool call should terminate the blocking run");
assert!(
format!("{err}").contains(ABORT_REASON),
"the abort reason should surface in the blocking error, got: {err}"
);
let streaming_model = MockCompletionModel::from_stream_turns(
turns
.iter()
.map(|turn| turn.as_stream_events(StreamShape::Complete)),
);
let mut stream = AgentBuilder::new(streaming_model)
.tool(MockAddTool)
.build()
.runner("do the sensitive thing")
.max_turns(3)
.add_hook(HumanApprovalHook::new([Decision::Abort(ABORT_REASON)]))
.stream()
.await;
let mut stream_error = None;
while let Some(item) = stream.next().await {
match item {
Err(err) => stream_error = Some(format!("{err}")),
Ok(MultiTurnStreamItem::FinalResponse(resp)) => {
panic!("aborted stream must not finalize, got: {}", resp.output())
}
Ok(_) => {}
}
}
let stream_error = stream_error.expect("an aborted tool call should error the stream");
assert!(
stream_error.contains(ABORT_REASON),
"the abort reason should surface in the streaming error, got: {stream_error}"
);
}
#[derive(Clone)]
struct PolicyHook {
auto_approve: std::collections::HashSet<&'static str>,
evaluated: Arc<Mutex<Vec<String>>>,
cache: Arc<Mutex<std::collections::HashMap<String, bool>>>,
}
impl PolicyHook {
fn new(auto_approve: impl IntoIterator<Item = &'static str>) -> Self {
Self {
auto_approve: auto_approve.into_iter().collect(),
evaluated: Arc::new(Mutex::new(Vec::new())),
cache: Arc::new(Mutex::new(std::collections::HashMap::new())),
}
}
fn evaluated(&self) -> Vec<String> {
self.evaluated.lock().unwrap().clone()
}
}
impl AgentHook for PolicyHook {
async fn on_tool_call(&self, _ctx: &HookContext, event: ToolCall<'_>) -> ToolCallAction {
let ToolCall { tool_name, .. } = event else {
return ToolCallAction::run();
};
let cached = self.cache.lock().unwrap().get(tool_name).copied();
let approved = match cached {
Some(decision) => decision, None => {
self.evaluated.lock().unwrap().push(tool_name.to_string());
let decision = self.auto_approve.contains(tool_name);
self.cache
.lock()
.unwrap()
.insert(tool_name.to_string(), decision);
decision
}
};
if approved {
ToolCallAction::run()
} else {
ToolCallAction::skip(format!("denied by policy: `{tool_name}` not allowed"))
}
}
}
#[tokio::test]
async fn approval_policy_allow_list_with_sticky_decisions() {
let turns = [
ScriptedTurn::ToolCalls(vec![
add_call("c1", 2, 3),
ScriptedToolCall {
id: "c2",
name: "subtract",
args: json!({ "x": 10, "y": 4 }),
},
add_call("c3", 2, 3),
]),
ScriptedTurn::Text("done"),
];
let model =
MockCompletionModel::from_turns(turns.iter().map(ScriptedTurn::as_blocking_turn));
let recorder = RecordingHook::default();
let policy = PolicyHook::new(["add"]);
let out = AgentBuilder::new(model)
.tool(MockAddTool)
.tool(MockSubtractTool)
.build()
.runner("go")
.max_turns(3)
.add_hook(recorder.clone())
.add_hook(policy.clone())
.run()
.await
.expect("policy run should succeed");
assert_eq!(out.output, "done");
assert_eq!(
recorder.tool_results(),
vec![
"5".to_string(),
"denied by policy: `subtract` not allowed".to_string(),
"5".to_string()
]
);
assert_eq!(
policy.evaluated(),
vec!["add".to_string(), "subtract".to_string()]
);
let messages = out.messages.expect("messages");
assert!(
tool_result_text_in_history(&messages, "denied by policy: `subtract` not allowed"),
"the policy denial reason must reach the model as the subtract tool result"
);
}
static NEXT_RESPONSE_RETRY_HOOK_ID: AtomicU64 = AtomicU64::new(1);
#[derive(Clone, Default)]
struct ResponseRetryAttempts(HashMap<u64, usize>);
#[derive(Clone)]
enum TestRetryMode {
Repeat,
Feedback(&'static str),
}
#[derive(Clone)]
struct BoundedResponseRetry {
id: u64,
rejected_text: &'static str,
max_retries: usize,
mode: TestRetryMode,
}
#[derive(Clone, Default)]
struct StatefulCompletionPatch {
calls: Arc<AtomicU32>,
}
impl StatefulCompletionPatch {
fn calls(&self) -> u32 {
self.calls.load(SeqCst)
}
}
impl AgentHook for StatefulCompletionPatch {
async fn on_completion_call(
&self,
_ctx: &HookContext,
_event: crate::agent::CompletionCallEvent<'_>,
) -> CompletionCallAction {
let call = self.calls.fetch_add(1, SeqCst);
CompletionCallAction::patch(RequestPatch::new().temperature(if call == 0 {
0.1
} else {
0.9
}))
}
}
impl BoundedResponseRetry {
fn new(rejected_text: &'static str, max_retries: usize, mode: TestRetryMode) -> Self {
Self {
id: NEXT_RESPONSE_RETRY_HOOK_ID.fetch_add(1, SeqCst),
rejected_text,
max_retries,
mode,
}
}
}
impl AgentHook for BoundedResponseRetry {
async fn on_model_turn_finished(
&self,
ctx: &HookContext,
event: ModelTurnFinished<'_>,
) -> ModelTurnAction {
let rejected = event.content.iter().any(
|content| matches!(content, AssistantContent::Text(text) if text.text == self.rejected_text),
);
if !rejected {
return ModelTurnAction::continue_run();
}
let attempt = ctx
.scratchpad()
.update::<ResponseRetryAttempts, _>(|attempts| {
let attempt = attempts.0.entry(self.id).or_default();
*attempt += 1;
*attempt
});
if attempt > self.max_retries {
return ModelTurnAction::stop(format!(
"response retry limit ({}) exceeded",
self.max_retries
));
}
match self.mode {
TestRetryMode::Repeat => ModelTurnAction::repeat(),
TestRetryMode::Feedback(feedback) => ModelTurnAction::retry_with_feedback(feedback),
}
}
}
fn retry_usage(input_tokens: u64, output_tokens: u64) -> Usage {
Usage {
input_tokens,
output_tokens,
total_tokens: input_tokens + output_tokens,
..Usage::new()
}
}
#[tokio::test]
async fn blocking_model_turn_repeat_preserves_prompt_history_with_fresh_preparation() {
let first_usage = retry_usage(10, 3);
let second_usage = retry_usage(7, 2);
let completion_patch = StatefulCompletionPatch::default();
let model = MockCompletionModel::from_turns([
MockTurn::text("rejected").with_usage(first_usage),
MockTurn::text("accepted").with_usage(second_usage),
]);
let response = AgentBuilder::new(model.clone())
.add_hook(completion_patch.clone())
.add_hook(BoundedResponseRetry::new(
"rejected",
1,
TestRetryMode::Repeat,
))
.build()
.runner("question")
.max_turns(2)
.run()
.await
.expect("repeat should recover");
assert_eq!(response.output, "accepted");
assert_eq!(response.usage, first_usage + second_usage);
assert_eq!(response.completion_calls.len(), 2);
let messages = response.messages.expect("response messages");
assert_eq!(
messages,
vec![Message::user("question"), Message::assistant("accepted")]
);
let requests = model.requests();
assert_eq!(requests.len(), 2);
let first = requests[0].chat_history.iter().cloned().collect::<Vec<_>>();
let second = requests[1].chat_history.iter().cloned().collect::<Vec<_>>();
assert_eq!(first, vec![Message::user("question")]);
assert_eq!(
second, first,
"Repeat must preserve the prompt and preceding history"
);
assert_eq!(requests[0].temperature, Some(0.1));
assert_eq!(requests[1].temperature, Some(0.9));
assert_eq!(completion_patch.calls(), 2);
}
#[tokio::test]
async fn blocking_model_turn_feedback_preserves_rejected_response() {
let model = MockCompletionModel::from_turns([
MockTurn::text("rejected"),
MockTurn::text("accepted"),
]);
let response = AgentBuilder::new(model.clone())
.add_hook(BoundedResponseRetry::new(
"rejected",
1,
TestRetryMode::Feedback("try another approach"),
))
.build()
.runner("question")
.max_turns(2)
.run()
.await
.expect("feedback retry should recover");
assert_eq!(response.output, "accepted");
assert_eq!(
response.messages.expect("response messages"),
vec![
Message::user("question"),
Message::assistant("rejected"),
Message::user("try another approach"),
Message::assistant("accepted"),
]
);
let second_request = &model.requests()[1];
assert_eq!(
second_request
.chat_history
.iter()
.cloned()
.collect::<Vec<_>>(),
vec![
Message::user("question"),
Message::assistant("rejected"),
Message::user("try another approach"),
]
);
}
#[tokio::test]
async fn blocking_empty_feedback_retry_omits_empty_assistant_history() {
let first_usage = retry_usage(5, 1);
let second_usage = retry_usage(7, 2);
let model = MockCompletionModel::from_turns([
MockTurn::text("").with_usage(first_usage),
MockTurn::text("accepted").with_usage(second_usage),
]);
let response = AgentBuilder::new(model.clone())
.add_hook(BoundedResponseRetry::new(
"",
1,
TestRetryMode::Feedback("provide an answer"),
))
.build()
.runner("question")
.max_turns(2)
.run()
.await
.expect("feedback retry should recover from an empty turn");
assert_eq!(response.output, "accepted");
assert_eq!(response.usage, first_usage + second_usage);
assert_eq!(response.completion_calls.len(), 2);
assert_eq!(
response.messages.expect("response messages"),
vec![
Message::user("question"),
Message::user("provide an answer"),
Message::assistant("accepted"),
]
);
assert_eq!(
model.requests()[1]
.chat_history
.iter()
.cloned()
.collect::<Vec<_>>(),
vec![
Message::user("question"),
Message::user("provide an answer"),
],
"the retry request must not contain an empty assistant message"
);
}
#[tokio::test]
async fn streaming_model_turn_retry_marks_rollback_and_matches_blocking_accounting() {
let first_usage = retry_usage(10, 3);
let second_usage = retry_usage(7, 2);
let model = MockCompletionModel::from_stream_turns([
[
MockStreamEvent::text("rejected"),
MockStreamEvent::final_response(first_usage),
],
[
MockStreamEvent::text("accepted"),
MockStreamEvent::final_response(second_usage),
],
]);
let mut stream = AgentBuilder::new(model.clone())
.add_hook(BoundedResponseRetry::new(
"rejected",
1,
TestRetryMode::Repeat,
))
.build()
.runner("question")
.max_turns(2)
.stream()
.await;
let mut retries = Vec::new();
let mut provider_finals = 0;
let mut completion_calls = 0;
let mut final_response = None;
while let Some(item) = stream.next().await {
match item.expect("stream item") {
MultiTurnStreamItem::ModelTurnRetried { turn } => retries.push(turn),
MultiTurnStreamItem::CompletionCall(_) => completion_calls += 1,
MultiTurnStreamItem::StreamAssistantItem(StreamedAssistantContent::Final(_)) => {
provider_finals += 1
}
MultiTurnStreamItem::FinalResponse(response) => final_response = Some(response),
_ => {}
}
}
assert_eq!(retries, vec![1]);
assert_eq!(
provider_finals, 1,
"the rejected provider final is suppressed"
);
assert_eq!(completion_calls, 2);
let response = final_response.expect("run final response");
assert_eq!(response.output, "accepted");
assert_eq!(response.usage, first_usage + second_usage);
assert_eq!(response.completion_calls.len(), 2);
assert_eq!(
response.messages.expect("response messages"),
vec![Message::user("question"), Message::assistant("accepted")]
);
assert_eq!(model.requests().len(), 2);
}
#[tokio::test]
async fn streaming_feedback_retry_matches_blocking_history_and_usage() {
let first_usage = retry_usage(5, 2);
let second_usage = retry_usage(8, 4);
let blocking = AgentBuilder::new(MockCompletionModel::from_turns([
MockTurn::text("rejected").with_usage(first_usage),
MockTurn::text("accepted").with_usage(second_usage),
]))
.add_hook(BoundedResponseRetry::new(
"rejected",
1,
TestRetryMode::Feedback("correct the answer"),
))
.build()
.runner("question")
.max_turns(2)
.run()
.await
.expect("blocking feedback retry");
let mut stream = AgentBuilder::new(MockCompletionModel::from_stream_turns([
[
MockStreamEvent::text("rejected"),
MockStreamEvent::final_response(first_usage),
],
[
MockStreamEvent::text("accepted"),
MockStreamEvent::final_response(second_usage),
],
]))
.add_hook(BoundedResponseRetry::new(
"rejected",
1,
TestRetryMode::Feedback("correct the answer"),
))
.build()
.runner("question")
.max_turns(2)
.stream()
.await;
let mut saw_retry = false;
let mut streaming = None;
while let Some(item) = stream.next().await {
match item.expect("stream item") {
MultiTurnStreamItem::ModelTurnRetried { turn: 1 } => saw_retry = true,
MultiTurnStreamItem::FinalResponse(response) => streaming = Some(response),
_ => {}
}
}
let streaming = streaming.expect("streaming final response");
assert!(saw_retry);
assert_eq!(streaming.output, blocking.output);
assert_eq!(streaming.usage, blocking.usage);
assert_eq!(streaming.completion_calls, blocking.completion_calls);
assert_eq!(
serde_json::to_value(streaming.messages).expect("streaming history"),
serde_json::to_value(blocking.messages).expect("blocking history")
);
}
#[tokio::test]
async fn streaming_empty_feedback_retry_omits_empty_assistant_history() {
let first_usage = retry_usage(5, 1);
let second_usage = retry_usage(7, 2);
let model = MockCompletionModel::from_stream_turns([
[
MockStreamEvent::text(""),
MockStreamEvent::final_response(first_usage),
],
[
MockStreamEvent::text("accepted"),
MockStreamEvent::final_response(second_usage),
],
]);
let mut stream = AgentBuilder::new(model.clone())
.add_hook(BoundedResponseRetry::new(
"",
1,
TestRetryMode::Feedback("provide an answer"),
))
.build()
.runner("question")
.max_turns(2)
.stream()
.await;
let mut retries = Vec::new();
let mut provider_finals = 0;
let mut completion_calls = 0;
let mut final_response = None;
while let Some(item) = stream.next().await {
match item.expect("stream item") {
MultiTurnStreamItem::ModelTurnRetried { turn } => retries.push(turn),
MultiTurnStreamItem::CompletionCall(_) => completion_calls += 1,
MultiTurnStreamItem::StreamAssistantItem(StreamedAssistantContent::Final(_)) => {
provider_finals += 1;
}
MultiTurnStreamItem::FinalResponse(response) => final_response = Some(response),
_ => {}
}
}
assert_eq!(retries, vec![1]);
assert_eq!(provider_finals, 1, "the rejected final is suppressed");
assert_eq!(completion_calls, 2);
let response = final_response.expect("run final response");
assert_eq!(response.output, "accepted");
assert_eq!(response.usage, first_usage + second_usage);
assert_eq!(response.completion_calls.len(), 2);
assert_eq!(
response.messages.expect("response messages"),
vec![
Message::user("question"),
Message::user("provide an answer"),
Message::assistant("accepted"),
]
);
assert_eq!(
model.requests()[1]
.chat_history
.iter()
.cloned()
.collect::<Vec<_>>(),
vec![
Message::user("question"),
Message::user("provide an answer"),
],
"the retry request must not contain an empty assistant message"
);
}
#[tokio::test]
async fn response_retry_preserves_model_turn_hook_order_across_surfaces() {
let blocking_events = RecordingHook::default();
AgentBuilder::new(MockCompletionModel::from_turns([
MockTurn::text("rejected"),
MockTurn::text("accepted"),
]))
.add_hook(blocking_events.clone())
.add_hook(BoundedResponseRetry::new(
"rejected",
1,
TestRetryMode::Repeat,
))
.build()
.runner("question")
.max_turns(2)
.run()
.await
.expect("blocking retry");
let streaming_events = RecordingHook::default();
let mut stream = AgentBuilder::new(MockCompletionModel::from_stream_turns([
[
MockStreamEvent::text("rejected"),
MockStreamEvent::final_response_with_default_usage(),
],
[
MockStreamEvent::text("accepted"),
MockStreamEvent::final_response_with_default_usage(),
],
]))
.add_hook(streaming_events.clone())
.add_hook(BoundedResponseRetry::new(
"rejected",
1,
TestRetryMode::Repeat,
))
.build()
.runner("question")
.max_turns(2)
.stream()
.await;
while let Some(item) = stream.next().await {
item.expect("streaming retry item");
}
let shared_order = |events: &RecordingHook| {
events
.events
.lock()
.expect("events")
.iter()
.copied()
.filter(|event| {
matches!(
event,
StepEventKind::CompletionCall | StepEventKind::ModelTurnFinished
)
})
.collect::<Vec<_>>()
};
let expected = vec![
StepEventKind::CompletionCall,
StepEventKind::ModelTurnFinished,
StepEventKind::CompletionCall,
StepEventKind::ModelTurnFinished,
];
assert_eq!(shared_order(&blocking_events), expected);
assert_eq!(shared_order(&streaming_events), expected);
let blocking_order = blocking_events.events.lock().expect("events").clone();
assert_eq!(
blocking_order,
vec![
StepEventKind::CompletionCall,
StepEventKind::CompletionResponse,
StepEventKind::ModelTurnFinished,
StepEventKind::CompletionCall,
StepEventKind::CompletionResponse,
StepEventKind::ModelTurnFinished,
]
);
let streaming_order = streaming_events.events.lock().expect("events").clone();
assert_eq!(
streaming_order,
vec![
StepEventKind::CompletionCall,
StepEventKind::TextDelta,
StepEventKind::StreamResponseFinish,
StepEventKind::ModelTurnFinished,
StepEventKind::CompletionCall,
StepEventKind::TextDelta,
StepEventKind::StreamResponseFinish,
StepEventKind::ModelTurnFinished,
]
);
}
#[tokio::test]
async fn streaming_model_turn_retry_respects_max_turns() {
let model = MockCompletionModel::from_stream_turns([[
MockStreamEvent::text("rejected"),
MockStreamEvent::final_response_with_default_usage(),
]]);
let mut stream = AgentBuilder::new(model)
.add_hook(BoundedResponseRetry::new(
"rejected",
1,
TestRetryMode::Repeat,
))
.build()
.runner("question")
.max_turns(1)
.stream()
.await;
let mut saw_rollback = false;
let mut error = None;
while let Some(item) = stream.next().await {
match item {
Ok(MultiTurnStreamItem::ModelTurnRetried { turn: 1 }) => saw_rollback = true,
Ok(_) => {}
Err(err) => error = Some(err),
}
}
assert!(saw_rollback);
assert!(matches!(
error,
Some(StreamingError::Prompt(error))
if matches!(error.as_ref(), PromptError::MaxTurnsError { max_turns: 1, .. })
));
}
struct AlwaysRepeatModelTurn;
impl AgentHook for AlwaysRepeatModelTurn {
async fn on_model_turn_finished(
&self,
_ctx: &HookContext,
_event: ModelTurnFinished<'_>,
) -> ModelTurnAction {
ModelTurnAction::repeat()
}
}
#[tokio::test]
async fn model_turn_retry_rejects_tool_turn_before_tool_hooks_or_execution() {
let recorder = RecordingHook::default();
let executions = Arc::new(AtomicU32::new(0));
let err = AgentBuilder::new(MockCompletionModel::from_turns([MockTurn::tool_call(
"tc1",
"add",
json!({"x": 1, "y": 2}),
)]))
.tool(CountingAddTool {
calls: executions.clone(),
})
.add_hook(recorder.clone())
.add_hook(AlwaysRepeatModelTurn)
.build()
.runner("add")
.max_turns(2)
.run()
.await
.expect_err("tool-bearing retry must fail closed");
let PromptError::PromptCancelled {
chat_history,
reason,
} = err
else {
panic!("tool-bearing retry should return PromptCancelled");
};
assert!(reason.contains("tool-bearing model turns"));
assert!(reason.contains("tool-call hooks"));
assert_eq!(chat_history, vec![Message::user("add")]);
assert_eq!(recorder.count(StepEventKind::ToolCall), 0);
assert_eq!(recorder.count(StepEventKind::ToolResult), 0);
assert_eq!(executions.load(SeqCst), 0);
}
#[tokio::test]
async fn streaming_model_turn_retry_rejects_tool_turn_without_committed_execution() {
let recorder = RecordingHook::default();
let executions = Arc::new(AtomicU32::new(0));
let mut stream = AgentBuilder::new(MockCompletionModel::from_stream_turns([[
MockStreamEvent::tool_call_name_delta("tc1", "ic1", "add"),
MockStreamEvent::tool_call_arguments_delta("tc1", "ic1", r#"{"x":1,"y":2}"#),
MockStreamEvent::tool_call("tc1", "add", json!({"x": 1, "y": 2})),
MockStreamEvent::final_response_with_default_usage(),
]]))
.tool(CountingAddTool {
calls: executions.clone(),
})
.add_hook(recorder.clone())
.add_hook(AlwaysRepeatModelTurn)
.build()
.runner("add")
.max_turns(2)
.stream()
.await;
let mut execution_commits = 0;
let mut tool_results = 0;
let mut provider_finals = 0;
let mut agent_finals = 0;
let mut retry_markers = 0;
let mut error = None;
while let Some(item) = stream.next().await {
match item {
Ok(MultiTurnStreamItem::ToolExecutionCommitted { .. }) => execution_commits += 1,
Ok(MultiTurnStreamItem::StreamUserItem(StreamedUserContent::ToolResult {
..
})) => tool_results += 1,
Ok(MultiTurnStreamItem::StreamAssistantItem(StreamedAssistantContent::Final(
_,
))) => provider_finals += 1,
Ok(MultiTurnStreamItem::FinalResponse(_)) => agent_finals += 1,
Ok(MultiTurnStreamItem::ModelTurnRetried { .. }) => retry_markers += 1,
Ok(_) => {}
Err(err) => error = Some(err),
}
}
let Some(StreamingError::Prompt(error)) = error else {
panic!("tool-bearing streaming retry should return PromptCancelled");
};
let PromptError::PromptCancelled {
chat_history,
reason,
} = error.as_ref()
else {
panic!("tool-bearing streaming retry should return PromptCancelled");
};
assert!(reason.contains("tool-bearing model turns"));
assert!(reason.contains("tool-call hooks"));
assert_eq!(chat_history, &[Message::user("add")]);
assert_eq!(execution_commits, 0);
assert_eq!(tool_results, 0);
assert_eq!(provider_finals, 0);
assert_eq!(agent_finals, 0);
assert_eq!(retry_markers, 0);
assert_eq!(recorder.count(StepEventKind::ToolCall), 0);
assert_eq!(recorder.count(StepEventKind::ToolResult), 0);
assert_eq!(executions.load(SeqCst), 0);
}
#[derive(Clone)]
struct BarrierResponseRetry {
inner: BoundedResponseRetry,
barrier: Arc<Barrier>,
}
impl AgentHook for BarrierResponseRetry {
async fn on_model_turn_finished(
&self,
ctx: &HookContext,
event: ModelTurnFinished<'_>,
) -> ModelTurnAction {
let rejected = event.content.iter().any(
|content| matches!(content, AssistantContent::Text(text) if text.text == "rejected"),
);
if rejected {
self.barrier.wait().await;
}
self.inner.on_model_turn_finished(ctx, event).await
}
}
#[tokio::test]
async fn concurrent_runs_of_same_agent_have_independent_retry_budgets() {
let hook = BarrierResponseRetry {
inner: BoundedResponseRetry::new("rejected", 1, TestRetryMode::Repeat),
barrier: Arc::new(Barrier::new(2)),
};
let agent = AgentBuilder::new(MockCompletionModel::from_turns([
MockTurn::text("rejected"),
MockTurn::text("rejected"),
MockTurn::text("accepted one"),
MockTurn::text("accepted two"),
]))
.add_hook(hook)
.build();
let first = agent.runner("first").max_turns(2).run();
let second = agent.runner("second").max_turns(2).run();
let (first, second) = tokio::join!(first, second);
let first = first.expect("first run");
let second = second.expect("second run");
let outputs = std::collections::HashSet::from([first.output, second.output]);
assert_eq!(
outputs,
std::collections::HashSet::from([
"accepted one".to_string(),
"accepted two".to_string(),
])
);
assert_eq!(first.completion_calls.len(), 2);
assert_eq!(second.completion_calls.len(), 2);
}
#[tokio::test]
async fn retry_scratchpad_state_is_isolated_by_run_and_hook_instance() {
let shared_hook = BoundedResponseRetry::new("rejected", 1, TestRetryMode::Repeat);
let first_ctx = HookContext::new(false, None);
let second_ctx = HookContext::new(false, None);
let content = OneOrMany::one(AssistantContent::text("rejected"));
let first_event = ModelTurnFinished {
turn: 1,
content: &content,
usage: Usage::new(),
};
let second_event = first_event;
let (first, second) = tokio::join!(
shared_hook.on_model_turn_finished(&first_ctx, first_event),
shared_hook.on_model_turn_finished(&second_ctx, second_event),
);
assert!(matches!(first, ModelTurnAction::Retry(_)));
assert!(matches!(second, ModelTurnAction::Retry(_)));
assert!(matches!(
shared_hook
.on_model_turn_finished(&first_ctx, first_event)
.await,
ModelTurnAction::Stop(_)
));
let same_run_ctx = HookContext::new(false, None);
let first_hook = BoundedResponseRetry::new("first", 1, TestRetryMode::Repeat);
let second_hook = BoundedResponseRetry::new("second", 1, TestRetryMode::Repeat);
let first_content = OneOrMany::one(AssistantContent::text("first"));
let second_content = OneOrMany::one(AssistantContent::text("second"));
let first_action = first_hook
.on_model_turn_finished(
&same_run_ctx,
ModelTurnFinished {
turn: 1,
content: &first_content,
usage: Usage::new(),
},
)
.await;
let second_action = second_hook
.on_model_turn_finished(
&same_run_ctx,
ModelTurnFinished {
turn: 2,
content: &second_content,
usage: Usage::new(),
},
)
.await;
assert!(matches!(first_action, ModelTurnAction::Retry(_)));
assert!(matches!(second_action, ModelTurnAction::Retry(_)));
}
#[derive(Clone)]
struct FixedModelTurnAction {
action: ModelTurnAction,
calls: Arc<AtomicU32>,
}
impl AgentHook for FixedModelTurnAction {
async fn on_model_turn_finished(
&self,
_ctx: &HookContext,
_event: ModelTurnFinished<'_>,
) -> ModelTurnAction {
self.calls.fetch_add(1, SeqCst);
self.action.clone()
}
}
#[tokio::test]
async fn model_turn_action_short_circuits_flat_and_nested_hook_stacks() {
let content = OneOrMany::one(AssistantContent::text("response"));
let event = ModelTurnFinished {
turn: 1,
content: &content,
usage: Usage::new(),
};
let ctx = HookContext::new(false, None);
let first_calls = Arc::new(AtomicU32::new(0));
let retry_calls = Arc::new(AtomicU32::new(0));
let skipped_calls = Arc::new(AtomicU32::new(0));
let mut flat = HookStack::new();
flat.push(FixedModelTurnAction {
action: ModelTurnAction::Continue,
calls: first_calls.clone(),
});
flat.push(FixedModelTurnAction {
action: ModelTurnAction::repeat(),
calls: retry_calls.clone(),
});
flat.push(FixedModelTurnAction {
action: ModelTurnAction::stop("unreachable"),
calls: skipped_calls.clone(),
});
assert!(matches!(
flat.on_model_turn_finished(&ctx, event).await,
ModelTurnAction::Retry(_)
));
assert_eq!(first_calls.load(SeqCst), 1);
assert_eq!(retry_calls.load(SeqCst), 1);
assert_eq!(skipped_calls.load(SeqCst), 0);
let nested_retry_calls = Arc::new(AtomicU32::new(0));
let outer_skipped_calls = Arc::new(AtomicU32::new(0));
let mut nested = HookStack::new();
nested.push(FixedModelTurnAction {
action: ModelTurnAction::retry_with_feedback("fix it"),
calls: nested_retry_calls.clone(),
});
let mut outer = HookStack::new();
outer.push(nested);
outer.push(FixedModelTurnAction {
action: ModelTurnAction::Continue,
calls: outer_skipped_calls.clone(),
});
assert!(matches!(
outer.on_model_turn_finished(&ctx, event).await,
ModelTurnAction::Retry(crate::agent::RetryRequest::Feedback(feedback))
if feedback == "fix it"
));
assert_eq!(nested_retry_calls.load(SeqCst), 1);
assert_eq!(outer_skipped_calls.load(SeqCst), 0);
let stop_calls = Arc::new(AtomicU32::new(0));
let after_stop_calls = Arc::new(AtomicU32::new(0));
let mut stopping = HookStack::new();
stopping.push(FixedModelTurnAction {
action: ModelTurnAction::stop("stop now"),
calls: stop_calls.clone(),
});
stopping.push(FixedModelTurnAction {
action: ModelTurnAction::Continue,
calls: after_stop_calls.clone(),
});
assert!(matches!(
stopping.on_model_turn_finished(&ctx, event).await,
ModelTurnAction::Stop(reason) if reason == "stop now"
));
assert_eq!(stop_calls.load(SeqCst), 1);
assert_eq!(after_stop_calls.load(SeqCst), 0);
}
}