tokens 表增加 token_text 字段,旧库自动补列并回填符号文本

This commit is contained in:
2026-08-14 20:02:11 +08:00
parent b25dfb9d55
commit 378a6e0d41
5 changed files with 193 additions and 18 deletions
+98 -5
View File
@@ -5,6 +5,8 @@ import (
"fmt"
"strings"
"time"
"myaibot/internal/tokens"
)
type Memory struct {
@@ -41,12 +43,14 @@ const createMemoriesIndex = "CREATE INDEX IF NOT EXISTS idx_memories_created_at
const createTokensSQLite = `
CREATE TABLE IF NOT EXISTS tokens (
token_id INTEGER PRIMARY KEY,
token_text TEXT NOT NULL DEFAULT '',
created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP
)`
const createTokensMySQL = `
CREATE TABLE IF NOT EXISTS tokens (
token_id BIGINT PRIMARY KEY,
token_text TEXT NOT NULL,
created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP
)`
@@ -66,6 +70,92 @@ CREATE TABLE IF NOT EXISTS memory_tokens (
const createMemoryTokensIndex = "CREATE INDEX IF NOT EXISTS idx_memory_tokens_token_id ON memory_tokens (token_id)"
// migrateTokens 创建 tokens 表;旧表缺失 token_text 列时补列并回填文本。
func migrateTokens(db *sql.DB, driver string) error {
switch driver {
case "sqlite3":
if _, err := db.Exec(createTokensSQLite); err != nil {
return fmt.Errorf("创建 tokens 表失败: %w", err)
}
case "mysql":
if _, err := db.Exec(createTokensMySQL); err != nil {
return fmt.Errorf("创建 tokens 表失败: %w", err)
}
}
has, err := hasTokenTextColumn(db, driver)
if err != nil {
return err
}
if !has {
ddl := "ALTER TABLE tokens ADD COLUMN token_text TEXT NOT NULL DEFAULT ''"
if driver == "mysql" {
ddl = "ALTER TABLE tokens ADD COLUMN token_text TEXT NOT NULL"
}
if _, err := db.Exec(ddl); err != nil {
return fmt.Errorf("tokens 表增加 token_text 列失败: %w", err)
}
}
return backfillTokenTexts(db)
}
func hasTokenTextColumn(db *sql.DB, driver string) (bool, error) {
var query string
switch driver {
case "sqlite3":
query = "SELECT name FROM pragma_table_info('tokens')"
case "mysql":
query = "SELECT column_name FROM information_schema.COLUMNS WHERE TABLE_SCHEMA = DATABASE() AND TABLE_NAME = 'tokens' AND COLUMN_NAME = 'token_text'"
default:
return false, fmt.Errorf("不支持的数据库驱动: %s", driver)
}
rows, err := db.Query(query)
if err != nil {
return false, fmt.Errorf("检查 tokens 表结构失败: %w", err)
}
defer rows.Close()
for rows.Next() {
var name string
if err := rows.Scan(&name); err != nil {
return false, fmt.Errorf("读取 tokens 表结构失败: %w", err)
}
if name == "token_text" {
return true, nil
}
}
return false, rows.Err()
}
// backfillTokenTexts 为 token_text 为空的 token 行回填符号文本。
func backfillTokenTexts(db *sql.DB) error {
rows, err := db.Query("SELECT token_id FROM tokens WHERE token_text = ''")
if err != nil {
return fmt.Errorf("查询待回填 token 失败: %w", err)
}
var ids []int64
for rows.Next() {
var id int64
if err := rows.Scan(&id); err != nil {
rows.Close()
return fmt.Errorf("读取 token id 失败: %w", err)
}
ids = append(ids, id)
}
rows.Close()
if err := rows.Err(); err != nil {
return err
}
for _, id := range ids {
text := tokens.Text(id)
if text == "" {
continue
}
if _, err := db.Exec("UPDATE tokens SET token_text = ? WHERE token_id = ? AND token_text = ''", text, id); err != nil {
return fmt.Errorf("回填 token 文本失败: %w", err)
}
}
return nil
}
func migrateMemories(db *sql.DB, driver string) error {
var ddl string
switch driver {
@@ -91,6 +181,9 @@ func migrateMemories(db *sql.DB, driver string) error {
if _, err := db.Exec(ddl); err != nil {
return fmt.Errorf("创建 tokens 表失败: %w", err)
}
if err := migrateTokens(db, driver); err != nil {
return err
}
switch driver {
case "sqlite3":
ddl = createMemoryTokensSQLite
@@ -130,8 +223,8 @@ func SaveMemories(db *sql.DB, memories []Memory) ([]int64, error) {
}
// SaveMemoryTokens 为记忆建立 token 索引:token 去重、关联幂等。
func SaveMemoryTokens(db *sql.DB, memoryID int64, tokenIDs []int64) error {
if len(tokenIDs) == 0 {
func SaveMemoryTokens(db *sql.DB, memoryID int64, tokenList []tokens.Token) error {
if len(tokenList) == 0 {
return nil
}
tx, err := db.Begin()
@@ -139,11 +232,11 @@ func SaveMemoryTokens(db *sql.DB, memoryID int64, tokenIDs []int64) error {
return fmt.Errorf("开启事务失败: %w", err)
}
defer tx.Rollback()
for _, tid := range tokenIDs {
if _, err := tx.Exec("INSERT OR IGNORE INTO tokens (token_id) VALUES (?)", tid); err != nil {
for _, t := range tokenList {
if _, err := tx.Exec("INSERT OR IGNORE INTO tokens (token_id, token_text) VALUES (?, ?)", t.ID, t.Text); err != nil {
return fmt.Errorf("保存 token 失败: %w", err)
}
if _, err := tx.Exec("INSERT OR IGNORE INTO memory_tokens (memory_id, token_id) VALUES (?, ?)", memoryID, tid); err != nil {
if _, err := tx.Exec("INSERT OR IGNORE INTO memory_tokens (memory_id, token_id) VALUES (?, ?)", memoryID, t.ID); err != nil {
return fmt.Errorf("建立 token 关联失败: %w", err)
}
}
+54 -5
View File
@@ -6,6 +6,7 @@ import (
"testing"
"myaibot/internal/config"
"myaibot/internal/tokens"
)
func openMemDB(t *testing.T) *sql.DB {
@@ -80,8 +81,9 @@ func TestMemoryTokens(t *testing.T) {
t.Fatalf("SaveMemories 出错: %v", err)
}
// 同一记忆建索引两次应幂等
tok := []tokens.Token{{ID: 100, Text: " 100"}, {ID: 200, Text: " 200"}}
for i := 0; i < 2; i++ {
if err := SaveMemoryTokens(db, ids[0], []int64{100, 200, 100}); err != nil {
if err := SaveMemoryTokens(db, ids[0], tok); err != nil {
t.Fatalf("SaveMemoryTokens 出错: %v", err)
}
}
@@ -92,6 +94,13 @@ func TestMemoryTokens(t *testing.T) {
if tokenCount != 2 {
t.Errorf("token 应去重为 2, got %d", tokenCount)
}
var text string
if err := db.QueryRow("SELECT token_text FROM tokens WHERE token_id = 100").Scan(&text); err != nil {
t.Fatal(err)
}
if text != " 100" {
t.Errorf("token_text 应为 \" 100\", got %q", text)
}
var linkCount int64
if err := db.QueryRow("SELECT COUNT(*) FROM memory_tokens").Scan(&linkCount); err != nil {
t.Fatal(err)
@@ -101,7 +110,7 @@ func TestMemoryTokens(t *testing.T) {
}
// 另一条记忆建索引
if err := SaveMemoryTokens(db, ids[1], []int64{200, 300}); err != nil {
if err := SaveMemoryTokens(db, ids[1], []tokens.Token{{ID: 200, Text: " 200"}, {ID: 300, Text: " 300"}}); err != nil {
t.Fatalf("SaveMemoryTokens 出错: %v", err)
}
@@ -128,6 +137,46 @@ func TestMemoryTokens(t *testing.T) {
}
}
func TestMigrateLegacyTokensTable(t *testing.T) {
db := openMemDB(t)
// 模拟旧库:删除新结构 tokens 表,重建无 token_text 的旧表并插入旧数据
if _, err := db.Exec("DROP TABLE tokens"); err != nil {
t.Fatal(err)
}
if _, err := db.Exec(`CREATE TABLE tokens (
token_id INTEGER PRIMARY KEY,
created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP
)`); err != nil {
t.Fatal(err)
}
if _, err := db.Exec("INSERT INTO tokens (token_id) VALUES (100), (200)"); err != nil {
t.Fatal(err)
}
// 重新执行迁移:应补列并回填文本
if err := Migrate(db, "sqlite3"); err != nil {
t.Fatalf("Migrate 出错: %v", err)
}
var count int64
if err := db.QueryRow("SELECT COUNT(*) FROM tokens WHERE token_text = ''").Scan(&count); err != nil {
t.Fatal(err)
}
if count != 0 {
t.Errorf("回填后不应有空的 token_text, 剩余 %d", count)
}
var text string
if err := db.QueryRow("SELECT token_text FROM tokens WHERE token_id = 100").Scan(&text); err != nil {
t.Fatal(err)
}
if text == "" {
t.Error("token_id 100 的文本应已回填")
}
// 幂等:再次迁移不报错
if err := Migrate(db, "sqlite3"); err != nil {
t.Fatalf("重复迁移出错: %v", err)
}
}
func TestLoadMemory(t *testing.T) {
db := openMemDB(t)
ids, err := SaveMemories(db, []Memory{{Content: "测试内容", Category: "fact", Importance: 5}})
@@ -158,13 +207,13 @@ func TestSearchMemoriesByTokens(t *testing.T) {
t.Fatal(err)
}
// 手工建索引:记忆1 含 token 100,200;记忆2 含 100,200,300;记忆3 含 400
if err := SaveMemoryTokens(db, ids[0], []int64{100, 200}); err != nil {
if err := SaveMemoryTokens(db, ids[0], []tokens.Token{{ID: 100, Text: " 100"}, {ID: 200, Text: " 200"}}); err != nil {
t.Fatal(err)
}
if err := SaveMemoryTokens(db, ids[1], []int64{100, 200, 300}); err != nil {
if err := SaveMemoryTokens(db, ids[1], []tokens.Token{{ID: 100, Text: " 100"}, {ID: 200, Text: " 200"}, {ID: 300, Text: " 300"}}); err != nil {
t.Fatal(err)
}
if err := SaveMemoryTokens(db, ids[2], []int64{400}); err != nil {
if err := SaveMemoryTokens(db, ids[2], []tokens.Token{{ID: 400, Text: " 400"}}); err != nil {
t.Fatal(err)
}