会话持久化与恢复、模型级上下文窗口配置与自动获取
This commit is contained in:
+58
-1
@@ -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")
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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("不同列表应不相等")
|
||||
}
|
||||
}
|
||||
@@ -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{}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user