//! 工具注册表 —— 管理工具注册、发现、调用。 use std::collections::HashMap; use std::sync::Arc; use std::time::Duration; use futures::future::join_all; use serde_json::Value; #[allow(deprecated)] use crate::llm::types::ToolDefinition; use crate::tools::base::{ToolContext, ToolRef}; use crate::tools::error::ToolError; use crate::tools::permission::PermissionChecker; /// 工具调用记录 —— 用于追踪和调试。 #[derive(Debug, Clone)] pub struct ToolInvocation { /// LLM 返回的 tool_call_id —— 用于回传 `Message::ToolResult` 时关联原始调用。 /// /// ponytail: Phase 2 引入。老的循环用 `tool_name` 冒充 tool_call_id, /// 对 OpenAI 碰巧可用,对 Anthropic 必然失败。Anthropic 协议要求 /// `tool_result.tool_use_id` 与上一轮 `tool_use.id` 严格一致。 pub tool_call_id: String, /// 被调用的工具名。 pub tool_name: String, /// 工具的入参。 pub input: Value, /// 工具的输出。 pub output: Result, } impl ToolInvocation { /// 创建一个新的工具调用记录。 pub fn new( tool_call_id: String, tool_name: String, input: Value, output: Result, ) -> Self { Self { tool_call_id, tool_name, input, output, } } } /// 工具注册表 —— 管理工具注册、发现、调用。 /// /// 通过 `Arc` 共享,方法签名 `&self`,可安全跨 task 并行调用。 /// 不支持运行时并发注册(应在 setup 阶段一次性构建后冻结)。 #[derive(Clone, Default)] pub struct ToolRegistry { inner: Arc, } #[derive(Clone, Default)] struct ToolRegistryInner { tools: HashMap, permission_checker: Option>, } impl std::fmt::Debug for ToolRegistry { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { f.debug_struct("ToolRegistry") .field("tool_names", &self.inner.tools.keys().collect::>()) .field("has_checker", &self.inner.permission_checker.is_some()) .finish() } } #[allow(deprecated)] impl ToolRegistry { /// 创建一个新的工具注册表。 pub fn new() -> Self { Self { inner: Arc::new(ToolRegistryInner { tools: HashMap::new(), permission_checker: None, }), } } /// 设置权限检查器(Builder 模式)。 pub fn with_permission_checker(mut self, checker: PermissionChecker) -> Self { let inner = Arc::make_mut(&mut self.inner); inner.permission_checker = Some(Arc::new(checker)); self } /// 注册一个工具。 /// /// 重复注册同名工具返回错误。 pub fn register(&mut self, tool: ToolRef) -> Result<(), ToolError> { let name = tool.name().to_string(); let inner = Arc::make_mut(&mut self.inner); if inner.tools.contains_key(&name) { return Err(ToolError::ExecutionFailed(name, "工具已存在".to_string())); } inner.tools.insert(name, tool); Ok(()) } /// 批量注册工具。 pub fn register_all(&mut self, tools: Vec) -> Result<(), ToolError> { for tool in tools { self.register(tool)?; } Ok(()) } /// 注销一个工具。 pub fn unregister(&mut self, name: &str) -> Option { let inner = Arc::make_mut(&mut self.inner); inner.tools.remove(name) } /// 按名称查找工具。 pub fn get(&self, name: &str) -> Option { self.inner.tools.get(name).cloned() } /// 获取所有已注册工具的名称列表。 pub fn list_tools(&self) -> Vec { self.inner.tools.keys().cloned().collect() } /// 获取所有工具的 `ToolDefinition` 列表(用于传递给 LLM)。 pub fn definitions(&self) -> Vec { self.inner .tools .values() .map(|tool| ToolDefinition { name: tool.name().to_string(), description: Some(tool.description().to_string()), parameters: tool.parameters(), }) .collect() } /// 调用单个工具(含权限检查)。 /// /// `tool_call_id` 来源于 LLM 流式响应中的 `tool_calls[i].id`,用于回传 /// 工具结果时与原始 `tool_use` block 关联。 pub async fn invoke( &self, tool_call_id: &str, name: &str, args: Value, ) -> Result { let tool = self .get(name) .ok_or_else(|| ToolError::NotFound(name.to_string()))?; if let Some(checker) = &self.inner.permission_checker { checker.check(name, &tool.required_permissions())?; } let ctx = ToolContext::new(name, ""); let output = tool.execute(args.clone(), &ctx).await; Ok(ToolInvocation::new( tool_call_id.to_string(), name.to_string(), args, output, )) } /// 并行执行多个工具调用(互不依赖的工具)。 /// /// 每个工具独立超时(`timeout_per_call_secs`,0 表示不超时)。 /// 单个工具超时不会影响其他工具的返回。 /// /// 入参元组为 `(tool_call_id, tool_name, args)` —— `tool_call_id` 来自 LLM 响应。 pub async fn invoke_all( &self, calls: Vec<(String, String, Value)>, timeout_per_call_secs: u64, ) -> Vec { let this = self.clone(); let futures = calls.into_iter().map(|(tool_call_id, name, args)| { let this = this.clone(); async move { match if timeout_per_call_secs == 0 { Ok(this.invoke(&tool_call_id, &name, args.clone()).await) } else { tokio::time::timeout( Duration::from_secs(timeout_per_call_secs), this.invoke(&tool_call_id, &name, args.clone()), ) .await } { Ok(result) => result.unwrap_or_else(|e| { ToolInvocation::new( tool_call_id.clone(), name.clone(), args.clone(), Err(e), ) }), Err(_) => ToolInvocation::new( tool_call_id, name, args, Err(ToolError::McpTimeout("timeout".into())), ), } } }); join_all(futures).await } } #[cfg(test)] mod tests { use super::*; use crate::tools::BaseTool; use async_trait::async_trait; use serde_json::json; struct AddTool { base: i64, } #[async_trait] impl BaseTool for AddTool { fn name(&self) -> &str { "add" } fn description(&self) -> &str { "加法" } fn parameters(&self) -> Value { json!({ "type": "object", "properties": { "n": { "type": "integer" } }, "required": ["n"] }) } async fn execute(&self, args: Value, _ctx: &ToolContext<'_>) -> Result { let n = args["n"].as_i64().unwrap_or(0); Ok(json!({ "result": self.base + n })) } } struct FailTool; #[async_trait] impl BaseTool for FailTool { fn name(&self) -> &str { "fail" } fn description(&self) -> &str { "总会失败" } fn parameters(&self) -> Value { json!({}) } async fn execute(&self, _args: Value, _ctx: &ToolContext<'_>) -> Result { Err(ToolError::ExecutionFailed("fail".into(), "boom".into())) } } struct ShellTool; #[async_trait] impl BaseTool for ShellTool { fn name(&self) -> &str { "shell" } fn description(&self) -> &str { "shell" } fn parameters(&self) -> Value { json!({}) } fn required_permissions(&self) -> Vec { vec![crate::tools::permission::Permission::Shell] } async fn execute(&self, _args: Value, _ctx: &ToolContext<'_>) -> Result { Ok(json!({})) } } #[test] fn test_register_and_get() { let mut reg = ToolRegistry::new(); reg.register(Arc::new(AddTool { base: 10 })).unwrap(); assert!(reg.get("add").is_some()); assert!(reg.get("nonexistent").is_none()); } #[test] fn test_register_duplicate() { let mut reg = ToolRegistry::new(); reg.register(Arc::new(AddTool { base: 0 })).unwrap(); let result = reg.register(Arc::new(AddTool { base: 1 })); assert!(result.is_err()); } #[test] fn test_register_all() { let mut reg = ToolRegistry::new(); let result = reg.register_all(vec![ Arc::new(AddTool { base: 1 }), Arc::new(AddTool { base: 2 }), ]); assert!(result.is_err()); // 重名 add → 失败 } #[test] fn test_unregister() { let mut reg = ToolRegistry::new(); reg.register(Arc::new(AddTool { base: 0 })).unwrap(); let removed = reg.unregister("add"); assert!(removed.is_some()); assert!(reg.get("add").is_none()); } #[test] fn test_list_tools() { let mut reg = ToolRegistry::new(); reg.register(Arc::new(AddTool { base: 0 })).unwrap(); reg.register(Arc::new(FailTool)).unwrap(); let names = reg.list_tools(); assert_eq!(names.len(), 2); assert!(names.contains(&"add".to_string())); assert!(names.contains(&"fail".to_string())); } #[test] fn test_definitions() { let mut reg = ToolRegistry::new(); reg.register(Arc::new(AddTool { base: 0 })).unwrap(); let defs = reg.definitions(); assert_eq!(defs.len(), 1); assert_eq!(defs[0].name, "add"); assert!(defs[0].description.is_some()); } #[tokio::test] async fn test_invoke_success() { let mut reg = ToolRegistry::new(); reg.register(Arc::new(AddTool { base: 100 })).unwrap(); let result = reg .invoke("call_1", "add", json!({ "n": 5 })) .await .unwrap(); let value = result.output.unwrap(); assert_eq!(value["result"], 105); assert_eq!(result.tool_call_id, "call_1"); assert_eq!(result.tool_name, "add"); } #[tokio::test] async fn test_invoke_not_found() { let reg = ToolRegistry::new(); let result = reg.invoke("call_x", "nope", json!({})).await; assert!(matches!(result, Err(ToolError::NotFound(_)))); } #[tokio::test] async fn test_invoke_execution_error() { let mut reg = ToolRegistry::new(); reg.register(Arc::new(FailTool)).unwrap(); let result = reg.invoke("call_y", "fail", json!({})).await.unwrap(); assert!(result.output.is_err()); } #[tokio::test] async fn test_invoke_with_permission_denied() { let mut reg = ToolRegistry::new().with_permission_checker(PermissionChecker::new(Default::default())); reg.register(Arc::new(ShellTool)).unwrap(); let result = reg.invoke("call_z", "shell", json!({})).await; assert!(matches!(result, Err(ToolError::PermissionDenied(_, _)))); } #[tokio::test] async fn test_invoke_all_parallel() { let mut reg = ToolRegistry::new(); reg.register(Arc::new(AddTool { base: 1 })).unwrap(); reg.register(Arc::new(FailTool)).unwrap(); let calls = vec![ ("c1".into(), "add".into(), json!({ "n": 1 })), ("c2".into(), "add".into(), json!({ "n": 2 })), ("c3".into(), "fail".into(), json!({})), ]; let results = reg.invoke_all(calls, 0).await; assert_eq!(results.len(), 3); assert_eq!(results[0].tool_call_id, "c1"); assert_eq!(results[1].tool_call_id, "c2"); assert_eq!(results[2].tool_call_id, "c3"); assert!(results[0].output.is_ok()); assert!(results[1].output.is_ok()); assert!(results[2].output.is_err()); } #[tokio::test] async fn test_invoke_all_with_timeout() { let mut reg = ToolRegistry::new(); reg.register(Arc::new(AddTool { base: 0 })).unwrap(); let calls = vec![("c1".into(), "add".into(), json!({ "n": 1 }))]; let results = reg.invoke_all(calls, 5).await; assert_eq!(results.len(), 1); assert_eq!(results[0].tool_call_id, "c1"); assert!(results[0].output.is_ok()); } }