263 lines
6.3 KiB
Go
263 lines
6.3 KiB
Go
// crypt_manager.go - 修复 ExecuteTool 返回结果格式
|
||
package crypt
|
||
|
||
import (
|
||
"context"
|
||
"encoding/json"
|
||
"fmt"
|
||
"sort"
|
||
"strings"
|
||
"sync"
|
||
|
||
"github.com/sashabaranov/go-openai"
|
||
)
|
||
|
||
// ToolResult 工具执行结果
|
||
type ToolResult struct {
|
||
Success bool `json:"success"`
|
||
Data string `json:"data"`
|
||
Error string `json:"error,omitempty"`
|
||
Tool string `json:"tool"`
|
||
}
|
||
|
||
// CryptoSkillManager 加密技能管理器 - 支持渐进式披露和缓存
|
||
type CryptoSkillManager struct {
|
||
registry *ToolRegistry
|
||
cache map[string]string // cacheKey -> result
|
||
cacheMu sync.RWMutex
|
||
// 记录已调用的工具,用于调试
|
||
calledTools map[string]int
|
||
calledMu sync.Mutex
|
||
}
|
||
|
||
func NewCryptoSkillManager() *CryptoSkillManager {
|
||
return &CryptoSkillManager{
|
||
registry: NewToolRegistry(),
|
||
cache: make(map[string]string),
|
||
calledTools: make(map[string]int),
|
||
}
|
||
}
|
||
|
||
// GetToolsForLLM 获取所有加密工具供 LLM 调用(AI 自主决策)
|
||
func (m *CryptoSkillManager) GetToolsForLLM() []openai.Tool {
|
||
tools := m.registry.GetToolManifest()
|
||
|
||
// 为每个工具添加分组信息到描述中
|
||
for i := range tools {
|
||
tool := &tools[i]
|
||
group := m.getToolGroup(tool.Function.Name)
|
||
if group != "" {
|
||
tool.Function.Description = fmt.Sprintf("[%s] %s", group, tool.Function.Description)
|
||
}
|
||
}
|
||
|
||
return tools
|
||
}
|
||
|
||
// GetToolDescriptions 获取工具描述(注入到 System Prompt)
|
||
func (m *CryptoSkillManager) GetToolDescriptions() string {
|
||
return m.registry.GetToolDescriptions()
|
||
}
|
||
|
||
// ExecuteTool 执行工具调用(由 LLM 发起的 Tool Call)- 带缓存
|
||
func (m *CryptoSkillManager) ExecuteTool(ctx context.Context, toolCall openai.ToolCall) (string, error) {
|
||
toolName := toolCall.Function.Name
|
||
|
||
// 记录调用次数
|
||
m.calledMu.Lock()
|
||
m.calledTools[toolName] = m.calledTools[toolName] + 1
|
||
callCount := m.calledTools[toolName]
|
||
m.calledMu.Unlock()
|
||
|
||
// 解析参数 - 支持多种类型
|
||
params, err := m.parseArguments(toolCall.Function.Arguments)
|
||
if err != nil {
|
||
return "", fmt.Errorf("解析参数失败: %v", err)
|
||
}
|
||
|
||
// 生成缓存 key
|
||
cacheKey := m.generateCacheKey(toolName, params)
|
||
|
||
// 检查缓存
|
||
m.cacheMu.RLock()
|
||
if cached, ok := m.cache[cacheKey]; ok {
|
||
m.cacheMu.RUnlock()
|
||
// 如果是重复调用,在返回结果中明确告知 AI
|
||
if callCount > 1 {
|
||
return fmt.Sprintf(`[已缓存] 工具 %s 已经被调用过 %d 次。
|
||
上次返回的结果如下(请直接使用,无需重复调用):
|
||
|
||
%s`, toolName, callCount, cached), nil
|
||
}
|
||
return cached, nil
|
||
}
|
||
m.cacheMu.RUnlock()
|
||
|
||
// 执行工具
|
||
result, err := m.registry.ExecuteTool(ctx, toolName, params)
|
||
|
||
// 包装结果
|
||
response := ToolResult{
|
||
Success: err == nil,
|
||
Tool: toolName,
|
||
Data: result,
|
||
}
|
||
|
||
if err != nil {
|
||
response.Error = err.Error()
|
||
}
|
||
|
||
jsonData, _ := json.MarshalIndent(response, "", " ")
|
||
resultStr := string(jsonData)
|
||
|
||
// 缓存结果(只有成功时才缓存)
|
||
if err == nil {
|
||
m.cacheMu.Lock()
|
||
m.cache[cacheKey] = resultStr
|
||
m.cacheMu.Unlock()
|
||
}
|
||
|
||
return resultStr, err
|
||
}
|
||
|
||
// parseArguments 解析参数,支持多种类型(string, bool, number 等)
|
||
func (m *CryptoSkillManager) parseArguments(arguments string) (map[string]string, error) {
|
||
if arguments == "" {
|
||
return map[string]string{}, nil
|
||
}
|
||
|
||
// 首先尝试解析为 map[string]interface{}
|
||
var raw map[string]interface{}
|
||
if err := json.Unmarshal([]byte(arguments), &raw); err != nil {
|
||
return nil, err
|
||
}
|
||
|
||
// 转换为 map[string]string,处理各种类型
|
||
result := make(map[string]string)
|
||
for key, value := range raw {
|
||
switch v := value.(type) {
|
||
case string:
|
||
result[key] = v
|
||
case bool:
|
||
if v {
|
||
result[key] = "true"
|
||
} else {
|
||
result[key] = "false"
|
||
}
|
||
case float64:
|
||
result[key] = fmt.Sprintf("%.0f", v)
|
||
case int:
|
||
result[key] = fmt.Sprintf("%d", v)
|
||
case int64:
|
||
result[key] = fmt.Sprintf("%d", v)
|
||
case nil:
|
||
// 忽略 nil
|
||
default:
|
||
if b, err := json.Marshal(v); err == nil {
|
||
result[key] = string(b)
|
||
}
|
||
}
|
||
}
|
||
|
||
return result, nil
|
||
}
|
||
|
||
// ExecuteToolByName 直接按名称执行工具(用于测试或直接调用)
|
||
func (m *CryptoSkillManager) ExecuteToolByName(ctx context.Context, toolName string, params map[string]string) (string, error) {
|
||
cacheKey := m.generateCacheKey(toolName, params)
|
||
|
||
// 检查缓存
|
||
m.cacheMu.RLock()
|
||
if cached, ok := m.cache[cacheKey]; ok {
|
||
m.cacheMu.RUnlock()
|
||
return cached, nil
|
||
}
|
||
m.cacheMu.RUnlock()
|
||
|
||
result, err := m.registry.ExecuteTool(ctx, toolName, params)
|
||
if err != nil {
|
||
return "", err
|
||
}
|
||
|
||
// 缓存结果
|
||
m.cacheMu.Lock()
|
||
m.cache[cacheKey] = result
|
||
m.cacheMu.Unlock()
|
||
|
||
return result, nil
|
||
}
|
||
|
||
// generateCacheKey 生成缓存 key
|
||
func (m *CryptoSkillManager) generateCacheKey(toolName string, params map[string]string) string {
|
||
keys := make([]string, 0, len(params))
|
||
for k := range params {
|
||
keys = append(keys, k)
|
||
}
|
||
sort.Strings(keys)
|
||
|
||
var sb strings.Builder
|
||
sb.WriteString(toolName)
|
||
for _, k := range keys {
|
||
sb.WriteString("|")
|
||
sb.WriteString(k)
|
||
sb.WriteString("=")
|
||
sb.WriteString(params[k])
|
||
}
|
||
return sb.String()
|
||
}
|
||
|
||
// getToolGroup 获取工具分组
|
||
func (m *CryptoSkillManager) getToolGroup(toolName string) string {
|
||
switch {
|
||
case strings.HasPrefix(toolName, "aes_"):
|
||
return "对称加密"
|
||
case strings.HasPrefix(toolName, "rsa_"):
|
||
return "非对称加密"
|
||
case strings.HasPrefix(toolName, "sm2_"):
|
||
return "国密"
|
||
case strings.HasPrefix(toolName, "sm3_"):
|
||
return "国密哈希"
|
||
case strings.HasPrefix(toolName, "sm4_"):
|
||
return "国密对称加密"
|
||
case strings.HasPrefix(toolName, "param_"):
|
||
return "辅助工具"
|
||
case strings.HasPrefix(toolName, "nonce_"):
|
||
return "辅助工具"
|
||
default:
|
||
return ""
|
||
}
|
||
}
|
||
|
||
// ClearCache 清空缓存
|
||
func (m *CryptoSkillManager) ClearCache() {
|
||
m.cacheMu.Lock()
|
||
defer m.cacheMu.Unlock()
|
||
m.cache = make(map[string]string)
|
||
m.calledMu.Lock()
|
||
defer m.calledMu.Unlock()
|
||
m.calledTools = make(map[string]int)
|
||
}
|
||
|
||
// GetCacheStats 获取缓存统计
|
||
func (m *CryptoSkillManager) GetCacheStats() map[string]int {
|
||
m.cacheMu.RLock()
|
||
defer m.cacheMu.RUnlock()
|
||
m.calledMu.Lock()
|
||
defer m.calledMu.Unlock()
|
||
return map[string]int{
|
||
"cache_size": len(m.cache),
|
||
"called_tools": len(m.calledTools),
|
||
}
|
||
}
|
||
|
||
// GetCalledTools 获取工具调用统计
|
||
func (m *CryptoSkillManager) GetCalledTools() map[string]int {
|
||
m.calledMu.Lock()
|
||
defer m.calledMu.Unlock()
|
||
result := make(map[string]int)
|
||
for k, v := range m.calledTools {
|
||
result[k] = v
|
||
}
|
||
return result
|
||
}
|