From bb5bc619c9398c1829922f2b2368f8ae6241797e Mon Sep 17 00:00:00 2001 From: kevin Date: Fri, 14 Aug 2026 18:50:11 +0800 Subject: [PATCH] =?UTF-8?q?=E4=BC=9A=E8=AF=9D=E6=8C=81=E4=B9=85=E5=8C=96?= =?UTF-8?q?=E4=B8=8E=E6=81=A2=E5=A4=8D=E3=80=81=E6=A8=A1=E5=9E=8B=E7=BA=A7?= =?UTF-8?q?=E4=B8=8A=E4=B8=8B=E6=96=87=E7=AA=97=E5=8F=A3=E9=85=8D=E7=BD=AE?= =?UTF-8?q?=E4=B8=8E=E8=87=AA=E5=8A=A8=E8=8E=B7=E5=8F=96?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- internal/bot/bot.go | 59 +++++++++- internal/bot/session_test.go | 87 +++++++++++++++ internal/cli/cli.go | 58 +++++++++- internal/cli/complete.go | 2 +- internal/config/config.go | 192 +++++++++++++++++++++++++++------ internal/config/config_test.go | 8 +- internal/config/models_test.go | 149 +++++++++++++++++++++++++ internal/store/session.go | 142 ++++++++++++++++++++++++ internal/store/session_test.go | 99 +++++++++++++++++ main.go | 64 ++++++++++- 10 files changed, 820 insertions(+), 40 deletions(-) create mode 100644 internal/bot/session_test.go create mode 100644 internal/config/models_test.go create mode 100644 internal/store/session.go create mode 100644 internal/store/session_test.go diff --git a/internal/bot/bot.go b/internal/bot/bot.go index b597d49..643e0aa 100644 --- a/internal/bot/bot.go +++ b/internal/bot/bot.go @@ -13,6 +13,7 @@ import ( "github.com/openai/openai-go/shared/constant" "github.com/tidwall/gjson" "myaibot/internal/config" + "myaibot/internal/store" "myaibot/internal/tools" "myaibot/internal/tools/builtin" ) @@ -43,7 +44,7 @@ func New(cfg *config.Config) (*Bot, error) { systemPrompt: cfg.SystemPrompt, toolRegistry: tools.NewRegistry(builtin.NewTimeTool(), builtin.NewCalculatorTool(), builtin.NewRandomTool()), } - b.provider = config.FindProvider(cfg.DefaultProvider) + b.provider = config.FindProviderIn(cfg, cfg.DefaultProvider) b.model = cfg.DefaultModel b.toolProvider, b.toolModel = b.provider, b.model if cfg.ToolModel != "" { @@ -126,6 +127,62 @@ func (b *Bot) Tools() []string { return b.toolRegistry.List() } +func (b *Bot) ContextWindow() int64 { + if m := config.FindModel(b.provider, b.model); m != nil { + return m.ContextWindow + } + return 0 +} + +func (b *Bot) SessionMessages() []store.Message { + out := make([]store.Message, 0, len(b.history)+1) + out = append(out, store.Message{Role: "system", Content: b.systemPrompt}) + for _, msg := range b.history { + var role, content string + switch { + case msg.OfUser != nil: + role, content = "user", msg.OfUser.Content.OfString.Value + case msg.OfAssistant != nil: + role, content = "assistant", msg.OfAssistant.Content.OfString.Value + case msg.OfSystem != nil: + role, content = "system", msg.OfSystem.Content.OfString.Value + default: + continue + } + if content == "" { + continue + } + out = append(out, store.Message{Role: role, Content: content}) + } + return out +} + +func (b *Bot) RestoreSession(s *store.Session) { + if s == nil { + return + } + if s.SystemPrompt != "" { + b.systemPrompt = s.SystemPrompt + } + var history []openai.ChatCompletionMessageParamUnion + for _, m := range s.Messages { + switch m.Role { + case "user": + history = append(history, openai.UserMessage(m.Content)) + case "assistant": + history = append(history, openai.AssistantMessage(m.Content)) + case "system": + if b.systemPrompt == "" || m.Content != b.systemPrompt { + b.systemPrompt = m.Content + } + } + } + if len(history) > maxHistory { + history = history[len(history)-maxHistory:] + } + b.history = history +} + func (b *Bot) ContextDump() string { var sb strings.Builder sb.WriteString("[系统] " + b.systemPrompt + "\n") diff --git a/internal/bot/session_test.go b/internal/bot/session_test.go new file mode 100644 index 0000000..e5dda1a --- /dev/null +++ b/internal/bot/session_test.go @@ -0,0 +1,87 @@ +package bot + +import ( + "testing" + + "github.com/openai/openai-go" + + "myaibot/internal/config" + "myaibot/internal/store" +) + +func newTestBot(t *testing.T) *Bot { + t.Helper() + t.Chdir(t.TempDir()) + for _, name := range []string{"get_current_time", "calculate", "random_number"} { + if err := config.WriteDefaultToolConfig(name, map[string]any{"enabled": true, "prompt": "p"}); err != nil { + t.Fatalf("写入工具配置失败: %v", err) + } + } + cfg := &config.Config{ + BotName: "test", + SystemPrompt: "测试系统提示", + DefaultProvider: "p", + DefaultModel: "m", + Providers: []config.Provider{ + {Name: "p", BaseURL: "x", Models: []config.ModelConfig{{Name: "m"}}}, + }, + } + b, err := New(cfg) + if err != nil { + t.Fatalf("New 出错: %v", err) + } + return b +} + +func TestSessionRoundtrip(t *testing.T) { + b := newTestBot(t) + b.history = append(b.history, + openai.UserMessage("你好"), + openai.AssistantMessage("你好!有什么可以帮你?"), + ) + msgs := b.SessionMessages() + if len(msgs) != 3 { + t.Fatalf("SessionMessages 数量 = %d, want 3", len(msgs)) + } + if msgs[0].Role != "system" || msgs[0].Content != "测试系统提示" { + t.Errorf("首条应为系统提示: %+v", msgs[0]) + } + + restored := &Bot{} + restored.RestoreSession(&store.Session{Messages: msgs}) + if restored.systemPrompt != "测试系统提示" { + t.Errorf("systemPrompt 未恢复: %q", restored.systemPrompt) + } + if len(restored.history) != 2 { + t.Fatalf("history 数量 = %d, want 2", len(restored.history)) + } + if u := restored.history[0].OfUser; u == nil || u.Content.OfString.Value != "你好" { + t.Errorf("用户消息未还原: %+v", restored.history[0]) + } + if a := restored.history[1].OfAssistant; a == nil || a.Content.OfString.Value != "你好!有什么可以帮你?" { + t.Errorf("助手消息未还原: %+v", restored.history[1]) + } +} + +func TestRestoreSessionOverridesPrompt(t *testing.T) { + b := newTestBot(t) + b.RestoreSession(&store.Session{ + SystemPrompt: "会话覆盖的系统提示", + Messages: []store.Message{{Role: "user", Content: "hi"}}, + }) + if b.systemPrompt != "会话覆盖的系统提示" { + t.Errorf("systemPrompt 未覆盖: %q", b.systemPrompt) + } +} + +func TestRestoreTrimsHistory(t *testing.T) { + b := newTestBot(t) + var msgs []store.Message + for i := 0; i < maxHistory+10; i++ { + msgs = append(msgs, store.Message{Role: "user", Content: "m"}) + } + b.RestoreSession(&store.Session{Messages: msgs}) + if len(b.history) != maxHistory { + t.Errorf("history 应截断到 %d, got %d", maxHistory, len(b.history)) + } +} diff --git a/internal/cli/cli.go b/internal/cli/cli.go index 8744d81..cc01146 100644 --- a/internal/cli/cli.go +++ b/internal/cli/cli.go @@ -1,18 +1,35 @@ package cli import ( + "database/sql" "fmt" + "strconv" "strings" "myaibot/internal/bot" + "myaibot/internal/store" ) type Handler struct { bot *bot.Bot + db *sql.DB } -func New(b *bot.Bot) *Handler { - return &Handler{bot: b} +func New(b *bot.Bot, db *sql.DB) *Handler { + return &Handler{bot: b, db: db} +} + +func formatWindow(n int64) string { + switch { + case n <= 0: + return "未配置" + case n >= 1048576: + return fmt.Sprintf("%dM tokens", n/1048576) + case n >= 1024: + return fmt.Sprintf("%dK tokens", n/1024) + default: + return fmt.Sprintf("%d tokens", n) + } } func (h *Handler) Handle(input string) bool { @@ -30,6 +47,8 @@ func (h *Handler) Handle(input string) bool { fmt.Println(" /effort 设置思考强度") fmt.Println(" /context 打印当前聊天上下文") fmt.Println(" /tools 列出可用工具") + fmt.Println(" /sessions 列出历史会话") + fmt.Println(" /session 切换到历史会话,如 /session 3") fmt.Println(" /info 显示当前供应商、模型和思考配置") fmt.Println(" /exit 退出") case "/models": @@ -74,6 +93,40 @@ func (h *Handler) Handle(input string) bool { for _, t := range h.bot.Tools() { fmt.Println(" " + t) } + case "/sessions": + list, err := store.ListSessions(h.db) + if err != nil { + fmt.Printf("⚠️ %v\n", err) + return true + } + if len(list) == 0 { + fmt.Println("暂无历史会话") + return true + } + for _, s := range list { + fmt.Printf(" #%d %s (%d 条消息)\n", s.ID, s.CreatedAt.Format("2006-01-02 15:04:05"), s.MessageCount) + } + case "/session": + if len(args) == 0 { + fmt.Println("用法: /session ,如 /session 3") + return true + } + id, err := strconv.ParseInt(args[0], 10, 64) + if err != nil || id <= 0 { + fmt.Printf("无效的会话 id: %s\n", args[0]) + return true + } + sess, err := store.LoadSession(h.db, id) + if err != nil { + fmt.Printf("⚠️ %v\n", err) + return true + } + if sess == nil { + fmt.Printf("会话 #%d 不存在\n", id) + return true + } + h.bot.RestoreSession(sess) + fmt.Printf("已切换到会话 #%d (%d 条消息)\n", id, len(sess.Messages)) case "/info": provider, model := h.bot.Current() thinking, effort := h.bot.ThinkingConfig() @@ -85,6 +138,7 @@ func (h *Handler) Handle(input string) bool { effort = "high(默认)" } fmt.Printf("供应商: %s, 模型: %s, 思考模式: %s, 思考强度: %s\n", provider, model, thinking, effort) + fmt.Printf("上下文窗口: %s\n", formatWindow(h.bot.ContextWindow())) fmt.Printf("工具调用AI: %s\n图片识别AI: %s\n", tool, vision) default: fmt.Printf("未知命令: %s,输入 /help 查看命令列表\n", cmd) diff --git a/internal/cli/complete.go b/internal/cli/complete.go index b04f7f4..78238c1 100644 --- a/internal/cli/complete.go +++ b/internal/cli/complete.go @@ -2,7 +2,7 @@ package cli import "strings" -var commands = []string{"/exit", "/quit", "/help", "/models", "/use", "/think", "/effort", "/context", "/tools", "/info"} +var commands = []string{"/exit", "/quit", "/help", "/models", "/use", "/think", "/effort", "/context", "/tools", "/sessions", "/session", "/info"} func Complete(line string, models []string) []string { fields := strings.Fields(line) diff --git a/internal/config/config.go b/internal/config/config.go index 7a5831e..d4dbef3 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -1,12 +1,16 @@ package config import ( + "context" "errors" "fmt" "os" "path/filepath" + "sort" "strings" + "github.com/openai/openai-go" + "github.com/openai/openai-go/option" "gopkg.in/yaml.v3" ) @@ -16,13 +20,43 @@ const ( toolConfigDir = "data/tools" ) +type ModelConfig struct { + Name string `yaml:"name"` + ContextWindow int64 `yaml:"context_window"` +} + +// UnmarshalYAML 兼容两种格式: +// +// models: [deepseek-v4-flash, deepseek-v4-pro] # 字符串列表(旧格式) +// models: +// - name: deepseek-v4-flash +// context_window: 1048576 # 对象列表 +func (m *ModelConfig) UnmarshalYAML(node *yaml.Node) error { + switch node.Kind { + case yaml.ScalarNode: + m.Name = node.Value + return nil + case yaml.MappingNode: + type raw ModelConfig + var r raw + if err := node.Decode(&r); err != nil { + return err + } + *m = ModelConfig(r) + return nil + default: + return fmt.Errorf("模型配置必须是字符串或对象") + } +} + type Provider struct { - Name string `yaml:"name"` - APIKey string `yaml:"api_key"` - BaseURL string `yaml:"base_url"` - Models []string `yaml:"models"` - Thinking string `yaml:"thinking"` - ReasoningEffort string `yaml:"reasoning_effort"` + Name string `yaml:"name"` + APIKey string `yaml:"api_key"` + BaseURL string `yaml:"base_url"` + Models []ModelConfig `yaml:"models"` + AutoFetchModels bool `yaml:"auto_fetch_models"` + Thinking string `yaml:"thinking"` + ReasoningEffort string `yaml:"reasoning_effort"` } type Config struct { @@ -110,17 +144,17 @@ func migrateLegacy(path string, data []byte) error { Name: "openai", APIKey: legacy.APIKey, BaseURL: legacy.BaseURL, - Models: []string{legacy.Model}, + Models: []ModelConfig{{Name: 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"} + if len(p.Models) == 0 || p.Models[0].Name == "" { + p.Models = []ModelConfig{{Name: "gpt-4o-mini"}} } cfg.Providers = []Provider{p} cfg.DefaultProvider = p.Name - cfg.DefaultModel = p.Models[0] + cfg.DefaultModel = p.Models[0].Name return writeFile(path, cfg) } @@ -183,13 +217,16 @@ func validate(c *Config) error { if p.BaseURL == "" { return fmt.Errorf("供应商 %s 缺少 base_url", p.Name) } - if len(p.Models) == 0 { + if len(p.Models) == 0 && !p.AutoFetchModels { return fmt.Errorf("供应商 %s 未配置 models", p.Name) } for _, m := range p.Models { - if m == "" { + if m.Name == "" { return fmt.Errorf("供应商 %s 包含空模型名", p.Name) } + if m.ContextWindow < 0 { + return fmt.Errorf("供应商 %s 的模型 %s context_window 无效: %d(不能为负数)", p.Name, m.Name, m.ContextWindow) + } } if p.Thinking != "" && !contains([]string{"enabled", "disabled"}, p.Thinking) { return fmt.Errorf("供应商 %s 的 thinking 无效: %q(可选 enabled/disabled)", p.Name, p.Thinking) @@ -201,17 +238,17 @@ func validate(c *Config) error { 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) + if err := validateModelRef("default_model", c.DefaultModel, c); err != nil { + return err } if c.ToolModel != "" { - if _, _, err := ResolveModel(c.ToolModel); err != nil { - return fmt.Errorf("tool_model 无效: %w", err) + if err := validateModelRef("tool_model", c.ToolModel, c); err != nil { + return err } } if c.VisionModel != "" { - if _, _, err := ResolveModel(c.VisionModel); err != nil { - return fmt.Errorf("vision_model 无效: %w", err) + if err := validateModelRef("vision_model", c.VisionModel, c); err != nil { + return err } } d := c.Database @@ -224,36 +261,69 @@ func validate(c *Config) error { return nil } +// validateModelRef 校验模型引用;若引用指向启用了 auto_fetch_models 的供应商, +// 则跳过存在性校验(模型列表将在启动时从 API 拉取)。 +func validateModelRef(field, id string, c *Config) error { + if _, _, err := ResolveModelIn(c, id); err == nil { + return nil + } + providerName, _, hasProvider := strings.Cut(id, "/") + if !hasProvider { + providerName = c.DefaultProvider + } + if p := FindProviderIn(c, providerName); p != nil && p.AutoFetchModels { + return nil + } + return fmt.Errorf("%s 无效: 模型 %q 不存在", field, id) +} + func FindProvider(name string) *Provider { - for i := range cfg.Providers { - if cfg.Providers[i].Name == name { - return &cfg.Providers[i] + return FindProviderIn(cfg, name) +} + +func FindProviderIn(c *Config, name string) *Provider { + for i := range c.Providers { + if c.Providers[i].Name == name { + return &c.Providers[i] + } + } + return nil +} + +func FindModel(p *Provider, name string) *ModelConfig { + for i := range p.Models { + if p.Models[i].Name == name { + return &p.Models[i] } } return nil } func ResolveModel(id string) (*Provider, string, error) { + return ResolveModelIn(cfg, id) +} + +func ResolveModelIn(c *Config, id string) (*Provider, string, error) { if id == "" { - id = cfg.DefaultModel + id = c.DefaultModel } if providerName, modelName, ok := strings.Cut(id, "/"); ok { - p := FindProvider(providerName) + p := FindProviderIn(c, providerName) if p == nil { return nil, "", fmt.Errorf("供应商 %q 不存在", providerName) } - if !contains(p.Models, modelName) { + if FindModel(p, modelName) == nil { 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) { + for i := range c.Providers { + if FindModel(&c.Providers[i], id) != nil { if found != nil { return nil, "", fmt.Errorf("模型 %q 在多个供应商中存在,请使用 provider/model 格式指定", id) } - found = &cfg.Providers[i] + found = &c.Providers[i] } } if found == nil { @@ -267,7 +337,7 @@ func AllModels() []string { for i := range cfg.Providers { p := &cfg.Providers[i] for _, m := range p.Models { - out = append(out, p.Name+"/"+m) + out = append(out, p.Name+"/"+m.Name) } } return out @@ -293,13 +363,19 @@ func writeDefault(path string) error { Name: "openai", APIKey: "", BaseURL: "https://api.openai.com/v1", - Models: []string{"gpt-4o-mini", "gpt-4o"}, + Models: []ModelConfig{ + {Name: "gpt-4o-mini", ContextWindow: 128000}, + {Name: "gpt-4o", ContextWindow: 128000}, + }, }, { Name: "deepseek", APIKey: "", BaseURL: "https://api.deepseek.com/v1", - Models: []string{"deepseek-chat", "deepseek-reasoner"}, + Models: []ModelConfig{ + {Name: "deepseek-v4-flash", ContextWindow: 1048576}, + {Name: "deepseek-v4-pro", ContextWindow: 1048576}, + }, }, }, DefaultProvider: "openai", @@ -359,3 +435,59 @@ func WriteDefaultToolConfig(name string, defaults map[string]any) error { func ToolConfigPath(name string) string { return filepath.Join(toolConfigDir, name+".yaml") } + +// Save 将配置写回 data/config.yaml。 +func Save(c *Config) error { + path := filepath.Join(configDir, configFile) + return writeFile(path, c) +} + +// ModelsEqual 比较两个模型的名称与上下文窗口是否完全一致。 +func ModelsEqual(a, b []ModelConfig) bool { + if len(a) != len(b) { + return false + } + for i := range a { + if a[i].Name != b[i].Name || a[i].ContextWindow != b[i].ContextWindow { + return false + } + } + return true +} + +// FetchModels 从供应商 API 拉取模型列表并合并进 p.Models。 +// 已配置模型的 context_window 保留,新模型默认 0。 +func FetchModels(ctx context.Context, p *Provider) error { + if p.APIKey == "" { + return errors.New("未配置 api_key") + } + client := openai.NewClient( + option.WithAPIKey(p.APIKey), + option.WithBaseURL(p.BaseURL), + ) + page, err := client.Models.List(ctx) + if err != nil { + return fmt.Errorf("请求模型列表失败: %w", err) + } + ids := make([]string, 0, len(page.Data)) + seen := make(map[string]bool, len(page.Data)) + for _, m := range page.Data { + if m.ID == "" || seen[m.ID] { + continue + } + seen[m.ID] = true + ids = append(ids, m.ID) + } + sort.Strings(ids) + + existing := make(map[string]int64, len(p.Models)) + for _, m := range p.Models { + existing[m.Name] = m.ContextWindow + } + merged := make([]ModelConfig, 0, len(ids)) + for _, id := range ids { + merged = append(merged, ModelConfig{Name: id, ContextWindow: existing[id]}) + } + p.Models = merged + return nil +} diff --git a/internal/config/config_test.go b/internal/config/config_test.go index a7b88bf..24b8977 100644 --- a/internal/config/config_test.go +++ b/internal/config/config_test.go @@ -8,7 +8,7 @@ import ( ) func TestApplyDatabaseDefaults(t *testing.T) { - c := &Config{Providers: []Provider{{Name: "p", BaseURL: "x", Models: []string{"m"}}}} + c := &Config{Providers: []Provider{{Name: "p", BaseURL: "x", Models: []ModelConfig{{Name: "m"}}}}} changed := applyDefaults(c) if !changed { t.Error("缺失字段应返回 changed=true") @@ -28,7 +28,7 @@ func TestApplyDefaultsNoChange(t *testing.T) { LogLevel: "debug", SystemPrompt: "sp", DefaultProvider: "p", - Providers: []Provider{{Name: "p", BaseURL: "x", Models: []string{"m"}}}, + Providers: []Provider{{Name: "p", BaseURL: "x", Models: []ModelConfig{{Name: "m"}}}}, Database: DatabaseConfig{Driver: "mysql", File: "f", Host: "h", Port: 3307, Name: "n"}, } if applyDefaults(c) { @@ -38,7 +38,7 @@ func TestApplyDefaultsNoChange(t *testing.T) { func TestApplyMySQLDefaults(t *testing.T) { c := &Config{ - Providers: []Provider{{Name: "p", BaseURL: "x", Models: []string{"m"}}}, + Providers: []Provider{{Name: "p", BaseURL: "x", Models: []ModelConfig{{Name: "m"}}}}, Database: DatabaseConfig{Driver: "mysql", Name: "memory"}, } changed := applyDefaults(c) @@ -117,7 +117,7 @@ func TestValidateDatabase(t *testing.T) { c := &Config{ DefaultProvider: "p", DefaultModel: "m", - Providers: []Provider{{Name: "p", BaseURL: "x", Models: []string{"m"}}}, + Providers: []Provider{{Name: "p", BaseURL: "x", Models: []ModelConfig{{Name: "m"}}}}, Database: DatabaseConfig{Driver: "oracle"}, } cfg = c diff --git a/internal/config/models_test.go b/internal/config/models_test.go new file mode 100644 index 0000000..6ce26e1 --- /dev/null +++ b/internal/config/models_test.go @@ -0,0 +1,149 @@ +package config + +import ( + "context" + "net/http" + "net/http/httptest" + "testing" + + "gopkg.in/yaml.v3" +) + +func TestModelConfigUnmarshalStringList(t *testing.T) { + var p struct { + Models []ModelConfig `yaml:"models"` + } + err := yaml.Unmarshal([]byte("models:\n - deepseek-v4-flash\n - deepseek-v4-pro\n"), &p) + if err != nil { + t.Fatalf("解析失败: %v", err) + } + if len(p.Models) != 2 || p.Models[0].Name != "deepseek-v4-flash" || p.Models[1].Name != "deepseek-v4-pro" { + t.Errorf("字符串列表解析异常: %+v", p.Models) + } +} + +func TestModelConfigUnmarshalObjectList(t *testing.T) { + var p struct { + Models []ModelConfig `yaml:"models"` + } + err := yaml.Unmarshal([]byte("models:\n - name: deepseek-v4-flash\n context_window: 1048576\n - name: deepseek-v4-pro\n"), &p) + if err != nil { + t.Fatalf("解析失败: %v", err) + } + if len(p.Models) != 2 { + t.Fatalf("数量 = %d, want 2", len(p.Models)) + } + if p.Models[0].Name != "deepseek-v4-flash" || p.Models[0].ContextWindow != 1048576 { + t.Errorf("对象列表解析异常: %+v", p.Models[0]) + } + if p.Models[1].Name != "deepseek-v4-pro" || p.Models[1].ContextWindow != 0 { + t.Errorf("缺省 context_window 应为 0: %+v", p.Models[1]) + } +} + +func TestValidateContextWindow(t *testing.T) { + base := &Config{ + DefaultProvider: "p", + DefaultModel: "m", + Database: DatabaseConfig{Driver: "sqlite3"}, + Providers: []Provider{{ + Name: "p", BaseURL: "x", + Models: []ModelConfig{{Name: "m", ContextWindow: -1}}, + }}, + } + cfg = base + if err := validate(base); err == nil { + t.Error("负 context_window 应报错") + } + base.Providers[0].Models[0].ContextWindow = 0 + if err := validate(base); err != nil { + t.Errorf("context_window 0 不应报错: %v", err) + } +} + +func TestValidateAutoFetch(t *testing.T) { + c := &Config{ + DefaultProvider: "p", + DefaultModel: "m", + Database: DatabaseConfig{Driver: "sqlite3"}, + Providers: []Provider{{ + Name: "p", BaseURL: "x", + AutoFetchModels: true, + }}, + } + cfg = c + if err := validate(c); err != nil { + t.Errorf("auto_fetch 空 models 不应报错: %v", err) + } + c.Providers[0].AutoFetchModels = false + if err := validate(c); err == nil { + t.Error("非 auto_fetch 空 models 应报错") + } +} + +func TestValidateAutoFetchDefaultModel(t *testing.T) { + c := &Config{ + DefaultProvider: "p", + DefaultModel: "future-model", + Database: DatabaseConfig{Driver: "sqlite3"}, + Providers: []Provider{{ + Name: "p", BaseURL: "x", + AutoFetchModels: true, + }}, + } + cfg = c + if err := validate(c); err != nil { + t.Errorf("auto_fetch 时 default_model 应跳过存在性校验: %v", err) + } +} + +func TestFetchModels(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/models" { + t.Errorf("请求路径 = %q, want /models", r.URL.Path) + } + w.Header().Set("Content-Type", "application/json") + w.Write([]byte(`{"object":"list","data":[ + {"id":"deepseek-v4-pro","object":"model","owned_by":"deepseek"}, + {"id":"deepseek-v4-flash","object":"model","owned_by":"deepseek"} + ]}`)) + })) + defer srv.Close() + + p := &Provider{ + Name: "deepseek", + APIKey: "sk-test", + BaseURL: srv.URL, + Models: []ModelConfig{{Name: "deepseek-v4-flash", ContextWindow: 1048576}}, + } + if err := FetchModels(context.Background(), p); err != nil { + t.Fatalf("FetchModels 出错: %v", err) + } + if len(p.Models) != 2 { + t.Fatalf("模型数量 = %d, want 2", len(p.Models)) + } + if p.Models[0].Name != "deepseek-v4-flash" || p.Models[0].ContextWindow != 1048576 { + t.Errorf("已有模型应保留 context_window: %+v", p.Models[0]) + } + if p.Models[1].Name != "deepseek-v4-pro" || p.Models[1].ContextWindow != 0 { + t.Errorf("新模型 context_window 应为 0: %+v", p.Models[1]) + } +} + +func TestFetchModelsNoAPIKey(t *testing.T) { + if err := FetchModels(context.Background(), &Provider{Name: "p"}); err == nil { + t.Error("无 api_key 应报错") + } +} + +func TestModelsEqual(t *testing.T) { + a := []ModelConfig{{Name: "x", ContextWindow: 1}} + b := []ModelConfig{{Name: "x", ContextWindow: 1}} + if !ModelsEqual(a, b) { + t.Error("相同列表应相等") + } + b[0].ContextWindow = 2 + if ModelsEqual(a, b) { + t.Error("不同列表应不相等") + } +} diff --git a/internal/store/session.go b/internal/store/session.go new file mode 100644 index 0000000..fbdeab1 --- /dev/null +++ b/internal/store/session.go @@ -0,0 +1,142 @@ +package store + +import ( + "database/sql" + "encoding/json" + "fmt" + "time" +) + +type Message struct { + Role string `json:"role"` + Content string `json:"content"` +} + +type Session struct { + ID int64 + CreatedAt time.Time + Provider string + Model string + SystemPrompt string + Messages []Message +} + +type SessionSummary struct { + ID int64 + CreatedAt time.Time + MessageCount int +} + +const createSessionsSQLite = ` +CREATE TABLE IF NOT EXISTS sessions ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP, + provider TEXT NOT NULL DEFAULT '', + model TEXT NOT NULL DEFAULT '', + system_prompt TEXT NOT NULL DEFAULT '', + messages TEXT NOT NULL, + message_count INTEGER NOT NULL DEFAULT 0 +)` + +const createSessionsMySQL = ` +CREATE TABLE IF NOT EXISTS sessions ( + id BIGINT AUTO_INCREMENT PRIMARY KEY, + created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP, + provider VARCHAR(255) NOT NULL DEFAULT '', + model VARCHAR(255) NOT NULL DEFAULT '', + system_prompt TEXT NOT NULL, + messages LONGTEXT NOT NULL, + message_count INT NOT NULL DEFAULT 0 +)` + +func Migrate(db *sql.DB, driver string) error { + var ddl string + switch driver { + case "sqlite3": + ddl = createSessionsSQLite + case "mysql": + ddl = createSessionsMySQL + default: + return fmt.Errorf("不支持的数据库驱动: %s", driver) + } + if _, err := db.Exec(ddl); err != nil { + return fmt.Errorf("创建 sessions 表失败: %w", err) + } + return nil +} + +func SaveSession(db *sql.DB, s *Session) (int64, error) { + if s.Messages == nil { + s.Messages = []Message{} + } + data, err := json.Marshal(s.Messages) + if err != nil { + return 0, fmt.Errorf("序列化消息失败: %w", err) + } + res, err := db.Exec( + "INSERT INTO sessions (provider, model, system_prompt, messages, message_count) VALUES (?, ?, ?, ?, ?)", + s.Provider, s.Model, s.SystemPrompt, string(data), len(s.Messages), + ) + if err != nil { + return 0, fmt.Errorf("保存会话失败: %w", err) + } + return res.LastInsertId() +} + +func LoadLatestSession(db *sql.DB) (*Session, error) { + return loadSession(db, "SELECT id, created_at, provider, model, system_prompt, messages FROM sessions ORDER BY id DESC LIMIT 1") +} + +func LoadSession(db *sql.DB, id int64) (*Session, error) { + return loadSession(db, "SELECT id, created_at, provider, model, system_prompt, messages FROM sessions WHERE id = ?", id) +} + +func loadSession(db *sql.DB, query string, args ...any) (*Session, error) { + row := db.QueryRow(query, args...) + var ( + s Session + created string + messages string + ) + if err := row.Scan(&s.ID, &created, &s.Provider, &s.Model, &s.SystemPrompt, &messages); err != nil { + if err == sql.ErrNoRows { + return nil, nil + } + return nil, fmt.Errorf("读取会话失败: %w", err) + } + s.CreatedAt = parseTime(created) + if err := json.Unmarshal([]byte(messages), &s.Messages); err != nil { + return nil, fmt.Errorf("解析会话消息失败: %w", err) + } + return &s, nil +} + +func ListSessions(db *sql.DB) ([]SessionSummary, error) { + rows, err := db.Query("SELECT id, created_at, message_count FROM sessions ORDER BY id DESC") + if err != nil { + return nil, fmt.Errorf("查询会话列表失败: %w", err) + } + defer rows.Close() + var out []SessionSummary + for rows.Next() { + var ( + sm SessionSummary + t string + ) + if err := rows.Scan(&sm.ID, &t, &sm.MessageCount); err != nil { + return nil, fmt.Errorf("读取会话列表失败: %w", err) + } + sm.CreatedAt = parseTime(t) + out = append(out, sm) + } + return out, rows.Err() +} + +func parseTime(s string) time.Time { + for _, layout := range []string{time.RFC3339, "2006-01-02 15:04:05"} { + if t, err := time.ParseInLocation(layout, s, time.Local); err == nil { + return t + } + } + return time.Time{} +} diff --git a/internal/store/session_test.go b/internal/store/session_test.go new file mode 100644 index 0000000..58b7064 --- /dev/null +++ b/internal/store/session_test.go @@ -0,0 +1,99 @@ +package store + +import ( + "database/sql" + "path/filepath" + "testing" + + "myaibot/internal/config" +) + +func openTestDB(t *testing.T) *sql.DB { + t.Helper() + cfg := &config.DatabaseConfig{ + Driver: "sqlite3", + File: filepath.Join(t.TempDir(), "memory.db"), + } + db, err := Open(cfg) + if err != nil { + t.Fatalf("Open 出错: %v", err) + } + t.Cleanup(func() { Close(db) }) + if err := Migrate(db, "sqlite3"); err != nil { + t.Fatalf("Migrate 出错: %v", err) + } + return db +} + +func TestSessionRoundtrip(t *testing.T) { + db := openTestDB(t) + first := &Session{ + Provider: "deepseek", + Model: "deepseek-v4-flash", + Messages: []Message{{Role: "user", Content: "你好"}, {Role: "assistant", Content: "你好!"}}, + } + id1, err := SaveSession(db, first) + if err != nil { + t.Fatalf("SaveSession 出错: %v", err) + } + second := &Session{ + Provider: "deepseek", + Model: "deepseek-v4-flash", + Messages: []Message{{Role: "user", Content: "现在几点"}}, + } + id2, err := SaveSession(db, second) + if err != nil { + t.Fatalf("SaveSession 出错: %v", err) + } + if id2 <= id1 { + t.Errorf("id2 (%d) 应大于 id1 (%d)", id2, id1) + } + + latest, err := LoadLatestSession(db) + if err != nil { + t.Fatalf("LoadLatestSession 出错: %v", err) + } + if latest == nil || latest.ID != id2 { + t.Errorf("最新会话应为 #%d, got %+v", id2, latest) + } + if len(latest.Messages) != 1 || latest.Messages[0].Content != "现在几点" { + t.Errorf("消息还原异常: %+v", latest.Messages) + } + + byID, err := LoadSession(db, id1) + if err != nil { + t.Fatalf("LoadSession 出错: %v", err) + } + if byID == nil || len(byID.Messages) != 2 { + t.Errorf("按 id 加载异常: %+v", byID) + } + + list, err := ListSessions(db) + if err != nil { + t.Fatalf("ListSessions 出错: %v", err) + } + if len(list) != 2 { + t.Errorf("列表数量 = %d, want 2", len(list)) + } + if list[0].ID != id2 || list[0].MessageCount != 1 { + t.Errorf("列表首条应为最新会话: %+v", list[0]) + } +} + +func TestLoadSessionMissing(t *testing.T) { + db := openTestDB(t) + sess, err := LoadSession(db, 999) + if err != nil { + t.Fatalf("LoadSession 出错: %v", err) + } + if sess != nil { + t.Errorf("不存在的会话应返回 nil, got %+v", sess) + } + latest, err := LoadLatestSession(db) + if err != nil { + t.Fatalf("LoadLatestSession 出错: %v", err) + } + if latest != nil { + t.Errorf("空库最新会话应为 nil, got %+v", latest) + } +} diff --git a/main.go b/main.go index 12648d3..5711298 100644 --- a/main.go +++ b/main.go @@ -2,6 +2,7 @@ package main import ( "context" + "database/sql" "errors" "fmt" "io" @@ -22,19 +23,32 @@ func main() { if err != nil { log.Fatalf("加载配置失败: %v", err) } + autoFetchModels(cfg) + db, err := store.Open(&cfg.Database) if err != nil { log.Fatalf("数据库连接失败: %v", err) } - defer store.Close(db) fmt.Printf("💾 数据库已连接 (%s)\n", cfg.Database.Driver) b, err := bot.New(cfg) if err != nil { fmt.Printf("⚠️ %v\n", err) fmt.Println("请填写工具配置文件后重新启动。") + store.Close(db) os.Exit(1) } + + if err := store.Migrate(db, cfg.Database.Driver); err != nil { + log.Fatalf("数据库迁移失败: %v", err) + } + if sess, err := store.LoadLatestSession(db); err != nil { + fmt.Printf("⚠️ 读取上次会话失败: %v\n", err) + } else if sess != nil && len(sess.Messages) > 0 { + b.RestoreSession(sess) + fmt.Printf("💬 已恢复上次会话 (%d 条消息)\n", len(sess.Messages)) + } + provider, model := b.Current() fmt.Printf("🤖 %s 已启动 (供应商: %s, 模型: %s)。输入问题开始对话,输入 /help 查看命令。\n", cfg.BotName, provider, model) @@ -46,7 +60,7 @@ func main() { return cli.Complete(s, b.Models()) }) - h := cli.New(b) + h := cli.New(b, db) for { input, err := line.Prompt("你: ") if errors.Is(err, io.EOF) || errors.Is(err, liner.ErrPromptAborted) { @@ -102,4 +116,50 @@ func main() { continue } } + + saveSession(db, b) + store.Close(db) +} + +func autoFetchModels(cfg *config.Config) { + for i := range cfg.Providers { + p := &cfg.Providers[i] + if !p.AutoFetchModels { + continue + } + if p.APIKey == "" { + fmt.Printf("⚠️ 供应商 %s 启用了 auto_fetch_models 但未配置 api_key,跳过自动获取\n", p.Name) + continue + } + before := append([]config.ModelConfig(nil), p.Models...) + if err := config.FetchModels(context.Background(), p); err != nil { + fmt.Printf("⚠️ 自动获取模型失败 (供应商 %s): %v,使用现有模型列表\n", p.Name, err) + continue + } + fmt.Printf("📚 已从 API 获取模型列表 (供应商 %s): %d 个模型\n", p.Name, len(p.Models)) + if !config.ModelsEqual(before, p.Models) { + if err := config.Save(cfg); err != nil { + fmt.Printf("⚠️ 模型列表写回配置失败: %v\n", err) + } + } + } +} + +func saveSession(db *sql.DB, b *bot.Bot) { + msgs := b.SessionMessages() + if len(msgs) == 0 { + return + } + provider, model := b.Current() + sess := &store.Session{ + Provider: provider, + Model: model, + SystemPrompt: "", + Messages: msgs, + } + if _, err := store.SaveSession(db, sess); err != nil { + fmt.Printf("⚠️ 保存会话失败: %v\n", err) + return + } + fmt.Printf("💾 会话已保存 (%d 条消息)\n", len(msgs)) }