use crate::engine::middleware::{Middleware, PostLlmCtx, PreLlmCtx, UserMessageCtx};
use crate::engine::{AgentBuilder, DenyAllApprovalHandler, RetryOnError};
use crate::llm::StreamChunk;
use crate::tool::{Content, Tool, ToolContext, ToolPolicy};
use crate::types::{
AgentError, AgentResult, ApprovalRequest, ChatMessage, CheckpointData, CheckpointStep,
RiskLevel, RunOutcome, RuntimeEvent, SessionId, TurnContext,
};
use async_trait::async_trait;
use futures_core::Stream;
use llm_trait::LlmProvider;
use serde_json::Value;
use std::pin::Pin;
use std::sync::Arc;
use std::sync::Mutex;
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use std::task::{Context, Poll};
struct DummyProvider;
#[async_trait]
impl LlmProvider for DummyProvider {
async fn stream(
&self,
_request: llm_trait::ChatRequest,
) -> Result<llm_trait::ChatStream, llm_trait::LlmError> {
unimplemented!("not used")
}
async fn chat(
&self,
_request: llm_trait::ChatRequest,
) -> Result<llm_trait::ChatResponse, llm_trait::LlmError> {
Ok(llm_trait::ChatResponse {
content: String::new(),
reasoning_content: None,
tool_calls: vec![],
usage: Default::default(),
finish_reason: llm_trait::response::FinishReason::Stop,
raw: None,
thinking_signature: None,
})
}
fn capabilities(&self) -> llm_trait::Capabilities {
Default::default()
}
fn info(&self) -> llm_trait::ProviderInfo {
llm_trait::ProviderInfo {
name: "test".to_string(),
model: "test".to_string(),
version: None,
}
}
}
#[tokio::test]
async fn run_turn_emits_run_finished_on_session_not_found() {
let runtime = AgentBuilder::new(Arc::new(DummyProvider))
.system_prompt("test")
.build()
.expect("build runtime");
let nonexistent = SessionId::new(99999);
let event_fired = Arc::new(AtomicBool::new(false));
let event_fired_clone = event_fired.clone();
let result = runtime
.run_turn(nonexistent.clone(), "test input", move |event| {
if let RuntimeEvent::RunFinished { .. } = &event {
event_fired_clone.store(true, Ordering::SeqCst);
}
Ok(())
})
.await;
assert!(
result.is_err(),
"run_turn should return Err for nonexistent session"
);
assert!(
event_fired.load(Ordering::SeqCst),
"run_turn must emit RunFinished before returning Err on session not found"
);
}
struct FailingMiddleware;
#[async_trait]
impl Middleware for FailingMiddleware {
async fn on_user_message(&self, _ctx: &mut UserMessageCtx) -> AgentResult<()> {
Err(AgentError::internal("middleware intentionally fails"))
}
}
#[tokio::test]
async fn run_turn_emits_run_finished_on_middleware_failure() {
let runtime = AgentBuilder::new(Arc::new(DummyProvider))
.system_prompt("test")
.middleware(FailingMiddleware)
.build()
.expect("build runtime");
let sid = runtime.create_session().await;
let event_fired = Arc::new(AtomicBool::new(false));
let event_fired_clone = event_fired.clone();
let result = runtime
.run_turn(sid, "test input", move |event| {
if let RuntimeEvent::RunFinished { .. } = &event {
event_fired_clone.store(true, Ordering::SeqCst);
}
Ok(())
})
.await;
assert!(
result.is_err(),
"run_turn should return Err when middleware fails"
);
assert!(
event_fired.load(Ordering::SeqCst),
"run_turn must emit RunFinished before returning Err on middleware failure"
);
}
struct ErrorStreamProvider;
#[async_trait]
impl LlmProvider for ErrorStreamProvider {
async fn stream(
&self,
_request: llm_trait::ChatRequest,
) -> Result<llm_trait::ChatStream, llm_trait::LlmError> {
struct ErrorStream;
impl Stream for ErrorStream {
type Item = Result<StreamChunk, llm_trait::LlmError>;
fn poll_next(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
Poll::Ready(Some(Err(llm_trait::LlmError::llm("simulated LLM error"))))
}
}
Ok(llm_trait::ChatStream::new(Box::pin(ErrorStream)))
}
async fn chat(
&self,
_request: llm_trait::ChatRequest,
) -> Result<llm_trait::ChatResponse, llm_trait::LlmError> {
Ok(llm_trait::ChatResponse {
content: String::new(),
reasoning_content: None,
tool_calls: vec![],
usage: Default::default(),
finish_reason: llm_trait::response::FinishReason::Stop,
raw: None,
thinking_signature: None,
})
}
fn capabilities(&self) -> llm_trait::Capabilities {
Default::default()
}
fn info(&self) -> llm_trait::ProviderInfo {
llm_trait::ProviderInfo {
name: "test".to_string(),
model: "test".to_string(),
version: None,
}
}
}
struct CancelledStreamProvider;
#[async_trait]
impl LlmProvider for CancelledStreamProvider {
async fn stream(
&self,
_request: llm_trait::ChatRequest,
) -> Result<llm_trait::ChatStream, llm_trait::LlmError> {
struct CancelledStream;
impl Stream for CancelledStream {
type Item = Result<StreamChunk, llm_trait::LlmError>;
fn poll_next(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
Poll::Ready(Some(Err(llm_trait::LlmError::llm("cancelled"))))
}
}
Ok(llm_trait::ChatStream::new(Box::pin(CancelledStream)))
}
async fn chat(
&self,
_request: llm_trait::ChatRequest,
) -> Result<llm_trait::ChatResponse, llm_trait::LlmError> {
Ok(llm_trait::ChatResponse {
content: String::new(),
reasoning_content: None,
tool_calls: vec![],
usage: Default::default(),
finish_reason: llm_trait::response::FinishReason::Stop,
raw: None,
thinking_signature: None,
})
}
fn capabilities(&self) -> llm_trait::Capabilities {
Default::default()
}
fn info(&self) -> llm_trait::ProviderInfo {
llm_trait::ProviderInfo {
name: "test".to_string(),
model: "test".to_string(),
version: None,
}
}
}
#[tokio::test]
async fn run_turn_emits_run_finished_on_llm_error() {
let runtime = AgentBuilder::new(Arc::new(ErrorStreamProvider))
.system_prompt("test")
.build()
.expect("build runtime");
let sid = runtime.create_session().await;
let event_fired = Arc::new(AtomicBool::new(false));
let event_fired_clone = event_fired.clone();
let result = runtime
.run_turn(sid, "test input", move |event| {
if let RuntimeEvent::RunFinished { .. } = &event {
event_fired_clone.store(true, Ordering::SeqCst);
}
Ok(())
})
.await;
assert!(
event_fired.load(Ordering::SeqCst),
"run_turn must emit RunFinished when LLM returns an error"
);
let _ = result;
}
struct ScriptedProvider {
script: Mutex<std::vec::IntoIter<Vec<StreamChunk>>>,
}
impl ScriptedProvider {
fn new(script: Vec<Vec<StreamChunk>>) -> Self {
Self {
script: Mutex::new(script.into_iter()),
}
}
}
#[async_trait]
impl LlmProvider for ScriptedProvider {
async fn stream(
&self,
_request: llm_trait::ChatRequest,
) -> Result<llm_trait::ChatStream, llm_trait::LlmError> {
let chunks: Vec<Result<StreamChunk, llm_trait::LlmError>> = self
.script
.lock()
.unwrap()
.next()
.unwrap_or_default()
.into_iter()
.map(Ok)
.collect();
Ok(llm_trait::ChatStream::new(Box::pin(
futures_util::stream::iter(chunks),
)))
}
async fn chat(
&self,
_request: llm_trait::ChatRequest,
) -> Result<llm_trait::ChatResponse, llm_trait::LlmError> {
Ok(llm_trait::ChatResponse {
content: String::new(),
reasoning_content: None,
tool_calls: vec![],
usage: Default::default(),
finish_reason: llm_trait::response::FinishReason::Stop,
raw: None,
thinking_signature: None,
})
}
fn capabilities(&self) -> llm_trait::Capabilities {
llm_trait::Capabilities {
supports_streaming: true,
supports_tools: true,
supports_vision: false,
supports_thinking: false,
max_context_tokens: None,
max_output_tokens: None,
}
}
fn info(&self) -> llm_trait::ProviderInfo {
llm_trait::ProviderInfo {
name: "test".to_string(),
model: "test".to_string(),
version: None,
}
}
}
#[tokio::test]
async fn truncation_guard_blocks_tool_calls_on_length_finish_reason() {
let runtime = AgentBuilder::new(Arc::new(ScriptedProvider::new(vec![
vec![
StreamChunk::ToolCall(serde_json::json!({
"delta": {
"tool_calls": [{
"id": "call_trunc",
"function": {
"name": "shell",
"arguments": "{\"cmd\": \"rm -rf /inco"
}
}]
}
})),
StreamChunk::Stop {
finish_reason: Some("length".to_string()),
},
],
vec![
StreamChunk::Text(
"I see the previous call was truncated. Let me re-issue it.".to_string(),
),
StreamChunk::Stop {
finish_reason: Some("stop".to_string()),
},
],
])))
.system_prompt("You are a careful assistant.")
.build()
.expect("build runtime");
let sid = runtime.create_session().await;
let events = Arc::new(Mutex::new(Vec::new()));
let events_clone = events.clone();
let result = runtime
.run_turn(sid.clone(), "run a command", move |event| {
events_clone.lock().unwrap().push(event);
Ok(())
})
.await;
assert!(
result.is_ok(),
"run_turn should complete: {:?}",
result.err()
);
let session = runtime.session(&sid).await.expect("session exists");
let messages = session.chat_messages().to_vec();
let has_truncation_error = messages.iter().any(|m| {
if let ChatMessage::Tool { content, .. } = m {
content.contains("Tool call was not executed") && content.contains("output token limit")
} else {
false
}
});
assert!(
has_truncation_error,
"session should contain truncation error tool result. Messages: {:#?}",
messages
.iter()
.map(|m| format!("{:?}", m))
.collect::<Vec<_>>()
);
}
#[tokio::test]
async fn truncation_guard_blocks_tool_calls_on_finish_tool_calls_mimo() {
let runtime = AgentBuilder::new(Arc::new(ScriptedProvider::new(vec![
vec![
StreamChunk::ToolCall(serde_json::json!({
"delta": {
"tool_calls": [{
"id": "call_trunc_1",
"function": {
"name": "spawn_agent",
"arguments": "{\"agent_path\": "
}
}]
}
})),
StreamChunk::Stop {
finish_reason: Some("tool_calls".to_string()),
},
],
vec![
StreamChunk::Text("Re-issuing the call with complete arguments.".to_string()),
StreamChunk::Stop {
finish_reason: Some("stop".to_string()),
},
],
])))
.system_prompt("You are a careful assistant.")
.build()
.expect("build runtime");
let sid = runtime.create_session().await;
let events = Arc::new(Mutex::new(Vec::new()));
let events_clone = events.clone();
let result = runtime
.run_turn(sid.clone(), "spawn a sub-agent", move |event| {
events_clone.lock().unwrap().push(event);
Ok(())
})
.await;
assert!(
result.is_ok(),
"run_turn should complete: {:?}",
result.err()
);
let session = runtime.session(&sid).await.expect("session exists");
let messages = session.chat_messages().to_vec();
let has_reissue = messages.iter().any(|m| {
if let ChatMessage::Tool { content, .. } = m {
content.contains("Tool call was not executed")
&& content.contains("provider truncated the argument stream")
} else {
false
}
});
assert!(
has_reissue,
"session should contain a re-issue tool result for the truncated call. Messages: {:#?}",
messages
.iter()
.map(|m| format!("{:?}", m))
.collect::<Vec<_>>()
);
let has_args_invalid = messages.iter().any(|m| {
let s = format!("{:?}", m);
s.contains("argument parsing failed") || s.contains("EOF while parsing")
});
assert!(
!has_args_invalid,
"truncated args must not surface as a ToolArgsInvalid failure. Messages: {:#?}",
messages
.iter()
.map(|m| format!("{:?}", m))
.collect::<Vec<_>>()
);
}
#[tokio::test]
async fn truncation_guard_recognizes_wrapper_echo() {
let wrapper = serde_json::json!({
"error": "tool_call_arguments_truncated",
"message": "The tool call arguments were truncated or invalid. Please retry with complete arguments.",
"original_args_preview": "{\"agent_path\": "
})
.to_string();
let runtime = AgentBuilder::new(Arc::new(ScriptedProvider::new(vec![
vec![
StreamChunk::ToolCall(serde_json::json!({
"delta": {
"tool_calls": [{
"id": "call_echo",
"function": {
"name": "send_message",
"arguments": wrapper
}
}]
}
})),
StreamChunk::Stop {
finish_reason: Some("tool_calls".to_string()),
},
],
vec![
StreamChunk::Text("Now sending a real message.".to_string()),
StreamChunk::Stop {
finish_reason: Some("stop".to_string()),
},
],
])))
.system_prompt("You are a careful assistant.")
.build()
.expect("build runtime");
let sid = runtime.create_session().await;
let events = Arc::new(Mutex::new(Vec::new()));
let events_clone = events.clone();
let result = runtime
.run_turn(sid.clone(), "continue", move |event| {
events_clone.lock().unwrap().push(event);
Ok(())
})
.await;
assert!(
result.is_ok(),
"run_turn should complete: {:?}",
result.err()
);
let session = runtime.session(&sid).await.expect("session exists");
let messages = session.chat_messages().to_vec();
let has_reissue = messages.iter().any(|m| {
if let ChatMessage::Tool { content, .. } = m {
content.contains("Tool call was not executed")
} else {
false
}
});
assert!(
has_reissue,
"echoed wrapper args must be routed to the re-issue guard. Messages: {:#?}",
messages
.iter()
.map(|m| format!("{:?}", m))
.collect::<Vec<_>>()
);
let has_args_invalid = messages.iter().any(|m| {
let s = format!("{:?}", m);
s.contains("argument parsing failed")
});
assert!(
!has_args_invalid,
"echoed wrapper must not reach tool execution. Messages: {:#?}",
messages
.iter()
.map(|m| format!("{:?}", m))
.collect::<Vec<_>>()
);
}
struct SpawnLikeTool;
#[async_trait]
impl Tool for SpawnLikeTool {
fn name(&self) -> &'static str {
"spawn_agent"
}
fn description(&self) -> &'static str {
"Spawn a sub-agent"
}
fn schema(&self) -> Value {
serde_json::json!({
"type": "object",
"properties": {
"task_name": { "type": "string" },
"message": { "type": "string" },
},
"required": ["task_name", "message"]
})
}
async fn call(&self, _args: &Value, _ctx: &ToolContext) -> AgentResult<Vec<Content>> {
Ok(vec![Content::text("spawned")])
}
}
struct NoArgTool;
#[async_trait]
impl Tool for NoArgTool {
fn name(&self) -> &'static str {
"list_agents"
}
fn description(&self) -> &'static str {
"List running agents"
}
fn schema(&self) -> Value {
serde_json::json!({ "type": "object", "properties": {} })
}
async fn call(&self, _args: &Value, _ctx: &ToolContext) -> AgentResult<Vec<Content>> {
Ok(vec![Content::text("no agents")])
}
}
#[tokio::test]
async fn truncation_guard_blocks_empty_args_for_required_field_tool() {
let runtime = AgentBuilder::new(Arc::new(ScriptedProvider::new(vec![
vec![
StreamChunk::ToolCall(serde_json::json!({
"delta": {
"tool_calls": [
{
"index": 0,
"id": "call_real",
"function": {
"name": "spawn_agent",
"arguments": "{\"task_name\":\"foo\",\"message\":\"do stuff\"}"
}
},
{
"index": 1,
"id": "call_empty",
"function": {
"name": "spawn_agent",
"arguments": "{}"
}
}
]
}
})),
StreamChunk::Stop {
finish_reason: Some("tool_calls".to_string()),
},
],
vec![
StreamChunk::Text("Re-issuing with complete arguments.".to_string()),
StreamChunk::Stop {
finish_reason: Some("stop".to_string()),
},
],
])))
.system_prompt("You are a careful assistant.")
.register_tool(SpawnLikeTool)
.build()
.expect("build runtime");
let sid = runtime.create_session().await;
let events = Arc::new(Mutex::new(Vec::new()));
let events_clone = events.clone();
let result = runtime
.run_turn(sid.clone(), "spawn two agents", move |event| {
events_clone.lock().unwrap().push(event);
Ok(())
})
.await;
assert!(
result.is_ok(),
"run_turn should complete: {:?}",
result.err()
);
let session = runtime.session(&sid).await.expect("session exists");
let messages = session.chat_messages().to_vec();
let has_reissue = messages.iter().any(|m| {
if let ChatMessage::Tool { content, .. } = m {
content.contains("Tool call was not executed")
&& content.contains("empty argument object")
&& content.contains("requires fields")
} else {
false
}
});
assert!(
has_reissue,
"empty {{}} for a required-field tool must be caught by case 4 guard. Messages: {:#?}",
messages
.iter()
.map(|m| format!("{:?}", m))
.collect::<Vec<_>>()
);
let has_args_invalid = messages.iter().any(|m| {
let s = format!("{:?}", m);
s.contains("argument parsing failed") || s.contains("EOF while parsing")
});
assert!(
!has_args_invalid,
"empty {{}} args must not surface as a ToolArgsInvalid failure. Messages: {:#?}",
messages
.iter()
.map(|m| format!("{:?}", m))
.collect::<Vec<_>>()
);
}
#[tokio::test]
async fn no_arg_tool_with_empty_object_args_executes_normally() {
let runtime = AgentBuilder::new(Arc::new(ScriptedProvider::new(vec![
vec![
StreamChunk::ToolCall(serde_json::json!({
"delta": {
"tool_calls": [{
"index": 0,
"id": "call_list",
"function": {
"name": "list_agents",
"arguments": "{}"
}
}]
}
})),
StreamChunk::Stop {
finish_reason: Some("tool_calls".to_string()),
},
],
vec![
StreamChunk::Text("No agents running.".to_string()),
StreamChunk::Stop {
finish_reason: Some("stop".to_string()),
},
],
])))
.system_prompt("test")
.register_tool(NoArgTool)
.build()
.expect("build runtime");
let sid = runtime.create_session().await;
let events = Arc::new(Mutex::new(Vec::new()));
let events_clone = events.clone();
let result = runtime
.run_turn(sid.clone(), "list agents", move |event| {
events_clone.lock().unwrap().push(event);
Ok(())
})
.await;
assert!(
result.is_ok(),
"run_turn should complete: {:?}",
result.err()
);
let session = runtime.session(&sid).await.expect("session exists");
let messages = session.chat_messages().to_vec();
let has_reissue = messages.iter().any(|m| {
if let ChatMessage::Tool { content, .. } = m {
content.contains("Tool call was not executed")
&& content.contains("empty argument object")
} else {
false
}
});
assert!(
!has_reissue,
"no-arg tool with {{}} must execute, not be re-issued. Messages: {:#?}",
messages
.iter()
.map(|m| format!("{:?}", m))
.collect::<Vec<_>>()
);
let has_output = messages.iter().any(|m| {
if let ChatMessage::Tool { content, .. } = m {
content.contains("no agents")
} else {
false
}
});
assert!(
has_output,
"no-arg tool must have executed and returned output. Messages: {:#?}",
messages
.iter()
.map(|m| format!("{:?}", m))
.collect::<Vec<_>>()
);
}
#[tokio::test]
async fn partial_execution_valid_calls_execute_invalid_get_guidance() {
let runtime = AgentBuilder::new(Arc::new(ScriptedProvider::new(vec![
vec![
StreamChunk::ToolCall(serde_json::json!({
"delta": {
"tool_calls": [
{
"index": 0,
"id": "call_1",
"function": { "name": "spawn_agent", "arguments": "{\"task_name\":\"a\",\"message\":\"first\"}" }
},
{
"index": 1,
"id": "call_2",
"function": { "name": "spawn_agent", "arguments": "{\"task_name\":\"b\",\"message\":\"tru" }
},
{
"index": 2,
"id": "call_3",
"function": { "name": "spawn_agent", "arguments": "{\"task_name\":\"c\",\"message\":\"third\"}" }
}
]
}
})),
StreamChunk::Stop { finish_reason: Some("tool_calls".to_string()) },
],
vec![
StreamChunk::Text("Done.".to_string()),
StreamChunk::Stop { finish_reason: Some("stop".to_string()) },
],
])))
.system_prompt("test")
.register_tool(SpawnLikeTool)
.build().expect("build");
let sid = runtime.create_session().await;
let events = Arc::new(Mutex::new(Vec::new()));
let ec = events.clone();
let result = runtime
.run_turn(sid.clone(), "spawn three", move |e| {
ec.lock().unwrap().push(e);
Ok(())
})
.await;
assert!(result.is_ok(), "run_turn failed: {:?}", result.err());
let session = runtime.session(&sid).await.expect("session");
let messages = session.chat_messages().to_vec();
let has_reissue = messages.iter().any(|m| {
if let ChatMessage::Tool { content, .. } = m {
content.contains("incomplete")
} else {
false
}
});
assert!(
has_reissue,
"truncated call must get guidance. Messages: {:#?}",
messages
.iter()
.map(|m| format!("{:?}", m))
.collect::<Vec<_>>()
);
let spawned_count = messages
.iter()
.filter(|m| {
if let ChatMessage::Tool { content, .. } = m {
content.contains("spawned")
} else {
false
}
})
.count();
assert_eq!(
spawned_count,
2,
"exactly 2 valid calls must execute. Messages: {:#?}",
messages
.iter()
.map(|m| format!("{:?}", m))
.collect::<Vec<_>>()
);
}
#[tokio::test]
async fn partial_execution_all_valid_no_guard() {
let runtime = AgentBuilder::new(Arc::new(ScriptedProvider::new(vec![
vec![
StreamChunk::ToolCall(serde_json::json!({
"delta": {
"tool_calls": [
{ "index": 0, "id": "c1", "function": { "name": "spawn_agent", "arguments": "{\"task_name\":\"a\",\"message\":\"first\"}" } },
{ "index": 1, "id": "c2", "function": { "name": "spawn_agent", "arguments": "{\"task_name\":\"b\",\"message\":\"second\"}" } },
{ "index": 2, "id": "c3", "function": { "name": "spawn_agent", "arguments": "{\"task_name\":\"c\",\"message\":\"third\"}" } }
]
}
})),
StreamChunk::Stop { finish_reason: Some("tool_calls".to_string()) },
],
vec![
StreamChunk::Text("All done.".to_string()),
StreamChunk::Stop { finish_reason: Some("stop".to_string()) },
],
])))
.system_prompt("test")
.register_tool(SpawnLikeTool)
.build().expect("build");
let sid = runtime.create_session().await;
let events = Arc::new(Mutex::new(Vec::new()));
let ec = events.clone();
let result = runtime
.run_turn(sid.clone(), "spawn three", move |e| {
ec.lock().unwrap().push(e);
Ok(())
})
.await;
assert!(result.is_ok(), "run_turn failed: {:?}", result.err());
let session = runtime.session(&sid).await.expect("session");
let messages = session.chat_messages().to_vec();
let has_reissue = messages.iter().any(|m| {
if let ChatMessage::Tool { content, .. } = m {
content.contains("Tool call was not executed")
} else {
false
}
});
assert!(
!has_reissue,
"all valid calls must not trigger guard. Messages: {:#?}",
messages
.iter()
.map(|m| format!("{:?}", m))
.collect::<Vec<_>>()
);
let spawned_count = messages
.iter()
.filter(|m| {
if let ChatMessage::Tool { content, .. } = m {
content.contains("spawned")
} else {
false
}
})
.count();
assert_eq!(
spawned_count,
3,
"all 3 must execute. Messages: {:#?}",
messages
.iter()
.map(|m| format!("{:?}", m))
.collect::<Vec<_>>()
);
}
#[tokio::test]
async fn partial_execution_mixed_truncation_and_empty_required() {
let runtime = AgentBuilder::new(Arc::new(ScriptedProvider::new(vec![
vec![
StreamChunk::ToolCall(serde_json::json!({
"delta": {
"tool_calls": [
{ "index": 0, "id": "c_trunc", "function": { "name": "spawn_agent", "arguments": "{\"task_name\":\"x\"," } },
{ "index": 1, "id": "c_empty", "function": { "name": "spawn_agent", "arguments": "{}" } },
{ "index": 2, "id": "c_ok", "function": { "name": "spawn_agent", "arguments": "{\"task_name\":\"z\",\"message\":\"ok\"}" } }
]
}
})),
StreamChunk::Stop { finish_reason: Some("tool_calls".to_string()) },
],
vec![
StreamChunk::Text("Fixed.".to_string()),
StreamChunk::Stop { finish_reason: Some("stop".to_string()) },
],
])))
.system_prompt("test")
.register_tool(SpawnLikeTool)
.build().expect("build");
let sid = runtime.create_session().await;
let events = Arc::new(Mutex::new(Vec::new()));
let ec = events.clone();
let result = runtime
.run_turn(sid.clone(), "spawn mixed", move |e| {
ec.lock().unwrap().push(e);
Ok(())
})
.await;
assert!(result.is_ok(), "run_turn failed: {:?}", result.err());
let session = runtime.session(&sid).await.expect("session");
let messages = session.chat_messages().to_vec();
let has_trunc_guidance = messages.iter().any(|m| {
if let ChatMessage::Tool { content, .. } = m {
content.contains("incomplete") || content.contains("truncated")
} else {
false
}
});
assert!(
has_trunc_guidance,
"must have truncation guidance. Messages: {:#?}",
messages
.iter()
.map(|m| format!("{:?}", m))
.collect::<Vec<_>>()
);
let spawned_count = messages
.iter()
.filter(|m| {
if let ChatMessage::Tool { content, .. } = m {
content.contains("spawned")
} else {
false
}
})
.count();
assert_eq!(
spawned_count,
1,
"only valid call executes. Messages: {:#?}",
messages
.iter()
.map(|m| format!("{:?}", m))
.collect::<Vec<_>>()
);
let guidance_count = messages
.iter()
.filter(|m| {
if let ChatMessage::Tool { content, .. } = m {
content.contains("Tool call was not executed")
} else {
false
}
})
.count();
assert_eq!(
guidance_count,
2,
"2 invalid calls get guidance. Messages: {:#?}",
messages
.iter()
.map(|m| format!("{:?}", m))
.collect::<Vec<_>>()
);
}
#[tokio::test]
async fn truncation_circuit_breaker_redirects_then_hard_stops() {
let runtime = AgentBuilder::new(Arc::new(ScriptedProvider::new(vec![
vec![
StreamChunk::ToolCall(serde_json::json!({
"delta": {
"tool_calls": [{
"id": "t1",
"function": { "name": "spawn_agent", "arguments": "{\"task_name\":\"x\",\"message\":\"tru" }
}]
}
})),
StreamChunk::Stop { finish_reason: Some("tool_calls".to_string()) },
],
vec![
StreamChunk::ToolCall(serde_json::json!({
"delta": {
"tool_calls": [{
"id": "t2",
"function": { "name": "spawn_agent", "arguments": "{\"task_name\":\"x\",\"message\":\"tru" }
}]
}
})),
StreamChunk::Stop { finish_reason: Some("tool_calls".to_string()) },
],
vec![
StreamChunk::ToolCall(serde_json::json!({
"delta": {
"tool_calls": [{
"id": "t3",
"function": { "name": "spawn_agent", "arguments": "{\"task_name\":\"x\",\"message\":\"tru" }
}]
}
})),
StreamChunk::Stop { finish_reason: Some("tool_calls".to_string()) },
],
vec![
StreamChunk::ToolCall(serde_json::json!({
"delta": {
"tool_calls": [{
"id": "t4",
"function": { "name": "spawn_agent", "arguments": "{\"task_name\":\"x\",\"message\":\"tru" }
}]
}
})),
StreamChunk::Stop { finish_reason: Some("tool_calls".to_string()) },
],
vec![
StreamChunk::ToolCall(serde_json::json!({
"delta": {
"tool_calls": [{
"id": "t5",
"function": { "name": "spawn_agent", "arguments": "{\"task_name\":\"x\",\"message\":\"tru" }
}]
}
})),
StreamChunk::Stop { finish_reason: Some("tool_calls".to_string()) },
],
])))
.system_prompt("test")
.register_tool(SpawnLikeTool)
.build()
.expect("build");
let sid = runtime.create_session().await;
let events = Arc::new(Mutex::new(Vec::new()));
let ec = events.clone();
let result = runtime
.run_turn(sid.clone(), "do something", move |e| {
ec.lock().unwrap().push(e);
Ok(())
})
.await;
assert!(
result.is_ok(),
"run_turn should not error: {:?}",
result.err()
);
let outcome = result.unwrap();
match &outcome {
RunOutcome::Failed { error } => {
assert!(
error.contains("repeatedly truncated"),
"error should mention repeated truncation: {}",
error,
);
}
other => panic!("expected RunOutcome::Failed at hard limit, got {:?}", other),
}
let session = runtime.session(&sid).await.expect("session");
assert_eq!(session.run_state.truncation_strikes, 5);
}
#[tokio::test]
async fn truncation_strikes_reset_on_successful_tool_call() {
let runtime = AgentBuilder::new(Arc::new(ScriptedProvider::new(vec![
vec![
StreamChunk::ToolCall(serde_json::json!({
"delta": {
"tool_calls": [{
"id": "bad1",
"function": { "name": "spawn_agent", "arguments": "{\"task_name\":\"x\",\"message\":\"tru" }
}]
}
})),
StreamChunk::Stop { finish_reason: Some("tool_calls".to_string()) },
],
vec![
StreamChunk::ToolCall(serde_json::json!({
"delta": {
"tool_calls": [{
"id": "good1",
"function": { "name": "spawn_agent", "arguments": "{\"task_name\":\"y\",\"message\":\"hello\"}" }
}]
}
})),
StreamChunk::Stop { finish_reason: Some("tool_calls".to_string()) },
],
vec![
StreamChunk::ToolCall(serde_json::json!({
"delta": {
"tool_calls": [{
"id": "bad2",
"function": { "name": "spawn_agent", "arguments": "{\"task_name\":\"z\",\"message\":\"tru" }
}]
}
})),
StreamChunk::Stop { finish_reason: Some("tool_calls".to_string()) },
],
vec![
StreamChunk::Text("Done.".to_string()),
StreamChunk::Stop { finish_reason: Some("stop".to_string()) },
],
])))
.system_prompt("test")
.register_tool(SpawnLikeTool)
.build()
.expect("build");
let sid = runtime.create_session().await;
let events = Arc::new(Mutex::new(Vec::new()));
let ec = events.clone();
let result = runtime
.run_turn(sid.clone(), "do something", move |e| {
ec.lock().unwrap().push(e);
Ok(())
})
.await;
assert!(
result.is_ok(),
"run_turn should not error: {:?}",
result.err()
);
let outcome = result.unwrap();
assert!(
matches!(outcome, RunOutcome::Completed),
"expected Completed after reset, got {:?}",
outcome,
);
let session = runtime.session(&sid).await.expect("session");
assert_eq!(session.run_state.truncation_strikes, 1);
}
#[tokio::test]
async fn empty_args_counted_as_truncation_strike() {
let runtime = AgentBuilder::new(Arc::new(ScriptedProvider::new(vec![
vec![
StreamChunk::ToolCall(serde_json::json!({
"delta": {
"tool_calls": [{
"id": "e1",
"function": { "name": "spawn_agent", "arguments": "" }
}]
}
})),
StreamChunk::Stop {
finish_reason: Some("tool_calls".to_string()),
},
],
vec![
StreamChunk::ToolCall(serde_json::json!({
"delta": {
"tool_calls": [{
"id": "e2",
"function": { "name": "spawn_agent", "arguments": "" }
}]
}
})),
StreamChunk::Stop {
finish_reason: Some("tool_calls".to_string()),
},
],
vec![
StreamChunk::ToolCall(serde_json::json!({
"delta": {
"tool_calls": [{
"id": "e3",
"function": { "name": "spawn_agent", "arguments": "" }
}]
}
})),
StreamChunk::Stop {
finish_reason: Some("tool_calls".to_string()),
},
],
vec![
StreamChunk::Text("Done.".to_string()),
StreamChunk::Stop {
finish_reason: Some("stop".to_string()),
},
],
])))
.system_prompt("test")
.register_tool(SpawnLikeTool)
.build()
.expect("build");
let sid = runtime.create_session().await;
let events = Arc::new(Mutex::new(Vec::new()));
let ec = events.clone();
let result = runtime
.run_turn(sid.clone(), "do something", move |e| {
ec.lock().unwrap().push(e);
Ok(())
})
.await;
assert!(
result.is_ok(),
"run_turn should not error: {:?}",
result.err()
);
let outcome = result.unwrap();
assert!(
matches!(outcome, RunOutcome::Completed),
"expected Completed after redirect, got {:?}",
outcome,
);
let session = runtime.session(&sid).await.expect("session");
assert_eq!(session.run_state.truncation_strikes, 3);
}
#[tokio::test]
async fn run_managed_processes_follow_up_messages() {
let scripted = Arc::new(ScriptedProvider::new(vec![
vec![
StreamChunk::Text("done".to_string()),
StreamChunk::Stop {
finish_reason: Some("stop".to_string()),
},
],
vec![
StreamChunk::Text("follow-up done".to_string()),
StreamChunk::Stop {
finish_reason: Some("stop".to_string()),
},
],
]));
let runtime = AgentBuilder::new(scripted)
.system_prompt("You are a helpful assistant.")
.build()
.expect("build runtime");
let sid = runtime.create_session().await;
runtime.follow_up("continue working on the task".to_string());
let events = Arc::new(Mutex::new(Vec::new()));
let events_clone = events.clone();
let result = runtime
.run_managed(sid.clone(), "do something", move |event| {
events_clone.lock().unwrap().push(event);
Ok(())
})
.await;
assert!(
result.is_ok(),
"run_managed should complete: {:?}",
result.err()
);
let session = runtime.session(&sid).await.expect("session exists");
let messages = session.chat_messages().to_vec();
let text_count = messages
.iter()
.filter(|m| matches!(m, ChatMessage::Assistant { content: Some(c), .. } if c == "done" || c == "follow-up done"))
.count();
assert_eq!(
text_count, 2,
"both initial and follow-up turns should produce assistant responses"
);
}
struct EchoTool;
#[async_trait]
impl Tool for EchoTool {
fn name(&self) -> &'static str {
"echo"
}
fn description(&self) -> &'static str {
"Echo back the provided text"
}
fn schema(&self) -> Value {
serde_json::json!({
"type": "object",
"properties": { "text": { "type": "string" } }
})
}
async fn call(&self, args: &Value, _ctx: &ToolContext) -> AgentResult<Vec<Content>> {
let text = args.get("text").and_then(Value::as_str).unwrap_or("");
Ok(vec![Content::text(format!("echo: {text}"))])
}
}
#[tokio::test]
async fn run_turn_executes_tool_call_and_returns_text() {
let runtime = AgentBuilder::new(Arc::new(ScriptedProvider::new(vec![
vec![
StreamChunk::ToolCall(serde_json::json!({
"delta": {
"tool_calls": [{
"id": "call_1",
"function": {
"name": "echo",
"arguments": "{\"text\":\"hello\"}"
}
}]
}
})),
StreamChunk::Stop {
finish_reason: Some("tool_calls".to_string()),
},
],
vec![
StreamChunk::Text("final answer".to_string()),
StreamChunk::Stop {
finish_reason: Some("stop".to_string()),
},
],
])))
.system_prompt("test")
.register_tool(EchoTool)
.build()
.expect("build runtime");
let sid = runtime.create_session().await;
let events = Arc::new(Mutex::new(Vec::new()));
let events_clone = events.clone();
let result = runtime
.run_turn(sid.clone(), "echo hello", move |event| {
events_clone.lock().unwrap().push(event);
Ok(())
})
.await;
assert!(result.is_ok(), "run_turn failed: {:?}", result.err());
let session = runtime.session(&sid).await.expect("session exists");
let messages = session.chat_messages().to_vec();
assert!(
messages.iter().any(
|m| matches!(m, ChatMessage::Tool { content, .. } if content.contains("echo: hello"))
),
"session should contain the echo tool result: {:#?}",
messages
.iter()
.map(|m| format!("{:?}", m))
.collect::<Vec<_>>()
);
assert!(
messages.iter().any(
|m| matches!(m, ChatMessage::Assistant { content: Some(c), .. } if c == "final answer")
),
"session should contain the final assistant text"
);
let events = events.lock().unwrap();
assert!(
events
.iter()
.any(|e| matches!(e, RuntimeEvent::ToolCallStarted { .. })),
"run_turn should emit ToolCallStarted"
);
assert!(
events
.iter()
.any(|e| matches!(e, RuntimeEvent::ToolCallFinished { .. })),
"run_turn should emit ToolCallFinished"
);
}
struct FailingTool;
#[async_trait]
impl Tool for FailingTool {
fn name(&self) -> &'static str {
"fail"
}
fn description(&self) -> &'static str {
"Always fails"
}
fn schema(&self) -> Value {
serde_json::json!({"type": "object", "properties": {}})
}
async fn call(&self, _args: &Value, _ctx: &ToolContext) -> AgentResult<Vec<Content>> {
Err(AgentError::internal("simulated tool failure"))
}
}
struct CancelTool;
#[async_trait]
impl Tool for CancelTool {
fn name(&self) -> &'static str {
"cancel"
}
fn description(&self) -> &'static str {
"Cancel the current run"
}
fn schema(&self) -> Value {
serde_json::json!({"type": "object", "properties": {}})
}
async fn call(&self, _args: &Value, ctx: &ToolContext) -> AgentResult<Vec<Content>> {
ctx.cancel_token.cancel();
Ok(vec![Content::text("cancelled")])
}
}
struct CancelPendingTool;
#[async_trait]
impl Tool for CancelPendingTool {
fn name(&self) -> &'static str {
"cancel_pending"
}
fn description(&self) -> &'static str {
"Cancel the current run mid-execution"
}
fn schema(&self) -> Value {
serde_json::json!({"type": "object", "properties": {}})
}
async fn call(&self, _args: &Value, ctx: &ToolContext) -> AgentResult<Vec<Content>> {
ctx.cancel_token.cancel();
std::future::pending::<()>().await;
unreachable!("cancelled tool should never resolve")
}
}
struct RequireApproval;
#[async_trait]
impl ToolPolicy for RequireApproval {
async fn evaluate_approval(&self, tool_name: &str, _args: &Value) -> Option<ApprovalRequest> {
Some(ApprovalRequest {
title: format!("Approve {tool_name}"),
message: "Approve this tool call?".to_string(),
action_key: Some(format!("approve:{tool_name}")),
risk_level: RiskLevel::Sensitive,
raw: None,
})
}
}
struct HookMiddleware {
pre_called: Arc<AtomicBool>,
post_called: Arc<AtomicBool>,
}
#[async_trait]
impl Middleware for HookMiddleware {
async fn on_pre_llm(&self, _ctx: &mut PreLlmCtx) -> AgentResult<()> {
self.pre_called.store(true, Ordering::SeqCst);
Ok(())
}
async fn on_post_llm(&self, ctx: &mut PostLlmCtx) -> AgentResult<()> {
self.post_called.store(true, Ordering::SeqCst);
ctx.nudge_count += 1;
Ok(())
}
}
#[tokio::test]
async fn run_emits_run_finished_on_session_not_found() {
let runtime = AgentBuilder::new(Arc::new(DummyProvider))
.system_prompt("test")
.build()
.expect("build runtime");
let nonexistent = SessionId::new(99999);
let event_fired = Arc::new(AtomicBool::new(false));
let event_fired_clone = event_fired.clone();
let result = runtime
.run(nonexistent.clone(), move |event| {
if matches!(event, RuntimeEvent::RunFinished { .. }) {
event_fired_clone.store(true, Ordering::SeqCst);
}
Ok(())
})
.await;
assert!(
result.is_err(),
"run should return Err for a nonexistent session"
);
assert!(
event_fired.load(Ordering::SeqCst),
"run must emit RunFinished before returning Err on session not found"
);
}
#[tokio::test]
async fn run_completes_with_text_response() {
let runtime = AgentBuilder::new(Arc::new(ScriptedProvider::new(vec![vec![
StreamChunk::Text("answer".to_string()),
StreamChunk::Stop {
finish_reason: Some("stop".to_string()),
},
]])))
.system_prompt("test")
.build()
.expect("build runtime");
let sid = runtime.create_session().await;
runtime.add_user_message(&sid, "question").await.unwrap();
let events = Arc::new(Mutex::new(Vec::new()));
let events_clone = events.clone();
let result = runtime
.run(sid.clone(), move |event| {
events_clone.lock().unwrap().push(event);
Ok(())
})
.await;
assert!(
matches!(result, Ok(RunOutcome::Completed)),
"run failed: {:?}",
result.err()
);
let session = runtime.session(&sid).await.expect("session exists");
assert!(
session
.chat_messages()
.iter()
.any(|m| matches!(m, ChatMessage::Assistant { content: Some(c), .. } if c == "answer")),
"run should push the assistant text response into the session"
);
}
#[tokio::test]
async fn run_turn_emits_run_cancelled_when_cancelled_mid_run() {
let runtime = AgentBuilder::new(Arc::new(ScriptedProvider::new(vec![vec![
StreamChunk::ToolCall(serde_json::json!({
"delta": {
"tool_calls": [{
"id": "call_cancel",
"function": { "name": "cancel", "arguments": "{}" }
}]
}
})),
StreamChunk::Stop {
finish_reason: Some("tool_calls".to_string()),
},
]])))
.system_prompt("test")
.register_tool(CancelTool)
.build()
.expect("build runtime");
let sid = runtime.create_session().await;
let events = Arc::new(Mutex::new(Vec::new()));
let events_clone = events.clone();
let result = runtime
.run_turn(sid.clone(), "cancel now", move |event| {
events_clone.lock().unwrap().push(event);
Ok(())
})
.await;
let events = events.lock().unwrap();
assert!(
matches!(result, Ok(RunOutcome::Cancelled)),
"run_turn should return Cancelled outcome: {:?}",
result
);
assert!(
events
.iter()
.any(|e| matches!(e, RuntimeEvent::RunCancelled { .. })),
"run_turn should emit RunCancelled when cancelled mid-run"
);
}
#[tokio::test]
async fn run_turn_failing_tool_stops_with_failed_outcome() {
let runtime = AgentBuilder::new(Arc::new(ScriptedProvider::new(vec![vec![
StreamChunk::ToolCall(serde_json::json!({
"delta": {
"tool_calls": [{
"id": "call_fail",
"function": { "name": "fail", "arguments": "{}" }
}]
}
})),
StreamChunk::Stop {
finish_reason: Some("tool_calls".to_string()),
},
]])))
.system_prompt("test")
.register_tool(FailingTool)
.build()
.expect("build runtime");
let sid = runtime.create_session().await;
let result = runtime
.run_turn(sid.clone(), "do something", |_| Ok(()))
.await;
assert!(
matches!(result, Ok(RunOutcome::Failed { .. })),
"failing tool with StopOnError should yield Failed outcome: {:?}",
result
);
}
#[tokio::test]
async fn run_turn_failing_tool_retries_then_completes() {
let runtime = AgentBuilder::new(Arc::new(ScriptedProvider::new(vec![
vec![
StreamChunk::ToolCall(serde_json::json!({
"delta": {
"tool_calls": [{
"id": "call_fail",
"function": { "name": "fail", "arguments": "{}" }
}]
}
})),
StreamChunk::Stop {
finish_reason: Some("tool_calls".to_string()),
},
],
vec![
StreamChunk::Text("recovered".to_string()),
StreamChunk::Stop {
finish_reason: Some("stop".to_string()),
},
],
])))
.system_prompt("test")
.register_tool(FailingTool)
.error_recovery(Arc::new(RetryOnError))
.build()
.expect("build runtime");
let sid = runtime.create_session().await;
let result = runtime
.run_turn(sid.clone(), "do something", |_| Ok(()))
.await;
assert!(
matches!(result, Ok(RunOutcome::Completed)),
"failing tool with RetryOnError should eventually complete: {:?}",
result
);
}
#[tokio::test]
async fn run_turn_approval_denied_completes_without_executing() {
let runtime = AgentBuilder::new(Arc::new(ScriptedProvider::new(vec![vec![
StreamChunk::ToolCall(serde_json::json!({
"delta": {
"tool_calls": [{
"id": "call_denied",
"function": { "name": "echo", "arguments": "{\"text\":\"hello\"}" }
}]
}
})),
StreamChunk::Stop {
finish_reason: Some("tool_calls".to_string()),
},
]])))
.system_prompt("test")
.register_tool(EchoTool)
.approval_handler(Arc::new(DenyAllApprovalHandler))
.tool_policy(Arc::new(RequireApproval))
.build()
.expect("build runtime");
let sid = runtime.create_session().await;
let result = runtime
.run_turn(sid.clone(), "echo hello", |_| Ok(()))
.await;
assert!(
matches!(result, Ok(RunOutcome::Completed)),
"approval denial should stop the run cleanly (Completed): {:?}",
result
);
let session = runtime.session(&sid).await.expect("session exists");
assert!(
!session.chat_messages().iter().any(
|m| matches!(m, ChatMessage::Tool { content, .. } if content.contains("echo: hello"))
),
"denied tool call must not execute"
);
}
#[tokio::test]
async fn resume_from_checkpoint_executes_pending_tool_calls() {
let runtime = AgentBuilder::new(Arc::new(ScriptedProvider::new(vec![vec![
StreamChunk::Text("resumed answer".to_string()),
StreamChunk::Stop {
finish_reason: Some("stop".to_string()),
},
]])))
.system_prompt("test")
.register_tool(EchoTool)
.build()
.expect("build runtime");
let sid = runtime.create_session().await;
let checkpoint = CheckpointData {
session_id: sid.clone(),
user_input: "resume".to_string(),
step: CheckpointStep::BeforeToolCalls {
tool_calls: vec![(
"call_resume".to_string(),
"echo".to_string(),
r#"{"text":"resume"}"#.to_string(),
)],
},
turn_count: 0,
};
let events = Arc::new(Mutex::new(Vec::new()));
let events_clone = events.clone();
let result = runtime
.runner
.resume_from_checkpoint(checkpoint, move |event| {
events_clone.lock().unwrap().push(event);
Ok(())
})
.await;
assert!(
matches!(result, Ok(RunOutcome::Completed)),
"resume_from_checkpoint should complete: {:?}",
result.err()
);
let session = runtime.session(&sid).await.expect("session exists");
assert!(
session.chat_messages().iter().any(
|m| matches!(m, ChatMessage::Tool { content, .. } if content.contains("echo: resume"))
),
"resumed tool call should have executed"
);
}
#[tokio::test]
async fn middleware_on_pre_and_post_llm_hooks_fire() {
let pre_called = Arc::new(AtomicBool::new(false));
let post_called = Arc::new(AtomicBool::new(false));
let mw = HookMiddleware {
pre_called: pre_called.clone(),
post_called: post_called.clone(),
};
let runtime = AgentBuilder::new(Arc::new(ScriptedProvider::new(vec![vec![
StreamChunk::Text("hello".to_string()),
StreamChunk::Stop {
finish_reason: Some("stop".to_string()),
},
]])))
.system_prompt("test")
.middleware(mw)
.build()
.expect("build runtime");
let sid = runtime.create_session().await;
let result = runtime.run_turn(sid.clone(), "hi", |_| Ok(())).await;
assert!(result.is_ok(), "run_turn failed: {:?}", result.err());
assert!(
pre_called.load(Ordering::SeqCst),
"on_pre_llm hook should fire"
);
assert!(
post_called.load(Ordering::SeqCst),
"on_post_llm hook should fire"
);
let session = runtime.session(&sid).await.expect("session exists");
assert_eq!(
session.run_state.nudge_count, 1,
"nudge_count should be written back"
);
}
#[tokio::test]
async fn run_emits_runfinished_exactly_once() {
let runtime = AgentBuilder::new(Arc::new(ScriptedProvider::new(vec![vec![
StreamChunk::Text("answer".to_string()),
StreamChunk::Stop {
finish_reason: Some("stop".to_string()),
},
]])))
.system_prompt("test")
.build()
.expect("build runtime");
let sid = runtime.create_session().await;
runtime.add_user_message(&sid, "question").await.unwrap();
let events = Arc::new(Mutex::new(Vec::new()));
let events_clone = events.clone();
let outcome = runtime
.run(sid.clone(), move |event| {
events_clone.lock().unwrap().push(event);
Ok(())
})
.await
.expect("run");
let events = events.lock().unwrap();
assert!(matches!(outcome, RunOutcome::Completed));
let finished = events
.iter()
.filter(|e| matches!(e, RuntimeEvent::RunFinished { .. }))
.count();
assert_eq!(finished, 1, "run() must emit exactly one RunFinished");
}
#[tokio::test]
async fn run_managed_emits_runfinished_exactly_once() {
let runtime = AgentBuilder::new(Arc::new(ScriptedProvider::new(vec![vec![
StreamChunk::Text("done".to_string()),
StreamChunk::Stop {
finish_reason: Some("stop".to_string()),
},
]])))
.system_prompt("test")
.build()
.expect("build runtime");
let sid = runtime.create_session().await;
let events = Arc::new(Mutex::new(Vec::new()));
let events_clone = events.clone();
let outcome = runtime
.run_managed(sid.clone(), "do something", move |event| {
events_clone.lock().unwrap().push(event);
Ok(())
})
.await
.expect("run_managed");
let events = events.lock().unwrap();
assert!(matches!(outcome, RunOutcome::Completed));
let finished = events
.iter()
.filter(|e| matches!(e, RuntimeEvent::RunFinished { .. }))
.count();
assert_eq!(
finished, 1,
"run_managed() must emit exactly one RunFinished"
);
}
#[tokio::test]
async fn run_managed_cancel_cleans_up_ephemeral_messages() {
let runtime = AgentBuilder::new(Arc::new(ScriptedProvider::new(vec![vec![
StreamChunk::ToolCall(serde_json::json!({
"delta": {
"tool_calls": [{
"id": "call_cancel",
"function": { "name": "cancel", "arguments": "{}" }
}]
}
})),
StreamChunk::Stop {
finish_reason: Some("tool_calls".to_string()),
},
]])))
.system_prompt("test")
.register_tool(CancelTool)
.build()
.expect("build runtime");
let sid = runtime.create_session().await;
runtime
.set_messages(&sid, vec![ChatMessage::user_ephemeral("temp nudge")])
.await
.unwrap();
let result = runtime
.run_managed(sid.clone(), "cancel now", |_| Ok(()))
.await;
assert!(
matches!(result, Ok(RunOutcome::Cancelled)),
"run_managed should return Cancelled: {:?}",
result
);
let session = runtime.session(&sid).await.expect("session exists");
assert!(
session.chat_messages().iter().all(|m| !m.is_ephemeral()),
"run_managed cancel must clean up ephemeral messages, got: {:?}",
session.chat_messages()
);
}
fn terminal_seq(events: &[RuntimeEvent]) -> Vec<&'static str> {
events
.iter()
.filter_map(|e| match e {
RuntimeEvent::Checkpoint { checkpoint, .. } => Some(match &checkpoint.step {
CheckpointStep::AfterUserInput => "AfterUserInput",
CheckpointStep::BeforeLlm { .. } => "BeforeLlm",
CheckpointStep::BeforeToolCalls { .. } => "BeforeToolCalls",
CheckpointStep::AfterToolCalls { .. } => "AfterToolCalls",
}),
RuntimeEvent::RunFinished { .. } => Some("RunFinished"),
RuntimeEvent::RunCancelled { .. } => Some("RunCancelled"),
_ => None,
})
.collect()
}
fn text_script(text: &str) -> Vec<Vec<StreamChunk>> {
vec![vec![
StreamChunk::Text(text.to_string()),
StreamChunk::Stop {
finish_reason: Some("stop".to_string()),
},
]]
}
#[tokio::test]
async fn turn_end_callback_receives_correct_context() {
let runtime = AgentBuilder::new(Arc::new(ScriptedProvider::new(text_script("answer"))))
.system_prompt("test")
.build()
.expect("build runtime");
let captured = Arc::new(Mutex::new(Vec::<TurnContext>::new()));
let captured_clone = captured.clone();
runtime.on_turn_end(move |ctx| {
captured_clone.lock().unwrap().push(ctx.clone());
});
let sid = runtime.create_session().await;
let result = runtime.run_turn(sid.clone(), "question", |_| Ok(())).await;
assert!(matches!(result, Ok(RunOutcome::Completed)));
let contexts = captured.lock().unwrap();
assert_eq!(contexts.len(), 1, "one turn → one turn-end callback");
let ctx = &contexts[0];
assert!(matches!(&ctx.outcome, RunOutcome::Completed));
assert_eq!(ctx.turn_number, 1);
assert_eq!(ctx.tool_call_count, 0);
assert_eq!(ctx.tool_success, 0);
assert_eq!(ctx.tool_failed, 0);
assert!(ctx.error_message.is_none());
assert_eq!(ctx.full_text_len, "answer".len() as u64);
assert!(!ctx.has_thinking);
assert_eq!(ctx.llm_calls, 1);
assert_eq!(ctx.user_input.as_str(), "question");
}
#[tokio::test]
async fn run_turn_text_only_event_order() {
let runtime = AgentBuilder::new(Arc::new(ScriptedProvider::new(text_script("answer"))))
.system_prompt("test")
.build()
.expect("build runtime");
let sid = runtime.create_session().await;
let (events, outcome) = runtime
.run_turn_collect(sid, "question")
.await
.expect("run_turn_collect");
assert!(matches!(outcome, RunOutcome::Completed));
assert_eq!(
terminal_seq(&events),
vec!["AfterUserInput", "BeforeLlm", "RunFinished"]
);
}
fn reasoning_script() -> Vec<Vec<StreamChunk>> {
(0..10)
.map(|_| {
vec![
StreamChunk::Thought("thinking but never committing".to_string()),
StreamChunk::Stop {
finish_reason: Some("stop".to_string()),
},
]
})
.collect()
}
#[tokio::test]
async fn reasoning_only_with_tools_fails_instead_of_fake_completing() {
let runtime = AgentBuilder::new(Arc::new(ScriptedProvider::new(reasoning_script())))
.system_prompt("test")
.register_tool(EchoTool)
.build()
.expect("build runtime");
let sid = runtime.create_session().await;
let result = runtime
.run_turn(sid.clone(), "do the thing", |_| Ok(()))
.await;
assert!(
matches!(result, Ok(RunOutcome::Failed { .. })),
"reasoning-only with tools must fail, not complete: {:?}",
result
);
}
#[tokio::test]
async fn reasoning_only_without_tools_fails_instead_of_promoting() {
let runtime = AgentBuilder::new(Arc::new(ScriptedProvider::new(reasoning_script())))
.system_prompt("test")
.build()
.expect("build runtime");
let sid = runtime.create_session().await;
let result = runtime
.run_turn(sid.clone(), "do the thing", |_| Ok(()))
.await;
assert!(
matches!(result, Ok(RunOutcome::Failed { .. })),
"reasoning-only without tools must fail, not promote reasoning to the answer: {:?}",
result
);
}
#[tokio::test]
async fn empty_response_fails_after_bounded_retries() {
let script = (0..5)
.map(|_| {
vec![StreamChunk::Stop {
finish_reason: Some("stop".to_string()),
}]
})
.collect();
let runtime = AgentBuilder::new(Arc::new(ScriptedProvider::new(script)))
.system_prompt("test")
.register_tool(EchoTool)
.build()
.expect("build runtime");
let sid = runtime.create_session().await;
let result = runtime
.run_turn(sid.clone(), "do the thing", |_| Ok(()))
.await;
assert!(
matches!(result, Ok(RunOutcome::Failed { .. })),
"empty response must fail after bounded retries, not loop forever: {:?}",
result
);
}
#[tokio::test]
async fn run_turn_tool_call_then_text_event_order() {
let runtime = AgentBuilder::new(Arc::new(ScriptedProvider::new(vec![
vec![
StreamChunk::ToolCall(serde_json::json!({
"delta": {
"tool_calls": [{
"id": "call_1",
"function": {
"name": "echo",
"arguments": "{\"text\":\"hello\"}"
}
}]
}
})),
StreamChunk::Stop {
finish_reason: Some("tool_calls".to_string()),
},
],
vec![
StreamChunk::Text("final answer".to_string()),
StreamChunk::Stop {
finish_reason: Some("stop".to_string()),
},
],
])))
.system_prompt("test")
.register_tool(EchoTool)
.build()
.expect("build runtime");
let sid = runtime.create_session().await;
let (events, outcome) = runtime
.run_turn_collect(sid, "echo hello")
.await
.expect("run_turn_collect");
assert!(matches!(outcome, RunOutcome::Completed));
assert_eq!(
terminal_seq(&events),
vec![
"AfterUserInput",
"BeforeLlm",
"BeforeToolCalls",
"AfterToolCalls",
"BeforeLlm",
"RunFinished"
]
);
}
#[tokio::test]
async fn run_turn_cancel_does_not_emit_runfinished() {
let runtime = AgentBuilder::new(Arc::new(ScriptedProvider::new(vec![vec![
StreamChunk::ToolCall(serde_json::json!({
"delta": {
"tool_calls": [{
"id": "call_cancel",
"function": { "name": "cancel", "arguments": "{}" }
}]
}
})),
StreamChunk::Stop {
finish_reason: Some("tool_calls".to_string()),
},
]])))
.system_prompt("test")
.register_tool(CancelTool)
.build()
.expect("build runtime");
let sid = runtime.create_session().await;
let events = Arc::new(Mutex::new(Vec::new()));
let events_clone = events.clone();
let result = runtime
.run_turn(sid.clone(), "cancel now", move |event| {
events_clone.lock().unwrap().push(event);
Ok(())
})
.await;
let events = events.lock().unwrap();
assert!(matches!(result, Ok(RunOutcome::Cancelled)));
assert!(
events
.iter()
.any(|e| matches!(e, RuntimeEvent::RunCancelled { .. })),
"cancel must emit RunCancelled"
);
assert_eq!(
events
.iter()
.filter(|e| matches!(e, RuntimeEvent::RunFinished { .. }))
.count(),
0,
"cancelled run_turn must not emit RunFinished"
);
}
#[tokio::test]
async fn run_text_only_event_order() {
let runtime = AgentBuilder::new(Arc::new(ScriptedProvider::new(text_script("answer"))))
.system_prompt("test")
.build()
.expect("build runtime");
let sid = runtime.create_session().await;
runtime.add_user_message(&sid, "question").await.unwrap();
let events = Arc::new(Mutex::new(Vec::new()));
let events_clone = events.clone();
let outcome = runtime
.run(sid.clone(), move |event| {
events_clone.lock().unwrap().push(event);
Ok(())
})
.await
.expect("run");
let events = events.lock().unwrap();
assert!(matches!(outcome, RunOutcome::Completed));
assert_eq!(terminal_seq(&events), vec!["BeforeLlm", "RunFinished"]);
}
#[tokio::test]
async fn run_managed_text_only_event_order() {
let runtime = AgentBuilder::new(Arc::new(ScriptedProvider::new(text_script("done"))))
.system_prompt("test")
.build()
.expect("build runtime");
let sid = runtime.create_session().await;
let events = Arc::new(Mutex::new(Vec::new()));
let events_clone = events.clone();
let outcome = runtime
.run_managed(sid.clone(), "do something", move |event| {
events_clone.lock().unwrap().push(event);
Ok(())
})
.await
.expect("run_managed");
let events = events.lock().unwrap();
assert!(matches!(outcome, RunOutcome::Completed));
assert_eq!(terminal_seq(&events), vec!["BeforeLlm", "RunFinished"]);
}
#[tokio::test]
async fn max_turns_exceeded_event_order() {
let runtime = AgentBuilder::new(Arc::new(ScriptedProvider::new(vec![
vec![
StreamChunk::ToolCall(serde_json::json!({
"delta": {
"tool_calls": [{
"id": "call_1",
"function": {
"name": "echo",
"arguments": "{\"text\":\"hello\"}"
}
}]
}
})),
StreamChunk::Stop {
finish_reason: Some("tool_calls".to_string()),
},
],
vec![
StreamChunk::Text("unused".to_string()),
StreamChunk::Stop {
finish_reason: Some("stop".to_string()),
},
],
])))
.system_prompt("test")
.register_tool(EchoTool)
.execution_max_turns(1)
.build()
.expect("build runtime");
let sid = runtime.create_session().await;
let (events, outcome) = runtime
.run_turn_collect(sid, "go")
.await
.expect("run_turn_collect");
assert!(matches!(outcome, RunOutcome::MaxTurnsExceeded { .. }));
assert_eq!(
terminal_seq(&events),
vec![
"AfterUserInput",
"BeforeLlm",
"BeforeToolCalls",
"AfterToolCalls",
"RunFinished"
]
);
}
#[tokio::test]
async fn run_turn_tool_cancelled_mid_execution_returns_cancelled() {
let runtime = AgentBuilder::new(Arc::new(ScriptedProvider::new(vec![vec![
StreamChunk::ToolCall(serde_json::json!({
"delta": {
"tool_calls": [{
"id": "call_cancel",
"function": { "name": "cancel_pending", "arguments": "{}" }
}]
}
})),
StreamChunk::Stop {
finish_reason: Some("tool_calls".to_string()),
},
]])))
.system_prompt("test")
.register_tool(CancelPendingTool)
.build()
.expect("build runtime");
let sid = runtime.create_session().await;
let events = Arc::new(Mutex::new(Vec::new()));
let events_clone = events.clone();
let result = runtime
.run_turn(sid.clone(), "cancel now", move |event| {
events_clone.lock().unwrap().push(event);
Ok(())
})
.await;
let events = events.lock().unwrap();
assert!(
matches!(&result, Err(e) if e.is_cancelled()),
"run_turn should surface a Cancelled error for mid-execution tool cancel, got: {result:?}"
);
assert!(
events
.iter()
.any(|e| matches!(e, RuntimeEvent::RunCancelled { .. })),
"cancelled tool execution must emit RunCancelled"
);
}
#[tokio::test]
async fn run_turn_llm_stream_cancelled_returns_cancelled() {
let runtime = AgentBuilder::new(Arc::new(CancelledStreamProvider))
.system_prompt("test")
.build()
.expect("build runtime");
let sid = runtime.create_session().await;
let events = Arc::new(Mutex::new(Vec::new()));
let events_clone = events.clone();
let result = runtime
.run_turn(sid.clone(), "test input", move |event| {
events_clone.lock().unwrap().push(event);
Ok(())
})
.await;
let events = events.lock().unwrap();
assert!(
matches!(&result, Ok(RunOutcome::Failed { .. }) | Err(_)),
"run_turn should surface the LLM stream error, got: {result:?}"
);
assert!(
events
.iter()
.any(|e| matches!(e, RuntimeEvent::RunFinished { .. })),
"LLM stream error must emit RunFinished"
);
}
struct FailThenSucceedProvider {
fail_count: Mutex<u32>,
remaining_fails: Mutex<u32>,
}
impl FailThenSucceedProvider {
fn new(fail_count: u32) -> Self {
Self {
fail_count: Mutex::new(fail_count),
remaining_fails: Mutex::new(fail_count),
}
}
fn calls_made(&self) -> u32 {
*self.fail_count.lock().unwrap() - *self.remaining_fails.lock().unwrap()
}
}
#[async_trait]
impl LlmProvider for FailThenSucceedProvider {
async fn stream(
&self,
_request: llm_trait::ChatRequest,
) -> Result<llm_trait::ChatStream, llm_trait::LlmError> {
let mut remaining = self.remaining_fails.lock().unwrap();
if *remaining > 0 {
*remaining -= 1;
struct ErrorStream;
impl Stream for ErrorStream {
type Item = Result<StreamChunk, llm_trait::LlmError>;
fn poll_next(
self: Pin<&mut Self>,
_cx: &mut Context<'_>,
) -> Poll<Option<Self::Item>> {
Poll::Ready(Some(Err(llm_trait::LlmError::llm(
"simulated SSE read error",
))))
}
}
Ok(llm_trait::ChatStream::new(Box::pin(ErrorStream)))
} else {
Ok(llm_trait::ChatStream::new(Box::pin(
futures_util::stream::iter(vec![
Ok(StreamChunk::Text("recovered answer".to_string())),
Ok(StreamChunk::Stop {
finish_reason: Some("stop".to_string()),
}),
]),
)))
}
}
async fn chat(
&self,
_request: llm_trait::ChatRequest,
) -> Result<llm_trait::ChatResponse, llm_trait::LlmError> {
Ok(llm_trait::ChatResponse {
content: String::new(),
reasoning_content: None,
tool_calls: vec![],
usage: Default::default(),
finish_reason: llm_trait::response::FinishReason::Stop,
raw: None,
thinking_signature: None,
})
}
fn capabilities(&self) -> llm_trait::Capabilities {
llm_trait::Capabilities {
supports_streaming: true,
supports_tools: true,
supports_vision: false,
supports_thinking: false,
max_context_tokens: None,
max_output_tokens: None,
}
}
fn info(&self) -> llm_trait::ProviderInfo {
llm_trait::ProviderInfo {
name: "test".to_string(),
model: "test".to_string(),
version: None,
}
}
}
#[tokio::test]
async fn mid_stream_retry_succeeds_after_transient_error() {
let provider = FailThenSucceedProvider::new(1);
let runtime = AgentBuilder::new(Arc::new(provider))
.system_prompt("test")
.build()
.expect("build runtime");
let sid = runtime.create_session().await;
let result = runtime.run_turn(sid, "test", |_| Ok(())).await;
assert!(
matches!(result, Ok(RunOutcome::Completed)),
"mid-stream retry should recover from transient error, got: {result:?}"
);
}
#[tokio::test]
async fn mid_stream_retry_exhausted_returns_error() {
let provider = Arc::new(FailThenSucceedProvider::new(4));
let runtime = AgentBuilder::new(provider.clone())
.system_prompt("test")
.build()
.expect("build runtime");
let sid = runtime.create_session().await;
let events = Arc::new(Mutex::new(Vec::new()));
let events_clone = events.clone();
let result = runtime
.run_turn(sid, "test", move |event| {
events_clone.lock().unwrap().push(event);
Ok(())
})
.await;
assert!(
matches!(&result, Ok(RunOutcome::Failed { .. }) | Err(_)),
"exhausted retries should return error, got: {result:?}"
);
assert_eq!(
provider.calls_made(),
4,
"should make 4 calls: 1 initial + 3 retries"
);
}
#[tokio::test]
async fn mid_stream_retry_does_not_retry_cancellation() {
let provider = Arc::new(FailThenSucceedProvider::new(10));
let runtime = AgentBuilder::new(provider.clone())
.system_prompt("test")
.build()
.expect("build runtime");
let sid = runtime.create_session().await;
let result = runtime.run_turn(sid, "test", |_| Ok(())).await;
assert!(
matches!(&result, Ok(RunOutcome::Failed { .. }) | Err(_)),
"10 consecutive failures should exhaust retries, got: {result:?}"
);
assert_eq!(
provider.calls_made(),
4,
"should make 4 calls (1 initial + 3 retries), not 10"
);
}
#[tokio::test]
async fn run_turn_with_llm_retry_completes() {
let runtime = AgentBuilder::new(Arc::new(ScriptedProvider::new(text_script("answer"))))
.system_prompt("test")
.llm_retry(crate::types::RetryConfig::default().max_retries(1))
.build()
.expect("build runtime");
let sid = runtime.create_session().await;
let result = runtime.run_turn(sid.clone(), "hi", |_| Ok(())).await;
assert!(
matches!(result, Ok(RunOutcome::Completed)),
"run_turn with llm_retry should complete: {:?}",
result
);
}
#[tokio::test]
async fn run_managed_tool_cancelled_mid_execution_returns_cancelled() {
let runtime = AgentBuilder::new(Arc::new(ScriptedProvider::new(vec![vec![
StreamChunk::ToolCall(serde_json::json!({
"delta": {
"tool_calls": [{
"id": "call_cancel",
"function": { "name": "cancel_pending", "arguments": "{}" }
}]
}
})),
StreamChunk::Stop {
finish_reason: Some("tool_calls".to_string()),
},
]])))
.system_prompt("test")
.register_tool(CancelPendingTool)
.build()
.expect("build runtime");
let sid = runtime.create_session().await;
let result = runtime
.run_managed(sid.clone(), "cancel now", |_| Ok(()))
.await;
assert!(
matches!(&result, Err(e) if e.is_cancelled()),
"run_managed should surface a Cancelled error for mid-execution tool cancel, got: {result:?}"
);
}
#[tokio::test]
async fn run_managed_max_turns_exceeded() {
let runtime = AgentBuilder::new(Arc::new(ScriptedProvider::new(text_script("done"))))
.system_prompt("test")
.execution_max_turns(1)
.build()
.expect("build runtime");
let sid = runtime.create_session().await;
runtime.follow_up("continue".to_string());
let result = runtime
.run_managed(sid.clone(), "do something", |_| Ok(()))
.await;
assert!(
matches!(result, Ok(RunOutcome::MaxTurnsExceeded { .. })),
"run_managed should exceed its cumulative turn budget: {:?}",
result
);
}
#[tokio::test]
async fn tool_use_with_no_tool_calls_returns_failed() {
let runtime = AgentBuilder::new(Arc::new(ScriptedProvider::new(vec![vec![
StreamChunk::ToolCall(serde_json::json!({ "no_delta": true })),
StreamChunk::Stop {
finish_reason: Some("tool_use".to_string()),
},
]])))
.system_prompt("test")
.build()
.expect("build runtime");
let sid = runtime.create_session().await;
let result = runtime.run_turn(sid, "do something", |_| Ok(())).await;
assert!(
matches!(result, Ok(RunOutcome::Failed { .. })),
"tool_use with no tool calls should return Failed: {:?}",
result
);
}
#[tokio::test]
async fn truncated_text_only_returns_failed() {
let runtime = AgentBuilder::new(Arc::new(ScriptedProvider::new(vec![vec![
StreamChunk::Text("Here is part of my answer before being cut".to_string()),
StreamChunk::Stop {
finish_reason: Some("max_tokens".to_string()),
},
]])))
.system_prompt("test")
.build()
.expect("build runtime");
let sid = runtime.create_session().await;
let result = runtime.run_turn(sid, "do something", |_| Ok(())).await;
assert!(
matches!(result, Ok(RunOutcome::Failed { .. })),
"truncated text-only should return Failed: {:?}",
result
);
}
#[tokio::test]
async fn openai_length_text_only_returns_failed() {
let runtime = AgentBuilder::new(Arc::new(ScriptedProvider::new(vec![vec![
StreamChunk::Text("Partial response before length limit".to_string()),
StreamChunk::Stop {
finish_reason: Some("length".to_string()),
},
]])))
.system_prompt("test")
.build()
.expect("build runtime");
let sid = runtime.create_session().await;
let result = runtime.run_turn(sid, "do something", |_| Ok(())).await;
assert!(
matches!(result, Ok(RunOutcome::Failed { .. })),
"OpenAI length with text-only should return Failed: {:?}",
result
);
}
#[tokio::test]
async fn tool_call_chunk_with_no_parsed_calls_and_stop_finish_fails() {
let guard = Arc::new(RecordingGuard::new(GuardDecision::Fail {
error: "incomplete tool call".to_string(),
}));
let guard_clone = guard.clone();
let runtime = AgentBuilder::new(Arc::new(ScriptedProvider::new(vec![vec![
StreamChunk::ToolCall(serde_json::json!({ "no_delta": true })),
StreamChunk::Stop {
finish_reason: Some("stop".to_string()),
},
]])))
.system_prompt("test")
.register_tool(EchoTool)
.guard_dyn(guard_clone)
.build()
.expect("build runtime");
let sid = runtime.create_session().await;
let result = runtime
.run_turn(sid, "帮我解读一下这个工程", |_| Ok(()))
.await;
assert!(
matches!(result, Ok(RunOutcome::Failed { .. })),
"incomplete tool call should fail, not silently complete. Got: {:?}",
result
);
}
use crate::engine::react_loop_guard::{GuardCtx, GuardDecision, ReactLoopGuard};
struct RecordingGuard {
calls: Mutex<Vec<(bool, String)>>,
decision: GuardDecision,
}
impl RecordingGuard {
fn new(decision: GuardDecision) -> Self {
Self {
calls: Mutex::new(Vec::new()),
decision,
}
}
}
#[async_trait]
impl ReactLoopGuard for RecordingGuard {
async fn on_turn(&self, ctx: &GuardCtx) -> GuardDecision {
if ctx.is_text_only {
self.calls
.lock()
.unwrap()
.push((ctx.run_has_tool_calls, ctx.model_response.clone()));
}
self.decision.clone()
}
}
#[tokio::test]
async fn branch4_guard_called_with_run_has_tool_calls() {
let guard = Arc::new(RecordingGuard::new(GuardDecision::Complete));
let guard_clone = guard.clone();
let runtime = AgentBuilder::new(Arc::new(ScriptedProvider::new(vec![
vec![
StreamChunk::ToolCall(serde_json::json!({
"delta": {
"tool_calls": [{
"id": "call_1",
"function": {
"name": "echo",
"arguments": "{\"text\":\"hello\"}"
}
}]
}
})),
StreamChunk::Stop {
finish_reason: Some("tool_calls".to_string()),
},
],
vec![
StreamChunk::Text("The answer is 42.".to_string()),
StreamChunk::Stop {
finish_reason: Some("stop".to_string()),
},
],
])))
.system_prompt("test")
.register_tool(EchoTool)
.guard_dyn(guard_clone)
.build()
.expect("build runtime");
let sid = runtime.create_session().await;
let result = runtime.run_turn(sid, "echo hello", |_| Ok(())).await;
assert!(
matches!(result, Ok(RunOutcome::Completed)),
"run_turn should complete: {:?}",
result
);
let calls = guard.calls.lock().unwrap();
assert_eq!(calls.len(), 1, "on_text_only should be called exactly once");
let (run_has_tools, response) = &calls[0];
assert!(
*run_has_tools,
"run_has_tool_calls must be true after tool use"
);
assert_eq!(response, "The answer is 42.");
}
#[tokio::test]
async fn branch4_guard_continue_nudges_and_retries() {
let guard = Arc::new(RecordingGuard::new(GuardDecision::Continue {
nudge: Some("please continue".to_string()),
}));
let guard_clone = guard.clone();
let runtime = AgentBuilder::new(Arc::new(ScriptedProvider::new(vec![
vec![
StreamChunk::ToolCall(serde_json::json!({
"delta": {
"tool_calls": [{
"id": "call_1",
"function": {
"name": "echo",
"arguments": "{\"text\":\"hello\"}"
}
}]
}
})),
StreamChunk::Stop {
finish_reason: Some("tool_calls".to_string()),
},
],
vec![
StreamChunk::Text("maybe done".to_string()),
StreamChunk::Stop {
finish_reason: Some("stop".to_string()),
},
],
vec![
StreamChunk::Text("here is the complete answer".to_string()),
StreamChunk::Stop {
finish_reason: Some("stop".to_string()),
},
],
])))
.system_prompt("test")
.register_tool(EchoTool)
.guard_dyn(guard_clone)
.build()
.expect("build runtime");
let sid = runtime.create_session().await;
let (_events, _outcome) = runtime
.run_turn_collect(sid.clone(), "echo hello")
.await
.expect("run_turn_collect");
let (call_count, all_have_tools) = {
let calls = guard.calls.lock().unwrap();
(calls.len(), calls.iter().all(|(has_tools, _)| *has_tools))
};
assert!(
call_count >= 2,
"guard should be called at least twice (initial + after nudge), got: {}",
call_count
);
let session = runtime.session(&sid).await.expect("session exists");
let has_nudge = session.chat_messages().iter().any(
|m| matches!(m, ChatMessage::User { content, .. } if content.contains("please continue")),
);
assert!(has_nudge, "session should contain the nudge message");
assert!(
all_have_tools,
"all on_text_only calls should have run_has_tool_calls = true"
);
}
struct EmptyResponseStrikeGuard {
max_strikes: usize,
calls: AtomicUsize,
}
impl EmptyResponseStrikeGuard {
fn new(max_strikes: usize) -> Self {
Self {
max_strikes,
calls: AtomicUsize::new(0),
}
}
}
#[async_trait]
impl ReactLoopGuard for EmptyResponseStrikeGuard {
async fn on_turn(&self, ctx: &GuardCtx) -> GuardDecision {
if ctx.is_empty_response {
let count = self.calls.fetch_add(1, Ordering::SeqCst) + 1;
if count >= self.max_strikes {
GuardDecision::Fail {
error: format!("exceeded max retries ({})", self.max_strikes),
}
} else {
GuardDecision::Continue {
nudge: Some("please retry".to_string()),
}
}
} else {
GuardDecision::Complete
}
}
}
#[tokio::test]
async fn incomplete_tool_calls_fail_after_max_strikes() {
let max_strikes = 3;
let guard = Arc::new(EmptyResponseStrikeGuard::new(max_strikes));
let guard_clone = guard.clone();
let script: Vec<Vec<StreamChunk>> = (0..max_strikes + 1)
.map(|_| {
vec![
StreamChunk::ToolCall(serde_json::json!({ "no_delta": true })),
StreamChunk::Stop {
finish_reason: Some("stop".to_string()),
},
]
})
.collect();
let runtime = AgentBuilder::new(Arc::new(ScriptedProvider::new(script)))
.system_prompt("test")
.register_tool(EchoTool)
.guard_dyn(guard_clone)
.build()
.expect("build runtime");
let sid = runtime.create_session().await;
let result = runtime
.run_turn(sid, "帮我解读一下这个工程", |_| Ok(()))
.await;
assert!(
matches!(result, Ok(RunOutcome::Failed { .. })),
"incomplete tool calls should fail after {} strikes, got: {:?}",
max_strikes,
result
);
let calls = guard.calls.load(Ordering::SeqCst);
assert_eq!(
calls, max_strikes,
"guard should be called {} times, got: {}",
max_strikes, calls
);
}