sdk_generate/internal/prompts/tools/crypt/crypt_manager.go

263 lines
6.3 KiB
Go
Raw Permalink 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.

// 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
}