topfans/backend/services/aiChatService/service/memory_service.go
2026-07-06 10:25:06 +08:00

321 lines
9.6 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

package service
import (
"context"
"strings"
"unicode"
"github.com/topfans/backend/pkg/logger"
"github.com/topfans/backend/services/aiChatService/model"
"github.com/topfans/backend/services/aiChatService/repository"
"go.uber.org/zap"
)
// MemoryService 记忆服务
type MemoryService struct {
shortTermRepo repository.ShortTermMemoryRepository
longTermRepo repository.LongTermMemoryRepository
}
// NewMemoryService 创建记忆服务
func NewMemoryService(shortTermRepo repository.ShortTermMemoryRepository, longTermRepo repository.LongTermMemoryRepository) *MemoryService {
return &MemoryService{
shortTermRepo: shortTermRepo,
longTermRepo: longTermRepo,
}
}
// SaveContext 保存短期上下文
func (s *MemoryService) SaveContext(ctx context.Context, sessionID string, messages []model.Message, personaID string) error {
return s.shortTermRepo.SaveContext(ctx, sessionID, messages, personaID)
}
// GetContext 获取短期上下文
func (s *MemoryService) GetContext(ctx context.Context, sessionID string) ([]model.Message, error) {
return s.shortTermRepo.GetContext(ctx, sessionID)
}
// RecallMemories 召回相关记忆
func (s *MemoryService) RecallMemories(ctx context.Context, userID int64, userInput string, limit int) (string, error) {
keywords := extractKeywords(userInput)
memories, err := s.longTermRepo.GetMemories(ctx, userID, keywords, limit)
if err != nil {
return "", err
}
if len(memories) == 0 {
logger.Logger.Debug("No memories matched for keywords",
zap.Int64("user_id", userID),
zap.Strings("keywords", keywords),
)
return "", nil
}
logger.Logger.Debug("Memories recalled",
zap.Int64("user_id", userID),
zap.Int("count", len(memories)),
)
var builder strings.Builder
builder.WriteString("# 用户核心记忆\n")
for _, m := range memories {
builder.WriteString("- ")
builder.WriteString(m.Content)
builder.WriteString("\n")
}
return builder.String(), nil
}
// ExtractMemory 提取并保存记忆(接收 LLM 提取结果)
func (s *MemoryService) ExtractMemory(ctx context.Context, userID int64, recentMessages []model.Message, extractedMemories []string) error {
if len(extractedMemories) == 0 {
return nil
}
// 从最近的用户消息中提取关键词
var userMessages []string
for i := len(recentMessages) - 1; i >= 0 && len(userMessages) < 5; i-- {
if recentMessages[i].Role == "user" {
userMessages = append(userMessages, recentMessages[i].Content)
}
}
// 如果没有用户消息(极端情况),从所有消息中提取关键词作为兜底
if len(userMessages) == 0 {
for i := len(recentMessages) - 1; i >= 0 && len(userMessages) < 5; i-- {
userMessages = append(userMessages, recentMessages[i].Content)
}
}
keywords := extractKeywordsFromMessages(userMessages)
for _, content := range extractedMemories {
memory := &model.UserMemory{
UserID: userID,
Content: content,
Keywords: keywords,
Weight: 50,
}
if err := s.longTermRepo.SaveMemory(ctx, memory); err != nil {
logger.Logger.Error("Failed to save extracted memory",
zap.Int64("user_id", userID),
zap.String("content", content),
zap.Error(err),
)
// 继续处理其他记忆,不因单条失败而中断
continue
}
}
logger.Logger.Info("Memories saved",
zap.Int64("user_id", userID),
zap.Int("count", len(extractedMemories)),
)
return nil
}
// GetUserMemories 获取用户所有长期记忆
func (s *MemoryService) GetUserMemories(ctx context.Context, userID int64) ([]model.UserMemory, error) {
return s.longTermRepo.GetMemoriesByUserID(ctx, userID)
}
// 中文停用词(常见的无语义虚词)
var chineseStopWords = map[string]bool{
"的": true, "了": true, "是": true, "我": true, "你": true,
"他": true, "她": true, "它": true, "们": true, "这": true,
"那": true, "在": true, "有": true, "不": true, "和": true,
"就": true, "都": true, "也": true, "还": true, "要": true,
"会": true, "可": true, "没": true, "很": true, "个": true,
"对": true, "与": true, "或": true, "但": true, "而": true,
"且": true, "所": true, "为": true, "以": true, "及": true,
"上": true, "中": true, "下": true, "着": true, "过": true,
"得": true, "地": true, "把": true, "被": true, "让": true,
"从": true, "到": true, "向": true, "将": true, "能": true,
"想": true, "说": true, "去": true, "来": true, "做": true,
"看": true, "听": true, "知": true, "觉": true,
"给": true, "用": true, "吧": true, "吗": true, "呢": true,
"啊": true, "哦": true, "嗯": true, "哈": true, "嘛": true,
"呀": true, "啦": true, "哟": true, "哇": true, "喔": true,
"么": true, "什": true, "怎": true, "哪": true, "谁": true,
}
// isCJK 判断是否为中日韩文字
func isCJK(r rune) bool {
return unicode.Is(unicode.Han, r) || unicode.Is(unicode.Hiragana, r) || unicode.Is(unicode.Katakana, r)
}
// extractKeywords 从用户输入提取关键词(中文 bigram + 英文分词)
func extractKeywords(text string) []string {
var keywords []string
seen := make(map[string]bool)
// 1. 中文 bigram 提取2-3 字组合)
runes := []rune(text)
for i := 0; i < len(runes)-1; i++ {
if isCJK(runes[i]) && isCJK(runes[i+1]) {
bigram := string(runes[i : i+2])
if !chineseStopWords[bigram] && !seen[bigram] && len(bigram) >= 2 {
seen[bigram] = true
keywords = append(keywords, bigram)
}
}
// 3-gram更长的短语
if i < len(runes)-2 && isCJK(runes[i]) && isCJK(runes[i+1]) && isCJK(runes[i+2]) {
trigram := string(runes[i : i+3])
if !seen[trigram] && len(trigram) >= 3 {
seen[trigram] = true
keywords = append(keywords, trigram)
}
}
}
// 2. 英文/数字分词:按空白和标点分割
words := strings.FieldsFunc(text, func(r rune) bool {
return unicode.IsPunct(r) || unicode.IsSpace(r) || r == '' || r == '。' || r == '' || r == '' || r == '\n'
})
for _, word := range words {
word = strings.TrimSpace(word)
// 保留长度 >= 2 的非中文词
if len(word) >= 2 && !isCJK([]rune(word)[0]) && !seen[word] {
seen[word] = true
keywords = append(keywords, word)
}
}
return keywords
}
// extractKeywordsFromMessages 从多条消息提取关键词
func extractKeywordsFromMessages(messages []string) []string {
var allKeywords []string
seen := make(map[string]bool)
for _, msg := range messages {
keywords := extractKeywords(msg)
for _, k := range keywords {
if !seen[k] {
seen[k] = true
allKeywords = append(allKeywords, k)
}
}
}
// 限制关键词数量,优先保留长的(更具体)
if len(allKeywords) > 10 {
// 按长度降序排列,长词更具体
sorted := make([]string, len(allKeywords))
copy(sorted, allKeywords)
for i := 0; i < len(sorted); i++ {
for j := i + 1; j < len(sorted); j++ {
if len(sorted[j]) > len(sorted[i]) {
sorted[i], sorted[j] = sorted[j], sorted[i]
}
}
}
allKeywords = sorted[:10]
}
return allKeywords
}
// ExtractMemoriesWithLLM 使用 LLM 从对话中提取用户记忆
// 返回提取到的记忆文本列表
func ExtractMemoriesWithLLM(ctx context.Context, llmChatFunc func(ctx context.Context, messages []model.Message) (string, error), recentMessages []model.Message) ([]string, error) {
// 只取最近 10 条消息5轮对话防止 prompt 过大
maxMessages := 10
if len(recentMessages) > maxMessages {
recentMessages = recentMessages[len(recentMessages)-maxMessages:]
}
// 构建记忆提取 prompt
var conversationBuilder strings.Builder
for _, msg := range recentMessages {
role := "用户"
if msg.Role == "assistant" {
role = "AI"
}
// 截断过长消息
content := msg.Content
if len([]rune(content)) > 200 {
content = string([]rune(content)[:200]) + "..."
}
conversationBuilder.WriteString(role)
conversationBuilder.WriteString(": ")
conversationBuilder.WriteString(content)
conversationBuilder.WriteString("\n")
}
extractionPrompt := []model.Message{
{
Role: "system",
Content: `你是一个记忆提取系统。从以下对话中提取关于用户的关键事实和记忆。
每条记忆用一句话概括,只输出记忆内容,每行一条。
重点关注:
- 用户的偏好、兴趣、爱好
- 用户分享的个人信息
- 用户表达的观点和态度
- 用户反复提及的话题
如果没有可提取的记忆,输出"无"。
示例输出格式:
用户喜欢打篮球,每周去健身房三次
用户在学编程,主要用 Go 语言
用户养了一只叫小白的猫`,
},
{
Role: "user",
Content: "对话内容:\n" + conversationBuilder.String() + "\n请提取用户的关键记忆",
},
}
response, err := llmChatFunc(ctx, extractionPrompt)
if err != nil {
return nil, err
}
// 解析响应,每行一条记忆
lines := strings.Split(strings.TrimSpace(response), "\n")
var memories []string
for _, line := range lines {
line = strings.TrimSpace(line)
if line == "" || line == "无" || strings.HasPrefix(line, "示例") {
continue
}
// 去掉可能的序号前缀1. 2. - 等)
line = strings.TrimLeft(line, "0123456789.-) ")
line = strings.TrimSpace(line)
if len([]rune(line)) >= 4 {
memories = append(memories, line)
}
}
logger.Logger.Info("LLM memory extraction completed",
zap.Int("extracted_count", len(memories)),
)
return memories, nil
}
// FormatMemoriesForPrompt 将记忆格式化为 prompt 片段
func FormatMemoriesForPrompt(memories []model.UserMemory) string {
if len(memories) == 0 {
return ""
}
var builder strings.Builder
builder.WriteString("# 用户核心记忆\n")
for _, m := range memories {
builder.WriteString("- ")
builder.WriteString(m.Content)
builder.WriteString("\n")
}
return builder.String()
}