use crate::interceptor::ToolCallInterceptor;
use crate::registry::ToolRegistry;
use pe_core::error::PeError;
use pe_core::message::{Message, ToolCall};
use pe_core::node::{NodeContext, NodeFn, NodeFuture, NodeResult};
use pe_core::state::{CoreState, CoreStateUpdate};
use std::marker::PhantomData;
use std::sync::Arc;
pub struct ToolNode<S: CoreState> {
registry: Arc<ToolRegistry>,
handle_tool_errors: bool,
interceptor: Option<Arc<dyn ToolCallInterceptor>>,
_phantom: PhantomData<S>,
}
impl<S: CoreState> ToolNode<S>
where
S::Update: CoreStateUpdate,
{
pub fn new(registry: Arc<ToolRegistry>) -> Self {
Self {
registry,
handle_tool_errors: true,
interceptor: None,
_phantom: PhantomData,
}
}
#[must_use = "builder methods return the modified builder"]
pub fn with_error_handling(mut self, handle: bool) -> Self {
self.handle_tool_errors = handle;
self
}
#[must_use = "builder methods return the modified builder"]
pub fn with_interceptor(mut self, interceptor: impl ToolCallInterceptor + 'static) -> Self {
self.interceptor = Some(Arc::new(interceptor));
self
}
}
impl<S: CoreState> NodeFn<S> for ToolNode<S>
where
S::Update: CoreStateUpdate,
{
fn call(&self, state: &S, ctx: &NodeContext) -> NodeFuture<S::Update> {
let registry = self.registry.clone();
let handle_errors = self.handle_tool_errors;
let interceptor = self.interceptor.clone();
let tool_observer = ctx.tool_observer.clone();
let tool_calls = extract_tool_calls(state.messages());
if tool_calls.is_empty() {
return Box::pin(async move { NodeResult::Update(S::Update::default()) });
}
Box::pin(async move {
let tool_calls = match &interceptor {
Some(int) => tool_calls.into_iter().map(|tc| int.intercept(tc)).collect(),
None => tool_calls,
};
let tasks: Vec<_> = tool_calls
.into_iter()
.map(|tc| {
let reg = registry.clone();
let obs = tool_observer.clone();
tokio::spawn(async move {
execute_single_tool_observed(®, tc, handle_errors, obs.as_deref()).await
})
})
.collect();
let results = futures::future::join_all(tasks).await;
let mut messages = Vec::new();
for join_result in results {
match join_result {
Ok(ToolExecResult::Message(msg)) => messages.push(msg),
Ok(ToolExecResult::PropagateError(err)) => {
return NodeResult::Error(err);
}
Err(join_err) => {
return NodeResult::Error(PeError::Internal {
details: format!("Tool task panicked: {join_err}"),
});
}
}
}
let update = <S::Update as CoreStateUpdate>::from_messages(messages);
NodeResult::Update(update)
})
}
fn name(&self) -> &str {
"tool_node"
}
}
enum ToolExecResult {
Message(Message),
PropagateError(PeError),
}
async fn execute_single_tool_observed(
registry: &ToolRegistry,
tc: ToolCall,
handle_errors: bool,
observer: Option<&dyn pe_core::node::ToolObserver>,
) -> ToolExecResult {
let input_summary = summarize_tool_input(&tc);
if let Some(obs) = observer {
obs.on_tool_start(&tc.name, &input_summary).await;
}
let start = std::time::Instant::now();
let result = match registry.get(&tc.name) {
None => {
tracing::warn!("Tool '{}' not found in registry", tc.name);
if let Some(obs) = observer {
obs.on_tool_error(&tc.name, "Tool not found in registry")
.await;
}
return ToolExecResult::Message(Message::tool(
format!("Tool '{}' not found in registry.", tc.name),
tc.id,
));
}
Some(tool) => tool.execute_structured(tc.args.clone()).await,
};
let elapsed = start.elapsed();
match result {
Ok(tool_result) => {
if let Some(obs) = observer {
obs.on_tool_complete(&tc.name, elapsed).await;
}
let metadata = if tool_result.metadata.is_empty() {
None
} else {
Some(tool_result.metadata)
};
ToolExecResult::Message(Message::tool_with_metadata(
tool_result.output.to_string(),
tc.id,
metadata,
Some(elapsed.as_millis() as u64),
))
}
Err(e) if handle_errors => {
tracing::warn!("Tool '{}' failed (handled): {}", tc.name, e);
if let Some(obs) = observer {
obs.on_tool_error(&tc.name, &e.to_string()).await;
}
ToolExecResult::Message(Message::tool_with_metadata(
format!("Tool error: {e}"),
tc.id,
None,
Some(elapsed.as_millis() as u64),
))
}
Err(e) => {
tracing::error!("Tool '{}' failed (propagating): {}", tc.name, e);
if let Some(obs) = observer {
obs.on_tool_error(&tc.name, &e.to_string()).await;
}
ToolExecResult::PropagateError(PeError::ToolExecution {
tool: tc.name,
reason: e.to_string(),
})
}
}
}
fn summarize_tool_input(tc: &ToolCall) -> String {
let s = tc.args.to_string();
if s.len() <= 100 {
s
} else {
format!("{}...", &s[..97])
}
}
fn extract_tool_calls(messages: &[Message]) -> Vec<ToolCall> {
messages
.iter()
.rev()
.find_map(|m| {
if let Message::Ai(ai) = m {
if !ai.tool_calls.is_empty() {
return Some(ai.tool_calls.clone());
}
}
None
})
.unwrap_or_default()
}
impl<S: CoreState> std::fmt::Debug for ToolNode<S>
where
S::Update: CoreStateUpdate,
{
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ToolNode")
.field("handle_tool_errors", &self.handle_tool_errors)
.field("has_interceptor", &self.interceptor.is_some())
.finish()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::interceptor::ToolCallInterceptor;
use crate::tool::FunctionTool;
use pe_core::message::{AiMessage, MessageContent};
use pe_core::state::{CoreStateUpdate, ExecutionContext, StateUpdate};
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::time::{Duration, Instant};
#[derive(Debug, Clone, Serialize, Deserialize)]
struct ToolTestState {
messages: Vec<Message>,
iterations: u32,
thread_id: String,
context: ExecutionContext,
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
struct ToolTestUpdate {
messages: Option<Vec<Message>>,
}
impl StateUpdate for ToolTestUpdate {}
impl CoreStateUpdate for ToolTestUpdate {
fn from_messages(msgs: Vec<Message>) -> Self {
Self {
messages: Some(msgs),
}
}
}
impl pe_core::state::State for ToolTestState {
type Update = ToolTestUpdate;
fn apply(&mut self, update: Self::Update) {
if let Some(msgs) = update.messages {
self.messages.extend(msgs);
}
}
}
impl CoreState for ToolTestState {
fn messages(&self) -> &[Message] {
&self.messages
}
fn messages_mut(&mut self) -> &mut Vec<Message> {
&mut self.messages
}
fn iterations(&self) -> u32 {
self.iterations
}
fn set_iterations(&mut self, n: u32) {
self.iterations = n;
}
fn thread_id(&self) -> &str {
&self.thread_id
}
fn context(&self) -> &ExecutionContext {
&self.context
}
fn context_mut(&mut self) -> &mut ExecutionContext {
&mut self.context
}
}
fn make_ctx() -> NodeContext {
NodeContext {
step: 1,
recursion_limit: 25,
node_name: "tool_node".into(),
activation: pe_core::node::ActivationReason::EntryPoint,
metadata: HashMap::new(),
phase_store: pe_core::phase_store::PhaseStateStore::new(),
stream_sender: None,
tool_observer: None,
lobe_runtime_service_factory: None,
}
}
fn make_state_with_tool_calls(calls: Vec<ToolCall>) -> ToolTestState {
ToolTestState {
messages: vec![Message::Ai(AiMessage {
content: MessageContent::Text("Calling tools...".into()),
tool_calls: calls,
invalid_tool_calls: vec![],
usage_metadata: None,
response_metadata: HashMap::new(),
id: None,
})],
iterations: 0,
thread_id: "t1".into(),
context: ExecutionContext::new("test"),
}
}
fn make_echo_registry() -> ToolRegistry {
let mut reg = ToolRegistry::new();
reg.register(FunctionTool::new(
"echo",
"Echoes input",
serde_json::json!({"type": "object"}),
|input| Box::pin(async move { Ok(input) }),
))
.unwrap();
reg
}
#[tokio::test]
async fn single_tool_call_produces_tool_message() {
let reg = make_echo_registry();
let node = ToolNode::<ToolTestState>::new(Arc::new(reg));
let state = make_state_with_tool_calls(vec![ToolCall {
id: "tc_1".into(),
name: "echo".into(),
args: serde_json::json!({"msg": "hello"}),
}]);
let result = node.call(&state, &make_ctx()).await;
match result {
NodeResult::Update(update) => {
let msgs = update.messages.expect("should have messages");
assert_eq!(msgs.len(), 1);
if let Message::Tool(tm) = &msgs[0] {
assert_eq!(tm.tool_call_id, "tc_1");
assert!(tm.content.contains("hello"));
} else {
panic!("expected ToolMessage, got {:?}", msgs[0]);
}
}
other => panic!("expected Update, got {:?}", other),
}
}
#[tokio::test]
async fn parallel_tool_calls_all_execute() {
let mut reg = ToolRegistry::new();
for name in ["tool_a", "tool_b", "tool_c"] {
let n = name.to_string();
reg.register(FunctionTool::new(
name,
format!("Tool {name}"),
serde_json::json!({"type": "object"}),
move |_| {
let n = n.clone();
Box::pin(async move {
tokio::time::sleep(Duration::from_millis(50)).await;
Ok(serde_json::json!({"from": n}))
})
},
))
.unwrap();
}
let node = ToolNode::<ToolTestState>::new(Arc::new(reg));
let state = make_state_with_tool_calls(vec![
ToolCall {
id: "tc_a".into(),
name: "tool_a".into(),
args: serde_json::json!({}),
},
ToolCall {
id: "tc_b".into(),
name: "tool_b".into(),
args: serde_json::json!({}),
},
ToolCall {
id: "tc_c".into(),
name: "tool_c".into(),
args: serde_json::json!({}),
},
]);
let start = Instant::now();
let result = node.call(&state, &make_ctx()).await;
let elapsed = start.elapsed();
assert!(
elapsed < Duration::from_millis(150),
"Expected parallel execution (<150ms), took {:?}",
elapsed
);
match result {
NodeResult::Update(update) => {
let msgs = update.messages.expect("should have messages");
assert_eq!(msgs.len(), 3);
for msg in &msgs {
assert!(matches!(msg, Message::Tool(_)));
}
}
other => panic!("expected Update, got {:?}", other),
}
}
#[tokio::test]
async fn tool_error_handled_becomes_tool_message() {
let mut reg = ToolRegistry::new();
reg.register(FunctionTool::new(
"fail",
"Always fails",
serde_json::json!({"type": "object"}),
|_| {
Box::pin(async {
Err(PeError::ToolExecution {
tool: "fail".into(),
reason: "boom".into(),
})
})
},
))
.unwrap();
let node = ToolNode::<ToolTestState>::new(Arc::new(reg)).with_error_handling(true);
let state = make_state_with_tool_calls(vec![ToolCall {
id: "tc_fail".into(),
name: "fail".into(),
args: serde_json::json!({}),
}]);
let result = node.call(&state, &make_ctx()).await;
match result {
NodeResult::Update(update) => {
let msgs = update.messages.expect("should have messages");
assert_eq!(msgs.len(), 1);
if let Message::Tool(tm) = &msgs[0] {
assert!(tm.content.contains("Tool error"));
assert!(tm.content.contains("boom"));
assert_eq!(tm.tool_call_id, "tc_fail");
} else {
panic!("expected ToolMessage");
}
}
other => panic!("expected Update, got {:?}", other),
}
}
#[tokio::test]
async fn tool_error_propagated_when_handling_disabled() {
let mut reg = ToolRegistry::new();
reg.register(FunctionTool::new(
"fail",
"Always fails",
serde_json::json!({"type": "object"}),
|_| {
Box::pin(async {
Err(PeError::ToolExecution {
tool: "fail".into(),
reason: "critical failure".into(),
})
})
},
))
.unwrap();
let node = ToolNode::<ToolTestState>::new(Arc::new(reg)).with_error_handling(false);
let state = make_state_with_tool_calls(vec![ToolCall {
id: "tc_fail".into(),
name: "fail".into(),
args: serde_json::json!({}),
}]);
let result = node.call(&state, &make_ctx()).await;
assert!(result.is_error(), "expected Error, got {:?}", result);
}
#[tokio::test]
async fn unknown_tool_always_produces_not_found_message() {
let reg = ToolRegistry::new(); let node = ToolNode::<ToolTestState>::new(Arc::new(reg));
let state = make_state_with_tool_calls(vec![ToolCall {
id: "tc_missing".into(),
name: "nonexistent".into(),
args: serde_json::json!({}),
}]);
let result = node.call(&state, &make_ctx()).await;
match result {
NodeResult::Update(update) => {
let msgs = update.messages.expect("should have messages");
assert_eq!(msgs.len(), 1);
if let Message::Tool(tm) = &msgs[0] {
assert!(tm.content.contains("not found"));
assert_eq!(tm.tool_call_id, "tc_missing");
} else {
panic!("expected ToolMessage");
}
}
other => panic!("expected Update, got {:?}", other),
}
}
#[tokio::test]
async fn interceptor_modifies_args_before_execution() {
struct InjectKeyInterceptor;
impl ToolCallInterceptor for InjectKeyInterceptor {
fn intercept(&self, mut call: ToolCall) -> ToolCall {
if let Some(obj) = call.args.as_object_mut() {
obj.insert("injected".into(), serde_json::json!(true));
}
call
}
}
let reg = make_echo_registry();
let node =
ToolNode::<ToolTestState>::new(Arc::new(reg)).with_interceptor(InjectKeyInterceptor);
let state = make_state_with_tool_calls(vec![ToolCall {
id: "tc_int".into(),
name: "echo".into(),
args: serde_json::json!({"original": "data"}),
}]);
let result = node.call(&state, &make_ctx()).await;
match result {
NodeResult::Update(update) => {
let msgs = update.messages.expect("should have messages");
assert_eq!(msgs.len(), 1);
if let Message::Tool(tm) = &msgs[0] {
assert!(tm.content.contains("injected"));
assert!(tm.content.contains("original"));
} else {
panic!("expected ToolMessage");
}
}
other => panic!("expected Update, got {:?}", other),
}
}
#[tokio::test]
async fn tool_message_includes_duration_ms() {
let reg = make_echo_registry();
let node = ToolNode::<ToolTestState>::new(Arc::new(reg));
let state = make_state_with_tool_calls(vec![ToolCall {
id: "tc_dur".into(),
name: "echo".into(),
args: serde_json::json!({"msg": "timing"}),
}]);
let result = node.call(&state, &make_ctx()).await;
match result {
NodeResult::Update(update) => {
let msgs = update.messages.expect("should have messages");
assert_eq!(msgs.len(), 1);
if let Message::Tool(tm) = &msgs[0] {
assert!(
tm.duration_ms.is_some(),
"ToolMessage should include duration_ms"
);
assert!(tm.duration_ms.unwrap() < 5000);
} else {
panic!("expected ToolMessage, got {:?}", msgs[0]);
}
}
other => panic!("expected Update, got {:?}", other),
}
}
#[tokio::test]
async fn tool_error_message_includes_duration_ms() {
let mut reg = ToolRegistry::new();
reg.register(FunctionTool::new(
"fail",
"Always fails",
serde_json::json!({"type": "object"}),
|_| {
Box::pin(async {
Err(PeError::ToolExecution {
tool: "fail".into(),
reason: "boom".into(),
})
})
},
))
.unwrap();
let node = ToolNode::<ToolTestState>::new(Arc::new(reg)).with_error_handling(true);
let state = make_state_with_tool_calls(vec![ToolCall {
id: "tc_fail_dur".into(),
name: "fail".into(),
args: serde_json::json!({}),
}]);
let result = node.call(&state, &make_ctx()).await;
match result {
NodeResult::Update(update) => {
let msgs = update.messages.expect("should have messages");
assert_eq!(msgs.len(), 1);
if let Message::Tool(tm) = &msgs[0] {
assert!(
tm.duration_ms.is_some(),
"Error ToolMessage should include duration_ms"
);
} else {
panic!("expected ToolMessage");
}
}
other => panic!("expected Update, got {:?}", other),
}
}
#[tokio::test]
async fn structured_tool_metadata_passes_through() {
let mut reg = ToolRegistry::new();
reg.register(crate::tool::StructuredFunctionTool::new(
"search",
"Search things",
serde_json::json!({"type": "object"}),
|_input| {
Box::pin(async move {
Ok(
crate::tool::ToolResult::ok(serde_json::json!({"results": [1, 2, 3]}))
.with_metadata("result_count", serde_json::json!(3))
.with_metadata("source", serde_json::json!("index")),
)
})
},
))
.unwrap();
let node = ToolNode::<ToolTestState>::new(Arc::new(reg));
let state = make_state_with_tool_calls(vec![ToolCall {
id: "tc_meta".into(),
name: "search".into(),
args: serde_json::json!({}),
}]);
let result = node.call(&state, &make_ctx()).await;
match result {
NodeResult::Update(update) => {
let msgs = update.messages.expect("should have messages");
assert_eq!(msgs.len(), 1);
if let Message::Tool(tm) = &msgs[0] {
let meta = tm.metadata.as_ref().expect("should have metadata");
assert_eq!(meta["result_count"], serde_json::json!(3));
assert_eq!(meta["source"], serde_json::json!("index"));
} else {
panic!("expected ToolMessage");
}
}
other => panic!("expected Update, got {:?}", other),
}
}
#[tokio::test]
async fn no_tool_calls_returns_empty_update() {
let reg = make_echo_registry();
let node = ToolNode::<ToolTestState>::new(Arc::new(reg));
let state = ToolTestState {
messages: vec![Message::ai("No tools needed")],
iterations: 0,
thread_id: "t1".into(),
context: ExecutionContext::new("test"),
};
let result = node.call(&state, &make_ctx()).await;
match result {
NodeResult::Update(update) => {
assert!(
update.messages.is_none(),
"empty update should have no messages"
);
}
other => panic!("expected Update, got {:?}", other),
}
}
}