From 69e13da145f9c4fa6de108200054fe9036ac4855 Mon Sep 17 00:00:00 2001 From: kevin Date: Fri, 14 Aug 2026 20:15:07 +0800 Subject: [PATCH] =?UTF-8?q?token=20=E7=B4=A2=E5=BC=95=E5=A4=A7=E5=B0=8F?= =?UTF-8?q?=E5=86=99=E8=A7=84=E8=8C=83=E5=8C=96=EF=BC=8C=E6=96=B0=E5=A2=9E?= =?UTF-8?q?/reindex=E9=87=8D=E5=BB=BA=E6=90=9C=E7=B4=A2=E7=B4=A2=E5=BC=95?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- internal/bot/stats_test.go | 17 +++++++++++ internal/cli/cli.go | 8 +++++ internal/cli/complete.go | 2 +- internal/store/memory.go | 37 +++++++++++++++++++++++ internal/store/memory_test.go | 57 +++++++++++++++++++++++++++++++++++ internal/tokens/tokens.go | 9 ++++-- 6 files changed, 127 insertions(+), 3 deletions(-) diff --git a/internal/bot/stats_test.go b/internal/bot/stats_test.go index a9423b1..fa4ceef 100644 --- a/internal/bot/stats_test.go +++ b/internal/bot/stats_test.go @@ -60,6 +60,23 @@ func TestTokenize(t *testing.T) { } } +func TestTokenizeCaseInsensitive(t *testing.T) { + upper := tokens.Tokenize("Kevin 的生日") + lower := tokens.Tokenize("kevin 的生日") + if len(upper) != len(lower) { + t.Fatalf("大小写 token 数量不一致: %v vs %v", upper, lower) + } + for i := range upper { + if upper[i].ID != lower[i].ID { + t.Errorf("token %d: %d != %d(大小写应编码为相同 id)", i, upper[i].ID, lower[i].ID) + } + } + chinese := tokens.Tokenize("用户喜欢喝咖啡") + if len(chinese) == 0 { + t.Error("中文文本应产生 token") + } +} + func TestContextStats(t *testing.T) { b := &Bot{ systemPrompt: "你是一个乐于助人的 AI 助手。", diff --git a/internal/cli/cli.go b/internal/cli/cli.go index c2ef0dc..a04fd4a 100644 --- a/internal/cli/cli.go +++ b/internal/cli/cli.go @@ -117,6 +117,7 @@ func (h *Handler) Handle(input string) bool { fmt.Println(" /dream 从对话中提取长期记忆并开启新对话") fmt.Println(" /forge 直接清空对话,不提取记忆不存档") fmt.Println(" /memories 列出已提取的记忆") + fmt.Println(" /reindex 重建全部记忆的搜索索引") fmt.Println(" /sessions 列出历史会话") fmt.Println(" /session 切换到历史会话,如 /session 3") fmt.Println(" /info 显示当前供应商、模型和思考配置") @@ -225,6 +226,13 @@ func (h *Handler) Handle(input string) bool { for _, m := range list { fmt.Printf(" #%d %s [%s %d] %s\n", m.ID, m.CreatedAt.Format("2006-01-02 15:04"), m.Category, m.Importance, m.Content) } + case "/reindex": + n, err := store.RebuildTokenIndex(h.db) + if err != nil { + fmt.Printf("⚠️ %v\n", err) + return true + } + fmt.Printf("🔗 已重建搜索索引 (%d 条记忆)\n", n) case "/sessions": list, err := store.ListSessions(h.db) if err != nil { diff --git a/internal/cli/complete.go b/internal/cli/complete.go index 41e1aba..7aeebfa 100644 --- a/internal/cli/complete.go +++ b/internal/cli/complete.go @@ -2,7 +2,7 @@ package cli import "strings" -var commands = []string{"/exit", "/quit", "/help", "/models", "/use", "/think", "/effort", "/context", "/tools", "/dream", "/forge", "/memories", "/sessions", "/session", "/info"} +var commands = []string{"/exit", "/quit", "/help", "/models", "/use", "/think", "/effort", "/context", "/tools", "/dream", "/forge", "/memories", "/reindex", "/sessions", "/session", "/info"} func Complete(line string, models []string) []string { fields := strings.Fields(line) diff --git a/internal/store/memory.go b/internal/store/memory.go index f3fcff0..225d31d 100644 --- a/internal/store/memory.go +++ b/internal/store/memory.go @@ -277,6 +277,43 @@ func UnindexedMemoryIDs(db *sql.DB) ([]int64, error) { return out, rows.Err() } +// RebuildTokenIndex 清空 token 索引并重新为全部记忆建立索引,返回处理的记忆数。 +func RebuildTokenIndex(db *sql.DB) (int, error) { + memories, err := ListMemories(db) + if err != nil { + return 0, err + } + tx, err := db.Begin() + if err != nil { + return 0, fmt.Errorf("开启事务失败: %w", err) + } + defer tx.Rollback() + if _, err := tx.Exec("DELETE FROM memory_tokens"); err != nil { + return 0, fmt.Errorf("清空 memory_tokens 失败: %w", err) + } + if _, err := tx.Exec("DELETE FROM tokens"); err != nil { + return 0, fmt.Errorf("清空 tokens 失败: %w", err) + } + for _, m := range memories { + tokenList := tokens.Tokenize(m.Content) + if len(tokenList) == 0 { + continue + } + 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 0, fmt.Errorf("保存 token 失败: %w", err) + } + if _, err := tx.Exec("INSERT OR IGNORE INTO memory_tokens (memory_id, token_id) VALUES (?, ?)", m.ID, t.ID); err != nil { + return 0, fmt.Errorf("建立 token 关联失败: %w", err) + } + } + } + if err := tx.Commit(); err != nil { + return 0, fmt.Errorf("提交事务失败: %w", err) + } + return len(memories), nil +} + // SearchMemoriesByTokens 按 token id 搜索相关记忆,按命中 token 数从多到少排序。 // limit 钳制在 1-10。 func SearchMemoriesByTokens(db *sql.DB, tokenIDs []int64, limit int) ([]Memory, error) { diff --git a/internal/store/memory_test.go b/internal/store/memory_test.go index af2d40e..ef35ff6 100644 --- a/internal/store/memory_test.go +++ b/internal/store/memory_test.go @@ -260,3 +260,60 @@ func TestSearchMemoriesByTokens(t *testing.T) { t.Errorf("空 token 应返回 nil, %v %v", res, err) } } + +// 回归:大小写规范化后,索引大写内容、小写查询应命中 +func TestSearchCaseInsensitive(t *testing.T) { + db := openMemDB(t) + ids, err := SaveMemories(db, []Memory{ + {Content: "用户朋友Kevin的生日是9月12日", Category: "fact"}, + }) + if err != nil { + t.Fatal(err) + } + if err := SaveMemoryTokens(db, ids[0], tokens.Tokenize("用户朋友Kevin的生日是9月12日")); err != nil { + t.Fatal(err) + } + res, err := SearchMemoriesByTokens(db, tokens.TokenIDs(tokens.Tokenize("kevin 是谁")), 5) + if err != nil { + t.Fatalf("搜索出错: %v", err) + } + if len(res) != 1 || res[0].ID != ids[0] { + t.Errorf("小写查询应命中大写索引的记忆: %+v", res) + } +} + +func TestRebuildTokenIndex(t *testing.T) { + db := openMemDB(t) + ids, err := SaveMemories(db, []Memory{ + {Content: "用户喜欢咖啡"}, + {Content: "用户喜欢茶"}, + }) + if err != nil { + t.Fatal(err) + } + if err := SaveMemoryTokens(db, ids[0], []tokens.Token{{ID: 1, Text: " 1"}}); err != nil { + t.Fatal(err) + } + n, err := RebuildTokenIndex(db) + if err != nil { + t.Fatalf("RebuildTokenIndex 出错: %v", err) + } + if n != 2 { + t.Errorf("重建数量 = %d, want 2", n) + } + var count int64 + if err := db.QueryRow("SELECT COUNT(*) FROM tokens").Scan(&count); err != nil { + t.Fatal(err) + } + if count == 0 { + t.Error("重建后 tokens 不应为空") + } + // 旧 token 1 已被清空重建 + var oldExists int64 + if err := db.QueryRow("SELECT COUNT(*) FROM tokens WHERE token_id = 1").Scan(&oldExists); err != nil { + t.Fatal(err) + } + if oldExists != 0 { + t.Error("旧 token 应被清空") + } +} diff --git a/internal/tokens/tokens.go b/internal/tokens/tokens.go index 7d2fb15..d6422d7 100644 --- a/internal/tokens/tokens.go +++ b/internal/tokens/tokens.go @@ -1,6 +1,10 @@ package tokens -import "github.com/pkoukk/tiktoken-go" +import ( + "strings" + + "github.com/pkoukk/tiktoken-go" +) var tke *tiktoken.Tiktoken @@ -22,11 +26,12 @@ type Token struct { } // Tokenize 将文本编码为 token 列表(去重),Text 为 token 对应的符号。 -// 编码器初始化失败时返回 nil。 +// 文本先统一转小写,使索引与查询大小写不敏感;编码器初始化失败时返回 nil。 func Tokenize(text string) []Token { if text == "" { return nil } + text = strings.ToLower(text) enc := getEncoding() if enc == nil { return nil