- 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>
382 lines
12 KiB
Go
382 lines
12 KiB
Go
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)
|
||
}
|
||
} |