package provider import ( "context" "fmt" "io" "github.com/topfans/backend/pkg/authctx" "github.com/topfans/backend/pkg/logger" "github.com/topfans/backend/services/aiChatService/model" "github.com/topfans/backend/services/aiChatService/service" pb "github.com/topfans/backend/pkg/proto/ai_chat" "go.uber.org/zap" "google.golang.org/grpc/codes" "google.golang.org/grpc/status" ) // memoryExtractionInterval 记忆提取的触发轮数间隔。 // 每 N 轮对话触发一次 LLM 记忆提取(第 N, 2N, 3N... 轮)。 const memoryExtractionInterval = 5 // AIChatProvider AI Chat 服务 Provider 实现 type AIChatProvider struct { chatService *service.ChatService personaService *service.PersonaService memoryService *service.MemoryService auditService *service.AuditService } // 确保 AIChatProvider 实现了 AIChatServiceHandler 接口 var _ pb.AIChatServiceHandler = (*AIChatProvider)(nil) // NewAIChatProvider 创建 AIChatProvider 实例 func NewAIChatProvider( chatService *service.ChatService, personaService *service.PersonaService, memoryService *service.MemoryService, auditService *service.AuditService, ) *AIChatProvider { return &AIChatProvider{ chatService: chatService, personaService: personaService, memoryService: memoryService, auditService: auditService, } } // InitSession 初始化会话,返回欢迎消息(同时创建当前明星的默认人设) func (p *AIChatProvider) InitSession(ctx context.Context, req *pb.InitSessionRequest) (*pb.InitSessionResponse, error) { // 身份必须来自 ctx 里的可信值(auth interceptor 从 JWT 解析后灌入)。 userID, starID, err := authctx.ExtractIdentity(ctx) if err != nil { logger.Logger.Warn("InitSession missing trusted identity", zap.Error(err)) return nil, status.Error(codes.Unauthenticated, "identity required") } sessionID := req.SessionId if sessionID == "" { sessionID = fmt.Sprintf("%d_%d", userID, starID) } logger.Logger.Info("Received InitSession request", zap.Int64("user_id", userID), zap.Int64("star_id", starID), zap.String("session_id", sessionID), ) // 进入聊天即创建当前明星的默认人设,防止出现 A 明星用 B 人设的情况 if _, err := p.personaService.GetPersonaOrDefault(ctx, userID, starID, ""); err != nil { logger.Logger.Warn("Failed to ensure default persona in InitSession", zap.Int64("user_id", userID), zap.Int64("star_id", starID), zap.Error(err), ) // 不阻塞进入聊天,SendMessage 时会再次尝试 } // 获取欢迎消息 welcomeMessage := p.chatService.GetWelcomeMessage(sessionID, userID, starID) return &pb.InitSessionResponse{ WelcomeMessage: welcomeMessage, SessionId: sessionID, }, nil } // SendMessage 发送消息(流式返回) func (p *AIChatProvider) SendMessage(ctx context.Context, req *pb.ChatMessageRequest, stream pb.AIChatService_SendMessageServer) error { // 身份必须来自 ctx 里的可信值(auth interceptor 从 JWT 解析后灌入)。 userID, starID, err := authctx.ExtractIdentity(ctx) if err != nil { logger.Logger.Warn("SendMessage missing trusted identity", zap.Error(err)) return status.Error(codes.Unauthenticated, "identity required") } sessionID := req.SessionId if sessionID == "" { sessionID = fmt.Sprintf("%d_%d", userID, starID) } message := req.Message personaID := req.PersonaId logger.Logger.Info("Received SendMessage request", zap.Int64("user_id", userID), zap.String("session_id", sessionID), zap.Int("message_len", len(message)), ) // 不打印 message 原文以防 PII 泄露;长度足够排查空包 / 超长包。 // 1. 前置审核 if !p.auditService.AuditText(message) { logger.Logger.Info("Message blocked by audit") stream.Send(&pb.ChatMessageResponse{ Content: p.auditService.DefaultSafeResponse(), SessionId: sessionID, IsEnd: false, }) stream.Send(&pb.ChatMessageResponse{ SessionId: sessionID, IsEnd: true, }) return nil } // 2. 获取人设(传入 starID 用于首次使用时创建明星专属人设) persona, err := p.personaService.GetPersonaOrDefault(ctx, userID, starID, personaID) if err != nil { logger.Logger.Error("Failed to get persona", zap.Error(err)) stream.Send(&pb.ChatMessageResponse{ Content: err.Error(), IsEnd: true, Error: err.Error(), }) return err } // 3. 记忆召回(失败降级到空, 记 WARN, 不阻断 chat) memoryText, err := p.memoryService.RecallMemories(ctx, userID, message, 5) if err != nil { logger.Logger.Warn("RecallMemories failed, continuing with empty memory", zap.Int64("user_id", userID), zap.Error(err)) memoryText = "" } // 4. 获取对话历史(失败降级到空, 记 WARN, 不阻断 chat) history, err := p.memoryService.GetContext(ctx, sessionID) if err != nil { logger.Logger.Warn("GetContext failed, continuing with empty history", zap.Int64("user_id", userID), zap.String("session_id", sessionID), zap.Error(err)) history = nil } // 5. 构建 Prompt tokenizer := &service.Tokenizer{} messages, _ := service.BuildPrompt( persona.SystemPrompt, memoryText, history, message, tokenizer, ) // 6. 检查是否需要调用大模型 if service.IsNoNeedLLMCall(message) { stream.Send(&pb.ChatMessageResponse{ Content: "好的,我听到了。", SessionId: sessionID, IsEnd: false, }) stream.Send(&pb.ChatMessageResponse{ SessionId: sessionID, IsEnd: true, }) return nil } // 7. 调用大模型(流式) streamReader, err := p.chatService.StreamChat(ctx, messages) if err != nil { logger.Logger.Error("AI call failed", zap.Error(err)) // 检查是否是敏感内容错误 if _, ok := err.(*service.SensitiveContentError); ok { logger.Logger.Info("Content blocked by safety filter") stream.Send(&pb.ChatMessageResponse{ Content: p.auditService.DefaultSafeResponse(), SessionId: sessionID, IsEnd: true, }) return nil } // 其他错误 - 不暴露原始 err.Error() 给客户端(可能含内部 URL / stack trace), // 详细错误仅留服务端日志, 客户端拿到稳定的内部错误码。 stream.Send(&pb.ChatMessageResponse{ Content: "抱歉,服务暂时不可用,请稍后重试", SessionId: sessionID, IsEnd: true, Error: "internal_error", }) return err } defer streamReader.Close() // 8. 流式处理 var fullResponse string var sentEnd = false for { content, done, err := streamReader.Next() if err != nil { if err == io.EOF { // 流结束,发送 is_end if !sentEnd { stream.Send(&pb.ChatMessageResponse{ SessionId: sessionID, IsEnd: true, }) sentEnd = true } break } logger.Logger.Error("Stream read error", zap.Error(err)) break } // 后置审核(逐 token) if !p.auditService.AuditResponse(content) { logger.Logger.Info("Response blocked by audit") streamReader.Close() // 发送安全回复作为替代 stream.Send(&pb.ChatMessageResponse{ Content: p.auditService.DefaultSafeResponse(), SessionId: sessionID, IsEnd: false, }) stream.Send(&pb.ChatMessageResponse{ SessionId: sessionID, IsEnd: true, }) sentEnd = true return nil } fullResponse += content // 发送 token 给客户端 if err := stream.Send(&pb.ChatMessageResponse{ Content: content, SessionId: sessionID, IsEnd: done, }); err != nil { logger.Logger.Error("Failed to send message to stream", zap.Error(err)) return err } if done { sentEnd = true } } // 9. 保存上下文 - 用解析后的 persona.ID 而非请求里的 personaID, // 保证历史记录关联到 GetPersonaOrDefault 实际返回的人设, // 避免请求给空 / 错误 personaID 时关联错乱。 resolvedPersonaID := resolvePersonaID(personaID, persona.ID.String()) newHistory := append(history, model.Message{Role: "user", Content: message}) newHistory = append(newHistory, model.Message{Role: "assistant", Content: fullResponse}) if err := p.memoryService.SaveContext(ctx, sessionID, newHistory, resolvedPersonaID); err != nil { logger.Logger.Warn("SaveContext failed", zap.Int64("user_id", userID), zap.String("session_id", sessionID), zap.Error(err), ) } // 10. 触发记忆提取(每5轮触发一次,在第5、10、15...轮提取) newTurns := len(newHistory) / 2 shouldExtract := newTurns >= memoryExtractionInterval && newTurns%memoryExtractionInterval == 0 logger.Logger.Info("Memory extraction check", zap.Int("message_count", len(newHistory)), zap.Int("turns", newTurns), zap.Bool("should_extract", shouldExtract), ) if shouldExtract { logger.Logger.Info("Triggering LLM memory extraction", zap.Int64("user_id", userID)) // 使用 LLM 提取记忆 extractedMemories, err := service.ExtractMemoriesWithLLM(ctx, p.chatService.Chat, newHistory) if err != nil { logger.Logger.Error("LLM memory extraction failed", zap.Error(err)) } else if len(extractedMemories) > 0 { if err := p.memoryService.ExtractMemory(ctx, userID, newHistory, extractedMemories); err != nil { logger.Logger.Error("Failed to save extracted memories", zap.Error(err)) } else { logger.Logger.Info("Memories extracted and saved successfully", zap.Int64("user_id", userID), zap.Int("count", len(extractedMemories)), ) } } } logger.Logger.Info("SendMessage completed", zap.Int64("user_id", userID), zap.String("session_id", sessionID), zap.Int("response_length", len(fullResponse)), ) return nil } // GetHistory 获取对话历史 func (p *AIChatProvider) GetHistory(ctx context.Context, req *pb.ChatHistoryRequest) (*pb.ChatHistoryResponse, error) { // 身份必须来自 ctx 里的可信值(auth interceptor 从 JWT 解析后灌入)。 userID, starID, err := authctx.ExtractIdentity(ctx) if err != nil { logger.Logger.Warn("GetHistory missing trusted identity", zap.Error(err)) return nil, status.Error(codes.Unauthenticated, "identity required") } sessionID := req.SessionId if sessionID == "" { sessionID = fmt.Sprintf("%d_%d", userID, starID) } logger.Logger.Info("Received GetHistory request", zap.Int64("user_id", userID), zap.String("session_id", sessionID), ) messages, err := p.memoryService.GetContext(ctx, sessionID) if err != nil { return nil, err } pbMessages := make([]*pb.Message, len(messages)) for i, m := range messages { pbMessages[i] = &pb.Message{ Role: m.Role, Content: m.Content, } } return &pb.ChatHistoryResponse{ History: pbMessages, }, nil } // GetPersonas 获取用户的所有人设 func (p *AIChatProvider) GetPersonas(ctx context.Context, req *pb.GetPersonasRequest) (*pb.PersonaListResponse, error) { // 身份必须来自 ctx 里的可信值(auth interceptor 从 JWT 解析后灌入)。 // GetPersonas 只用 user_id(service.GetPersonas 不绑定 star),用 ExtractUserID。 userID, err := authctx.ExtractUserID(ctx) if err != nil { logger.Logger.Warn("GetPersonas missing trusted identity", zap.Error(err)) return nil, status.Error(codes.Unauthenticated, "identity required") } logger.Logger.Info("Received GetPersonas request", zap.Int64("user_id", userID), ) personas, err := p.personaService.GetPersonas(ctx, userID) if err != nil { return nil, err } pbPersonas := make([]*pb.PersonaInfo, len(personas)) for i, persona := range personas { pbPersonas[i] = &pb.PersonaInfo{ Id: persona.ID, StarId: persona.StarID, Name: persona.Name, Description: persona.Description, AvatarUrl: persona.AvatarURL, TalkStyle: persona.TalkStyle, IsDefault: persona.IsDefault, CreatedAt: persona.CreatedAt, UpdatedAt: persona.UpdatedAt, } } return &pb.PersonaListResponse{ Personas: pbPersonas, }, nil } // resolvePersonaID 决定 SaveContext 应该用哪个 personaID。 // 优先使用 GetPersonaOrDefault 解析后的 persona.ID(保证关联到真实可用的人设), // 仅在解析结果为空时回落到请求里的 personaID(向后兼容)。 func resolvePersonaID(reqID, resolvedID string) string { if resolvedID != "" { return resolvedID } return reqID }