diff --git a/internal/bot/bot.go b/internal/bot/bot.go index ebe6d8f..ceae387 100644 --- a/internal/bot/bot.go +++ b/internal/bot/bot.go @@ -2,7 +2,6 @@ package bot import ( "context" - "errors" "fmt" "github.com/openai/openai-go" @@ -13,33 +12,66 @@ import ( const maxHistory = 20 type Bot struct { - client *openai.Client - cfg *config.Config - history []openai.ChatCompletionMessageParamUnion + clients map[string]*openai.Client + cfg *config.Config + provider *config.Provider + model string + history []openai.ChatCompletionMessageParamUnion + systemPrompt string } func New(cfg *config.Config) *Bot { - client := openai.NewClient( - option.WithAPIKey(cfg.APIKey), - option.WithBaseURL(cfg.BaseURL), - ) - return &Bot{ - client: &client, - cfg: cfg, + b := &Bot{ + clients: make(map[string]*openai.Client), + cfg: cfg, + systemPrompt: cfg.SystemPrompt, } + b.provider = config.FindProvider(cfg.DefaultProvider) + b.model = cfg.DefaultModel + return b +} + +func (b *Bot) client() *openai.Client { + if c, ok := b.clients[b.provider.Name]; ok { + return c + } + c := openai.NewClient( + option.WithAPIKey(b.provider.APIKey), + option.WithBaseURL(b.provider.BaseURL), + ) + b.clients[b.provider.Name] = &c + return &c +} + +func (b *Bot) Models() []string { + return config.AllModels() +} + +func (b *Bot) Current() (string, string) { + return b.provider.Name, b.model +} + +func (b *Bot) SwitchModel(id string) error { + p, modelName, err := config.ResolveModel(id) + if err != nil { + return err + } + b.provider = p + b.model = modelName + return nil } func (b *Bot) Chat(ctx context.Context, userMsg string) (string, error) { - if b.cfg.APIKey == "" { - return "", errors.New("未配置 api_key,请编辑 data/config.yaml") + if b.provider.APIKey == "" { + return "", fmt.Errorf("供应商 %s 未配置 api_key,请编辑 data/config.yaml", b.provider.Name) } history := make([]openai.ChatCompletionMessageParamUnion, 0, len(b.history)+2) - history = append(history, openai.SystemMessage(b.cfg.SystemPrompt)) + history = append(history, openai.SystemMessage(b.systemPrompt)) history = append(history, b.history...) history = append(history, openai.UserMessage(userMsg)) - stream := b.client.Chat.Completions.NewStreaming(ctx, openai.ChatCompletionNewParams{ - Model: b.cfg.Model, + stream := b.client().Chat.Completions.NewStreaming(ctx, openai.ChatCompletionNewParams{ + Model: b.model, Messages: history, }) answer := "" diff --git a/internal/config/config.go b/internal/config/config.go index c8632db..373d6e5 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -1,8 +1,11 @@ package config import ( + "errors" + "fmt" "os" "path/filepath" + "strings" "gopkg.in/yaml.v3" ) @@ -12,14 +15,27 @@ const ( configFile = "config.yaml" ) +type Provider struct { + Name string `yaml:"name"` + APIKey string `yaml:"api_key"` + BaseURL string `yaml:"base_url"` + Models []string `yaml:"models"` +} + type Config struct { - BotName string `yaml:"bot_name"` - Port int `yaml:"port"` - LogLevel string `yaml:"log_level"` - APIKey string `yaml:"api_key"` - BaseURL string `yaml:"base_url"` - Model string `yaml:"model"` - SystemPrompt string `yaml:"system_prompt"` + BotName string `yaml:"bot_name"` + Port int `yaml:"port"` + LogLevel string `yaml:"log_level"` + SystemPrompt string `yaml:"system_prompt"` + Providers []Provider `yaml:"providers"` + DefaultProvider string `yaml:"default_provider"` + DefaultModel string `yaml:"default_model"` +} + +type legacyConfig struct { + APIKey string `yaml:"api_key"` + BaseURL string `yaml:"base_url"` + Model string `yaml:"model"` } var cfg *Config @@ -45,54 +61,190 @@ func Load() (*Config, error) { if err := yaml.Unmarshal(data, cfg); err != nil { return nil, err } + if len(cfg.Providers) == 0 { + if err := migrateLegacy(path, data); err != nil { + return nil, err + } + } applyDefaults(cfg) + if err := validate(cfg); err != nil { + return nil, err + } return cfg, nil } -func applyDefaults(c *Config) { - def := &Config{ - BotName: "ai-bot", - Port: 8080, - LogLevel: "info", - BaseURL: "https://api.openai.com/v1", - Model: "gpt-4o-mini", - SystemPrompt: "你是一个乐于助人的 AI 助手。", - } - if c.BotName == "" { - c.BotName = def.BotName - } - if c.Port == 0 { - c.Port = def.Port - } - if c.LogLevel == "" { - c.LogLevel = def.LogLevel - } - if c.BaseURL == "" { - c.BaseURL = def.BaseURL - } - if c.Model == "" { - c.Model = def.Model - } - if c.SystemPrompt == "" { - c.SystemPrompt = def.SystemPrompt - } -} - func GetConfig() *Config { return cfg } +func migrateLegacy(path string, data []byte) error { + legacy := &legacyConfig{} + if err := yaml.Unmarshal(data, legacy); err != nil { + return err + } + if legacy.APIKey == "" && legacy.BaseURL == "" && legacy.Model == "" { + return errors.New("配置文件中没有 providers,请检查 data/config.yaml") + } + p := Provider{ + Name: "openai", + APIKey: legacy.APIKey, + BaseURL: legacy.BaseURL, + Models: []string{legacy.Model}, + } + if p.BaseURL == "" { + p.BaseURL = "https://api.openai.com/v1" + } + if len(p.Models) == 0 || p.Models[0] == "" { + p.Models = []string{"gpt-4o-mini"} + } + cfg.Providers = []Provider{p} + cfg.DefaultProvider = p.Name + cfg.DefaultModel = p.Models[0] + return writeFile(path, cfg) +} + +func applyDefaults(c *Config) { + if c.BotName == "" { + c.BotName = "ai-bot" + } + if c.Port == 0 { + c.Port = 8080 + } + if c.LogLevel == "" { + c.LogLevel = "info" + } + if c.SystemPrompt == "" { + c.SystemPrompt = "你是一个乐于助人的 AI 助手。" + } + if c.DefaultProvider == "" && len(c.Providers) > 0 { + c.DefaultProvider = c.Providers[0].Name + } +} + +func validate(c *Config) error { + if len(c.Providers) == 0 { + return errors.New("至少需要一个供应商 (providers)") + } + names := make(map[string]bool, len(c.Providers)) + for i := range c.Providers { + p := &c.Providers[i] + if p.Name == "" { + return fmt.Errorf("providers[%d] 缺少 name", i) + } + if names[p.Name] { + return fmt.Errorf("供应商名称重复: %s", p.Name) + } + names[p.Name] = true + if p.BaseURL == "" { + return fmt.Errorf("供应商 %s 缺少 base_url", p.Name) + } + if len(p.Models) == 0 { + return fmt.Errorf("供应商 %s 未配置 models", p.Name) + } + for _, m := range p.Models { + if m == "" { + return fmt.Errorf("供应商 %s 包含空模型名", p.Name) + } + } + } + if _, ok := names[c.DefaultProvider]; !ok { + return fmt.Errorf("default_provider %q 不存在", c.DefaultProvider) + } + if _, _, err := ResolveModel(c.DefaultModel); err != nil { + return fmt.Errorf("default_model 无效: %w", err) + } + return nil +} + +func FindProvider(name string) *Provider { + for i := range cfg.Providers { + if cfg.Providers[i].Name == name { + return &cfg.Providers[i] + } + } + return nil +} + +func ResolveModel(id string) (*Provider, string, error) { + if id == "" { + id = cfg.DefaultModel + } + if providerName, modelName, ok := strings.Cut(id, "/"); ok { + p := FindProvider(providerName) + if p == nil { + return nil, "", fmt.Errorf("供应商 %q 不存在", providerName) + } + if !contains(p.Models, modelName) { + return nil, "", fmt.Errorf("供应商 %s 没有模型 %q", p.Name, modelName) + } + return p, modelName, nil + } + var found *Provider + for i := range cfg.Providers { + if contains(cfg.Providers[i].Models, id) { + if found != nil { + return nil, "", fmt.Errorf("模型 %q 在多个供应商中存在,请使用 provider/model 格式指定", id) + } + found = &cfg.Providers[i] + } + } + if found == nil { + return nil, "", fmt.Errorf("模型 %q 不存在", id) + } + return found, id, nil +} + +func AllModels() []string { + var out []string + for i := range cfg.Providers { + p := &cfg.Providers[i] + for _, m := range p.Models { + out = append(out, p.Name+"/"+m) + } + } + return out +} + +func contains(list []string, s string) bool { + for _, v := range list { + if v == s { + return true + } + } + return false +} + func writeDefault(path string) error { cfg = &Config{ BotName: "ai-bot", Port: 8080, LogLevel: "info", - APIKey: "", - BaseURL: "https://api.openai.com/v1", - Model: "gpt-4o-mini", SystemPrompt: "你是一个乐于助人的 AI 助手。", + Providers: []Provider{ + { + Name: "openai", + APIKey: "", + BaseURL: "https://api.openai.com/v1", + Models: []string{"gpt-4o-mini", "gpt-4o"}, + }, + { + Name: "deepseek", + APIKey: "", + BaseURL: "https://api.deepseek.com/v1", + Models: []string{"deepseek-chat", "deepseek-reasoner"}, + }, + }, + DefaultProvider: "openai", + DefaultModel: "gpt-4o-mini", } - data, err := yaml.Marshal(cfg) + if err := validate(cfg); err != nil { + return err + } + return writeFile(path, cfg) +} + +func writeFile(path string, c *Config) error { + data, err := yaml.Marshal(c) if err != nil { return err } diff --git a/main.go b/main.go index 0284ca1..0709e41 100644 --- a/main.go +++ b/main.go @@ -17,12 +17,11 @@ func main() { if err != nil { log.Fatalf("加载配置失败: %v", err) } - if cfg.APIKey == "" { - log.Println("提示: 未配置 api_key,请编辑 data/config.yaml") - } - fmt.Printf("🤖 %s 已启动 (模型: %s)。输入问题开始对话,输入 /exit 退出。\n", cfg.BotName, cfg.Model) - b := bot.New(cfg) + provider, model := b.Current() + fmt.Printf("🤖 %s 已启动 (供应商: %s, 模型: %s)。输入问题开始对话,输入 /help 查看命令。\n", + cfg.BotName, provider, model) + scanner := bufio.NewScanner(os.Stdin) for { fmt.Print("你: ") @@ -33,9 +32,11 @@ func main() { if input == "" { continue } - if input == "/exit" || input == "/quit" { - fmt.Println("再见!") - break + if strings.HasPrefix(input, "/") { + if !handleCommand(b, input) { + break + } + continue } answer, err := b.Chat(context.Background(), input) if err != nil { @@ -48,3 +49,40 @@ func main() { log.Fatal(err) } } + +func handleCommand(b *bot.Bot, input string) bool { + fields := strings.Fields(input) + cmd, args := fields[0], fields[1:] + switch cmd { + case "/exit", "/quit": + fmt.Println("再见!") + return false + case "/help": + fmt.Println("命令列表:") + fmt.Println(" /models 列出所有供应商和模型") + fmt.Println(" /use <模型> 切换模型,如 /use deepseek-chat 或 /use deepseek/deepseek-chat") + fmt.Println(" /info 显示当前供应商和模型") + fmt.Println(" /exit 退出") + case "/models": + for _, m := range b.Models() { + fmt.Println(" " + m) + } + case "/use": + if len(args) == 0 { + fmt.Println("用法: /use <模型>,如 /use deepseek-chat") + return true + } + if err := b.SwitchModel(args[0]); err != nil { + fmt.Printf("⚠️ %v\n", err) + return true + } + provider, model := b.Current() + fmt.Printf("已切换到 %s/%s (对话历史已保留)\n", provider, model) + case "/info": + provider, model := b.Current() + fmt.Printf("供应商: %s, 模型: %s\n", provider, model) + default: + fmt.Printf("未知命令: %s,输入 /help 查看命令列表\n", cmd) + } + return true +}