diff --git a/internal/bot/bot.go b/internal/bot/bot.go index bb8d4f3..b03c8c0 100644 --- a/internal/bot/bot.go +++ b/internal/bot/bot.go @@ -22,8 +22,6 @@ const ( maxToolRounds = 5 ) -var toolRegistry = tools.NewRegistry(builtin.TimeTool{}, builtin.CalculatorTool{}, builtin.RandomTool{}) - type Bot struct { clients map[string]*openai.Client cfg *config.Config @@ -35,13 +33,15 @@ type Bot struct { visionModel string history []openai.ChatCompletionMessageParamUnion systemPrompt string + toolRegistry *tools.Registry } -func New(cfg *config.Config) *Bot { +func New(cfg *config.Config) (*Bot, error) { b := &Bot{ clients: make(map[string]*openai.Client), cfg: cfg, systemPrompt: cfg.SystemPrompt, + toolRegistry: tools.NewRegistry(builtin.TimeTool{}, builtin.CalculatorTool{}, builtin.RandomTool{}), } b.provider = config.FindProvider(cfg.DefaultProvider) b.model = cfg.DefaultModel @@ -57,7 +57,10 @@ func New(cfg *config.Config) *Bot { b.visionProvider, b.visionModel = p, m } } - return b + if err := b.toolRegistry.InitConfigs(); err != nil { + return nil, err + } + return b, nil } func (b *Bot) clientFor(p *config.Provider) *openai.Client { @@ -120,7 +123,7 @@ func (b *Bot) CurrentRoles() (tool, vision string) { } func (b *Bot) Tools() []string { - return toolRegistry.List() + return b.toolRegistry.List() } func (b *Bot) ContextDump() string { @@ -202,7 +205,7 @@ func (b *Bot) toolRound(ctx context.Context, userMsg string, onReasoning func(st params := openai.ChatCompletionNewParams{ Model: b.toolModel, Messages: history, - Tools: toolRegistry.ParamList(), + Tools: b.toolRegistry.ParamList(), } b.applyThinkingParams(¶ms, b.toolProvider) stream := b.clientFor(b.toolProvider).Chat.Completions.NewStreaming(ctx, params) @@ -254,7 +257,7 @@ func (b *Bot) toolRound(ctx context.Context, userMsg string, onReasoning func(st if onTool != nil { onTool(tc.Function.Name, tc.Function.Arguments) } - result, err := toolRegistry.Execute(tc.Function.Name, json.RawMessage(tc.Function.Arguments)) + result, err := b.toolRegistry.Execute(tc.Function.Name, json.RawMessage(tc.Function.Arguments)) if err != nil { result = "工具执行失败: " + err.Error() } diff --git a/internal/config/config.go b/internal/config/config.go index d5a97dd..fcbc142 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -11,8 +11,9 @@ import ( ) const ( - configDir = "data" - configFile = "config.yaml" + configDir = "data" + configFile = "config.yaml" + toolConfigDir = "data/tools" ) type Provider struct { @@ -270,3 +271,36 @@ func writeFile(path string, c *Config) error { } return os.WriteFile(path, data, 0o644) } + +// LoadToolConfig 读取 data/tools/.yaml,文件不存在时返回 ok=false。 +func LoadToolConfig(name string) (cfg map[string]any, ok bool, err error) { + path := filepath.Join(toolConfigDir, name+".yaml") + data, err := os.ReadFile(path) + if os.IsNotExist(err) { + return nil, false, nil + } + if err != nil { + return nil, false, err + } + if err := yaml.Unmarshal(data, &cfg); err != nil { + return nil, false, fmt.Errorf("解析工具配置 %s 失败: %w", path, err) + } + return cfg, true, nil +} + +// WriteDefaultToolConfig 生成工具默认配置模板到 data/tools/.yaml。 +func WriteDefaultToolConfig(name string, defaults map[string]any) error { + if err := os.MkdirAll(toolConfigDir, 0o755); err != nil { + return err + } + data, err := yaml.Marshal(defaults) + if err != nil { + return err + } + path := filepath.Join(toolConfigDir, name+".yaml") + return os.WriteFile(path, data, 0o644) +} + +func ToolConfigPath(name string) string { + return filepath.Join(toolConfigDir, name+".yaml") +} diff --git a/internal/tools/initconfigs_test.go b/internal/tools/initconfigs_test.go new file mode 100644 index 0000000..9be5c6e --- /dev/null +++ b/internal/tools/initconfigs_test.go @@ -0,0 +1,89 @@ +package tools + +import ( + "encoding/json" + "fmt" + "os" + "path/filepath" + "strings" + "testing" +) + +type stubConfigurable struct { + configured map[string]any + fail bool +} + +func (s *stubConfigurable) Name() string { return "db" } +func (s *stubConfigurable) Description() string { return "测试数据库工具" } +func (s *stubConfigurable) Parameters() map[string]any { + return map[string]any{"type": "object"} +} +func (s *stubConfigurable) Execute(args json.RawMessage) (string, error) { + return "ok", nil +} +func (s *stubConfigurable) DefaultConfig() map[string]any { + return map[string]any{"host": "127.0.0.1", "password": "请填写"} +} +func (s *stubConfigurable) Configure(cfg map[string]any) error { + if s.fail { + return fmt.Errorf("密码为空") + } + s.configured = cfg + return nil +} + +func TestInitConfigsMissing(t *testing.T) { + t.Chdir(t.TempDir()) + stub := &stubConfigurable{} + err := NewRegistry(stub).InitConfigs() + if err == nil || !strings.Contains(err.Error(), "db") { + t.Fatalf("缺少配置应报错, got %v", err) + } + data, rerr := os.ReadFile(filepath.Join("data", "tools", "db.yaml")) + if rerr != nil { + t.Fatalf("默认配置未生成: %v", rerr) + } + if !strings.Contains(string(data), "请填写") { + t.Errorf("默认模板内容异常: %s", data) + } + if stub.configured != nil { + t.Error("缺少配置时不应调用 Configure") + } +} + +func TestInitConfigsOK(t *testing.T) { + t.Chdir(t.TempDir()) + if err := os.MkdirAll(filepath.Join("data", "tools"), 0o755); err != nil { + t.Fatal(err) + } + path := filepath.Join("data", "tools", "db.yaml") + if err := os.WriteFile(path, []byte("host: localhost\npassword: secret\n"), 0o644); err != nil { + t.Fatal(err) + } + stub := &stubConfigurable{} + if err := NewRegistry(stub).InitConfigs(); err != nil { + t.Fatalf("InitConfigs 出错: %v", err) + } + if stub.configured == nil || stub.configured["host"] != "localhost" || stub.configured["password"] != "secret" { + t.Errorf("Configure 未收到配置: %v", stub.configured) + } +} + +func TestInitConfigsInvalid(t *testing.T) { + t.Chdir(t.TempDir()) + os.MkdirAll(filepath.Join("data", "tools"), 0o755) + os.WriteFile(filepath.Join("data", "tools", "db.yaml"), []byte("host: localhost\n"), 0o644) + stub := &stubConfigurable{fail: true} + err := NewRegistry(stub).InitConfigs() + if err == nil || !strings.Contains(err.Error(), "密码为空") { + t.Fatalf("非法配置应报错, got %v", err) + } +} + +func TestInitConfigsSkipsPlain(t *testing.T) { + t.Chdir(t.TempDir()) + if err := NewRegistry(stubTool{}).InitConfigs(); err != nil { + t.Fatalf("普通工具不应报错: %v", err) + } +} diff --git a/internal/tools/tools.go b/internal/tools/tools.go index a8de274..cc0c601 100644 --- a/internal/tools/tools.go +++ b/internal/tools/tools.go @@ -4,10 +4,12 @@ import ( "encoding/json" "fmt" "sort" + "strings" "github.com/openai/openai-go" "github.com/openai/openai-go/packages/param" "github.com/openai/openai-go/shared" + "myaibot/internal/config" ) type Tool interface { @@ -17,6 +19,14 @@ type Tool interface { Execute(args json.RawMessage) (string, error) } +// Configurable 是可配置工具的接口:配置文件 data/tools/.yaml +// 缺失时生成默认模板并提醒,存在时注入配置。 +type Configurable interface { + Tool + DefaultConfig() map[string]any + Configure(cfg map[string]any) error +} + type Registry struct { tools map[string]Tool } @@ -68,3 +78,33 @@ func (r *Registry) Execute(name string, args json.RawMessage) (string, error) { } return t.Execute(args) } + +// InitConfigs 为可配置工具注入配置。 +// 配置文件缺失时生成默认模板并返回提醒;配置非法时返回错误。 +func (r *Registry) InitConfigs() error { + var missing []string + for name, t := range r.tools { + c, ok := t.(Configurable) + if !ok { + continue + } + cfg, found, err := config.LoadToolConfig(name) + if err != nil { + return fmt.Errorf("工具 %s 配置加载失败: %w", name, err) + } + if !found { + if err := config.WriteDefaultToolConfig(name, c.DefaultConfig()); err != nil { + return fmt.Errorf("工具 %s 默认配置生成失败: %w", name, err) + } + missing = append(missing, fmt.Sprintf("工具 %s 缺少配置,已生成模板 %s,请填写后重启", name, config.ToolConfigPath(name))) + continue + } + if err := c.Configure(cfg); err != nil { + return fmt.Errorf("工具 %s 配置无效: %w", name, err) + } + } + if len(missing) > 0 { + return fmt.Errorf("%s", strings.Join(missing, "\n")) + } + return nil +} diff --git a/main.go b/main.go index fbe4616..6ad8ac0 100644 --- a/main.go +++ b/main.go @@ -6,6 +6,7 @@ import ( "fmt" "io" "log" + "os" "strings" "github.com/peterh/liner" @@ -20,7 +21,12 @@ func main() { if err != nil { log.Fatalf("加载配置失败: %v", err) } - b := bot.New(cfg) + b, err := bot.New(cfg) + if err != nil { + fmt.Printf("⚠️ %v\n", err) + fmt.Println("请填写工具配置文件后重新启动。") + os.Exit(1) + } provider, model := b.Current() fmt.Printf("🤖 %s 已启动 (供应商: %s, 模型: %s)。输入问题开始对话,输入 /help 查看命令。\n", cfg.BotName, provider, model)