会话持久化与恢复、模型级上下文窗口配置与自动获取
This commit is contained in:
@@ -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