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
+9 -3
View File
@@ -39,11 +39,14 @@ func TestTokenize(t *testing.T) {
t.Errorf("应有多个 token, got %v", ids) t.Errorf("应有多个 token, got %v", ids)
} }
seen := make(map[int64]bool) seen := make(map[int64]bool)
for _, id := range ids { for _, tok := range ids {
if seen[id] { if tok.Text == "" {
t.Errorf("token 文本不应为空: %+v", tok)
}
if seen[tok.ID] {
t.Errorf("token 应去重: %v", ids) t.Errorf("token 应去重: %v", ids)
} }
seen[id] = true seen[tok.ID] = true
} }
if len(tokens.Tokenize("")) != 0 { if len(tokens.Tokenize("")) != 0 {
t.Error("空串应返回空") t.Error("空串应返回空")
@@ -52,6 +55,9 @@ func TestTokenize(t *testing.T) {
if len(chinese) == 0 { if len(chinese) == 0 {
t.Error("中文文本应产生 token") 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) { func TestContextStats(t *testing.T) {
+98 -5
View File
@@ -5,6 +5,8 @@ import (
"fmt" "fmt"
"strings" "strings"
"time" "time"
"myaibot/internal/tokens"
) )
type Memory struct { type Memory struct {
@@ -41,12 +43,14 @@ const createMemoriesIndex = "CREATE INDEX IF NOT EXISTS idx_memories_created_at
const createTokensSQLite = ` const createTokensSQLite = `
CREATE TABLE IF NOT EXISTS tokens ( CREATE TABLE IF NOT EXISTS tokens (
token_id INTEGER PRIMARY KEY, token_id INTEGER PRIMARY KEY,
token_text TEXT NOT NULL DEFAULT '',
created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP
)` )`
const createTokensMySQL = ` const createTokensMySQL = `
CREATE TABLE IF NOT EXISTS tokens ( CREATE TABLE IF NOT EXISTS tokens (
token_id BIGINT PRIMARY KEY, token_id BIGINT PRIMARY KEY,
token_text TEXT NOT NULL,
created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP 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)" 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 { func migrateMemories(db *sql.DB, driver string) error {
var ddl string var ddl string
switch driver { switch driver {
@@ -91,6 +181,9 @@ func migrateMemories(db *sql.DB, driver string) error {
if _, err := db.Exec(ddl); err != nil { if _, err := db.Exec(ddl); err != nil {
return fmt.Errorf("创建 tokens 表失败: %w", err) return fmt.Errorf("创建 tokens 表失败: %w", err)
} }
if err := migrateTokens(db, driver); err != nil {
return err
}
switch driver { switch driver {
case "sqlite3": case "sqlite3":
ddl = createMemoryTokensSQLite ddl = createMemoryTokensSQLite
@@ -130,8 +223,8 @@ func SaveMemories(db *sql.DB, memories []Memory) ([]int64, error) {
} }
// SaveMemoryTokens 为记忆建立 token 索引:token 去重、关联幂等。 // SaveMemoryTokens 为记忆建立 token 索引:token 去重、关联幂等。
func SaveMemoryTokens(db *sql.DB, memoryID int64, tokenIDs []int64) error { func SaveMemoryTokens(db *sql.DB, memoryID int64, tokenList []tokens.Token) error {
if len(tokenIDs) == 0 { if len(tokenList) == 0 {
return nil return nil
} }
tx, err := db.Begin() tx, err := db.Begin()
@@ -139,11 +232,11 @@ func SaveMemoryTokens(db *sql.DB, memoryID int64, tokenIDs []int64) error {
return fmt.Errorf("开启事务失败: %w", err) return fmt.Errorf("开启事务失败: %w", err)
} }
defer tx.Rollback() defer tx.Rollback()
for _, tid := range tokenIDs { for _, t := range tokenList {
if _, err := tx.Exec("INSERT OR IGNORE INTO tokens (token_id) VALUES (?)", tid); err != nil { 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) 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) return fmt.Errorf("建立 token 关联失败: %w", err)
} }
} }
+54 -5
View File
@@ -6,6 +6,7 @@ import (
"testing" "testing"
"myaibot/internal/config" "myaibot/internal/config"
"myaibot/internal/tokens"
) )
func openMemDB(t *testing.T) *sql.DB { func openMemDB(t *testing.T) *sql.DB {
@@ -80,8 +81,9 @@ func TestMemoryTokens(t *testing.T) {
t.Fatalf("SaveMemories 出错: %v", err) t.Fatalf("SaveMemories 出错: %v", err)
} }
// 同一记忆建索引两次应幂等 // 同一记忆建索引两次应幂等
tok := []tokens.Token{{ID: 100, Text: " 100"}, {ID: 200, Text: " 200"}}
for i := 0; i < 2; i++ { 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) t.Fatalf("SaveMemoryTokens 出错: %v", err)
} }
} }
@@ -92,6 +94,13 @@ func TestMemoryTokens(t *testing.T) {
if tokenCount != 2 { if tokenCount != 2 {
t.Errorf("token 应去重为 2, got %d", tokenCount) 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 var linkCount int64
if err := db.QueryRow("SELECT COUNT(*) FROM memory_tokens").Scan(&linkCount); err != nil { if err := db.QueryRow("SELECT COUNT(*) FROM memory_tokens").Scan(&linkCount); err != nil {
t.Fatal(err) 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) 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) { func TestLoadMemory(t *testing.T) {
db := openMemDB(t) db := openMemDB(t)
ids, err := SaveMemories(db, []Memory{{Content: "测试内容", Category: "fact", Importance: 5}}) ids, err := SaveMemories(db, []Memory{{Content: "测试内容", Category: "fact", Importance: 5}})
@@ -158,13 +207,13 @@ func TestSearchMemoriesByTokens(t *testing.T) {
t.Fatal(err) t.Fatal(err)
} }
// 手工建索引:记忆1 含 token 100,200;记忆2 含 100,200,300;记忆3 含 400 // 手工建索引:记忆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) 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) 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) t.Fatal(err)
} }
+31 -4
View File
@@ -15,9 +15,15 @@ func getEncoding() *tiktoken.Tiktoken {
return tke return tke
} }
// Tokenize 将文本编码为 o200k_base token id 列表(去重) // Token 表示一个 o200k_base tokenID 与对应的符号文本
type Token struct {
ID int64
Text string
}
// Tokenize 将文本编码为 token 列表(去重),Text 为 token 对应的符号。
// 编码器初始化失败时返回 nil。 // 编码器初始化失败时返回 nil。
func Tokenize(text string) []int64 { func Tokenize(text string) []Token {
if text == "" { if text == "" {
return nil return nil
} }
@@ -26,18 +32,39 @@ func Tokenize(text string) []int64 {
return nil return nil
} }
seen := make(map[int64]bool) seen := make(map[int64]bool)
var out []int64 var out []Token
for _, id := range enc.Encode(text, nil, nil) { for _, id := range enc.Encode(text, nil, nil) {
tid := int64(id) tid := int64(id)
if seen[tid] { if seen[tid] {
continue continue
} }
seen[tid] = true seen[tid] = true
out = append(out, tid) out = append(out, Token{ID: tid, Text: enc.Decode([]int{int(tid)})})
} }
return out 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 估算。 // Count 统计文本的 token 数量;编码器初始化失败时回退为字符数/2 估算。
func Count(text string) int64 { func Count(text string) int64 {
if text == "" { if text == "" {
+1 -1
View File
@@ -68,7 +68,7 @@ func (t *recallTool) Execute(args json.RawMessage) (string, error) {
if p.Limit == 0 { if p.Limit == 0 {
p.Limit = 5 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 { if err != nil {
return "", err return "", err
} }