会话持久化与恢复、模型级上下文窗口配置与自动获取

This commit is contained in:
2026-08-14 18:50:11 +08:00
parent a4f07cc546
commit bb5bc619c9
10 changed files with 820 additions and 40 deletions
+58 -1
View File
@@ -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")
+87
View File
@@ -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))
}
}
+56 -2
View File
@@ -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 <low|high|max> 设置思考强度")
fmt.Println(" /context 打印当前聊天上下文")
fmt.Println(" /tools 列出可用工具")
fmt.Println(" /sessions 列出历史会话")
fmt.Println(" /session <id> 切换到历史会话,如 /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 <id>,如 /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)
+1 -1
View File
@@ -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)
+162 -30
View File
@@ -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
}
+4 -4
View File
@@ -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
+149
View File
@@ -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("不同列表应不相等")
}
}
+142
View File
@@ -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{}
}
+99
View File
@@ -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)
}
}
+62 -2
View File
@@ -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))
}