token 索引大小写规范化,新增/reindex重建搜索索引
This commit is contained in:
@@ -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 助手。",
|
||||
|
||||
@@ -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 <id> 切换到历史会话,如 /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 {
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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 应被清空")
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user