feat(tools): 添加工具系统框架与 MCP 协议客户端
This commit is contained in:
@@ -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());
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user