From 4e4c6b32fc7ea99d8cf9b1cceb521def7a739eb7 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E6=97=A0=E9=97=BB=E9=A3=8E?= Date: Wed, 17 Jun 2026 13:13:06 +0800 Subject: [PATCH] =?UTF-8?q?=E5=8A=A8=E6=80=81=E5=8A=A0=E8=BD=BD=E5=B7=A5?= =?UTF-8?q?=E5=85=B7=E9=85=8D=E7=BD=AE?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-Authored-By: Claude --- agents/calculator/calculator.go | 90 +++++++++++++++++ agents/search/search.go | 80 +++++++++++++++ agents/sql/sql_query.go | 66 +++++++++++++ agents/time/time.go | 94 ++++++++++++++++++ agenttool/registry.go | 100 +++++++++++++++++++ config/config.go | 59 ++--------- main.go | 27 ++--- main_test.go | 36 ++++--- server/handlers.go | 29 +++++- server/server.go | 11 +-- toolmanager/manager.go | 125 ++++++++++++++++++++++++ toolrouter/loop.go | 7 +- toolrouter/tools.go | 168 +++++++++++++++----------------- 13 files changed, 704 insertions(+), 188 deletions(-) create mode 100644 agenttool/registry.go create mode 100644 toolmanager/manager.go diff --git a/agents/calculator/calculator.go b/agents/calculator/calculator.go index 125564b..564ada8 100644 --- a/agents/calculator/calculator.go +++ b/agents/calculator/calculator.go @@ -1,15 +1,21 @@ package calculator import ( + "context" "encoding/json" "errors" "fmt" "math" + "os" + "path/filepath" "strconv" "strings" "unicode" + agents "aichat/agenttool" + "github.com/volcengine/volcengine-go-sdk/service/arkruntime/model" + "gopkg.in/yaml.v3" ) const ( @@ -17,11 +23,95 @@ const ( ActivationPrompt = "执行简单、确定性的数学四则运算。当用户询问加减乘除、括号表达式、小数运算或需要准确计算表达式结果时,应直接调用此工具;不用于代数推导、方程求解、统计分析或复杂数学证明。" ) +type Config struct { + Enabled bool `yaml:"enabled" json:"enabled"` + ActivationPrompt string `yaml:"activation_prompt" json:"activation_prompt"` +} + type ToolArgs struct { Expression string `json:"expression"` Reason string `json:"reason"` } +type LoadedTool struct { + cfg *Config +} + +func NewLoadedTool(cfg *Config) *LoadedTool { + if cfg == nil { + defaultCfg := defaultConfig() + cfg = &defaultCfg + } + return &LoadedTool{cfg: cfg} +} + +func init() { + agents.Register(agents.Descriptor{Name: ToolName, Load: func(path string, options agents.LoadOptions) (agents.LoadedTool, error) { + cfg, err := LoadConfig(path) + if err != nil { + return nil, err + } + return NewLoadedTool(cfg), nil + }}) +} + +func defaultConfig() Config { + return Config{Enabled: true, ActivationPrompt: ActivationPrompt} +} + +func LoadConfig(path string) (*Config, error) { + if _, err := os.Stat(path); err != nil { + if !os.IsNotExist(err) { + return nil, fmt.Errorf("检查计算器工具配置失败: %w", err) + } + cfg := defaultConfig() + if err := os.MkdirAll(filepath.Dir(path), 0755); err != nil { + return nil, fmt.Errorf("创建计算器工具配置目录失败: %w", err) + } + data, err := yaml.Marshal(&cfg) + if err != nil { + return nil, fmt.Errorf("生成计算器工具配置失败: %w", err) + } + if err := os.WriteFile(path, data, 0644); err != nil { + return nil, fmt.Errorf("写入计算器工具配置失败: %w", err) + } + } + + data, err := os.ReadFile(path) + if err != nil { + return nil, fmt.Errorf("读取计算器工具配置失败: %w", err) + } + var cfg Config + if err := yaml.Unmarshal(data, &cfg); err != nil { + return nil, fmt.Errorf("解析计算器工具配置失败: %w", err) + } + if strings.TrimSpace(cfg.ActivationPrompt) == "" { + cfg.ActivationPrompt = ActivationPrompt + } + return &cfg, nil +} + +func (t *LoadedTool) Name() string { return ToolName } + +func (t *LoadedTool) Enabled() bool { return t != nil && t.cfg != nil && t.cfg.Enabled } + +func (t *LoadedTool) ToolDefinition(description string) *model.Tool { + if strings.TrimSpace(description) == "" && t != nil && t.cfg != nil { + description = t.cfg.ActivationPrompt + } + return ToolDefinition(description) +} + +func (t *LoadedTool) Execute(ctx context.Context, args string, runtime agents.Runtime) (string, error) { + result, err := ExecuteTool(args) + if err == nil && runtime.Emit != nil { + runtime.Emit(agents.Frame{Type: "trace", Tool: ToolName, Stage: "calculate", Status: "success", Message: "四则运算完成"}) + } + return result, err +} + +func (t *LoadedTool) RawState() any { return nil } + func ToolDefinition(description string) *model.Tool { description = strings.TrimSpace(description) if description == "" { diff --git a/agents/search/search.go b/agents/search/search.go index c1334ad..124e9f2 100644 --- a/agents/search/search.go +++ b/agents/search/search.go @@ -14,6 +14,8 @@ import ( "sync" "time" + agents "aichat/agenttool" + "github.com/volcengine/volcengine-go-sdk/service/arkruntime/model" "gopkg.in/yaml.v3" ) @@ -82,6 +84,10 @@ type State struct { activeName string } +type LoadedTool struct { + state *State +} + type Result struct { Title string `json:"title"` URL string `json:"url"` @@ -98,6 +104,80 @@ type ToolArgs struct { Reason string `json:"reason"` } +func init() { + agents.Register(agents.Descriptor{Name: ToolName, Load: func(path string, options agents.LoadOptions) (agents.LoadedTool, error) { + legacy, err := legacyProfilesFromOption(options.Value(agents.LegacySearchProfilesKey)) + if err != nil { + return nil, err + } + cfg, err := LoadConfig(path, legacy) + if err != nil { + return nil, err + } + state, err := NewState(cfg) + if err != nil { + return nil, err + } + return &LoadedTool{state: state}, nil + }}) +} + +func legacyProfilesFromOption(value any) ([]ProfileConfig, error) { + switch profiles := value.(type) { + case nil: + return nil, nil + case []ProfileConfig: + return profiles, nil + case []map[string]any: + data, err := yaml.Marshal(profiles) + if err != nil { + return nil, fmt.Errorf("转换旧搜索配置失败: %w", err) + } + var parsed []ProfileConfig + if err := yaml.Unmarshal(data, &parsed); err != nil { + return nil, fmt.Errorf("解析旧搜索配置失败: %w", err) + } + return parsed, nil + default: + return nil, nil + } +} + +func (t *LoadedTool) Name() string { return ToolName } + +func (t *LoadedTool) Enabled() bool { return t != nil && t.state != nil && t.state.Enabled() } + +func (t *LoadedTool) ToolDefinition(description string) *model.Tool { + if t == nil || t.state == nil { + return nil + } + return t.state.ToolDefinition(description) +} + +func (t *LoadedTool) Execute(ctx context.Context, args string, runtime agents.Runtime) (string, error) { + if runtime.Emit != nil { + runtime.Emit(agents.Frame{Type: "trace", Tool: ToolName, Stage: "request", Status: "running", Message: "正在联网搜索"}) + } + result, err := t.state.ExecuteTool(ctx, args) + if runtime.Emit != nil { + status := "success" + messageText := "联网搜索完成" + if err != nil { + status = "error" + messageText = "联网搜索失败" + } + runtime.Emit(agents.Frame{Type: "trace", Tool: ToolName, Stage: "results", Status: status, Message: messageText}) + } + return result, err +} + +func (t *LoadedTool) RawState() any { + if t == nil { + return nil + } + return t.state +} + type braveSearchResponse struct { Web struct { Results []Result `json:"results"` diff --git a/agents/sql/sql_query.go b/agents/sql/sql_query.go index d8b5c70..f89fbba 100644 --- a/agents/sql/sql_query.go +++ b/agents/sql/sql_query.go @@ -15,6 +15,8 @@ import ( "time" "unicode/utf8" + agents "aichat/agenttool" + _ "github.com/go-sql-driver/mysql" "github.com/volcengine/volcengine-go-sdk/service/arkruntime/model" "gopkg.in/yaml.v3" @@ -69,6 +71,10 @@ type State struct { cacheAt time.Time } +type LoadedTool struct { + state *State +} + type database struct { cfg DatabaseConfig db *sql.DB @@ -90,6 +96,66 @@ type ToolArgs struct { type SQLGenerator func(ctx context.Context, prompt string, maxTokens int) (string, error) +func init() { + agents.Register(agents.Descriptor{Name: ToolName, Load: func(path string, options agents.LoadOptions) (agents.LoadedTool, error) { + cfg, err := LoadConfig(path) + if err != nil { + return nil, err + } + state, err := NewState(cfg) + if err != nil { + return nil, err + } + return &LoadedTool{state: state}, nil + }}) +} + +func (t *LoadedTool) Name() string { return ToolName } + +func (t *LoadedTool) Enabled() bool { return t != nil && t.state != nil && t.state.Enabled() } + +func (t *LoadedTool) ToolDefinition(description string) *model.Tool { + if t == nil || t.state == nil { + return nil + } + return t.state.ToolDefinition(description) +} + +func (t *LoadedTool) Execute(ctx context.Context, args string, runtime agents.Runtime) (string, error) { + if runtime.Emit != nil { + runtime.Emit(agents.Frame{Type: "trace", Tool: ToolName, Stage: "execute", Status: "running", Message: "正在查询数据库"}) + } + generator := runtime.CompleteText + if generator == nil { + return "", errors.New("SQL 查询工具缺少文本生成器") + } + result, err := t.state.ExecuteTool(ctx, args, generator) + if runtime.Emit != nil { + status := "success" + messageText := "数据库查询完成" + if err != nil { + status = "error" + messageText = "数据库查询失败" + } + runtime.Emit(agents.Frame{Type: "trace", Tool: ToolName, Stage: "execute", Status: status, Message: messageText}) + } + return result, err +} + +func (t *LoadedTool) RawState() any { + if t == nil { + return nil + } + return t.state +} + +func (t *LoadedTool) Close() error { + if t == nil || t.state == nil { + return nil + } + return t.state.Close() +} + type GenerationResult struct { Database string `json:"database"` SQL string `json:"sql"` diff --git a/agents/time/time.go b/agents/time/time.go index 2f6b716..115d528 100644 --- a/agents/time/time.go +++ b/agents/time/time.go @@ -1,12 +1,18 @@ package timeagent import ( + "context" "encoding/json" "fmt" + "os" + "path/filepath" "strings" "time" + agents "aichat/agenttool" + "github.com/volcengine/volcengine-go-sdk/service/arkruntime/model" + "gopkg.in/yaml.v3" ) const ( @@ -14,10 +20,98 @@ const ( ActivationPrompt = "提供当前日期、时间和常用时间范围。当用户问题包含今天、今日、明天、昨天、本周、本月、本年、最近、历史上的今天、日程安排等相对时间表达时,应先调用此工具;如果后续还需要联网搜索或查数据库,可继续调用 search 或 sql。" ) +type Config struct { + Enabled bool `yaml:"enabled" json:"enabled"` + ActivationPrompt string `yaml:"activation_prompt" json:"activation_prompt"` +} + type ToolArgs struct { Reason string `json:"reason"` } +type LoadedTool struct { + cfg *Config +} + +func NewLoadedTool(cfg *Config) *LoadedTool { + if cfg == nil { + defaultCfg := defaultConfig() + cfg = &defaultCfg + } + return &LoadedTool{cfg: cfg} +} + +func init() { + agents.Register(agents.Descriptor{Name: ToolName, Load: func(path string, options agents.LoadOptions) (agents.LoadedTool, error) { + cfg, err := LoadConfig(path) + if err != nil { + return nil, err + } + return NewLoadedTool(cfg), nil + }}) +} + +func defaultConfig() Config { + return Config{Enabled: true, ActivationPrompt: ActivationPrompt} +} + +func LoadConfig(path string) (*Config, error) { + if _, err := os.Stat(path); err != nil { + if !os.IsNotExist(err) { + return nil, fmt.Errorf("检查时间工具配置失败: %w", err) + } + cfg := defaultConfig() + if err := os.MkdirAll(filepath.Dir(path), 0755); err != nil { + return nil, fmt.Errorf("创建时间工具配置目录失败: %w", err) + } + data, err := yaml.Marshal(&cfg) + if err != nil { + return nil, fmt.Errorf("生成时间工具配置失败: %w", err) + } + if err := os.WriteFile(path, data, 0644); err != nil { + return nil, fmt.Errorf("写入时间工具配置失败: %w", err) + } + } + + data, err := os.ReadFile(path) + if err != nil { + return nil, fmt.Errorf("读取时间工具配置失败: %w", err) + } + var cfg Config + if err := yaml.Unmarshal(data, &cfg); err != nil { + return nil, fmt.Errorf("解析时间工具配置失败: %w", err) + } + if strings.TrimSpace(cfg.ActivationPrompt) == "" { + cfg.ActivationPrompt = ActivationPrompt + } + return &cfg, nil +} + +func (t *LoadedTool) Name() string { return ToolName } + +func (t *LoadedTool) Enabled() bool { return t != nil && t.cfg != nil && t.cfg.Enabled } + +func (t *LoadedTool) ToolDefinition(description string) *model.Tool { + if strings.TrimSpace(description) == "" && t != nil && t.cfg != nil { + description = t.cfg.ActivationPrompt + } + return ToolDefinition(description) +} + +func (t *LoadedTool) Execute(ctx context.Context, args string, runtime agents.Runtime) (string, error) { + now := runtime.Now + if now.IsZero() { + now = time.Now() + } + result, err := ExecuteTool(args, now) + if err == nil && runtime.Emit != nil { + runtime.Emit(agents.Frame{Type: "trace", Tool: ToolName, Stage: "resolve", Status: "success", Message: "已获取当前时间上下文"}) + } + return result, err +} + +func (t *LoadedTool) RawState() any { return nil } + type Range struct { Start time.Time End time.Time diff --git a/agenttool/registry.go b/agenttool/registry.go new file mode 100644 index 0000000..444a7f2 --- /dev/null +++ b/agenttool/registry.go @@ -0,0 +1,100 @@ +package agenttool + +import ( + "context" + "fmt" + "sort" + "strings" + "sync" + "time" + + "github.com/volcengine/volcengine-go-sdk/service/arkruntime/model" +) + +const LegacySearchProfilesKey = "legacy_search_profiles" + +type LoadOptions struct { + Values map[string]any +} + +func (o LoadOptions) Value(key string) any { + if o.Values == nil { + return nil + } + return o.Values[key] +} + +type Frame struct { + Type string + Tool string + Stage string + Status string + Message string + Data map[string]any + Error string + Text string +} + +type EmitFunc func(frame any) + +type Runtime struct { + Profile any + CompleteText func(context.Context, string, int) (string, error) + Emit EmitFunc + Now time.Time +} + +type LoadedTool interface { + Name() string + Enabled() bool + ToolDefinition(description string) *model.Tool + Execute(context.Context, string, Runtime) (string, error) + RawState() any +} + +type Descriptor struct { + Name string + Load func(path string, options LoadOptions) (LoadedTool, error) +} + +var ( + registryMu sync.RWMutex + registry = map[string]Descriptor{} +) + +func Register(descriptor Descriptor) { + name := strings.ToLower(strings.TrimSpace(descriptor.Name)) + if name == "" { + panic("agenttool: tool name is empty") + } + if descriptor.Load == nil { + panic(fmt.Sprintf("agenttool: %s load function is nil", name)) + } + descriptor.Name = name + + registryMu.Lock() + defer registryMu.Unlock() + if _, ok := registry[name]; ok { + panic(fmt.Sprintf("agenttool: tool %s already registered", name)) + } + registry[name] = descriptor +} + +func Lookup(name string) (Descriptor, bool) { + registryMu.RLock() + defer registryMu.RUnlock() + descriptor, ok := registry[strings.ToLower(strings.TrimSpace(name))] + return descriptor, ok +} + +func Names() []string { + registryMu.RLock() + defer registryMu.RUnlock() + + names := make([]string, 0, len(registry)) + for name := range registry { + names = append(names, name) + } + sort.Strings(names) + return names +} diff --git a/config/config.go b/config/config.go index 66bb234..0234b1a 100644 --- a/config/config.go +++ b/config/config.go @@ -5,8 +5,6 @@ import ( "os" "strings" - searchagent "aichat/agents/search" - "gopkg.in/yaml.v3" ) @@ -17,10 +15,7 @@ const ( defaultToolRouterTimeout = 30 defaultToolRouterMaxTokens = 512 defaultToolRouterSystemText = `你可以按需直接调用可用工具来回答用户问题。 -用户询问简单数学计算、四则运算、加减乘除、括号表达式或明确要求算出数值结果时,调用 calculator。 -如果用户问题包含今天、今日、明天、昨天、本周、本月、本年、最近等相对时间,且后续需要搜索或查询数据库,应先调用 time 获取绝对日期范围。 -需要实时网页资料、新闻、当前版本、近期事件、网页核验或用户明确要求联网时,调用 search。 -需要查询本地业务数据、日程、会议、待办、记录、统计或时间范围内数据时,调用 sql。 +每个工具的 description 描述了它的适用场景和调用条件。 工具结果优先于模型内置知识;工具失败时必须如实说明,不要编造结果。 只调用确实必要的工具。` ) @@ -44,7 +39,7 @@ type ToolRouterConfig struct { Timeout int `yaml:"timeout" json:"timeout"` MaxTokens int `yaml:"max_tokens" json:"max_tokens"` SystemPrompt string `yaml:"system_prompt" json:"system_prompt"` - Tools []ToolRouteConfig `yaml:"tools" json:"tools"` + Tools []ToolRouteConfig `yaml:"tools,omitempty" json:"tools,omitempty"` } type ToolRouteConfig struct { @@ -105,12 +100,6 @@ func DefaultToolRouterConfig() ToolRouterConfig { Timeout: defaultToolRouterTimeout, MaxTokens: defaultToolRouterMaxTokens, SystemPrompt: defaultToolRouterSystemText, - Tools: []ToolRouteConfig{ - {Name: "calculator", Enabled: true, Description: ""}, - {Name: "time", Enabled: true, Description: ""}, - {Name: "search", Enabled: true, Description: ""}, - {Name: "sql", Enabled: true, Description: ""}, - }, } } @@ -123,7 +112,7 @@ func Default() Config { return cfg } -func Load(path string) (*Config, []searchagent.ProfileConfig, error) { +func Load(path string) (*Config, []map[string]any, error) { if err := ensureFile(path); err != nil { return nil, nil, err } @@ -285,7 +274,6 @@ func NormalizeToolRouterConfig(cfg *Config) (bool, error) { func normalizeToolRouterConfig(cfg *Config) (bool, error) { changed := false - defaults := DefaultToolRouterConfig() cfg.ToolRouter.OpenAIName = strings.TrimSpace(cfg.ToolRouter.OpenAIName) if cfg.ToolRouter.Timeout <= 0 { cfg.ToolRouter.Timeout = defaultToolRouterTimeout @@ -303,10 +291,6 @@ func normalizeToolRouterConfig(cfg *Config) (bool, error) { cfg.ToolRouter.SystemPrompt = systemPrompt changed = true } - if len(cfg.ToolRouter.Tools) == 0 { - cfg.ToolRouter.Tools = defaults.Tools - changed = true - } seen := map[string]bool{} for i := range cfg.ToolRouter.Tools { tool := &cfg.ToolRouter.Tools[i] @@ -324,48 +308,17 @@ func normalizeToolRouterConfig(cfg *Config) (bool, error) { } seen[name] = true } - byName := map[string]ToolRouteConfig{} - for _, tool := range cfg.ToolRouter.Tools { - byName[tool.Name] = tool - } - merged := make([]ToolRouteConfig, 0, len(cfg.ToolRouter.Tools)+len(defaults.Tools)) - used := map[string]bool{} - for _, tool := range defaults.Tools { - if existing, ok := byName[tool.Name]; ok { - merged = append(merged, existing) - } else { - merged = append(merged, tool) - changed = true - } - used[tool.Name] = true - } - for _, tool := range cfg.ToolRouter.Tools { - if !used[tool.Name] { - merged = append(merged, tool) - } - } - if len(merged) != len(cfg.ToolRouter.Tools) { - changed = true - } else { - for i := range merged { - if merged[i].Name != cfg.ToolRouter.Tools[i].Name { - changed = true - break - } - } - } - cfg.ToolRouter.Tools = merged return changed, nil } -func readLegacySearchProfiles(data []byte) []searchagent.ProfileConfig { +func readLegacySearchProfiles(data []byte) []map[string]any { var legacy struct { - Search searchagent.ProfileConfigs `yaml:"search"` + Search []map[string]any `yaml:"search"` } if err := yaml.Unmarshal(data, &legacy); err != nil { return nil } - return []searchagent.ProfileConfig(legacy.Search) + return legacy.Search } func Write(path string, cfg Config) error { diff --git a/main.go b/main.go index 63dd7a9..1cb999b 100644 --- a/main.go +++ b/main.go @@ -5,12 +5,12 @@ import ( "os" "strings" - searchagent "aichat/agents/search" - sqlquery "aichat/agents/sql" + agents "aichat/agenttool" "aichat/config" "aichat/conversation" "aichat/llm" "aichat/server" + "aichat/toolmanager" "aichat/toolrouter" ) @@ -27,27 +27,12 @@ func main() { fmt.Fprintln(os.Stderr, "OpenAI 配置初始化失败:", err) os.Exit(1) } - searchConfig, err := searchagent.LoadConfig("agents/search/config.yaml", legacySearchProfiles) + toolManager, err := toolmanager.Load("agents", agents.LoadOptions{Values: map[string]any{agents.LegacySearchProfilesKey: legacySearchProfiles}}) if err != nil { - fmt.Fprintln(os.Stderr, "联网搜索配置加载失败:", err) + fmt.Fprintln(os.Stderr, "工具加载失败:", err) os.Exit(1) } - searchState, err := searchagent.NewState(searchConfig) - if err != nil { - fmt.Fprintln(os.Stderr, "联网搜索初始化失败:", err) - os.Exit(1) - } - sqlConfig, err := sqlquery.LoadConfig("agents/sql/config.yaml") - if err != nil { - fmt.Fprintln(os.Stderr, "SQL 查询插件配置加载失败:", err) - os.Exit(1) - } - sqlState, err := sqlquery.NewState(sqlConfig) - if err != nil { - fmt.Fprintln(os.Stderr, "SQL 查询插件初始化失败:", err) - os.Exit(1) - } - defer sqlState.Close() + defer toolManager.Close() toolRouterState, err := toolrouter.NewState(&cfg.ToolRouter, aiState) if err != nil { fmt.Fprintln(os.Stderr, "工具路由配置初始化失败:", err) @@ -55,7 +40,7 @@ func main() { } store := conversation.NewStore("conversations") - app := server.New(cfg, aiState, searchState, sqlState, toolRouterState, store) + app := server.New(cfg, aiState, toolManager, toolRouterState, store) cfg.Server.Mode = strings.ToLower(cfg.Server.Mode) if err := app.Run(); err != nil { fmt.Fprintln(os.Stderr, "服务异常退出:", err) diff --git a/main_test.go b/main_test.go index 1a69333..f39403a 100644 --- a/main_test.go +++ b/main_test.go @@ -7,10 +7,13 @@ import ( "testing" "time" + "aichat/agents/time" + agents "aichat/agenttool" "aichat/config" "aichat/llm" "aichat/message" "aichat/stream" + "aichat/toolmanager" "aichat/toolrouter" "github.com/volcengine/volcengine-go-sdk/service/arkruntime/model" @@ -18,6 +21,13 @@ import ( const testOpenAIBaseURL = "https://ark.cn-beijing.volces.com/api/v3" +func newTestToolManager(tools ...agents.LoadedTool) *toolmanager.Manager { + if len(tools) == 0 { + tools = []agents.LoadedTool{timeagent.NewLoadedTool(nil)} + } + return toolmanager.NewForTest(tools...) +} + func newTestAI(t *testing.T, configs []config.OpenAIConfig) *llm.State { t.Helper() ai, err := llm.NewState(configs) @@ -46,7 +56,7 @@ func TestNormalizeToolRouterConfigDefaults(t *testing.T) { if strings.TrimSpace(cfg.ToolRouter.SystemPrompt) == "" { t.Fatal("system prompt should be defaulted") } - if len(cfg.ToolRouter.Tools) != 4 || cfg.ToolRouter.Tools[0].Name != "calculator" || cfg.ToolRouter.Tools[1].Name != "time" || cfg.ToolRouter.Tools[2].Name != "search" || cfg.ToolRouter.Tools[3].Name != "sql" || !cfg.ToolRouter.Tools[0].Enabled || !cfg.ToolRouter.Tools[1].Enabled || !cfg.ToolRouter.Tools[2].Enabled || !cfg.ToolRouter.Tools[3].Enabled { + if len(cfg.ToolRouter.Tools) != 0 { t.Fatalf("unexpected tools: %#v", cfg.ToolRouter.Tools) } } @@ -66,14 +76,14 @@ func TestNormalizeOpenAIConfigDefaultsContextWindow(t *testing.T) { } } -func TestNormalizeToolRouterConfigAddsCalculatorAndTimeBeforeSQL(t *testing.T) { +func TestNormalizeToolRouterConfigKeepsConfiguredTools(t *testing.T) { cfg := &config.Config{ToolRouter: config.ToolRouterConfig{ Enabled: true, Timeout: 1, MaxTokens: 1, SystemPrompt: "tools", Tools: []config.ToolRouteConfig{ - {Name: "search", Enabled: true}, + {Name: " Search ", Enabled: true}, {Name: "sql", Enabled: true}, }, }} @@ -82,9 +92,9 @@ func TestNormalizeToolRouterConfigAddsCalculatorAndTimeBeforeSQL(t *testing.T) { t.Fatal(err) } if !changed { - t.Fatal("expected calculator and time tools to be added") + t.Fatal("expected configured tools to be normalized") } - if len(cfg.ToolRouter.Tools) < 4 || cfg.ToolRouter.Tools[0].Name != "calculator" || cfg.ToolRouter.Tools[1].Name != "time" || cfg.ToolRouter.Tools[3].Name != "sql" { + if len(cfg.ToolRouter.Tools) != 2 || cfg.ToolRouter.Tools[0].Name != "search" || cfg.ToolRouter.Tools[1].Name != "sql" { t.Fatalf("unexpected tool order: %#v", cfg.ToolRouter.Tools) } } @@ -121,15 +131,15 @@ func TestAvailableAgentToolsUsesConfigOrderAndEnabled(t *testing.T) { t.Fatal(err) } - tools := toolrouter.AvailableAgentTools(router, ai.ActiveProfile(), nil, nil, nil) - if len(tools) != 2 { + tools := toolrouter.AvailableAgentTools(router, ai.ActiveProfile(), newTestToolManager(), nil) + if len(tools) != 1 { t.Fatalf("tools length = %d", len(tools)) } - if tools[0].Name() != "calculator" || tools[1].Name() != "time" { + if tools[0].Name() != "time" { t.Fatalf("unexpected tools: %#v", tools) } definition := tools[0].Definition() - if definition.Function == nil || definition.Function.Description != "custom calculator" { + if definition.Function == nil || definition.Function.Description != "custom time" { t.Fatalf("unexpected definition: %#v", definition) } } @@ -160,7 +170,7 @@ func TestRunAgentToolLoopAppendsToolMessages(t *testing.T) { t.Fatal(err) } - messages, err := toolrouter.RunAgentToolLoop(context.Background(), router, ai.ActiveProfile(), []message.ChatMessage{{Role: "user", Content: "今天几号"}}, nil, nil, nil) + messages, err := toolrouter.RunAgentToolLoop(context.Background(), router, ai.ActiveProfile(), []message.ChatMessage{{Role: "user", Content: "今天几号"}}, newTestToolManager(), nil) if err != nil { t.Fatal(err) } @@ -208,7 +218,7 @@ func TestRunAgentToolLoopMaxIterations(t *testing.T) { t.Fatal(err) } - messages, err := toolrouter.RunAgentToolLoop(context.Background(), router, ai.ActiveProfile(), []message.ChatMessage{{Role: "user", Content: "今天"}}, nil, nil, nil) + messages, err := toolrouter.RunAgentToolLoop(context.Background(), router, ai.ActiveProfile(), []message.ChatMessage{{Role: "user", Content: "今天"}}, newTestToolManager(), nil) if err != nil { t.Fatal(err) } @@ -344,7 +354,7 @@ func TestRunAgentToolLoopImageUsesTextOnlyDecisionMessages(t *testing.T) { t.Fatal(err) } - messages, err := toolrouter.RunAgentToolLoop(context.Background(), router, ai.ActiveProfile(), []message.ChatMessage{{Role: "user", Content: "描述这张图", ImageURL: "data:image/png;base64,aGVsbG8="}}, nil, nil, nil) + messages, err := toolrouter.RunAgentToolLoop(context.Background(), router, ai.ActiveProfile(), []message.ChatMessage{{Role: "user", Content: "描述这张图", ImageURL: "data:image/png;base64,aGVsbG8="}}, newTestToolManager(), nil) if err != nil { t.Fatal(err) } @@ -381,7 +391,7 @@ func TestRunAgentToolLoopUsesConfiguredRouterProfile(t *testing.T) { t.Fatal(err) } - _, err = toolrouter.RunAgentToolLoop(context.Background(), router, ai.ActiveProfile(), []message.ChatMessage{{Role: "user", Content: "今天"}}, nil, nil, nil) + _, err = toolrouter.RunAgentToolLoop(context.Background(), router, ai.ActiveProfile(), []message.ChatMessage{{Role: "user", Content: "今天"}}, newTestToolManager(), nil) if err != nil { t.Fatal(err) } diff --git a/server/handlers.go b/server/handlers.go index 9be2bfc..8268351 100644 --- a/server/handlers.go +++ b/server/handlers.go @@ -10,6 +10,7 @@ import ( "strings" "time" + searchagent "aichat/agents/search" "aichat/contextwindow" "aichat/conversation" "aichat/llm" @@ -61,17 +62,39 @@ func (s *Server) switchOpenAIHandler(c *gin.Context) { }) } +func (s *Server) searchState() (*searchagent.State, bool) { + if s == nil || s.toolManager == nil { + return nil, false + } + raw, ok := s.toolManager.RawState(searchagent.ToolName) + if !ok { + return nil, false + } + state, ok := raw.(*searchagent.State) + return state, ok && state != nil +} + func (s *Server) listSearchHandler(c *gin.Context) { - c.JSON(http.StatusOK, s.searchState.ListProfiles()) + state, ok := s.searchState() + if !ok { + c.JSON(http.StatusOK, searchagent.ListResponse{}) + return + } + c.JSON(http.StatusOK, state.ListProfiles()) } func (s *Server) switchSearchHandler(c *gin.Context) { + state, ok := s.searchState() + if !ok { + c.JSON(http.StatusBadRequest, gin.H{"error": "搜索工具未加载"}) + return + } var req activeProfileRequest if err := c.ShouldBindJSON(&req); err != nil { c.JSON(http.StatusBadRequest, gin.H{"error": "请求格式错误: " + err.Error()}) return } - profile, err := s.searchState.SwitchActive(req.Name) + profile, err := state.SwitchActive(req.Name) if err != nil { c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) return @@ -219,7 +242,7 @@ func (s *Server) chatHandler(c *gin.Context) { contextMessages = chatWindow.ChatMessages // 用 Function Calling 工具循环替代旧的路由+隐藏上下文机制 - messages, err := toolrouter.RunAgentToolLoop(ctx, s.toolRouterState, profile, contextMessages, s.searchState, s.sqlState, emit) + messages, err := toolrouter.RunAgentToolLoop(ctx, s.toolRouterState, profile, contextMessages, s.toolManager, emit) if err != nil { fmt.Fprintln(os.Stderr, "Agent 工具循环失败:", err) messages, err = message.BuildArkMessages(contextMessages) diff --git a/server/server.go b/server/server.go index 4dacf74..f9482da 100644 --- a/server/server.go +++ b/server/server.go @@ -6,11 +6,10 @@ import ( "net/http" "os" - searchagent "aichat/agents/search" - sqlquery "aichat/agents/sql" "aichat/config" "aichat/conversation" "aichat/llm" + "aichat/toolmanager" "aichat/toolrouter" "github.com/gin-gonic/gin" @@ -19,19 +18,17 @@ import ( type Server struct { cfg *config.Config aiState *llm.State - searchState *searchagent.State - sqlState *sqlquery.State + toolManager *toolmanager.Manager toolRouterState *toolrouter.State store *conversation.Store router *gin.Engine } -func New(cfg *config.Config, aiState *llm.State, searchState *searchagent.State, sqlState *sqlquery.State, toolRouterState *toolrouter.State, store *conversation.Store) *Server { +func New(cfg *config.Config, aiState *llm.State, toolManager *toolmanager.Manager, toolRouterState *toolrouter.State, store *conversation.Store) *Server { s := &Server{ cfg: cfg, aiState: aiState, - searchState: searchState, - sqlState: sqlState, + toolManager: toolManager, toolRouterState: toolRouterState, store: store, router: gin.Default(), diff --git a/toolmanager/manager.go b/toolmanager/manager.go new file mode 100644 index 0000000..2b0c6b2 --- /dev/null +++ b/toolmanager/manager.go @@ -0,0 +1,125 @@ +package toolmanager + +import ( + "errors" + "fmt" + "os" + "path/filepath" + "sort" + "strings" + + _ "aichat/agents/calculator" + _ "aichat/agents/search" + _ "aichat/agents/sql" + _ "aichat/agents/time" + agents "aichat/agenttool" +) + +type Manager struct { + tools map[string]agents.LoadedTool + order []string +} + +func Load(root string, options agents.LoadOptions) (*Manager, error) { + entries, err := os.ReadDir(root) + if err != nil { + return nil, fmt.Errorf("扫描工具目录失败: %w", err) + } + + manager := &Manager{tools: map[string]agents.LoadedTool{}} + for _, entry := range entries { + if !entry.IsDir() { + continue + } + name := strings.ToLower(strings.TrimSpace(entry.Name())) + descriptor, ok := agents.Lookup(name) + if !ok { + continue + } + tool, err := descriptor.Load(filepath.Join(root, entry.Name(), "config.yaml"), options) + if err != nil { + manager.Close() + return nil, fmt.Errorf("加载工具 %s 失败: %w", name, err) + } + if tool == nil { + continue + } + toolName := strings.ToLower(strings.TrimSpace(tool.Name())) + if toolName == "" { + toolName = name + } + if _, ok := manager.tools[toolName]; ok { + manager.Close() + return nil, fmt.Errorf("工具名称重复: %s", toolName) + } + manager.tools[toolName] = tool + manager.order = append(manager.order, toolName) + } + return manager, nil +} + +func NewForTest(tools ...agents.LoadedTool) *Manager { + manager := &Manager{tools: map[string]agents.LoadedTool{}} + for _, tool := range tools { + if tool == nil { + continue + } + name := strings.ToLower(strings.TrimSpace(tool.Name())) + if name == "" { + continue + } + if _, ok := manager.tools[name]; !ok { + manager.order = append(manager.order, name) + } + manager.tools[name] = tool + } + return manager +} + +func (m *Manager) Tools() []agents.LoadedTool { + if m == nil { + return nil + } + tools := make([]agents.LoadedTool, 0, len(m.order)) + for _, name := range m.order { + if tool := m.tools[name]; tool != nil { + tools = append(tools, tool) + } + } + return tools +} + +func (m *Manager) Get(name string) (agents.LoadedTool, bool) { + if m == nil { + return nil, false + } + tool, ok := m.tools[strings.ToLower(strings.TrimSpace(name))] + return tool, ok +} + +func (m *Manager) RawState(name string) (any, bool) { + tool, ok := m.Get(name) + if !ok || tool == nil { + return nil, false + } + return tool.RawState(), true +} + +func (m *Manager) Close() error { + if m == nil { + return nil + } + var errs []string + for _, tool := range m.Tools() { + if closer, ok := tool.(interface{ Close() error }); ok { + if err := closer.Close(); err != nil { + errs = append(errs, err.Error()) + } + } + } + if len(errs) > 0 { + sort.Strings(errs) + return errors.New(strings.Join(errs, "; ")) + } + return nil +} diff --git a/toolrouter/loop.go b/toolrouter/loop.go index a568e33..0df8181 100644 --- a/toolrouter/loop.go +++ b/toolrouter/loop.go @@ -6,11 +6,10 @@ import ( "strings" "time" - searchagent "aichat/agents/search" - sqlquery "aichat/agents/sql" "aichat/llm" "aichat/message" "aichat/stream" + "aichat/toolmanager" "aichat/utils" "github.com/volcengine/volcengine-go-sdk/service/arkruntime/model" @@ -18,7 +17,7 @@ import ( const maxAgentToolIterations = 6 -func RunAgentToolLoop(ctx context.Context, state *State, profile *llm.Profile, chatMessages []message.ChatMessage, searchState *searchagent.State, sqlState *sqlquery.State, emit stream.EmitFunc) ([]*model.ChatCompletionMessage, error) { +func RunAgentToolLoop(ctx context.Context, state *State, profile *llm.Profile, chatMessages []message.ChatMessage, manager *toolmanager.Manager, emit stream.EmitFunc) ([]*model.ChatCompletionMessage, error) { finalMessages, err := message.BuildArkMessages(chatMessages) if err != nil { return nil, err @@ -27,7 +26,7 @@ func RunAgentToolLoop(ctx context.Context, state *State, profile *llm.Profile, c if state != nil { routerProfile = state.RouterProfile(profile) } - tools := availableAgentTools(state, routerProfile, searchState, sqlState, emit) + tools := availableAgentTools(state, routerProfile, manager, emit) if len(tools) == 0 { return finalMessages, nil } diff --git a/toolrouter/tools.go b/toolrouter/tools.go index 0c16733..a39f42f 100644 --- a/toolrouter/tools.go +++ b/toolrouter/tools.go @@ -6,14 +6,13 @@ import ( "strings" "time" - "aichat/agents/calculator" - searchagent "aichat/agents/search" - sqlquery "aichat/agents/sql" - timeagent "aichat/agents/time" + agents "aichat/agenttool" "aichat/completion" + "aichat/config" "aichat/llm" "aichat/message" "aichat/stream" + "aichat/toolmanager" "aichat/utils" "github.com/volcengine/volcengine-go-sdk/service/arkruntime/model" @@ -33,101 +32,96 @@ func (t AgentTool) Name() string { return t.name } func (t AgentTool) Definition() *model.Tool { return t.definition } -func AvailableAgentTools(state *State, profile *llm.Profile, searchState *searchagent.State, sqlState *sqlquery.State, emit stream.EmitFunc) []AgentTool { - return availableAgentTools(state, profile, searchState, sqlState, emit) +func AvailableAgentTools(state *State, profile *llm.Profile, manager *toolmanager.Manager, emit stream.EmitFunc) []AgentTool { + return availableAgentTools(state, profile, manager, emit) } -func availableAgentTools(state *State, profile *llm.Profile, searchState *searchagent.State, sqlState *sqlquery.State, emit stream.EmitFunc) []AgentTool { - if state == nil || state.cfg == nil || !state.cfg.Enabled { +func availableAgentTools(state *State, profile *llm.Profile, manager *toolmanager.Manager, emit stream.EmitFunc) []AgentTool { + if state == nil || state.cfg == nil || !state.cfg.Enabled || manager == nil { return nil } - tools := make([]AgentTool, 0, len(state.cfg.Tools)) + overrides := map[string]config.ToolRouteConfig{} for _, item := range state.cfg.Tools { - if !item.Enabled { + name := strings.ToLower(strings.TrimSpace(item.Name)) + if name != "" { + overrides[name] = item + } + } + + if len(state.cfg.Tools) > 0 { + tools := make([]AgentTool, 0, len(state.cfg.Tools)) + for _, item := range state.cfg.Tools { + name := strings.ToLower(strings.TrimSpace(item.Name)) + if name == "" || !item.Enabled { + continue + } + tool, ok := manager.Get(name) + if !ok { + continue + } + if agentTool, ok := buildAgentTool(tool, strings.TrimSpace(item.Description), profile, emit); ok { + tools = append(tools, agentTool) + } + } + return tools + } + + loaded := manager.Tools() + tools := make([]AgentTool, 0, len(loaded)) + for _, tool := range loaded { + if tool == nil { continue } - description := strings.TrimSpace(item.Description) - switch item.Name { - case calculator.ToolName: - tools = append(tools, AgentTool{ - name: calculator.ToolName, - definition: calculator.ToolDefinition(description), - execute: func(ctx context.Context, args string) (string, error) { - result, err := calculator.ExecuteTool(args) - if err == nil && emit != nil { - emit(stream.Frame{Type: "trace", Tool: calculator.ToolName, Stage: "calculate", Status: "success", Message: "四则运算完成"}) - } - return result, err - }, - }) - case timeagent.ToolName: - tools = append(tools, AgentTool{ - name: timeagent.ToolName, - definition: timeagent.ToolDefinition(description), - execute: func(ctx context.Context, args string) (string, error) { - result, err := timeagent.ExecuteTool(args, time.Now()) - if err == nil && emit != nil { - emit(stream.Frame{Type: "trace", Tool: timeagent.ToolName, Stage: "resolve", Status: "success", Message: "已获取当前时间上下文"}) - } - return result, err - }, - }) - case searchagent.ToolName: - if searchState == nil || !searchState.Enabled() { - continue - } - tools = append(tools, AgentTool{ - name: searchagent.ToolName, - definition: searchState.ToolDefinition(description), - execute: func(ctx context.Context, args string) (string, error) { - if emit != nil { - emit(stream.Frame{Type: "trace", Tool: searchagent.ToolName, Stage: "request", Status: "running", Message: "正在联网搜索"}) - } - result, err := searchState.ExecuteTool(ctx, args) - if emit != nil { - status := "success" - messageText := "联网搜索完成" - if err != nil { - status = "error" - messageText = "联网搜索失败" - } - emit(stream.Frame{Type: "trace", Tool: searchagent.ToolName, Stage: "results", Status: status, Message: messageText}) - } - return result, err - }, - }) - case sqlquery.ToolName: - if sqlState == nil || !sqlState.Enabled() { - continue - } - tools = append(tools, AgentTool{ - name: sqlquery.ToolName, - definition: sqlState.ToolDefinition(description), - execute: func(ctx context.Context, args string) (string, error) { - if emit != nil { - emit(stream.Frame{Type: "trace", Tool: sqlquery.ToolName, Stage: "execute", Status: "running", Message: "正在查询数据库"}) - } - generator := func(ctx context.Context, prompt string, maxTokens int) (string, error) { - return completion.CompleteText(ctx, profile, []message.ChatMessage{{Role: "system", Content: prompt}}, maxTokens) - } - result, err := sqlState.ExecuteTool(ctx, args, generator) - if emit != nil { - status := "success" - messageText := "数据库查询完成" - if err != nil { - status = "error" - messageText = "数据库查询失败" - } - emit(stream.Frame{Type: "trace", Tool: sqlquery.ToolName, Stage: "execute", Status: status, Message: messageText}) - } - return result, err - }, - }) + item := overrides[tool.Name()] + if agentTool, ok := buildAgentTool(tool, strings.TrimSpace(item.Description), profile, emit); ok { + tools = append(tools, agentTool) } } return tools } +func buildAgentTool(tool agents.LoadedTool, description string, profile *llm.Profile, emit stream.EmitFunc) (AgentTool, bool) { + if tool == nil || !tool.Enabled() { + return AgentTool{}, false + } + definition := tool.ToolDefinition(description) + if definition == nil || definition.Function == nil { + return AgentTool{}, false + } + name := tool.Name() + return AgentTool{ + name: name, + definition: definition, + execute: func(ctx context.Context, args string) (string, error) { + runtime := agents.Runtime{ + Profile: profile, + Now: time.Now(), + Emit: wrapAgentEmit(emit), + } + if profile != nil { + runtime.CompleteText = func(ctx context.Context, prompt string, maxTokens int) (string, error) { + return completion.CompleteText(ctx, profile, []message.ChatMessage{{Role: "system", Content: prompt}}, maxTokens) + } + } + return tool.Execute(ctx, args, runtime) + }, + }, true +} + +func wrapAgentEmit(emit stream.EmitFunc) agents.EmitFunc { + if emit == nil { + return nil + } + return func(frame any) { + switch value := frame.(type) { + case agents.Frame: + emit(stream.Frame{Type: value.Type, Tool: value.Tool, Stage: value.Stage, Status: value.Status, Message: value.Message, Data: value.Data, Error: value.Error, Text: value.Text}) + case stream.Frame: + emit(value) + } + } +} + func ExecuteAgentToolCall(ctx context.Context, call *model.ToolCall, tools map[string]AgentTool, emit stream.EmitFunc) string { if call == nil || call.Type != model.ToolTypeFunction { result := "工具调用无效:仅支持 function 类型工具。"