diff --git a/internal/bot/bot.go b/internal/bot/bot.go index 7646544..867efce 100644 --- a/internal/bot/bot.go +++ b/internal/bot/bot.go @@ -2,6 +2,7 @@ package bot import ( "context" + "database/sql" "encoding/json" "errors" "fmt" @@ -12,10 +13,10 @@ import ( "github.com/openai/openai-go/packages/param" "github.com/openai/openai-go/shared" "github.com/openai/openai-go/shared/constant" - "github.com/pkoukk/tiktoken-go" "github.com/tidwall/gjson" "myaibot/internal/config" "myaibot/internal/store" + "myaibot/internal/tokens" "myaibot/internal/tools" "myaibot/internal/tools/builtin" ) @@ -41,12 +42,12 @@ type Bot struct { toolRegistry *tools.Registry } -func New(cfg *config.Config) (*Bot, error) { +func New(cfg *config.Config, db *sql.DB) (*Bot, error) { b := &Bot{ clients: make(map[string]*openai.Client), cfg: cfg, systemPrompt: cfg.SystemPrompt, - toolRegistry: tools.NewRegistry(builtin.NewTimeTool(), builtin.NewCalculatorTool(), builtin.NewRandomTool()), + toolRegistry: tools.NewRegistry(builtin.NewTimeTool(), builtin.NewCalculatorTool(), builtin.NewRandomTool(), builtin.NewRecallTool(db)), } b.provider = config.FindProviderIn(cfg, cfg.DefaultProvider) b.model = cfg.DefaultModel @@ -153,7 +154,7 @@ func (b *Bot) ContextWindow() int64 { // ContextStats 统计当前上下文的 token 使用量与窗口总大小(0 表示未配置)。 func (b *Bot) ContextStats() (used, total int64) { total = b.ContextWindow() - used += estimateTokens(b.systemPrompt) + used += tokens.Count(b.systemPrompt) for _, msg := range b.history { var content string switch { @@ -164,54 +165,11 @@ func (b *Bot) ContextStats() (used, total int64) { case msg.OfSystem != nil: content = msg.OfSystem.Content.OfString.Value } - used += estimateTokens(content) + used += tokens.Count(content) } return used, total } -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 { - if s == "" { - return 0 - } - if tke == nil { - t, err := tiktoken.GetEncoding("o200k_base") - if err != nil { - return int64(len([]rune(s)) / 2) - } - tke = t - } - return int64(len(tke.Encode(s, nil, nil))) -} - const memoryExtractPrompt = `你是记忆提取器。从下面的对话中提取值得长期记住的信息,包括: - 用户的个人偏好、习惯、兴趣 - 用户的个人事实(职业、所在地、家庭等) diff --git a/internal/bot/session_test.go b/internal/bot/session_test.go index 8dee110..a13f332 100644 --- a/internal/bot/session_test.go +++ b/internal/bot/session_test.go @@ -12,7 +12,7 @@ import ( func newTestBot(t *testing.T) *Bot { t.Helper() t.Chdir(t.TempDir()) - for _, name := range []string{"get_current_time", "calculate", "random_number"} { + for _, name := range []string{"get_current_time", "calculate", "random_number", "recall_memory"} { if err := config.WriteDefaultToolConfig(name, map[string]any{"enabled": true, "prompt": "p"}); err != nil { t.Fatalf("写入工具配置失败: %v", err) } @@ -26,7 +26,7 @@ func newTestBot(t *testing.T) *Bot { {Name: "p", BaseURL: "x", Models: []config.ModelConfig{{Name: "m"}}}, }, } - b, err := New(cfg) + b, err := New(cfg, nil) if err != nil { t.Fatalf("New 出错: %v", err) } diff --git a/internal/bot/stats_test.go b/internal/bot/stats_test.go index 4be5409..a955bad 100644 --- a/internal/bot/stats_test.go +++ b/internal/bot/stats_test.go @@ -3,13 +3,15 @@ package bot import ( "testing" + "myaibot/internal/tokens" + "github.com/openai/openai-go" "myaibot/internal/config" ) func TestEstimateTokensEmpty(t *testing.T) { - if n := estimateTokens(""); n != 0 { + if n := tokens.Count(""); n != 0 { t.Errorf("空串应为 0, got %d", n) } } @@ -25,14 +27,14 @@ func TestEstimateTokensKnown(t *testing.T) { {"你是一个乐于助人的 AI 助手。", 12}, } for _, c := range cases { - if n := estimateTokens(c.text); n != c.want { - t.Errorf("estimateTokens(%q) = %d, want %d", c.text, n, c.want) + if n := tokens.Count(c.text); n != c.want { + t.Errorf("tokens.Count(%q) = %d, want %d", c.text, n, c.want) } } } func TestTokenize(t *testing.T) { - ids := Tokenize("hello hello world") + ids := tokens.Tokenize("hello hello world") if len(ids) < 2 { t.Errorf("应有多个 token, got %v", ids) } @@ -43,10 +45,10 @@ func TestTokenize(t *testing.T) { } seen[id] = true } - if len(Tokenize("")) != 0 { + if len(tokens.Tokenize("")) != 0 { t.Error("空串应返回空") } - chinese := Tokenize("用户喜欢喝咖啡") + chinese := tokens.Tokenize("用户喜欢喝咖啡") if len(chinese) == 0 { t.Error("中文文本应产生 token") } diff --git a/internal/cli/cli.go b/internal/cli/cli.go index b781084..c2ef0dc 100644 --- a/internal/cli/cli.go +++ b/internal/cli/cli.go @@ -9,6 +9,7 @@ import ( "myaibot/internal/bot" "myaibot/internal/store" + "myaibot/internal/tokens" ) type Handler struct { @@ -62,7 +63,7 @@ func (h *Handler) indexMemories(newIDs []int64) error { if m == nil { continue } - if err := store.SaveMemoryTokens(h.db, id, bot.Tokenize(m.Content)); err != nil { + if err := store.SaveMemoryTokens(h.db, id, tokens.Tokenize(m.Content)); err != nil { return fmt.Errorf("记忆 #%d 建索引失败: %w", id, err) } } diff --git a/internal/store/memory.go b/internal/store/memory.go index 6d94d47..4608414 100644 --- a/internal/store/memory.go +++ b/internal/store/memory.go @@ -3,6 +3,7 @@ package store import ( "database/sql" "fmt" + "strings" "time" ) @@ -183,6 +184,51 @@ func UnindexedMemoryIDs(db *sql.DB) ([]int64, error) { return out, rows.Err() } +// SearchMemoriesByTokens 按 token id 搜索相关记忆,按命中 token 数从多到少排序。 +// limit 钳制在 1-10。 +func SearchMemoriesByTokens(db *sql.DB, tokenIDs []int64, limit int) ([]Memory, error) { + if len(tokenIDs) == 0 { + return nil, nil + } + if limit < 1 { + limit = 1 + } + if limit > 10 { + limit = 10 + } + placeholders := make([]string, len(tokenIDs)) + args := make([]any, 0, len(tokenIDs)+1) + for i, tid := range tokenIDs { + placeholders[i] = "?" + args = append(args, tid) + } + args = append(args, limit) + query := `SELECT m.id, m.created_at, m.source_session_id, m.content, m.category, m.importance +FROM memory_tokens mt JOIN memories m ON m.id = mt.memory_id +WHERE mt.token_id IN (` + strings.Join(placeholders, ", ") + `) +GROUP BY m.id +ORDER BY COUNT(*) DESC, m.id DESC +LIMIT ?` + rows, err := db.Query(query, args...) + if err != nil { + return nil, fmt.Errorf("搜索记忆失败: %w", err) + } + defer rows.Close() + var out []Memory + for rows.Next() { + var ( + m Memory + created string + ) + if err := rows.Scan(&m.ID, &created, &m.SourceSessionID, &m.Content, &m.Category, &m.Importance); err != nil { + return nil, fmt.Errorf("读取记忆失败: %w", err) + } + m.CreatedAt = parseTime(created) + out = append(out, m) + } + return out, rows.Err() +} + func ListMemories(db *sql.DB) ([]Memory, error) { rows, err := db.Query("SELECT id, created_at, source_session_id, content, category, importance FROM memories ORDER BY id DESC") if err != nil { diff --git a/internal/store/memory_test.go b/internal/store/memory_test.go index 8e67e3f..f2f4f9e 100644 --- a/internal/store/memory_test.go +++ b/internal/store/memory_test.go @@ -146,3 +146,68 @@ func TestLoadMemory(t *testing.T) { t.Errorf("不存在应返回 nil: %v, %v", missing, err) } } + +func TestSearchMemoriesByTokens(t *testing.T) { + db := openMemDB(t) + ids, err := SaveMemories(db, []Memory{ + {Content: "用户喜欢喝咖啡", Category: "preference"}, + {Content: "用户喜欢喝咖啡和茶", Category: "preference"}, + {Content: "用户是 Go 开发者", Category: "fact"}, + }) + if err != nil { + t.Fatal(err) + } + // 手工建索引:记忆1 含 token 100,200;记忆2 含 100,200,300;记忆3 含 400 + if err := SaveMemoryTokens(db, ids[0], []int64{100, 200}); err != nil { + t.Fatal(err) + } + if err := SaveMemoryTokens(db, ids[1], []int64{100, 200, 300}); err != nil { + t.Fatal(err) + } + if err := SaveMemoryTokens(db, ids[2], []int64{400}); err != nil { + t.Fatal(err) + } + + // 命中数优先:记忆2 命中 3 个 token 排第一 + res, err := SearchMemoriesByTokens(db, []int64{100, 200, 300, 400}, 5) + if err != nil { + t.Fatalf("搜索出错: %v", err) + } + if len(res) != 3 { + t.Fatalf("命中数量 = %d, want 3", len(res)) + } + if res[0].ID != ids[1] { + t.Errorf("命中数最多的应排第一: %v", res[0]) + } + + // limit 钳制与命中子集 + res, err = SearchMemoriesByTokens(db, []int64{100}, 10) + if err != nil { + t.Fatal(err) + } + if len(res) != 2 { + t.Errorf("limit 10 应命中 2 条, got %d", len(res)) + } + res, err = SearchMemoriesByTokens(db, []int64{100}, 1) + if err != nil { + t.Fatal(err) + } + if len(res) != 1 { + t.Errorf("limit 1 应只返回 1 条, got %d", len(res)) + } + + // 无命中 + res, err = SearchMemoriesByTokens(db, []int64{9999}, 5) + if err != nil { + t.Fatal(err) + } + if len(res) != 0 { + t.Errorf("无命中应返回空, got %d", len(res)) + } + + // 空 token + res, err = SearchMemoriesByTokens(db, nil, 5) + if err != nil || res != nil { + t.Errorf("空 token 应返回 nil, %v %v", res, err) + } +} diff --git a/internal/tokens/tokens.go b/internal/tokens/tokens.go new file mode 100644 index 0000000..7c03114 --- /dev/null +++ b/internal/tokens/tokens.go @@ -0,0 +1,51 @@ +package tokens + +import "github.com/pkoukk/tiktoken-go" + +var tke *tiktoken.Tiktoken + +func getEncoding() *tiktoken.Tiktoken { + if tke == nil { + t, err := tiktoken.GetEncoding("o200k_base") + if err != nil { + return nil + } + tke = t + } + return tke +} + +// Tokenize 将文本编码为 o200k_base token id 列表(去重)。 +// 编码器初始化失败时返回 nil。 +func Tokenize(text string) []int64 { + if text == "" { + return nil + } + enc := getEncoding() + if enc == nil { + return nil + } + seen := make(map[int64]bool) + var out []int64 + for _, id := range enc.Encode(text, nil, nil) { + tid := int64(id) + if seen[tid] { + continue + } + seen[tid] = true + out = append(out, tid) + } + return out +} + +// Count 统计文本的 token 数量;编码器初始化失败时回退为字符数/2 估算。 +func Count(text string) int64 { + if text == "" { + return 0 + } + enc := getEncoding() + if enc == nil { + return int64(len([]rune(text)) / 2) + } + return int64(len(enc.Encode(text, nil, nil))) +} diff --git a/internal/tools/builtin/builtin_test.go b/internal/tools/builtin/builtin_test.go new file mode 100644 index 0000000..b054def --- /dev/null +++ b/internal/tools/builtin/builtin_test.go @@ -0,0 +1,84 @@ +package builtin + +import ( + "database/sql" + "encoding/json" + "path/filepath" + "strings" + "testing" + + "myaibot/internal/config" + "myaibot/internal/store" + "myaibot/internal/tokens" +) + +func recallTestDB(t *testing.T) *sql.DB { + t.Helper() + cfg := &config.DatabaseConfig{ + Driver: "sqlite3", + File: filepath.Join(t.TempDir(), "memory.db"), + } + db, err := store.Open(cfg) + if err != nil { + t.Fatalf("Open 出错: %v", err) + } + t.Cleanup(func() { store.Close(db) }) + if err := store.Migrate(db, "sqlite3"); err != nil { + t.Fatalf("Migrate 出错: %v", err) + } + ids, err := store.SaveMemories(db, []store.Memory{ + {Content: "用户喜欢喝咖啡", Category: "preference", Importance: 7}, + }) + if err != nil { + t.Fatal(err) + } + if err := store.SaveMemoryTokens(db, ids[0], tokens.Tokenize("用户喜欢喝咖啡")); err != nil { + t.Fatal(err) + } + return db +} + +func TestRecallToolFound(t *testing.T) { + tool := NewRecallTool(recallTestDB(t)) + args, _ := json.Marshal(map[string]any{"query": "咖啡"}) + out, err := tool.Execute(args) + if err != nil { + t.Fatalf("Execute 出错: %v", err) + } + if !strings.Contains(out, "找到 1 条相关记忆") || !strings.Contains(out, "用户喜欢喝咖啡") { + t.Errorf("输出异常: %q", out) + } +} + +func TestRecallToolNotFound(t *testing.T) { + tool := NewRecallTool(recallTestDB(t)) + args, _ := json.Marshal(map[string]any{"query": "不存在的关键词"}) + out, err := tool.Execute(args) + if err != nil { + t.Fatalf("Execute 出错: %v", err) + } + if out != "未找到相关记忆" { + t.Errorf("应返回未找到, got %q", out) + } +} + +func TestRecallToolEmptyQuery(t *testing.T) { + tool := NewRecallTool(recallTestDB(t)) + args, _ := json.Marshal(map[string]any{"query": ""}) + if _, err := tool.Execute(args); err == nil { + t.Error("空 query 应报错") + } +} + +func TestRecallToolConfigure(t *testing.T) { + tool := NewRecallTool(nil) + if err := tool.Configure(map[string]any{"enabled": false, "prompt": "自定义提示"}); err != nil { + t.Fatalf("Configure 出错: %v", err) + } + if tool.Enabled() { + t.Error("应被禁用") + } + if tool.Description() != "自定义提示" { + t.Errorf("Description = %q", tool.Description()) + } +} diff --git a/internal/tools/builtin/recall_tool.go b/internal/tools/builtin/recall_tool.go new file mode 100644 index 0000000..0fd0c9f --- /dev/null +++ b/internal/tools/builtin/recall_tool.go @@ -0,0 +1,84 @@ +package builtin + +import ( + "database/sql" + "encoding/json" + "fmt" + "strings" + + "myaibot/internal/store" + "myaibot/internal/tokens" +) + +type recallTool struct { + db *sql.DB + enabled bool + prompt string +} + +func NewRecallTool(db *sql.DB) *recallTool { + return &recallTool{ + db: db, + enabled: true, + prompt: "从长期记忆中回忆与用户提问相关的信息。当用户询问个人偏好、个人信息、之前聊过的话题或需要回顾历史对话时调用此工具", + } +} + +func (t *recallTool) Name() string { return "recall_memory" } +func (t *recallTool) Description() string { return t.prompt } +func (t *recallTool) Enabled() bool { return t.enabled } +func (t *recallTool) DefaultConfig() map[string]any { + return map[string]any{"enabled": true, "prompt": t.prompt} +} + +func (t *recallTool) Configure(cfg map[string]any) error { + var err error + if t.enabled, err = parseEnabled(cfg); err != nil { + return err + } + if p, ok := cfg["prompt"].(string); ok && p != "" { + t.prompt = p + } + return nil +} + +func (t *recallTool) Parameters() map[string]any { + return map[string]any{ + "type": "object", + "properties": map[string]any{ + "query": map[string]any{"type": "string", "description": "要回忆的内容或用户的提问"}, + "limit": map[string]any{"type": "integer", "description": "最多返回的记忆条数,默认 5,最大 10"}, + }, + "required": []string{"query"}, + } +} + +func (t *recallTool) Execute(args json.RawMessage) (string, error) { + var p struct { + Query string `json:"query"` + Limit int `json:"limit"` + } + if err := json.Unmarshal(args, &p); err != nil { + return "", err + } + query := strings.TrimSpace(p.Query) + if query == "" { + return "", fmt.Errorf("query 不能为空") + } + if p.Limit == 0 { + p.Limit = 5 + } + memories, err := store.SearchMemoriesByTokens(t.db, tokens.Tokenize(query), p.Limit) + if err != nil { + return "", err + } + if len(memories) == 0 { + return "未找到相关记忆", nil + } + var sb strings.Builder + fmt.Fprintf(&sb, "找到 %d 条相关记忆:\n", len(memories)) + for _, m := range memories { + fmt.Fprintf(&sb, "- [%s %d] %s\n", m.Category, m.Importance, m.Content) + } + return strings.TrimSuffix(sb.String(), "\n"), nil +} diff --git a/internal/tools/builtin_test.go b/internal/tools/builtin_test.go deleted file mode 100644 index 0836e52..0000000 --- a/internal/tools/builtin_test.go +++ /dev/null @@ -1,115 +0,0 @@ -package tools_test - -import ( - "encoding/json" - "strings" - "testing" - - "myaibot/internal/tools/builtin" -) - -func TestCalculator(t *testing.T) { - cases := []struct { - name string - expr string - want string - }{ - {"四则运算", "(12 + 3) * 4", "60"}, - {"幂运算", "2^10", "1024"}, - {"浮点", "7 / 2", "3.5"}, - {"小数精度", "1 / 3", "0.33333333"}, - {"布尔", "2 > 1", "true"}, - } - for _, c := range cases { - t.Run(c.name, func(t *testing.T) { - args, _ := json.Marshal(map[string]string{"expression": c.expr}) - got, err := builtin.NewCalculatorTool().Execute(args) - if err != nil { - t.Fatalf("Execute 出错: %v", err) - } - if got != c.want { - t.Errorf("Execute(%q) = %q, want %q", c.expr, got, c.want) - } - }) - } -} - -func TestCalculatorInvalid(t *testing.T) { - args, _ := json.Marshal(map[string]string{"expression": "1 +"}) - if _, err := builtin.NewCalculatorTool().Execute(args); err == nil { - t.Error("非法表达式应返回错误") - } -} - -func TestRandom(t *testing.T) { - args, _ := json.Marshal(map[string]any{"min": 1, "max": 10}) - for i := 0; i < 100; i++ { - out, err := builtin.NewRandomTool().Execute(args) - if err != nil { - t.Fatalf("Execute 出错: %v", err) - } - var v int - if err := json.Unmarshal([]byte(out), &v); err != nil { - t.Fatalf("结果 %q 解析失败: %v", out, err) - } - if v < 1 || v > 10 { - t.Fatalf("结果 %q 不在 [1,10] 内", out) - } - } - bad, _ := json.Marshal(map[string]any{"min": 10, "max": 1}) - if _, err := builtin.NewRandomTool().Execute(bad); err == nil { - t.Error("min>max 应返回错误") - } -} - -func TestTimeTool(t *testing.T) { - out, err := builtin.NewTimeTool().Execute(nil) - if err != nil { - t.Fatalf("Execute 出错: %v", err) - } - if !strings.Contains(out, "20") { - t.Errorf("时间输出异常: %q", out) - } -} - -func TestConfigurePrompt(t *testing.T) { - tool := builtin.NewTimeTool() - err := tool.Configure(map[string]any{"enabled": true, "prompt": "自定义提示词"}) - if err != nil { - t.Fatalf("Configure 出错: %v", err) - } - if tool.Description() != "自定义提示词" { - t.Errorf("Description = %q, want 自定义提示词", tool.Description()) - } - if !tool.Enabled() { - t.Error("enabled 应为 true") - } -} - -func TestConfigureDisable(t *testing.T) { - tool := builtin.NewCalculatorTool() - if err := tool.Configure(map[string]any{"enabled": false}); err != nil { - t.Fatalf("Configure 出错: %v", err) - } - if tool.Enabled() { - t.Error("enabled 应为 false") - } -} - -func TestConfigureInvalid(t *testing.T) { - tool := builtin.NewRandomTool() - if err := tool.Configure(map[string]any{"enabled": "yes"}); err == nil { - t.Error("enabled 非布尔值应报错") - } -} - -func TestDefaultConfig(t *testing.T) { - tool := builtin.NewTimeTool() - cfg := tool.DefaultConfig() - if cfg["enabled"] != true { - t.Errorf("默认 enabled 应为 true, got %v", cfg["enabled"]) - } - if p, ok := cfg["prompt"].(string); !ok || p == "" { - t.Errorf("默认 prompt 缺失: %v", cfg["prompt"]) - } -} diff --git a/main.go b/main.go index 0552d5f..3758de6 100644 --- a/main.go +++ b/main.go @@ -33,7 +33,7 @@ func main() { } fmt.Printf("💾 数据库已连接 (%s)\n", cfg.Database.Driver) - b, err := bot.New(cfg) + b, err := bot.New(cfg, db) if err != nil { fmt.Printf("⚠️ %v\n", err) fmt.Println("请填写工具配置文件后重新启动。")