topfans/backend/services/aiChatService/provider/ai_chat_provider_test.go
zerosaturation e1326acaf9 fix(backend): service stability — bcrypt off-txn / login anti-enum / MQ stub / aiChat / event reliability / gateway aggregate (batch 3)
- 3.1 bcrypt 移出事务 (Register): repository.HashPassword 前移到 db.Transaction 之前,消除连接池占用。
- 3.2 Login 消除用户枚举 + 限流 + timing 抹平: pkg/errors 加 ErrInvalidCredential
  /ErrTooManyLoginAttempts; 用户不存在/密码错/密码空 三路径统一返回同一错误;
  mobile 5次/ip 20次 per 15min 限流 (Redis, fail-open 降级); user-not-found 走
  dummy bcrypt 抹平 ~100ms 时序差,完全消除枚举侧信道;空密码分支已核实无时序 leak。
- 3.3 MQ streams adapter 停用 → stub: 0 业务调用方, 新 stub EventProducer.Publish no-op;
  pkg/mq/mq.go Init 不再装配 streams; 全仓 grep 验证 11 处硬编码
  'gallery'/'default' 集中到 pkg/queue/consts (值不变, 仅消漂移)。
- 3.5 JWT 密钥治理: pkg/jwt MustInit fail-fast + atomic.Value (见上一个 commit 293c7b1)。
- 3.6 aiChat 健壮性: SaveContext 用 persona.ID(非 req.PersonaId); Redis/memory 错误
  记 WARN 不静默; Dify err 映射稳定用户文案,原始 err 仅服务端日志。
- 3.7 statistic.Client 重构: TrackEvent 改 buffered channel (cap 1024) + dispatchLoop
  worker; 失败 ERROR 日志带字段; drop 记 WARN; Close 可重复调用。
- 3.8 网关聚合: StarCache (60s TTL, singleflight) 替换 5+ 处 GetFanIdentities 链式调用;
  DeleteAccount 改网关直调 userService.DeleteAccount(避免改 hand-written triple.go
  风险,见报告 §5 proto 风险复盘); 铸造双写改异步 channel+consumer (3 retry)。
- 大量单测: 各子项 TDD (RED→GREEN), 关键并发 race_test (50 goroutine)。

Co-Authored-By: Claude <noreply@anthropic.com>
2026-07-23 18:50:40 +08:00

382 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"
"errors"
"io"
"net/http"
"strings"
"sync"
"testing"
"github.com/google/uuid"
"go.uber.org/zap"
triple_protocol "dubbo.apache.org/dubbo-go/v3/protocol/triple/triple_protocol"
"github.com/topfans/backend/pkg/authctx"
"github.com/topfans/backend/pkg/logger"
pb "github.com/topfans/backend/pkg/proto/ai_chat"
"github.com/topfans/backend/services/aiChatService/model"
"github.com/topfans/backend/services/aiChatService/service"
)
// TestMain / init: 在测试启动前确保 logger.Logger 已初始化,
// 避免 provider 中 logger.Logger.Warn/Error 调用触发 nil pointer.
func init() {
if logger.Logger == nil {
logger.Logger = zap.NewNop()
}
}
// ========== mock 仓库层 ==========
// mockShortTermRepo 实现 repository.ShortTermMemoryRepository
type mockShortTermRepo struct {
mu sync.Mutex
saveCalls int
savedMsgs []model.Message
savedPersID string
saveErr error
getResult []model.Message
getErr error
}
func (m *mockShortTermRepo) SaveContext(_ context.Context, _ string, msgs []model.Message, personaID string) error {
m.mu.Lock()
defer m.mu.Unlock()
m.saveCalls++
m.savedMsgs = msgs
m.savedPersID = personaID
return m.saveErr
}
func (m *mockShortTermRepo) GetContext(_ context.Context, _ string) ([]model.Message, error) {
if m.getErr != nil {
return nil, m.getErr
}
return m.getResult, nil
}
func (m *mockShortTermRepo) DeleteContext(_ context.Context, _ string) error {
return nil
}
// mockLongTermRepo 实现 repository.LongTermMemoryRepository
type mockLongTermRepo struct {
recallResult string
recallErr error
}
func (m *mockLongTermRepo) SaveMemory(_ context.Context, _ *model.UserMemory) error {
return nil
}
func (m *mockLongTermRepo) GetMemories(_ context.Context, _ int64, _ []string, _ int) ([]model.UserMemory, error) {
if m.recallErr != nil {
return nil, m.recallErr
}
return nil, nil
}
func (m *mockLongTermRepo) GetMemoriesByUserID(_ context.Context, _ int64) ([]model.UserMemory, error) {
return nil, nil
}
// mockPersonaRepo 实现 repository.PersonaRepository
type mockPersonaRepo struct {
ensureResult *model.Persona
ensureErr error
getByIDFn func(ctx context.Context, id uuid.UUID) (*model.Persona, error)
}
func (m *mockPersonaRepo) Create(_ context.Context, _ *model.Persona) error { return nil }
func (m *mockPersonaRepo) Update(_ context.Context, _ *model.Persona) error { return nil }
func (m *mockPersonaRepo) Delete(_ context.Context, _ uuid.UUID) error { return nil }
func (m *mockPersonaRepo) GetByUserID(_ context.Context, _ int64) ([]model.Persona, error) {
return nil, nil
}
func (m *mockPersonaRepo) GetDefaultByUserIDAndStarID(_ context.Context, _ int64, _ int64) (*model.Persona, error) {
return nil, nil
}
func (m *mockPersonaRepo) GetByID(ctx context.Context, id uuid.UUID) (*model.Persona, error) {
if m.getByIDFn != nil {
return m.getByIDFn(ctx, id)
}
return nil, errors.New("not found")
}
func (m *mockPersonaRepo) EnsureDefaultPersona(_ context.Context, _ int64, _ int64) (*model.Persona, error) {
if m.ensureErr != nil {
return nil, m.ensureErr
}
return m.ensureResult, nil
}
// ========== fake AIProvider用于 ChatService ==========
type fakeAIProvider struct {
streamErr error
}
func (f *fakeAIProvider) StreamChat(_ context.Context, _ []model.Message) (service.StreamReader, error) {
if f.streamErr != nil {
return nil, f.streamErr
}
return &fakeStreamReader{}, nil
}
type fakeStreamReader struct {
mu sync.Mutex
called bool
content string
}
func (f *fakeStreamReader) Next() (string, bool, error) {
f.mu.Lock()
defer f.mu.Unlock()
if f.called {
return "", true, io.EOF
}
f.called = true
return f.content, true, nil
}
func (f *fakeStreamReader) Close() error { return nil }
// ========== stub gRPC server stream ==========
// stubSendServer 实现 pb.AIChatService_SendMessageServer,
// 把每次 Send 的响应收集起来供断言使用。
// 注意: ChatMessageResponse 内含 protobuf 锁, 所以存指针避免 go vet 报 copy lock.
type stubSendServer struct {
pb.AIChatService_SendMessageServer
mu sync.Mutex
sent []*pb.ChatMessageResponse
err error
}
func (s *stubSendServer) Send(r *pb.ChatMessageResponse) error {
s.mu.Lock()
defer s.mu.Unlock()
s.sent = append(s.sent, r)
if s.err != nil {
return s.err
}
return nil
}
func (s *stubSendServer) ResponseHeader() http.Header { return http.Header{} }
func (s *stubSendServer) ResponseTrailer() http.Header { return http.Header{} }
func (s *stubSendServer) Conn() triple_protocol.StreamingHandlerConn {
return stubStreamingHandlerConn{}
}
// stubStreamingHandlerConn 实现 triple_protocol.StreamingHandlerConn.
// provider 的业务路径不会调用, 这里仅满足接口签名.
type stubStreamingHandlerConn struct{}
func (stubStreamingHandlerConn) Spec() triple_protocol.Spec { return triple_protocol.Spec{} }
func (stubStreamingHandlerConn) Peer() triple_protocol.Peer { return triple_protocol.Peer{} }
func (stubStreamingHandlerConn) Receive(any) error { return nil }
func (stubStreamingHandlerConn) RequestHeader() http.Header { return http.Header{} }
func (stubStreamingHandlerConn) ExportableHeader() http.Header { return http.Header{} }
func (stubStreamingHandlerConn) Send(any) error { return nil }
func (stubStreamingHandlerConn) ResponseHeader() http.Header { return http.Header{} }
func (stubStreamingHandlerConn) ResponseTrailer() http.Header { return http.Header{} }
// ========== helper构造一个装配好的 AIChatProvider ==========
func buildProvider(
shortRepo *mockShortTermRepo,
longRepo *mockLongTermRepo,
personaRepo *mockPersonaRepo,
aiProvider service.AIProvider,
) *AIChatProvider {
memSvc := service.NewMemoryService(shortRepo, longRepo)
personaSvc := service.NewPersonaService(personaRepo)
auditSvc := service.NewAuditService()
chatSvc := service.NewChatService(aiProvider, personaSvc, memSvc, auditSvc)
return NewAIChatProvider(chatSvc, personaSvc, memSvc, auditSvc)
}
// resolvePersonaID 是 provider 中提取的 helper 函数。
// 单元测试目标: 解析后的 persona.ID 非空时, 必须用它而不是 req.PersonaId。
func TestResolvePersonaID_PrefersResolved(t *testing.T) {
if got := resolvePersonaID("from-request", "from-persona-object"); got != "from-persona-object" {
t.Fatalf("should prefer resolved persona.ID, got %q", got)
}
if got := resolvePersonaID("", "from-persona-object"); got != "from-persona-object" {
t.Fatalf("empty request should fall back to persona.ID, got %q", got)
}
// resolved 为空时回落到 request 值(向后兼容: 请求显式给空且 DB 没找到人设的情况)
if got := resolvePersonaID("from-request", ""); got != "from-request" {
t.Fatalf("empty resolved should fall back to req, got %q", got)
}
}
// TestSendMessage_UsesResolvedPersonaID 验证 L250-252 fix:
// SaveContext 必须用 persona.ID(由 GetPersonaOrDefault 返回), 不是 req.PersonaId。
func TestSendMessage_UsesResolvedPersonaID(t *testing.T) {
resolvedID := uuid.New()
// 用固定的 UUID 让 GetByID 一定能找到, 即便 req.PersonaId 传了别的值。
targetUUID, _ := uuid.Parse("11111111-1111-1111-1111-111111111111")
personaRepo := &mockPersonaRepo{
ensureResult: &model.Persona{
ID: resolvedID,
UserID: 42,
StarID: 7,
Name: "test-persona",
SystemPrompt: "you are a test persona",
},
// 让 GetByID 返回 resolvedID 对应的人设
getByIDFn: func(_ context.Context, id uuid.UUID) (*model.Persona, error) {
if id == targetUUID {
return &model.Persona{
ID: resolvedID,
UserID: 42,
StarID: 7,
Name: "resolved",
SystemPrompt: "resolved prompt",
}, nil
}
return nil, errors.New("not found")
},
}
shortRepo := &mockShortTermRepo{}
longRepo := &mockLongTermRepo{}
prov := buildProvider(shortRepo, longRepo, personaRepo, &fakeAIProvider{})
stream := &stubSendServer{}
ctx := authctx.WithIdentity(context.Background(), 42, 7)
req := &pb.ChatMessageRequest{
SessionId: "test-session",
Message: "你好",
PersonaId: targetUUID.String(), // 请求里给一个 UUID
UserId: 42,
}
if err := prov.SendMessage(ctx, req, stream); err != nil {
t.Fatalf("SendMessage returned error: %v", err)
}
shortRepo.mu.Lock()
defer shortRepo.mu.Unlock()
if shortRepo.saveCalls == 0 {
t.Fatalf("expected SaveContext to be called, got 0 calls")
}
if shortRepo.savedPersID != resolvedID.String() {
t.Fatalf("expected SaveContext called with resolved persona.ID=%q, got %q",
resolvedID.String(), shortRepo.savedPersID)
}
}
// TestSendMessage_RedisErrorsLoggedNotFatal 验证 L138/141 fix:
// RecallMemories / GetContext 返回 error 时 provider 不应崩溃,
// 也不应阻塞后续 LLM 调用。
func TestSendMessage_RedisErrorsLoggedNotFatal(t *testing.T) {
resolvedID := uuid.New()
personaRepo := &mockPersonaRepo{
ensureResult: &model.Persona{
ID: resolvedID,
UserID: 42,
StarID: 7,
Name: "test",
SystemPrompt: "test prompt",
},
}
shortRepo := &mockShortTermRepo{
getErr: errors.New("redis: connection refused"),
}
longRepo := &mockLongTermRepo{
recallErr: errors.New("redis: connection refused"),
}
prov := buildProvider(shortRepo, longRepo, personaRepo, &fakeAIProvider{})
stream := &stubSendServer{}
ctx := authctx.WithIdentity(context.Background(), 42, 7)
req := &pb.ChatMessageRequest{
SessionId: "redis-err-session",
Message: "你好",
UserId: 42,
}
// 即便 RecallMemories / GetContext 都失败, SendMessage 仍应正常返回 nil,
// 因为内存/历史是辅助数据, 不应阻断主流程。
if err := prov.SendMessage(ctx, req, stream); err != nil {
t.Fatalf("SendMessage should not return error when Redis fails, got: %v", err)
}
// 流应该至少包含 IsEnd=true 的最终消息(成功 chat 路径)
if len(stream.sent) == 0 {
t.Fatalf("expected at least one streamed response")
}
var sawEnd bool
for _, r := range stream.sent {
if r.IsEnd {
sawEnd = true
break
}
}
if !sawEnd {
t.Fatalf("expected IsEnd=true in stream, sent=%+v", stream.sent)
}
}
// TestSendMessage_RedactsErrorToClient 验证 L188 fix:
// 当 ChatService.StreamChat 返回非敏感内容错误时,
// 客户端响应里的 Error 字段必须是固定枚举(不能泄露 err.Error())。
func TestSendMessage_RedactsErrorToClient(t *testing.T) {
resolvedID := uuid.New()
personaRepo := &mockPersonaRepo{
ensureResult: &model.Persona{
ID: resolvedID,
UserID: 42,
StarID: 7,
Name: "test",
SystemPrompt: "test prompt",
},
}
shortRepo := &mockShortTermRepo{}
longRepo := &mockLongTermRepo{}
// 让 StreamChat 返回一个含敏感 URL / 内部信息的 error
leakErr := errors.New("dial tcp 10.0.0.1:443: i/o timeout with internal stack trace")
prov := buildProvider(shortRepo, longRepo, personaRepo, &fakeAIProvider{streamErr: leakErr})
stream := &stubSendServer{}
ctx := authctx.WithIdentity(context.Background(), 42, 7)
req := &pb.ChatMessageRequest{
SessionId: "err-redact-session",
Message: "你好",
UserId: 42,
}
err := prov.SendMessage(ctx, req, stream)
if err == nil {
t.Fatalf("expected SendMessage to return underlying error for grpc layer, got nil")
}
// 在所有发送的响应里, Error 字段都不应包含原始 err.Error() 中的敏感字符串
// 且应包含固定文案("internal_error" 等)。
var errorFieldSeen string
for _, r := range stream.sent {
if r.Error != "" {
errorFieldSeen = r.Error
}
if strings.Contains(r.Error, "10.0.0.1") || strings.Contains(r.Error, "i/o timeout") ||
strings.Contains(r.Error, "dial tcp") {
t.Fatalf("client-visible Error field leaks internal error: %q", r.Error)
}
}
if errorFieldSeen == "" {
t.Fatalf("expected Error field on terminal response, got nothing in stream=%+v", stream.sent)
}
}