记忆建立 token 搜索索引:tokens/memory_tokens 表与自动补索引
This commit is contained in:
@@ -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 {
|
||||
|
||||
@@ -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
@@ -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
@@ -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) {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user