topfans/backend/services/aiChatService/provider/ai_chat_provider.go
zerosaturation 8a767fb400 fix(backend): service stability — batch 3 accumulated (bcrypt off-txn / login anti-enum / MQ stub / aiChat / event reliability / gateway aggregate)
- 3.1 bcrypt 移出事务 (Register): repository.HashPassword 前移到 db.Transaction 之前。
- 3.2 Login 消除用户枚举 + 限流 + timing 抹平: pkg/errors 加 ErrInvalidCredential / ErrTooManyLoginAttempts;
  user-not-found 跑 dummy bcrypt 抹平 ~100ms 时序差; mobile 5次/ip 20次 per 15min 限流 (Redis, fail-open 降级)。
- 3.3 MQ streams adapter 停用 → stub: 0 业务调用方, noop EventProducer.Publish; Init 不再装配 streams;
  11 处硬编码 'gallery'/'default' 抽常量到 pkg/queue/consts (值不变, 消漂移)。
- 3.5 JWT 密钥治理: pkg/jwt MustInit fail-fast + atomic.Value, 50-goroutine race_test 零告警;
  MustInit 调用点 gateway main + auth_provider + loadgen 同步更新。
- 3.6 aiChat 健壮性: SaveContext 用 persona.ID(非 req.PersonaId); Redis/memory 错误 记 WARN 不静默;
  Dify err 映射稳定用户文案。
- 3.7 statistic.Client 重构: TrackEvent 改 buffered channel (cap 1024) + dispatchLoop worker。
- 3.8 网关聚合: StarCache (60s TTL, singleflight) 替换 GetFanIdentities 链式调用;
  DeleteAccount 改网关直调 userService.DeleteAccount(避免改 hand-written triple.go);
  铸造双写改异步 channel+consumer (3 retry)。
- 大量单测: 各子项 TDD (RED→GREEN), 关键并发 race_test (50 goroutine)。
- .env.example JWT_SECRET 改为 ≥32 字节 base64 示例(原为空, 被 MustInit 立即拒)。

Co-Authored-By: Claude <noreply@anthropic.com>
2026-07-24 14:03:21 +08:00

395 lines
12 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.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_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
}
// resolvePersonaID 决定 SaveContext 应该用哪个 personaID。
// 优先使用 GetPersonaOrDefault 解析后的 persona.ID保证关联到真实可用的人设
// 仅在解析结果为空时回落到请求里的 personaID向后兼容
func resolvePersonaID(reqID, resolvedID string) string {
if resolvedID != "" {
return resolvedID
}
return reqID
}