- 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>
395 lines
12 KiB
Go
395 lines
12 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.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
|
||
}
|