//! AI 工具注册 — 将现有 CRUD 操作注册为 LLM 可调用的 Tool //! //! 风险分级: //! - Low: 只读操作(list/get),AI 自动执行 //! - Medium: 创建操作(create),显示意图,可配置自动批准 //! - High: 破坏性操作(delete / run_workflow),必须人工批准 use std::collections::HashMap; use std::future::Future; use std::pin::Pin; use serde::{Deserialize, Serialize}; use serde_json::Value; use crate::provider::ToolDefinition; // ============================================================ // 风险级别 // ============================================================ /// 工具调用风险级别 #[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq, Eq)] #[serde(rename_all = "lowercase")] pub enum RiskLevel { /// 只读操作 — AI 自动执行 Low, /// 创建操作 — 可配置自动批准 Medium, /// 破坏性操作 — 必须人工批准 High, } // ============================================================ // 工具定义 // ============================================================ /// 已注册的 AI 工具 pub struct AiTool { /// 工具定义(发送给 LLM 的 JSON Schema) pub definition: ToolDefinition, /// 风险级别 pub risk_level: RiskLevel, /// 执行处理器 pub handler: AiToolHandler, } /// 工具处理器类型 — 异步函数,接收 JSON 参数,返回 JSON 结果 pub type AiToolHandler = Box Pin> + Send>> + Send + Sync>; // ============================================================ // 工具注册表 // ============================================================ /// AI 工具注册表 pub struct AiToolRegistry { tools: HashMap, } impl AiToolRegistry { /// 创建空注册表 pub fn new() -> Self { Self { tools: HashMap::new(), } } /// 注册一个工具 pub fn register( &mut self, name: impl Into, description: impl Into, parameters: Value, risk_level: RiskLevel, handler: AiToolHandler, ) { let name_str = name.into(); let definition = ToolDefinition::function(&name_str, description, parameters); self.tools.insert( name_str, AiTool { definition, risk_level, handler, }, ); } /// 获取所有工具定义(发送给 LLM 的 tools 参数) pub fn tool_definitions(&self) -> Vec { self.tools.values().map(|t| t.definition.clone()).collect() } /// 根据名称获取工具 pub fn get(&self, name: &str) -> Option<&AiTool> { self.tools.get(name) } /// 执行指定工具 — handler 是唯一执行路径(schema+risk+实现同源,消除双轨) pub async fn execute( &self, name: &str, args: serde_json::Value, ) -> anyhow::Result { let tool = self .tools .get(name) .ok_or_else(|| anyhow::anyhow!("未知工具: {}", name))?; (tool.handler)(args).await } /// 获取所有已注册工具名称 pub fn tool_names(&self) -> Vec { self.tools.keys().cloned().collect() } /// 已注册工具数量 pub fn len(&self) -> usize { self.tools.len() } /// 是否为空 pub fn is_empty(&self) -> bool { self.tools.is_empty() } } impl Default for AiToolRegistry { fn default() -> Self { Self::new() } } // ============================================================ // 辅助: 构建 JSON Schema 参数 // ============================================================ /// 构建一个简单的 object JSON Schema pub fn object_schema(properties: Vec<(&str, &str, bool)>) -> Value { let mut props = serde_json::Map::new(); let mut required = Vec::new(); for (name, type_str, is_required) in properties { props.insert( name.to_string(), serde_json::json!({ "type": type_str, "description": "" }), ); if is_required { required.push(name.to_string()); } } serde_json::json!({ "type": "object", "properties": props, "required": required, }) } #[cfg(test)] mod tests { use super::*; /// 构造一个空参恒返 Ok 的 handler(注册测试用·不实际执行) fn dummy_handler() -> AiToolHandler { Box::new(|_args: Value| { Box::pin(async { Ok(serde_json::json!({ "ok": true })) }) as Pin> + Send>> }) } #[test] fn risk_level_serde_lowercase() { // serde rename_all="lowercase": Low→"low" / Medium→"medium" / High→"high" assert_eq!(serde_json::to_string(&RiskLevel::Low).unwrap(), "\"low\""); assert_eq!(serde_json::to_string(&RiskLevel::Medium).unwrap(), "\"medium\""); assert_eq!(serde_json::to_string(&RiskLevel::High).unwrap(), "\"high\""); assert_eq!( serde_json::from_str::("\"high\"").unwrap(), RiskLevel::High ); } #[test] fn register_and_get_tool() { let mut reg = AiToolRegistry::new(); assert!(reg.is_empty()); reg.register( "list_tasks", "列出任务", object_schema(vec![]), RiskLevel::Low, dummy_handler(), ); assert!(!reg.is_empty()); assert_eq!(reg.len(), 1); assert_eq!(reg.get("list_tasks").unwrap().risk_level, RiskLevel::Low); assert!(reg.get("nonexistent").is_none()); } #[test] fn register_same_name_overwrites() { // 同名二次注册覆盖(HashMap insert 语义)·非新增·后注册的 risk_level 胜出 let mut reg = AiToolRegistry::new(); reg.register("t", "v1", object_schema(vec![]), RiskLevel::Low, dummy_handler()); reg.register("t", "v2", object_schema(vec![]), RiskLevel::High, dummy_handler()); assert_eq!(reg.len(), 1); assert_eq!(reg.get("t").unwrap().risk_level, RiskLevel::High); } #[test] fn tool_definitions_and_names() { let mut reg = AiToolRegistry::new(); reg.register("a", "desc a", object_schema(vec![]), RiskLevel::Low, dummy_handler()); reg.register("b", "desc b", object_schema(vec![]), RiskLevel::Medium, dummy_handler()); assert_eq!(reg.tool_definitions().len(), 2); let names = reg.tool_names(); assert_eq!(names.len(), 2); assert!(names.contains(&"a".to_string())); assert!(names.contains(&"b".to_string())); } #[test] fn default_is_empty() { let reg = AiToolRegistry::default(); assert!(reg.is_empty()); assert_eq!(reg.len(), 0); } #[test] fn object_schema_collects_required() { let schema = object_schema(vec![("name", "string", true), ("age", "integer", false)]); assert_eq!(schema["type"], "object"); assert_eq!(schema["properties"]["name"]["type"], "string"); let required = schema["required"].as_array().unwrap(); assert_eq!(required.len(), 1); assert_eq!(required[0], "name"); } }