Files
DevFlow/crates/df-ai/src/ai_tools.rs
T

240 lines
7.2 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.
//! 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");
}
}