tokens 表增加 token_text 字段,旧库自动补列并回填符号文本
This commit is contained in:
@@ -39,11 +39,14 @@ func TestTokenize(t *testing.T) {
|
|||||||
t.Errorf("应有多个 token, got %v", ids)
|
t.Errorf("应有多个 token, got %v", ids)
|
||||||
}
|
}
|
||||||
seen := make(map[int64]bool)
|
seen := make(map[int64]bool)
|
||||||
for _, id := range ids {
|
for _, tok := range ids {
|
||||||
if seen[id] {
|
if tok.Text == "" {
|
||||||
|
t.Errorf("token 文本不应为空: %+v", tok)
|
||||||
|
}
|
||||||
|
if seen[tok.ID] {
|
||||||
t.Errorf("token 应去重: %v", ids)
|
t.Errorf("token 应去重: %v", ids)
|
||||||
}
|
}
|
||||||
seen[id] = true
|
seen[tok.ID] = true
|
||||||
}
|
}
|
||||||
if len(tokens.Tokenize("")) != 0 {
|
if len(tokens.Tokenize("")) != 0 {
|
||||||
t.Error("空串应返回空")
|
t.Error("空串应返回空")
|
||||||
@@ -52,6 +55,9 @@ func TestTokenize(t *testing.T) {
|
|||||||
if len(chinese) == 0 {
|
if len(chinese) == 0 {
|
||||||
t.Error("中文文本应产生 token")
|
t.Error("中文文本应产生 token")
|
||||||
}
|
}
|
||||||
|
if got := tokens.TokenIDs(ids); len(got) != len(ids) {
|
||||||
|
t.Errorf("TokenIDs 数量 = %d, want %d", len(got), len(ids))
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestContextStats(t *testing.T) {
|
func TestContextStats(t *testing.T) {
|
||||||
|
|||||||
@@ -5,6 +5,8 @@ import (
|
|||||||
"fmt"
|
"fmt"
|
||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"myaibot/internal/tokens"
|
||||||
)
|
)
|
||||||
|
|
||||||
type Memory struct {
|
type Memory struct {
|
||||||
@@ -41,12 +43,14 @@ const createMemoriesIndex = "CREATE INDEX IF NOT EXISTS idx_memories_created_at
|
|||||||
const createTokensSQLite = `
|
const createTokensSQLite = `
|
||||||
CREATE TABLE IF NOT EXISTS tokens (
|
CREATE TABLE IF NOT EXISTS tokens (
|
||||||
token_id INTEGER PRIMARY KEY,
|
token_id INTEGER PRIMARY KEY,
|
||||||
|
token_text TEXT NOT NULL DEFAULT '',
|
||||||
created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP
|
created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP
|
||||||
)`
|
)`
|
||||||
|
|
||||||
const createTokensMySQL = `
|
const createTokensMySQL = `
|
||||||
CREATE TABLE IF NOT EXISTS tokens (
|
CREATE TABLE IF NOT EXISTS tokens (
|
||||||
token_id BIGINT PRIMARY KEY,
|
token_id BIGINT PRIMARY KEY,
|
||||||
|
token_text TEXT NOT NULL,
|
||||||
created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP
|
created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP
|
||||||
)`
|
)`
|
||||||
|
|
||||||
@@ -66,6 +70,92 @@ CREATE TABLE IF NOT EXISTS memory_tokens (
|
|||||||
|
|
||||||
const createMemoryTokensIndex = "CREATE INDEX IF NOT EXISTS idx_memory_tokens_token_id ON memory_tokens (token_id)"
|
const createMemoryTokensIndex = "CREATE INDEX IF NOT EXISTS idx_memory_tokens_token_id ON memory_tokens (token_id)"
|
||||||
|
|
||||||
|
// migrateTokens 创建 tokens 表;旧表缺失 token_text 列时补列并回填文本。
|
||||||
|
func migrateTokens(db *sql.DB, driver string) error {
|
||||||
|
switch driver {
|
||||||
|
case "sqlite3":
|
||||||
|
if _, err := db.Exec(createTokensSQLite); err != nil {
|
||||||
|
return fmt.Errorf("创建 tokens 表失败: %w", err)
|
||||||
|
}
|
||||||
|
case "mysql":
|
||||||
|
if _, err := db.Exec(createTokensMySQL); err != nil {
|
||||||
|
return fmt.Errorf("创建 tokens 表失败: %w", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
has, err := hasTokenTextColumn(db, driver)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if !has {
|
||||||
|
ddl := "ALTER TABLE tokens ADD COLUMN token_text TEXT NOT NULL DEFAULT ''"
|
||||||
|
if driver == "mysql" {
|
||||||
|
ddl = "ALTER TABLE tokens ADD COLUMN token_text TEXT NOT NULL"
|
||||||
|
}
|
||||||
|
if _, err := db.Exec(ddl); err != nil {
|
||||||
|
return fmt.Errorf("tokens 表增加 token_text 列失败: %w", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return backfillTokenTexts(db)
|
||||||
|
}
|
||||||
|
|
||||||
|
func hasTokenTextColumn(db *sql.DB, driver string) (bool, error) {
|
||||||
|
var query string
|
||||||
|
switch driver {
|
||||||
|
case "sqlite3":
|
||||||
|
query = "SELECT name FROM pragma_table_info('tokens')"
|
||||||
|
case "mysql":
|
||||||
|
query = "SELECT column_name FROM information_schema.COLUMNS WHERE TABLE_SCHEMA = DATABASE() AND TABLE_NAME = 'tokens' AND COLUMN_NAME = 'token_text'"
|
||||||
|
default:
|
||||||
|
return false, fmt.Errorf("不支持的数据库驱动: %s", driver)
|
||||||
|
}
|
||||||
|
rows, err := db.Query(query)
|
||||||
|
if err != nil {
|
||||||
|
return false, fmt.Errorf("检查 tokens 表结构失败: %w", err)
|
||||||
|
}
|
||||||
|
defer rows.Close()
|
||||||
|
for rows.Next() {
|
||||||
|
var name string
|
||||||
|
if err := rows.Scan(&name); err != nil {
|
||||||
|
return false, fmt.Errorf("读取 tokens 表结构失败: %w", err)
|
||||||
|
}
|
||||||
|
if name == "token_text" {
|
||||||
|
return true, nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return false, rows.Err()
|
||||||
|
}
|
||||||
|
|
||||||
|
// backfillTokenTexts 为 token_text 为空的 token 行回填符号文本。
|
||||||
|
func backfillTokenTexts(db *sql.DB) error {
|
||||||
|
rows, err := db.Query("SELECT token_id FROM tokens WHERE token_text = ''")
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("查询待回填 token 失败: %w", err)
|
||||||
|
}
|
||||||
|
var ids []int64
|
||||||
|
for rows.Next() {
|
||||||
|
var id int64
|
||||||
|
if err := rows.Scan(&id); err != nil {
|
||||||
|
rows.Close()
|
||||||
|
return fmt.Errorf("读取 token id 失败: %w", err)
|
||||||
|
}
|
||||||
|
ids = append(ids, id)
|
||||||
|
}
|
||||||
|
rows.Close()
|
||||||
|
if err := rows.Err(); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
for _, id := range ids {
|
||||||
|
text := tokens.Text(id)
|
||||||
|
if text == "" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if _, err := db.Exec("UPDATE tokens SET token_text = ? WHERE token_id = ? AND token_text = ''", text, id); err != nil {
|
||||||
|
return fmt.Errorf("回填 token 文本失败: %w", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
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 {
|
||||||
@@ -91,6 +181,9 @@ func migrateMemories(db *sql.DB, driver string) error {
|
|||||||
if _, err := db.Exec(ddl); err != nil {
|
if _, err := db.Exec(ddl); err != nil {
|
||||||
return fmt.Errorf("创建 tokens 表失败: %w", err)
|
return fmt.Errorf("创建 tokens 表失败: %w", err)
|
||||||
}
|
}
|
||||||
|
if err := migrateTokens(db, driver); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
switch driver {
|
switch driver {
|
||||||
case "sqlite3":
|
case "sqlite3":
|
||||||
ddl = createMemoryTokensSQLite
|
ddl = createMemoryTokensSQLite
|
||||||
@@ -130,8 +223,8 @@ func SaveMemories(db *sql.DB, memories []Memory) ([]int64, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// SaveMemoryTokens 为记忆建立 token 索引:token 去重、关联幂等。
|
// SaveMemoryTokens 为记忆建立 token 索引:token 去重、关联幂等。
|
||||||
func SaveMemoryTokens(db *sql.DB, memoryID int64, tokenIDs []int64) error {
|
func SaveMemoryTokens(db *sql.DB, memoryID int64, tokenList []tokens.Token) error {
|
||||||
if len(tokenIDs) == 0 {
|
if len(tokenList) == 0 {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
tx, err := db.Begin()
|
tx, err := db.Begin()
|
||||||
@@ -139,11 +232,11 @@ func SaveMemoryTokens(db *sql.DB, memoryID int64, tokenIDs []int64) error {
|
|||||||
return fmt.Errorf("开启事务失败: %w", err)
|
return fmt.Errorf("开启事务失败: %w", err)
|
||||||
}
|
}
|
||||||
defer tx.Rollback()
|
defer tx.Rollback()
|
||||||
for _, tid := range tokenIDs {
|
for _, t := range tokenList {
|
||||||
if _, err := tx.Exec("INSERT OR IGNORE INTO tokens (token_id) VALUES (?)", tid); err != nil {
|
if _, err := tx.Exec("INSERT OR IGNORE INTO tokens (token_id, token_text) VALUES (?, ?)", t.ID, t.Text); err != nil {
|
||||||
return fmt.Errorf("保存 token 失败: %w", err)
|
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 {
|
if _, err := tx.Exec("INSERT OR IGNORE INTO memory_tokens (memory_id, token_id) VALUES (?, ?)", memoryID, t.ID); err != nil {
|
||||||
return fmt.Errorf("建立 token 关联失败: %w", err)
|
return fmt.Errorf("建立 token 关联失败: %w", err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -6,6 +6,7 @@ import (
|
|||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"myaibot/internal/config"
|
"myaibot/internal/config"
|
||||||
|
"myaibot/internal/tokens"
|
||||||
)
|
)
|
||||||
|
|
||||||
func openMemDB(t *testing.T) *sql.DB {
|
func openMemDB(t *testing.T) *sql.DB {
|
||||||
@@ -80,8 +81,9 @@ func TestMemoryTokens(t *testing.T) {
|
|||||||
t.Fatalf("SaveMemories 出错: %v", err)
|
t.Fatalf("SaveMemories 出错: %v", err)
|
||||||
}
|
}
|
||||||
// 同一记忆建索引两次应幂等
|
// 同一记忆建索引两次应幂等
|
||||||
|
tok := []tokens.Token{{ID: 100, Text: " 100"}, {ID: 200, Text: " 200"}}
|
||||||
for i := 0; i < 2; i++ {
|
for i := 0; i < 2; i++ {
|
||||||
if err := SaveMemoryTokens(db, ids[0], []int64{100, 200, 100}); err != nil {
|
if err := SaveMemoryTokens(db, ids[0], tok); err != nil {
|
||||||
t.Fatalf("SaveMemoryTokens 出错: %v", err)
|
t.Fatalf("SaveMemoryTokens 出错: %v", err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -92,6 +94,13 @@ func TestMemoryTokens(t *testing.T) {
|
|||||||
if tokenCount != 2 {
|
if tokenCount != 2 {
|
||||||
t.Errorf("token 应去重为 2, got %d", tokenCount)
|
t.Errorf("token 应去重为 2, got %d", tokenCount)
|
||||||
}
|
}
|
||||||
|
var text string
|
||||||
|
if err := db.QueryRow("SELECT token_text FROM tokens WHERE token_id = 100").Scan(&text); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if text != " 100" {
|
||||||
|
t.Errorf("token_text 应为 \" 100\", got %q", text)
|
||||||
|
}
|
||||||
var linkCount int64
|
var linkCount int64
|
||||||
if err := db.QueryRow("SELECT COUNT(*) FROM memory_tokens").Scan(&linkCount); err != nil {
|
if err := db.QueryRow("SELECT COUNT(*) FROM memory_tokens").Scan(&linkCount); err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
@@ -101,7 +110,7 @@ func TestMemoryTokens(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// 另一条记忆建索引
|
// 另一条记忆建索引
|
||||||
if err := SaveMemoryTokens(db, ids[1], []int64{200, 300}); err != nil {
|
if err := SaveMemoryTokens(db, ids[1], []tokens.Token{{ID: 200, Text: " 200"}, {ID: 300, Text: " 300"}}); err != nil {
|
||||||
t.Fatalf("SaveMemoryTokens 出错: %v", err)
|
t.Fatalf("SaveMemoryTokens 出错: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -128,6 +137,46 @@ func TestMemoryTokens(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestMigrateLegacyTokensTable(t *testing.T) {
|
||||||
|
db := openMemDB(t)
|
||||||
|
// 模拟旧库:删除新结构 tokens 表,重建无 token_text 的旧表并插入旧数据
|
||||||
|
if _, err := db.Exec("DROP TABLE tokens"); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if _, err := db.Exec(`CREATE TABLE tokens (
|
||||||
|
token_id INTEGER PRIMARY KEY,
|
||||||
|
created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP
|
||||||
|
)`); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if _, err := db.Exec("INSERT INTO tokens (token_id) VALUES (100), (200)"); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// 重新执行迁移:应补列并回填文本
|
||||||
|
if err := Migrate(db, "sqlite3"); err != nil {
|
||||||
|
t.Fatalf("Migrate 出错: %v", err)
|
||||||
|
}
|
||||||
|
var count int64
|
||||||
|
if err := db.QueryRow("SELECT COUNT(*) FROM tokens WHERE token_text = ''").Scan(&count); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if count != 0 {
|
||||||
|
t.Errorf("回填后不应有空的 token_text, 剩余 %d", count)
|
||||||
|
}
|
||||||
|
var text string
|
||||||
|
if err := db.QueryRow("SELECT token_text FROM tokens WHERE token_id = 100").Scan(&text); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if text == "" {
|
||||||
|
t.Error("token_id 100 的文本应已回填")
|
||||||
|
}
|
||||||
|
// 幂等:再次迁移不报错
|
||||||
|
if err := Migrate(db, "sqlite3"); err != nil {
|
||||||
|
t.Fatalf("重复迁移出错: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestLoadMemory(t *testing.T) {
|
func TestLoadMemory(t *testing.T) {
|
||||||
db := openMemDB(t)
|
db := openMemDB(t)
|
||||||
ids, err := SaveMemories(db, []Memory{{Content: "测试内容", Category: "fact", Importance: 5}})
|
ids, err := SaveMemories(db, []Memory{{Content: "测试内容", Category: "fact", Importance: 5}})
|
||||||
@@ -158,13 +207,13 @@ func TestSearchMemoriesByTokens(t *testing.T) {
|
|||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
// 手工建索引:记忆1 含 token 100,200;记忆2 含 100,200,300;记忆3 含 400
|
// 手工建索引:记忆1 含 token 100,200;记忆2 含 100,200,300;记忆3 含 400
|
||||||
if err := SaveMemoryTokens(db, ids[0], []int64{100, 200}); err != nil {
|
if err := SaveMemoryTokens(db, ids[0], []tokens.Token{{ID: 100, Text: " 100"}, {ID: 200, Text: " 200"}}); err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
if err := SaveMemoryTokens(db, ids[1], []int64{100, 200, 300}); err != nil {
|
if err := SaveMemoryTokens(db, ids[1], []tokens.Token{{ID: 100, Text: " 100"}, {ID: 200, Text: " 200"}, {ID: 300, Text: " 300"}}); err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
if err := SaveMemoryTokens(db, ids[2], []int64{400}); err != nil {
|
if err := SaveMemoryTokens(db, ids[2], []tokens.Token{{ID: 400, Text: " 400"}}); err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -15,9 +15,15 @@ func getEncoding() *tiktoken.Tiktoken {
|
|||||||
return tke
|
return tke
|
||||||
}
|
}
|
||||||
|
|
||||||
// Tokenize 将文本编码为 o200k_base token id 列表(去重)。
|
// Token 表示一个 o200k_base token:ID 与对应的符号文本。
|
||||||
|
type Token struct {
|
||||||
|
ID int64
|
||||||
|
Text string
|
||||||
|
}
|
||||||
|
|
||||||
|
// Tokenize 将文本编码为 token 列表(去重),Text 为 token 对应的符号。
|
||||||
// 编码器初始化失败时返回 nil。
|
// 编码器初始化失败时返回 nil。
|
||||||
func Tokenize(text string) []int64 {
|
func Tokenize(text string) []Token {
|
||||||
if text == "" {
|
if text == "" {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
@@ -26,18 +32,39 @@ func Tokenize(text string) []int64 {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
seen := make(map[int64]bool)
|
seen := make(map[int64]bool)
|
||||||
var out []int64
|
var out []Token
|
||||||
for _, id := range enc.Encode(text, nil, nil) {
|
for _, id := range enc.Encode(text, nil, nil) {
|
||||||
tid := int64(id)
|
tid := int64(id)
|
||||||
if seen[tid] {
|
if seen[tid] {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
seen[tid] = true
|
seen[tid] = true
|
||||||
out = append(out, tid)
|
out = append(out, Token{ID: tid, Text: enc.Decode([]int{int(tid)})})
|
||||||
}
|
}
|
||||||
return out
|
return out
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// TokenIDs 提取 Token 列表的 id。
|
||||||
|
func TokenIDs(ts []Token) []int64 {
|
||||||
|
if len(ts) == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
out := make([]int64, len(ts))
|
||||||
|
for i, t := range ts {
|
||||||
|
out[i] = t.ID
|
||||||
|
}
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
|
// Text 返回 token id 对应的符号文本;未知 id 返回空串。
|
||||||
|
func Text(id int64) string {
|
||||||
|
enc := getEncoding()
|
||||||
|
if enc == nil {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
return enc.Decode([]int{int(id)})
|
||||||
|
}
|
||||||
|
|
||||||
// Count 统计文本的 token 数量;编码器初始化失败时回退为字符数/2 估算。
|
// Count 统计文本的 token 数量;编码器初始化失败时回退为字符数/2 估算。
|
||||||
func Count(text string) int64 {
|
func Count(text string) int64 {
|
||||||
if text == "" {
|
if text == "" {
|
||||||
|
|||||||
@@ -68,7 +68,7 @@ func (t *recallTool) Execute(args json.RawMessage) (string, error) {
|
|||||||
if p.Limit == 0 {
|
if p.Limit == 0 {
|
||||||
p.Limit = 5
|
p.Limit = 5
|
||||||
}
|
}
|
||||||
memories, err := store.SearchMemoriesByTokens(t.db, tokens.Tokenize(query), p.Limit)
|
memories, err := store.SearchMemoriesByTokens(t.db, tokens.TokenIDs(tokens.Tokenize(query)), p.Limit)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return "", err
|
return "", err
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user