Files
agcore/src/tools/registry.rs
T
徐涛 9da9b83167 feat(llm): 切换至 ToolDef IR 并适配 OpenAI Provider
将 MessageRequest.tools 切换为 Provider 无关的 ToolDef IR;
ToolDefinition 别名切换指向 ToolDef,registry/mcp 构造去掉
strict 字段;OpenAI 适配层通过 From 转换将 ToolDef 转为
OpenaiToolDefinition wire format。Anthropic 适配层字段名一致
无需改动,openai_compat/ollama 委托 GenericOpenaiProvider 无影响。
2026-07-05 10:14:32 +08:00

415 lines
13 KiB
Rust
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
//! 工具注册表 —— 管理工具注册、发现、调用。
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<Value, ToolError>,
}
impl ToolInvocation {
/// 创建一个新的工具调用记录。
pub fn new(
tool_call_id: String,
tool_name: String,
input: Value,
output: Result<Value, ToolError>,
) -> Self {
Self {
tool_call_id,
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()
}
}
#[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<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(),
})
.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<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(
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<ToolInvocation> {
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<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("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());
}
}