feat(tools): 添加工具系统框架与 MCP 协议客户端

This commit is contained in:
徐涛
2026-06-07 10:57:15 +08:00
parent e598f6d3ee
commit b6e7acfb0f
9 changed files with 2034 additions and 1 deletions
+371
View File
@@ -0,0 +1,371 @@
//! 工具注册表 —— 管理工具注册、发现、调用。
use std::collections::HashMap;
use std::sync::Arc;
use std::time::Duration;
use futures::future::join_all;
use serde_json::Value;
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 {
/// 被调用的工具名。
pub tool_name: String,
/// 工具的入参。
pub input: Value,
/// 工具的输出。
pub output: Result<Value, ToolError>,
}
impl ToolInvocation {
/// 创建一个新的工具调用记录。
pub fn new(tool_name: String, input: Value, output: Result<Value, ToolError>) -> Self {
Self {
tool_name,
input,
output,
}
}
}
/// 工具注册表 —— 管理工具注册、发现、调用。
///
/// 通过 `Arc` 共享,方法签名 `&self`,可安全跨 task 并行调用。
/// 不支持运行时并发注册(应在 setup 阶段一次性构建后冻结)。
#[derive(Clone, Default)]
pub struct ToolRegistry {
inner: Arc<ToolRegistryInner>,
}
#[derive(Clone, Default)]
struct ToolRegistryInner {
tools: HashMap<String, ToolRef>,
permission_checker: Option<Arc<PermissionChecker>>,
}
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::<Vec<_>>())
.field("has_checker", &self.inner.permission_checker.is_some())
.finish()
}
}
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<ToolRef>) -> Result<(), ToolError> {
for tool in tools {
self.register(tool)?;
}
Ok(())
}
/// 注销一个工具。
pub fn unregister(&mut self, name: &str) -> Option<ToolRef> {
let inner = Arc::make_mut(&mut self.inner);
inner.tools.remove(name)
}
/// 按名称查找工具。
pub fn get(&self, name: &str) -> Option<ToolRef> {
self.inner.tools.get(name).cloned()
}
/// 获取所有已注册工具的名称列表。
pub fn list_tools(&self) -> Vec<String> {
self.inner.tools.keys().cloned().collect()
}
/// 获取所有工具的 `ToolDefinition` 列表(用于传递给 LLM)。
pub fn definitions(&self) -> Vec<ToolDefinition> {
self.inner
.tools
.values()
.map(|tool| ToolDefinition {
name: tool.name().to_string(),
description: Some(tool.description().to_string()),
parameters: tool.parameters(),
strict: None,
})
.collect()
}
/// 调用单个工具(含权限检查)。
pub async fn invoke(&self, name: &str, args: Value) -> Result<ToolInvocation, ToolError> {
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(name.to_string(), args, output))
}
/// 并行执行多个工具调用(互不依赖的工具)。
///
/// 每个工具独立超时(`timeout_per_call_secs`0 表示不超时)。
/// 单个工具超时不会影响其他工具的返回。
pub async fn invoke_all(
&self,
calls: Vec<(String, Value)>,
timeout_per_call_secs: u64,
) -> Vec<ToolInvocation> {
let this = self.clone();
let futures = calls.into_iter().map(|(name, args)| {
let this = this.clone();
async move {
match if timeout_per_call_secs == 0 {
Ok(this.invoke(&name, args.clone()).await)
} else {
tokio::time::timeout(
Duration::from_secs(timeout_per_call_secs),
this.invoke(&name, args.clone()),
)
.await
} {
Ok(result) => result.unwrap_or_else(|e| {
ToolInvocation::new(name.clone(), args.clone(), Err(e))
}),
Err(_) => ToolInvocation::new(
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<Value, ToolError> {
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<Value, ToolError> {
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<crate::tools::permission::Permission> {
vec![crate::tools::permission::Permission::Shell]
}
async fn execute(&self, _args: Value, _ctx: &ToolContext<'_>) -> Result<Value, ToolError> {
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("add", json!({ "n": 5 })).await.unwrap();
let value = result.output.unwrap();
assert_eq!(value["result"], 105);
}
#[tokio::test]
async fn test_invoke_not_found() {
let reg = ToolRegistry::new();
let result = reg.invoke("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("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("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![
("add".into(), json!({ "n": 1 })),
("add".into(), json!({ "n": 2 })),
("fail".into(), json!({})),
];
let results = reg.invoke_all(calls, 0).await;
assert_eq!(results.len(), 3);
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![("add".into(), json!({ "n": 1 }))];
let results = reg.invoke_all(calls, 5).await;
assert_eq!(results.len(), 1);
assert!(results[0].output.is_ok());
}
}