use std::collections::HashMap;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
use async_trait::async_trait;
use lc_agents::types::{AgentAction, AgentFinish, AgentOutput, AgentStep, ToolInput};
use lc_agents::{AgentError, AgentExecutor, AgentStreamEvent, BaseAgent};
use lc_core::tools::{BaseTool, ToolError};
struct FailingTool {
calls: Arc<AtomicUsize>,
}
#[async_trait]
impl BaseTool for FailingTool {
fn name(&self) -> &str {
"failing"
}
fn description(&self) -> &str {
"always fails"
}
async fn run(&self, _input: String) -> Result<String, ToolError> {
self.calls.fetch_add(1, Ordering::SeqCst);
Err(ToolError::ExecutionFailed("boom".to_string()))
}
}
struct FailThenFinishAgent;
#[async_trait]
impl BaseAgent for FailThenFinishAgent {
async fn plan(
&self,
intermediate_steps: &[AgentStep],
_inputs: &HashMap<String, String>,
) -> Result<AgentOutput, AgentError> {
if intermediate_steps.is_empty() {
return Ok(AgentOutput::Action(AgentAction {
tool: "failing".to_string(),
tool_input: ToolInput::Object {
value: serde_json::json!({"x": 1}),
},
log: "call_fail".to_string(),
}));
}
Ok(AgentOutput::Finish(AgentFinish::new(
"done".to_string(),
String::new(),
)))
}
}
fn failing_harness() -> (AgentExecutor, Arc<AtomicUsize>) {
let calls = Arc::new(AtomicUsize::new(0));
let tool = FailingTool {
calls: calls.clone(),
};
let executor = AgentExecutor::new(Arc::new(FailThenFinishAgent), vec![Arc::new(tool)]);
(executor, calls)
}
#[tokio::test]
async fn stream_tool_error_continues() {
use futures_util::StreamExt;
let (executor, calls) = failing_harness();
let mut stream = executor.stream("go".to_string());
let mut events: Vec<Result<AgentStreamEvent, AgentError>> = Vec::new();
while let Some(event) = stream.next().await {
events.push(event);
}
assert_eq!(calls.load(Ordering::SeqCst), 1, "必败工具应执行 1 次");
assert!(
!events
.iter()
.any(|e| matches!(e, Err(_) | Ok(AgentStreamEvent::Error { .. }))),
"流式单工具失败应转 observation 继续,got: {:?}",
events
);
let tool_end = events.iter().find_map(|e| match e {
Ok(AgentStreamEvent::ToolEnd { output, .. }) => Some(output.clone()),
_ => None,
});
let tool_end = tool_end.expect("应发 ToolEnd 事件");
assert!(
tool_end.contains("[Tool execution error"),
"ToolEnd 应为错误 observation,got: {tool_end}"
);
assert!(
matches!(
events.last(),
Some(Ok(AgentStreamEvent::FinalAnswer { content })) if content == "done"
),
"流应继续到 FinalAnswer,got: {:?}",
events.last()
);
}
#[tokio::test]
async fn sequential_tool_error_stays_hard() {
let (executor, calls) = failing_harness();
let err = executor.invoke("go".to_string()).await.unwrap_err();
assert_eq!(calls.load(Ordering::SeqCst), 1);
assert!(
matches!(err, AgentError::ToolExecutionError(_)),
"顺序路径工具失败应硬停返回 Err(ToolExecutionError),got: {err:?}"
);
}