diff --git a/internal/bot/bot.go b/internal/bot/bot.go index 7c21e05..7646544 100644 --- a/internal/bot/bot.go +++ b/internal/bot/bot.go @@ -171,6 +171,31 @@ func (b *Bot) ContextStats() (used, total int64) { var tke *tiktoken.Tiktoken +// Tokenize 将文本编码为 o200k_base token id 列表(去重)。 +func Tokenize(text string) []int64 { + if text == "" { + return nil + } + if tke == nil { + t, err := tiktoken.GetEncoding("o200k_base") + if err != nil { + return nil + } + tke = t + } + seen := make(map[int64]bool) + var out []int64 + for _, id := range tke.Encode(text, nil, nil) { + tid := int64(id) + if seen[tid] { + continue + } + seen[tid] = true + out = append(out, tid) + } + return out +} + // estimateTokens 用 o200k_base 词表精确统计 token; // 初始化失败(如无法下载词表)时回退为字符数/2 估算。 func estimateTokens(s string) int64 { diff --git a/internal/bot/stats_test.go b/internal/bot/stats_test.go index cb52e38..4be5409 100644 --- a/internal/bot/stats_test.go +++ b/internal/bot/stats_test.go @@ -31,6 +31,27 @@ func TestEstimateTokensKnown(t *testing.T) { } } +func TestTokenize(t *testing.T) { + ids := Tokenize("hello hello world") + if len(ids) < 2 { + t.Errorf("应有多个 token, got %v", ids) + } + seen := make(map[int64]bool) + for _, id := range ids { + if seen[id] { + t.Errorf("token 应去重: %v", ids) + } + seen[id] = true + } + if len(Tokenize("")) != 0 { + t.Error("空串应返回空") + } + chinese := 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 9c36e3e..b781084 100644 --- a/internal/cli/cli.go +++ b/internal/cli/cli.go @@ -40,6 +40,36 @@ func (h *Handler) saveSessionAndReset() { fmt.Printf("💾 会话已存档 (%d 条消息),已开启新对话\n", len(msgs)) } +// indexMemories 为新记忆及历史未索引记忆建立 token 搜索索引(幂等)。 +func (h *Handler) indexMemories(newIDs []int64) error { + oldIDs, err := store.UnindexedMemoryIDs(h.db) + if err != nil { + return fmt.Errorf("查询未索引记忆失败: %w", err) + } + seen := make(map[int64]bool, len(newIDs)+len(oldIDs)) + var ids []int64 + for _, id := range append(newIDs, oldIDs...) { + if !seen[id] { + seen[id] = true + ids = append(ids, id) + } + } + for _, id := range ids { + m, err := store.LoadMemory(h.db, id) + if err != nil { + return fmt.Errorf("读取记忆 #%d 失败: %w", id, err) + } + if m == nil { + continue + } + if err := store.SaveMemoryTokens(h.db, id, bot.Tokenize(m.Content)); err != nil { + return fmt.Errorf("记忆 #%d 建索引失败: %w", id, err) + } + } + fmt.Printf("🔗 已建立搜索索引 (%d 条记忆)\n", len(ids)) + return nil +} + func formatWindow(n int64) string { switch { case n <= 0: @@ -162,16 +192,20 @@ func (h *Handler) Handle(input string) bool { } if len(ms) == 0 { fmt.Println("🧠 没有新的记忆") - h.saveSessionAndReset() - return true - } - if _, err := store.SaveMemories(h.db, ms); err != nil { - fmt.Printf("⚠️ %v\n", err) - return true - } - fmt.Printf("🧠 已提取 %d 条新记忆\n", len(ms)) - for _, m := range ms { - fmt.Printf(" [%s %d] %s\n", m.Category, m.Importance, m.Content) + } else { + ids, err := store.SaveMemories(h.db, ms) + if err != nil { + fmt.Printf("⚠️ %v\n", err) + return true + } + fmt.Printf("🧠 已提取 %d 条新记忆\n", len(ms)) + for _, m := range ms { + fmt.Printf(" [%s %d] %s\n", m.Category, m.Importance, m.Content) + } + if err := h.indexMemories(ids); err != nil { + fmt.Printf("⚠️ %v\n", err) + return true + } } h.saveSessionAndReset() case "/forge": diff --git a/internal/store/memory.go b/internal/store/memory.go index 163be75..6d94d47 100644 --- a/internal/store/memory.go +++ b/internal/store/memory.go @@ -3,7 +3,6 @@ package store import ( "database/sql" "fmt" - "strings" "time" ) @@ -38,6 +37,34 @@ CREATE TABLE IF NOT EXISTS memories ( const createMemoriesIndex = "CREATE INDEX IF NOT EXISTS idx_memories_created_at ON memories (created_at)" +const createTokensSQLite = ` +CREATE TABLE IF NOT EXISTS tokens ( + token_id INTEGER PRIMARY KEY, + created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP +)` + +const createTokensMySQL = ` +CREATE TABLE IF NOT EXISTS tokens ( + token_id BIGINT PRIMARY KEY, + created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP +)` + +const createMemoryTokensSQLite = ` +CREATE TABLE IF NOT EXISTS memory_tokens ( + memory_id INTEGER NOT NULL, + token_id INTEGER NOT NULL, + PRIMARY KEY (memory_id, token_id) +)` + +const createMemoryTokensMySQL = ` +CREATE TABLE IF NOT EXISTS memory_tokens ( + memory_id BIGINT NOT NULL, + token_id BIGINT NOT NULL, + PRIMARY KEY (memory_id, token_id) +)` + +const createMemoryTokensIndex = "CREATE INDEX IF NOT EXISTS idx_memory_tokens_token_id ON memory_tokens (token_id)" + func migrateMemories(db *sql.DB, driver string) error { var ddl string switch driver { @@ -54,26 +81,106 @@ func migrateMemories(db *sql.DB, driver string) error { if _, err := db.Exec(createMemoriesIndex); err != nil { return fmt.Errorf("创建 memories 索引失败: %w", err) } + switch driver { + case "sqlite3": + ddl = createTokensSQLite + case "mysql": + ddl = createTokensMySQL + } + if _, err := db.Exec(ddl); err != nil { + return fmt.Errorf("创建 tokens 表失败: %w", err) + } + switch driver { + case "sqlite3": + ddl = createMemoryTokensSQLite + case "mysql": + ddl = createMemoryTokensMySQL + } + if _, err := db.Exec(ddl); err != nil { + return fmt.Errorf("创建 memory_tokens 表失败: %w", err) + } + if _, err := db.Exec(createMemoryTokensIndex); err != nil { + return fmt.Errorf("创建 memory_tokens 索引失败: %w", err) + } return nil } -func SaveMemories(db *sql.DB, memories []Memory) (int64, error) { +// SaveMemories 逐条插入记忆,返回每条记忆的 id。 +func SaveMemories(db *sql.DB, memories []Memory) ([]int64, error) { if len(memories) == 0 { - return 0, nil + return nil, nil } - placeholders := make([]string, 0, len(memories)) - args := make([]any, 0, len(memories)*4) + ids := make([]int64, 0, len(memories)) for _, m := range memories { - placeholders = append(placeholders, "(?, ?, ?, ?)") - args = append(args, m.SourceSessionID, m.Content, m.Category, m.Importance) + res, err := db.Exec( + "INSERT INTO memories (source_session_id, content, category, importance) VALUES (?, ?, ?, ?)", + m.SourceSessionID, m.Content, m.Category, m.Importance, + ) + if err != nil { + return nil, fmt.Errorf("保存记忆失败: %w", err) + } + id, err := res.LastInsertId() + if err != nil { + return nil, fmt.Errorf("读取记忆 id 失败: %w", err) + } + ids = append(ids, id) } - query := "INSERT INTO memories (source_session_id, content, category, importance) VALUES " + - strings.Join(placeholders, ", ") - res, err := db.Exec(query, args...) + return ids, nil +} + +// SaveMemoryTokens 为记忆建立 token 索引:token 去重、关联幂等。 +func SaveMemoryTokens(db *sql.DB, memoryID int64, tokenIDs []int64) error { + if len(tokenIDs) == 0 { + return nil + } + tx, err := db.Begin() if err != nil { - return 0, fmt.Errorf("保存记忆失败: %w", err) + return fmt.Errorf("开启事务失败: %w", err) } - return res.LastInsertId() + defer tx.Rollback() + for _, tid := range tokenIDs { + if _, err := tx.Exec("INSERT OR IGNORE INTO tokens (token_id) VALUES (?)", tid); 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 { + return fmt.Errorf("建立 token 关联失败: %w", err) + } + } + return tx.Commit() +} + +func LoadMemory(db *sql.DB, id int64) (*Memory, error) { + row := db.QueryRow("SELECT id, created_at, source_session_id, content, category, importance FROM memories WHERE id = ?", id) + var ( + m Memory + created string + ) + if err := row.Scan(&m.ID, &created, &m.SourceSessionID, &m.Content, &m.Category, &m.Importance); err != nil { + if err == sql.ErrNoRows { + return nil, nil + } + return nil, fmt.Errorf("读取记忆失败: %w", err) + } + m.CreatedAt = parseTime(created) + return &m, nil +} + +// UnindexedMemoryIDs 返回尚未建立 token 索引的记忆 id。 +func UnindexedMemoryIDs(db *sql.DB) ([]int64, error) { + rows, err := db.Query("SELECT m.id FROM memories m LEFT JOIN memory_tokens mt ON m.id = mt.memory_id WHERE mt.memory_id IS NULL") + if err != nil { + return nil, fmt.Errorf("查询未索引记忆失败: %w", err) + } + defer rows.Close() + var out []int64 + for rows.Next() { + var id int64 + if err := rows.Scan(&id); err != nil { + return nil, fmt.Errorf("读取记忆 id 失败: %w", err) + } + out = append(out, id) + } + return out, rows.Err() } func ListMemories(db *sql.DB) ([]Memory, error) { diff --git a/internal/store/memory_test.go b/internal/store/memory_test.go index 3258364..8e67e3f 100644 --- a/internal/store/memory_test.go +++ b/internal/store/memory_test.go @@ -31,9 +31,16 @@ func TestMemoriesRoundtrip(t *testing.T) { {Content: "用户喜欢喝咖啡", Category: "preference", Importance: 7}, {Content: "用户是 Go 开发者", Category: "fact", Importance: 9, SourceSessionID: 3}, } - if _, err := SaveMemories(db, ms); err != nil { + ids, err := SaveMemories(db, ms) + if err != nil { t.Fatalf("SaveMemories 出错: %v", err) } + if len(ids) != 2 { + t.Fatalf("返回 ids 数量 = %d, want 2", len(ids)) + } + if ids[1] <= ids[0] { + t.Errorf("ids 应递增: %v", ids) + } if n, err := MemoryCount(db); err != nil || n != 2 { t.Errorf("MemoryCount = %d, %v; want 2", n, err) } @@ -54,7 +61,88 @@ func TestMemoriesRoundtrip(t *testing.T) { func TestSaveMemoriesEmpty(t *testing.T) { db := openMemDB(t) - if _, err := SaveMemories(db, nil); err != nil { + ids, err := SaveMemories(db, nil) + if err != nil { t.Fatalf("空列表不应报错: %v", err) } + if ids != nil { + t.Errorf("空列表应返回 nil, got %v", ids) + } +} + +func TestMemoryTokens(t *testing.T) { + db := openMemDB(t) + ids, err := SaveMemories(db, []Memory{ + {Content: "用户喜欢喝咖啡"}, + {Content: "用户是 Go 开发者"}, + }) + if err != nil { + t.Fatalf("SaveMemories 出错: %v", err) + } + // 同一记忆建索引两次应幂等 + for i := 0; i < 2; i++ { + if err := SaveMemoryTokens(db, ids[0], []int64{100, 200, 100}); err != nil { + t.Fatalf("SaveMemoryTokens 出错: %v", err) + } + } + var tokenCount int64 + if err := db.QueryRow("SELECT COUNT(*) FROM tokens").Scan(&tokenCount); err != nil { + t.Fatal(err) + } + if tokenCount != 2 { + t.Errorf("token 应去重为 2, got %d", tokenCount) + } + var linkCount int64 + if err := db.QueryRow("SELECT COUNT(*) FROM memory_tokens").Scan(&linkCount); err != nil { + t.Fatal(err) + } + if linkCount != 2 { + t.Errorf("关联应去重为 2, got %d", linkCount) + } + + // 另一条记忆建索引 + if err := SaveMemoryTokens(db, ids[1], []int64{200, 300}); err != nil { + t.Fatalf("SaveMemoryTokens 出错: %v", err) + } + + // 未索引记忆查询:全部已索引 + unindexed, err := UnindexedMemoryIDs(db) + if err != nil { + t.Fatalf("UnindexedMemoryIDs 出错: %v", err) + } + if len(unindexed) != 0 { + t.Errorf("应无未索引记忆, got %v", unindexed) + } + + // 新加一条未索引记忆 + ids2, err := SaveMemories(db, []Memory{{Content: "未索引"}}) + if err != nil { + t.Fatal(err) + } + unindexed, err = UnindexedMemoryIDs(db) + if err != nil { + t.Fatal(err) + } + if len(unindexed) != 1 || unindexed[0] != ids2[0] { + t.Errorf("未索引应只有新记忆: %v", unindexed) + } +} + +func TestLoadMemory(t *testing.T) { + db := openMemDB(t) + ids, err := SaveMemories(db, []Memory{{Content: "测试内容", Category: "fact", Importance: 5}}) + if err != nil { + t.Fatal(err) + } + m, err := LoadMemory(db, ids[0]) + if err != nil { + t.Fatalf("LoadMemory 出错: %v", err) + } + if m == nil || m.Content != "测试内容" || m.Category != "fact" { + t.Errorf("LoadMemory 异常: %+v", m) + } + missing, err := LoadMemory(db, 9999) + if err != nil || missing != nil { + t.Errorf("不存在应返回 nil: %v, %v", missing, err) + } }