- 新增 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>
364 lines
10 KiB
Go
364 lines
10 KiB
Go
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_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
|
||
}
|