记忆建立 token 搜索索引:tokens/memory_tokens 表与自动补索引

This commit is contained in:
2026-08-14 19:46:55 +08:00
parent 1973b9d516
commit 3c320986e0
5 changed files with 299 additions and 24 deletions
+25
View File
@@ -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 {
+21
View File
@@ -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 助手。",
+38 -4
View File
@@ -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,10 +192,9 @@ 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 {
} else {
ids, err := store.SaveMemories(h.db, ms)
if err != nil {
fmt.Printf("⚠️ %v\n", err)
return true
}
@@ -173,6 +202,11 @@ func (h *Handler) Handle(input string) bool {
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":
h.bot.ClearHistory()
+120 -13
View File
@@ -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)
}
query := "INSERT INTO memories (source_session_id, content, category, importance) VALUES " +
strings.Join(placeholders, ", ")
res, err := db.Exec(query, args...)
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 0, fmt.Errorf("保存记忆失败: %w", err)
return nil, fmt.Errorf("保存记忆失败: %w", err)
}
return res.LastInsertId()
id, err := res.LastInsertId()
if err != nil {
return nil, fmt.Errorf("读取记忆 id 失败: %w", err)
}
ids = append(ids, id)
}
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 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 {
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) {
+90 -2
View File
@@ -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)
}
}