记忆建立 token 搜索索引:tokens/memory_tokens 表与自动补索引
This commit is contained in:
@@ -171,6 +171,31 @@ func (b *Bot) ContextStats() (used, total int64) {
|
|||||||
|
|
||||||
var tke *tiktoken.Tiktoken
|
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;
|
// estimateTokens 用 o200k_base 词表精确统计 token;
|
||||||
// 初始化失败(如无法下载词表)时回退为字符数/2 估算。
|
// 初始化失败(如无法下载词表)时回退为字符数/2 估算。
|
||||||
func estimateTokens(s string) int64 {
|
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) {
|
func TestContextStats(t *testing.T) {
|
||||||
b := &Bot{
|
b := &Bot{
|
||||||
systemPrompt: "你是一个乐于助人的 AI 助手。",
|
systemPrompt: "你是一个乐于助人的 AI 助手。",
|
||||||
|
|||||||
+44
-10
@@ -40,6 +40,36 @@ func (h *Handler) saveSessionAndReset() {
|
|||||||
fmt.Printf("💾 会话已存档 (%d 条消息),已开启新对话\n", len(msgs))
|
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 {
|
func formatWindow(n int64) string {
|
||||||
switch {
|
switch {
|
||||||
case n <= 0:
|
case n <= 0:
|
||||||
@@ -162,16 +192,20 @@ func (h *Handler) Handle(input string) bool {
|
|||||||
}
|
}
|
||||||
if len(ms) == 0 {
|
if len(ms) == 0 {
|
||||||
fmt.Println("🧠 没有新的记忆")
|
fmt.Println("🧠 没有新的记忆")
|
||||||
h.saveSessionAndReset()
|
} else {
|
||||||
return true
|
ids, err := store.SaveMemories(h.db, ms)
|
||||||
}
|
if err != nil {
|
||||||
if _, err := store.SaveMemories(h.db, ms); err != nil {
|
fmt.Printf("⚠️ %v\n", err)
|
||||||
fmt.Printf("⚠️ %v\n", err)
|
return true
|
||||||
return true
|
}
|
||||||
}
|
fmt.Printf("🧠 已提取 %d 条新记忆\n", len(ms))
|
||||||
fmt.Printf("🧠 已提取 %d 条新记忆\n", len(ms))
|
for _, m := range ms {
|
||||||
for _, m := range ms {
|
fmt.Printf(" [%s %d] %s\n", m.Category, m.Importance, m.Content)
|
||||||
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()
|
h.saveSessionAndReset()
|
||||||
case "/forge":
|
case "/forge":
|
||||||
|
|||||||
+119
-12
@@ -3,7 +3,6 @@ package store
|
|||||||
import (
|
import (
|
||||||
"database/sql"
|
"database/sql"
|
||||||
"fmt"
|
"fmt"
|
||||||
"strings"
|
|
||||||
"time"
|
"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 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 {
|
func migrateMemories(db *sql.DB, driver string) error {
|
||||||
var ddl string
|
var ddl string
|
||||||
switch driver {
|
switch driver {
|
||||||
@@ -54,26 +81,106 @@ func migrateMemories(db *sql.DB, driver string) error {
|
|||||||
if _, err := db.Exec(createMemoriesIndex); err != nil {
|
if _, err := db.Exec(createMemoriesIndex); err != nil {
|
||||||
return fmt.Errorf("创建 memories 索引失败: %w", err)
|
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
|
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 {
|
if len(memories) == 0 {
|
||||||
return 0, nil
|
return nil, nil
|
||||||
}
|
}
|
||||||
placeholders := make([]string, 0, len(memories))
|
ids := make([]int64, 0, len(memories))
|
||||||
args := make([]any, 0, len(memories)*4)
|
|
||||||
for _, m := range memories {
|
for _, m := range memories {
|
||||||
placeholders = append(placeholders, "(?, ?, ?, ?)")
|
res, err := db.Exec(
|
||||||
args = append(args, m.SourceSessionID, m.Content, m.Category, m.Importance)
|
"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 " +
|
return ids, nil
|
||||||
strings.Join(placeholders, ", ")
|
}
|
||||||
res, err := db.Exec(query, args...)
|
|
||||||
|
// 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 {
|
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) {
|
func ListMemories(db *sql.DB) ([]Memory, error) {
|
||||||
|
|||||||
@@ -31,9 +31,16 @@ func TestMemoriesRoundtrip(t *testing.T) {
|
|||||||
{Content: "用户喜欢喝咖啡", Category: "preference", Importance: 7},
|
{Content: "用户喜欢喝咖啡", Category: "preference", Importance: 7},
|
||||||
{Content: "用户是 Go 开发者", Category: "fact", Importance: 9, SourceSessionID: 3},
|
{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)
|
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 {
|
if n, err := MemoryCount(db); err != nil || n != 2 {
|
||||||
t.Errorf("MemoryCount = %d, %v; want 2", n, err)
|
t.Errorf("MemoryCount = %d, %v; want 2", n, err)
|
||||||
}
|
}
|
||||||
@@ -54,7 +61,88 @@ func TestMemoriesRoundtrip(t *testing.T) {
|
|||||||
|
|
||||||
func TestSaveMemoriesEmpty(t *testing.T) {
|
func TestSaveMemoriesEmpty(t *testing.T) {
|
||||||
db := openMemDB(t)
|
db := openMemDB(t)
|
||||||
if _, err := SaveMemories(db, nil); err != nil {
|
ids, err := SaveMemories(db, nil)
|
||||||
|
if err != nil {
|
||||||
t.Fatalf("空列表不应报错: %v", err)
|
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