321 lines
9.6 KiB
Go
321 lines
9.6 KiB
Go
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()
|
||
}
|
||
|