tokens 表增加 token_text 字段,旧库自动补列并回填符号文本
This commit is contained in:
@@ -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) {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
|
||||
@@ -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 == "" {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user