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() }