增加多个ai
This commit is contained in:
+45
-13
@@ -2,7 +2,6 @@ package bot
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"errors"
|
|
||||||
"fmt"
|
"fmt"
|
||||||
|
|
||||||
"github.com/openai/openai-go"
|
"github.com/openai/openai-go"
|
||||||
@@ -13,33 +12,66 @@ import (
|
|||||||
const maxHistory = 20
|
const maxHistory = 20
|
||||||
|
|
||||||
type Bot struct {
|
type Bot struct {
|
||||||
client *openai.Client
|
clients map[string]*openai.Client
|
||||||
cfg *config.Config
|
cfg *config.Config
|
||||||
|
provider *config.Provider
|
||||||
|
model string
|
||||||
history []openai.ChatCompletionMessageParamUnion
|
history []openai.ChatCompletionMessageParamUnion
|
||||||
|
systemPrompt string
|
||||||
}
|
}
|
||||||
|
|
||||||
func New(cfg *config.Config) *Bot {
|
func New(cfg *config.Config) *Bot {
|
||||||
client := openai.NewClient(
|
b := &Bot{
|
||||||
option.WithAPIKey(cfg.APIKey),
|
clients: make(map[string]*openai.Client),
|
||||||
option.WithBaseURL(cfg.BaseURL),
|
|
||||||
)
|
|
||||||
return &Bot{
|
|
||||||
client: &client,
|
|
||||||
cfg: cfg,
|
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) {
|
func (b *Bot) Chat(ctx context.Context, userMsg string) (string, error) {
|
||||||
if b.cfg.APIKey == "" {
|
if b.provider.APIKey == "" {
|
||||||
return "", errors.New("未配置 api_key,请编辑 data/config.yaml")
|
return "", fmt.Errorf("供应商 %s 未配置 api_key,请编辑 data/config.yaml", b.provider.Name)
|
||||||
}
|
}
|
||||||
history := make([]openai.ChatCompletionMessageParamUnion, 0, len(b.history)+2)
|
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, b.history...)
|
||||||
history = append(history, openai.UserMessage(userMsg))
|
history = append(history, openai.UserMessage(userMsg))
|
||||||
|
|
||||||
stream := b.client.Chat.Completions.NewStreaming(ctx, openai.ChatCompletionNewParams{
|
stream := b.client().Chat.Completions.NewStreaming(ctx, openai.ChatCompletionNewParams{
|
||||||
Model: b.cfg.Model,
|
Model: b.model,
|
||||||
Messages: history,
|
Messages: history,
|
||||||
})
|
})
|
||||||
answer := ""
|
answer := ""
|
||||||
|
|||||||
+185
-33
@@ -1,8 +1,11 @@
|
|||||||
package config
|
package config
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
|
||||||
"gopkg.in/yaml.v3"
|
"gopkg.in/yaml.v3"
|
||||||
)
|
)
|
||||||
@@ -12,14 +15,27 @@ const (
|
|||||||
configFile = "config.yaml"
|
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 {
|
type Config struct {
|
||||||
BotName string `yaml:"bot_name"`
|
BotName string `yaml:"bot_name"`
|
||||||
Port int `yaml:"port"`
|
Port int `yaml:"port"`
|
||||||
LogLevel string `yaml:"log_level"`
|
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"`
|
APIKey string `yaml:"api_key"`
|
||||||
BaseURL string `yaml:"base_url"`
|
BaseURL string `yaml:"base_url"`
|
||||||
Model string `yaml:"model"`
|
Model string `yaml:"model"`
|
||||||
SystemPrompt string `yaml:"system_prompt"`
|
|
||||||
}
|
}
|
||||||
|
|
||||||
var cfg *Config
|
var cfg *Config
|
||||||
@@ -45,54 +61,190 @@ func Load() (*Config, error) {
|
|||||||
if err := yaml.Unmarshal(data, cfg); err != nil {
|
if err := yaml.Unmarshal(data, cfg); err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
if len(cfg.Providers) == 0 {
|
||||||
|
if err := migrateLegacy(path, data); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
}
|
||||||
applyDefaults(cfg)
|
applyDefaults(cfg)
|
||||||
|
if err := validate(cfg); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
return cfg, nil
|
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 {
|
func GetConfig() *Config {
|
||||||
return cfg
|
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 {
|
func writeDefault(path string) error {
|
||||||
cfg = &Config{
|
cfg = &Config{
|
||||||
BotName: "ai-bot",
|
BotName: "ai-bot",
|
||||||
Port: 8080,
|
Port: 8080,
|
||||||
LogLevel: "info",
|
LogLevel: "info",
|
||||||
|
SystemPrompt: "你是一个乐于助人的 AI 助手。",
|
||||||
|
Providers: []Provider{
|
||||||
|
{
|
||||||
|
Name: "openai",
|
||||||
APIKey: "",
|
APIKey: "",
|
||||||
BaseURL: "https://api.openai.com/v1",
|
BaseURL: "https://api.openai.com/v1",
|
||||||
Model: "gpt-4o-mini",
|
Models: []string{"gpt-4o-mini", "gpt-4o"},
|
||||||
SystemPrompt: "你是一个乐于助人的 AI 助手。",
|
},
|
||||||
|
{
|
||||||
|
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 {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -17,12 +17,11 @@ func main() {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
log.Fatalf("加载配置失败: %v", err)
|
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)
|
b := bot.New(cfg)
|
||||||
|
provider, model := b.Current()
|
||||||
|
fmt.Printf("🤖 %s 已启动 (供应商: %s, 模型: %s)。输入问题开始对话,输入 /help 查看命令。\n",
|
||||||
|
cfg.BotName, provider, model)
|
||||||
|
|
||||||
scanner := bufio.NewScanner(os.Stdin)
|
scanner := bufio.NewScanner(os.Stdin)
|
||||||
for {
|
for {
|
||||||
fmt.Print("你: ")
|
fmt.Print("你: ")
|
||||||
@@ -33,10 +32,12 @@ func main() {
|
|||||||
if input == "" {
|
if input == "" {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
if input == "/exit" || input == "/quit" {
|
if strings.HasPrefix(input, "/") {
|
||||||
fmt.Println("再见!")
|
if !handleCommand(b, input) {
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
|
continue
|
||||||
|
}
|
||||||
answer, err := b.Chat(context.Background(), input)
|
answer, err := b.Chat(context.Background(), input)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
fmt.Printf("⚠️ %v\n", err)
|
fmt.Printf("⚠️ %v\n", err)
|
||||||
@@ -48,3 +49,40 @@ func main() {
|
|||||||
log.Fatal(err)
|
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
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user