From 378a6e0d418e4d7a732110335dcaaf48734834ab Mon Sep 17 00:00:00 2001 From: kevin Date: Fri, 14 Aug 2026 20:02:11 +0800 Subject: [PATCH] =?UTF-8?q?tokens=20=E8=A1=A8=E5=A2=9E=E5=8A=A0=20token=5F?= =?UTF-8?q?text=20=E5=AD=97=E6=AE=B5=EF=BC=8C=E6=97=A7=E5=BA=93=E8=87=AA?= =?UTF-8?q?=E5=8A=A8=E8=A1=A5=E5=88=97=E5=B9=B6=E5=9B=9E=E5=A1=AB=E7=AC=A6?= =?UTF-8?q?=E5=8F=B7=E6=96=87=E6=9C=AC?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- internal/bot/stats_test.go | 12 ++- internal/store/memory.go | 103 ++++++++++++++++++++++++-- internal/store/memory_test.go | 59 +++++++++++++-- internal/tokens/tokens.go | 35 ++++++++- internal/tools/builtin/recall_tool.go | 2 +- 5 files changed, 193 insertions(+), 18 deletions(-) diff --git a/internal/bot/stats_test.go b/internal/bot/stats_test.go index a955bad..a9423b1 100644 --- a/internal/bot/stats_test.go +++ b/internal/bot/stats_test.go @@ -39,11 +39,14 @@ func TestTokenize(t *testing.T) { t.Errorf("应有多个 token, got %v", ids) } seen := make(map[int64]bool) - for _, id := range ids { - if seen[id] { + for _, tok := range ids { + if tok.Text == "" { + t.Errorf("token 文本不应为空: %+v", tok) + } + if seen[tok.ID] { t.Errorf("token 应去重: %v", ids) } - seen[id] = true + seen[tok.ID] = true } if len(tokens.Tokenize("")) != 0 { t.Error("空串应返回空") @@ -52,6 +55,9 @@ func TestTokenize(t *testing.T) { if len(chinese) == 0 { t.Error("中文文本应产生 token") } + if got := tokens.TokenIDs(ids); len(got) != len(ids) { + t.Errorf("TokenIDs 数量 = %d, want %d", len(got), len(ids)) + } } func TestContextStats(t *testing.T) { diff --git a/internal/store/memory.go b/internal/store/memory.go index 4608414..f3fcff0 100644 --- a/internal/store/memory.go +++ b/internal/store/memory.go @@ -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) } } diff --git a/internal/store/memory_test.go b/internal/store/memory_test.go index f2f4f9e..af2d40e 100644 --- a/internal/store/memory_test.go +++ b/internal/store/memory_test.go @@ -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) } diff --git a/internal/tokens/tokens.go b/internal/tokens/tokens.go index 7c03114..7d2fb15 100644 --- a/internal/tokens/tokens.go +++ b/internal/tokens/tokens.go @@ -15,9 +15,15 @@ func getEncoding() *tiktoken.Tiktoken { return tke } -// Tokenize 将文本编码为 o200k_base token id 列表(去重)。 +// Token 表示一个 o200k_base token:ID 与对应的符号文本。 +type Token struct { + ID int64 + Text string +} + +// Tokenize 将文本编码为 token 列表(去重),Text 为 token 对应的符号。 // 编码器初始化失败时返回 nil。 -func Tokenize(text string) []int64 { +func Tokenize(text string) []Token { if text == "" { return nil } @@ -26,18 +32,39 @@ func Tokenize(text string) []int64 { return nil } seen := make(map[int64]bool) - var out []int64 + var out []Token for _, id := range enc.Encode(text, nil, nil) { tid := int64(id) if seen[tid] { continue } seen[tid] = true - out = append(out, tid) + out = append(out, Token{ID: tid, Text: enc.Decode([]int{int(tid)})}) } return out } +// TokenIDs 提取 Token 列表的 id。 +func TokenIDs(ts []Token) []int64 { + if len(ts) == 0 { + return nil + } + out := make([]int64, len(ts)) + for i, t := range ts { + out[i] = t.ID + } + return out +} + +// Text 返回 token id 对应的符号文本;未知 id 返回空串。 +func Text(id int64) string { + enc := getEncoding() + if enc == nil { + return "" + } + return enc.Decode([]int{int(id)}) +} + // Count 统计文本的 token 数量;编码器初始化失败时回退为字符数/2 估算。 func Count(text string) int64 { if text == "" { diff --git a/internal/tools/builtin/recall_tool.go b/internal/tools/builtin/recall_tool.go index 0fd0c9f..2590331 100644 --- a/internal/tools/builtin/recall_tool.go +++ b/internal/tools/builtin/recall_tool.go @@ -68,7 +68,7 @@ func (t *recallTool) Execute(args json.RawMessage) (string, error) { if p.Limit == 0 { p.Limit = 5 } - memories, err := store.SearchMemoriesByTokens(t.db, tokens.Tokenize(query), p.Limit) + memories, err := store.SearchMemoriesByTokens(t.db, tokens.TokenIDs(tokens.Tokenize(query)), p.Limit) if err != nil { return "", err }