240 lines
7.2 KiB
Rust
240 lines
7.2 KiB
Rust
//! 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<dyn Fn(Value) -> Pin<Box<dyn Future<Output = anyhow::Result<Value>> + Send>> + Send + Sync>;
|
||
|
||
// ============================================================
|
||
// 工具注册表
|
||
// ============================================================
|
||
|
||
/// AI 工具注册表
|
||
pub struct AiToolRegistry {
|
||
tools: HashMap<String, AiTool>,
|
||
}
|
||
|
||
impl AiToolRegistry {
|
||
/// 创建空注册表
|
||
pub fn new() -> Self {
|
||
Self {
|
||
tools: HashMap::new(),
|
||
}
|
||
}
|
||
|
||
/// 注册一个工具
|
||
pub fn register(
|
||
&mut self,
|
||
name: impl Into<String>,
|
||
description: impl Into<String>,
|
||
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<ToolDefinition> {
|
||
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<serde_json::Value> {
|
||
let tool = self
|
||
.tools
|
||
.get(name)
|
||
.ok_or_else(|| anyhow::anyhow!("未知工具: {}", name))?;
|
||
(tool.handler)(args).await
|
||
}
|
||
|
||
/// 获取所有已注册工具名称
|
||
pub fn tool_names(&self) -> Vec<String> {
|
||
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<Box<dyn Future<Output = anyhow::Result<Value>> + 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::<RiskLevel>("\"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");
|
||
}
|
||
}
|