topfans/backend/services/aiChatService/provider/ai_chat_provider.go
zerosaturation b7022d2dc1 fix(backend): auth boundary — trusted identity from ctx (batch 2)
- 新增 pkg/authctx: 从 Dubbo attachment/gRPC metadata 提取可信 user_id/star_id,
  统一覆盖 req 同名字段, 缺身份返 Unauthenticated。
- 各 provider 接入(堵身份伪造/越权):
  * moderation SubmitReport 等 6 RPC(举报人伪造)
  * asset CheckAssetLike/GetAssetQrcode/TrackShare(点赞/分享归因伪造)
  * social CheckFriendship(修 starID=0 隐私预言机)
  * activity PurchaseItem/BatchPurchaseItem(水晶扣费伪造)等 5 RPC
  * gallery/aiChat/task/notification 迁移 authctx, 删散落 extractUserInfo*
- social 正确性: GetUserLikedAssets OR 显式分组(防御); GetRandomUsersByStar 真随机(去 rand.Seed)。
- gateway: /auth/validate 移入 AuthMiddleware 保护组(/refresh 保留,依赖注入身份)。
- 删 userService 已迁移死函数; notification 缺身份错误码统一为 Unauthenticated。
- 各 provider 单测(伪造身份被覆盖 + 缺身份拒绝)。

Co-Authored-By: Claude <noreply@anthropic.com>
2026-07-23 01:12:35 +08:00

364 lines
10 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 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.String("message", message),
)
// 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. 记忆召回
memoryText, _ := p.memoryService.RecallMemories(ctx, userID, message, 5)
// 4. 获取对话历史
history, _ := p.memoryService.GetContext(ctx, sessionID)
// 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
}
// 其他错误
stream.Send(&pb.ChatMessageResponse{
Content: "抱歉,服务暂时不可用",
SessionId: sessionID,
IsEnd: true,
Error: err.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. 保存上下文
newHistory := append(history, model.Message{Role: "user", Content: message})
newHistory = append(newHistory, model.Message{Role: "assistant", Content: fullResponse})
p.memoryService.SaveContext(ctx, sessionID, newHistory, personaID)
// 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_idservice.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
}