token 索引大小写规范化,新增/reindex重建搜索索引
This commit is contained in:
@@ -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) {
|
||||
|
||||
@@ -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 应被清空")
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user