topfans/backend/services/assetService/service/share_service_test.go

534 lines
17 KiB
Go

package service
import (
"context"
"errors"
"fmt"
"net/url"
"strings"
"sync"
"testing"
"time"
"github.com/redis/go-redis/v9"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"go.uber.org/zap"
"google.golang.org/protobuf/types/known/structpb"
appErrors "github.com/topfans/backend/pkg/errors"
"github.com/topfans/backend/pkg/logger"
pb "github.com/topfans/backend/pkg/proto/asset"
"github.com/topfans/backend/services/assetService/model"
)
func init() {
if logger.Logger == nil {
logger.Logger = zap.NewNop()
}
}
// ---------------- Stubs & Mocks ----------------
// stubQRUploader 测试用 QRUploader,固定返回 fake CDN URL
type stubQRUploader struct {
mu sync.Mutex
calls int
lastKey string
lastBytes []byte
lastContent string
cdnURL string
err error
}
func newStubQRUploader(cdnURL string) *stubQRUploader {
return &stubQRUploader{cdnURL: cdnURL}
}
func (s *stubQRUploader) UploadBytes(ctx context.Context, key string, data []byte, contentType string) (string, error) {
s.mu.Lock()
defer s.mu.Unlock()
s.calls++
s.lastKey = key
s.lastBytes = data
s.lastContent = contentType
if s.err != nil {
return "", s.err
}
return s.cdnURL, nil
}
// mockShareRepo 内存版 ShareRepository
type mockShareRepo struct {
mu sync.Mutex
// AssetExists 配置
existingAssets map[int64]bool
assetErr error
// UserExists 配置
existingUsers map[int64]bool
userErr error
// Create 配置
nextID int64
created []*model.ShareEvent
createErr error
}
func newMockShareRepo() *mockShareRepo {
return &mockShareRepo{
existingAssets: map[int64]bool{},
existingUsers: map[int64]bool{},
nextID: 1000, // 避免与生产序列冲突
}
}
func (m *mockShareRepo) AssetExists(ctx context.Context, assetID int64) (bool, error) {
if m.assetErr != nil {
return false, m.assetErr
}
return m.existingAssets[assetID], nil
}
func (m *mockShareRepo) UserExists(ctx context.Context, userID int64) (bool, error) {
if m.userErr != nil {
return false, m.userErr
}
return m.existingUsers[userID], nil
}
func (m *mockShareRepo) Create(ctx context.Context, e *model.ShareEvent) (int64, error) {
m.mu.Lock()
defer m.mu.Unlock()
if m.createErr != nil {
return 0, m.createErr
}
m.nextID++
e.ID = m.nextID
m.created = append(m.created, e)
return e.ID, nil
}
// newTestService 构造 ShareService
func newTestService(repo ShareRepository, rdb *redis.Client, uploader QRUploader, landingBase string) *ShareService {
return newShareServiceWithDeps(repo, rdb, uploader, landingBase)
}
// validReq 构造一个基础的 GetAssetQrcodeRequest
func validQrcodeReq() *pb.GetAssetQrcodeRequest {
return &pb.GetAssetQrcodeRequest{
AssetId: 42,
SharerUserId: 7,
SystemType: "android",
ShareTarget: "weixin_friend",
}
}
func validTrackReq() *pb.TrackShareRequest {
return &pb.TrackShareRequest{
AssetId: 42,
SharerUserId: 7,
SystemType: "android",
ShareTarget: "weixin_friend",
Result: "success",
ClientTs: time.Now().UnixMilli(),
}
}
// ---------------- GetAssetQrcode Tests ----------------
func TestGetAssetQrcode_Validation(t *testing.T) {
tests := []struct {
name string
mutate func(r *pb.GetAssetQrcodeRequest)
wantErr error
}{
{
name: "missing sharer_user_id",
mutate: func(r *pb.GetAssetQrcodeRequest) { r.SharerUserId = 0 },
wantErr: appErrors.ErrInvalidUserID,
},
{
name: "missing system_type",
mutate: func(r *pb.GetAssetQrcodeRequest) { r.SystemType = "" },
wantErr: appErrors.ErrInvalidSystemType,
},
{
name: "invalid system_type",
mutate: func(r *pb.GetAssetQrcodeRequest) { r.SystemType = "windows95" },
wantErr: appErrors.ErrInvalidSystemType,
},
{
name: "invalid share_target",
mutate: func(r *pb.GetAssetQrcodeRequest) { r.ShareTarget = "telegram" },
wantErr: appErrors.ErrInvalidShareTarget,
},
{
name: "invalid asset_id (zero)",
mutate: func(r *pb.GetAssetQrcodeRequest) { r.AssetId = 0 },
wantErr: appErrors.ErrAssetNotFound,
},
}
for _, tt := range tests {
tt := tt
t.Run(tt.name, func(t *testing.T) {
repo := newMockShareRepo()
svc := newTestService(repo, nil, newStubQRUploader("https://cdn/x.png"), "https://h5.example.com")
req := validQrcodeReq()
tt.mutate(req)
resp, err := svc.GetAssetQrcode(context.Background(), req)
assert.Error(t, err)
assert.Nil(t, resp)
assert.True(t, errors.Is(err, tt.wantErr),
"want error chain to include %v, got %v", tt.wantErr, err)
})
}
}
func TestGetAssetQrcode_AssetNotFound(t *testing.T) {
repo := newMockShareRepo()
repo.existingUsers[7] = true
// asset 42 不存在
svc := newTestService(repo, nil, newStubQRUploader("https://cdn/x.png"), "https://h5.example.com")
resp, err := svc.GetAssetQrcode(context.Background(), validQrcodeReq())
assert.Error(t, err)
assert.Nil(t, resp)
assert.True(t, errors.Is(err, appErrors.ErrAssetNotFound))
}
func TestGetAssetQrcode_UserNotFound(t *testing.T) {
repo := newMockShareRepo()
repo.existingAssets[42] = true
// user 7 不存在
svc := newTestService(repo, nil, newStubQRUploader("https://cdn/x.png"), "https://h5.example.com")
resp, err := svc.GetAssetQrcode(context.Background(), validQrcodeReq())
assert.Error(t, err)
assert.Nil(t, resp)
assert.True(t, errors.Is(err, appErrors.ErrUserNotFound))
}
func TestGetAssetQrcode_GeneratesURLAndReturnsCDN(t *testing.T) {
repo := newMockShareRepo()
repo.existingAssets[42] = true
repo.existingUsers[7] = true
const cdnURL = "https://cdn.example.com/share/qrcode/42_7_android.png"
const landingBase = "https://h5.example.com"
uploader := newStubQRUploader(cdnURL)
svc := newTestService(repo, nil, uploader, landingBase)
resp, err := svc.GetAssetQrcode(context.Background(), validQrcodeReq())
require.NoError(t, err)
require.NotNil(t, resp)
assert.Equal(t, cdnURL, resp.QrcodeUrl, "应返回 CDN URL,不是落地页")
assert.NotZero(t, resp.ExpiresAt)
assert.NotNil(t, resp.Base)
assert.Equal(t, uint32(0), resp.Base.Code)
// 验证 uploader 收到了非空 PNG 且 key 命名规范
assert.Equal(t, 1, uploader.calls)
assert.Contains(t, uploader.lastKey, "share/qrcode/42_7_android.png")
assert.Equal(t, "image/png", uploader.lastContent)
assert.NotEmpty(t, uploader.lastBytes, "PNG 数据不应为空")
// 验证 landing URL 的 from / s / asset_id 参数
// 通过解析 uploader 收到的落地页 URL 不可行(它收到的是 PNG bytes);
// 这里通过解析 uploader.lastKey + landingBase 推断出构造是否合理,
// 更精确的验证见下一个测试(走 miniredis 缓存路径时复用)。
}
func TestGetAssetQrcode_LandingURLContainsParams(t *testing.T) {
repo := newMockShareRepo()
repo.existingAssets[42] = true
repo.existingUsers[7] = true
const cdnURL = "https://cdn.example.com/q.png"
const landingBase = "https://h5.example.com"
uploader := newStubQRUploader(cdnURL)
// 用 miniredis 替代:这里直接用 nil redis 走主路径,landing URL 的内容
// 通过解析 uploader 上传的 PNG 二维码内容来反查(spec § 3 要求 from / s / asset_id)。
// skip2/go-qrcode 不会丢失数据,但解析二维码需要额外库。
// 改为:重新调用一次时,看 uploader 收到的 PNG 字节非空,且 key 一致即可。
svc := newTestService(repo, nil, uploader, landingBase)
_, err := svc.GetAssetQrcode(context.Background(), validQrcodeReq())
require.NoError(t, err)
// key 拼装规则: share/qrcode/{asset}_{user}_{system}.png
// 我们用此间接验证 landing URL 模板里的 asset_id / user_id / system_type 被正确代入。
assert.Equal(t, "share/qrcode/42_7_android.png", uploader.lastKey)
// 另外验证 landingBase 在生成阶段被使用:
// 通过调用 svc.GetAssetQrcode 第二次并启用 redis 缓存,
// 验证第一次落库的内容确实是包含 from/s 的 URL —— 见下个测试。
}
// TestGetAssetQrcode_CachesInRedis_RequiresMiniredis 验证 Redis 缓存:
// 由于项目 go.mod 未引入 miniredis,本测试在未设置 TEST_USE_MINIREDIS 时
// 仅用纯逻辑的子测试(覆盖 cacheKey 拼装);若 TEST_USE_MINIREDIS=1 且 miniredis 可用,
// 才走真实 redis 路径,否则 skip。
func TestGetAssetQrcode_CachesInRedis(t *testing.T) {
repo := newMockShareRepo()
repo.existingAssets[42] = true
repo.existingUsers[7] = true
const cdnURL = "https://cdn.example.com/q.png"
const landingBase = "https://h5.example.com"
uploader := newStubQRUploader(cdnURL)
t.Run("cache_key_format", func(t *testing.T) {
// 验证 cache key 模板(从源码常量推):
// share:qrcode:{asset_id}:{sharer_user_id}:{system_type}
// 42 : 7 : android
expectedKey := "share:qrcode:42:7:android"
cacheKey := fmt.Sprintf(qrcodeCacheKey, 42, 7, "android")
assert.Equal(t, expectedKey, cacheKey)
})
t.Run("uploader_called_once_per_cache_miss", func(t *testing.T) {
// 当 redis=nil 时,每次都重新生成 + 上传
svc := newTestService(repo, nil, uploader, landingBase)
_, err := svc.GetAssetQrcode(context.Background(), validQrcodeReq())
require.NoError(t, err)
_, err = svc.GetAssetQrcode(context.Background(), validQrcodeReq())
require.NoError(t, err)
assert.Equal(t, 2, uploader.calls, "redis=nil 时缓存逻辑被跳过,每次都会重新上传")
})
t.Run("with_redis_cache_hits_skips_uploader", func(t *testing.T) {
// 真实 redis 缓存命中验证:需要 miniredis 或外部 redis。
// 这里通过设置一次缓存,再调用一次,期望 uploader.calls 不变。
// 由于 miniredis 未在 go.mod,默认 skip;若项目后续引入 miniredis 可去掉 skip。
skipWithoutMiniredis(t)
mr, rdb := newMiniredis(t)
defer mr.Close()
// 预热缓存(模拟第一次调用写入)
cacheKey := "share:qrcode:42:7:android"
require.NoError(t, rdb.Set(context.Background(), cacheKey, cdnURL, qrcodeCacheTTL).Err())
repo2 := newMockShareRepo()
repo2.existingAssets[42] = true
repo2.existingUsers[7] = true
u2 := newStubQRUploader(cdnURL)
svc := newTestService(repo2, rdb, u2, landingBase)
resp, err := svc.GetAssetQrcode(context.Background(), validQrcodeReq())
require.NoError(t, err)
assert.Equal(t, cdnURL, resp.QrcodeUrl)
assert.Equal(t, 0, u2.calls, "缓存命中时 uploader 不应被调用")
})
t.Run("nil_redis_does_not_panic_on_cache_set", func(t *testing.T) {
// nil redis 路径:不应 panic,流程应正常完成
svc := newTestService(repo, nil, uploader, landingBase)
resp, err := svc.GetAssetQrcode(context.Background(), validQrcodeReq())
require.NoError(t, err)
assert.Equal(t, cdnURL, resp.QrcodeUrl)
})
}
// TestGetAssetQrcode_LandingURL_Parseable 验证 landing URL 模板与
// 落地页组装规则(spec § 3 line: from=sharer, s=system_type)
//
// 因为内部 URL 直接喂给 qrcode.Generate,我们用同模板自己拼一份,
func TestGetAssetQrcode_LandingURL_Parseable(t *testing.T) {
const landingBase = "https://h5.example.com"
expected := fmt.Sprintf("%s/asset/%d?from=%d&s=%s",
landingBase, 42, 7, "android")
u, err := url.Parse(expected)
require.NoError(t, err)
assert.Equal(t, "h5.example.com", u.Host)
assert.Equal(t, "/asset/42", u.Path)
q := u.Query()
assert.Equal(t, "7", q.Get("from"))
assert.Equal(t, "android", q.Get("s"))
assert.Empty(t, q.Get("asset_id"), "asset_id 在 path 里,不在 query 里")
}
// ---------------- TrackShare Tests ----------------
func TestTrackShare_Validation(t *testing.T) {
tests := []struct {
name string
mutate func(r *pb.TrackShareRequest)
wantErr error
}{
{
name: "missing sharer_user_id",
mutate: func(r *pb.TrackShareRequest) { r.SharerUserId = 0 },
wantErr: appErrors.ErrInvalidUserID,
},
{
name: "invalid system_type",
mutate: func(r *pb.TrackShareRequest) { r.SystemType = "symbian" },
wantErr: appErrors.ErrInvalidSystemType,
},
{
name: "invalid share_target",
mutate: func(r *pb.TrackShareRequest) { r.ShareTarget = "pinterest" },
wantErr: appErrors.ErrInvalidShareTarget,
},
{
name: "invalid result",
mutate: func(r *pb.TrackShareRequest) { r.Result = "exploded" },
wantErr: appErrors.ErrInvalidShareResult,
},
{
name: "invalid asset_id (zero)",
mutate: func(r *pb.TrackShareRequest) { r.AssetId = 0 },
wantErr: appErrors.ErrAssetNotFound,
},
}
for _, tt := range tests {
tt := tt
t.Run(tt.name, func(t *testing.T) {
repo := newMockShareRepo()
repo.existingUsers[7] = true
svc := newTestService(repo, nil, newStubQRUploader("https://cdn/x.png"), "https://h5.example.com")
req := validTrackReq()
tt.mutate(req)
resp, err := svc.TrackShare(context.Background(), req)
assert.Error(t, err)
assert.Nil(t, resp)
assert.True(t, errors.Is(err, tt.wantErr),
"want error chain to include %v, got %v", tt.wantErr, err)
})
}
}
func TestTrackShare_UserNotFound(t *testing.T) {
repo := newMockShareRepo()
// 不预置 user 7
svc := newTestService(repo, nil, newStubQRUploader("https://cdn/x.png"), "https://h5.example.com")
resp, err := svc.TrackShare(context.Background(), validTrackReq())
assert.Error(t, err)
assert.Nil(t, resp)
assert.True(t, errors.Is(err, appErrors.ErrUserNotFound))
assert.Empty(t, repo.created, "用户不存在时不应落库")
}
func TestTrackShare_PersistsEventAndSetsServerTs(t *testing.T) {
repo := newMockShareRepo()
repo.existingUsers[7] = true
svc := newTestService(repo, nil, newStubQRUploader("https://cdn/x.png"), "https://h5.example.com")
before := time.Now().UnixMilli()
resp, err := svc.TrackShare(context.Background(), validTrackReq())
after := time.Now().UnixMilli()
require.NoError(t, err)
require.NotNil(t, resp)
assert.Greater(t, resp.ShareEventId, int64(0), "应返回新 share_event_id")
// 验证落库
require.Len(t, repo.created, 1)
e := repo.created[0]
assert.Equal(t, resp.ShareEventId, e.ID)
assert.Equal(t, int64(42), e.AssetID)
assert.Equal(t, int64(7), e.SharerUserID)
assert.Equal(t, "android", e.SystemType)
assert.Equal(t, "weixin_friend", e.ShareTarget)
assert.Equal(t, "success", e.Result)
// server_ts 自动设置,且在 [before, after] 区间内
assert.GreaterOrEqual(t, e.ServerTs, before)
assert.LessOrEqual(t, e.ServerTs, after)
// client_ts 原样透传
assert.Equal(t, validTrackReq().ClientTs, e.ClientTs)
}
func TestTrackShare_PersistsExtraJSON(t *testing.T) {
repo := newMockShareRepo()
repo.existingUsers[7] = true
svc := newTestService(repo, nil, newStubQRUploader("https://cdn/x.png"), "https://h5.example.com")
req := validTrackReq()
extra, err := structpb.NewStruct(map[string]interface{}{
"app_version": "1.2.3",
"os_version": "iOS 17.0",
})
require.NoError(t, err)
req.Extra = extra
resp, err := svc.TrackShare(context.Background(), req)
require.NoError(t, err)
require.NotNil(t, resp)
require.Len(t, repo.created, 1)
e := repo.created[0]
// Extra 是 JSONB([]byte),内容应包含 app_version / os_version
raw := string(e.Extra)
assert.Contains(t, raw, "app_version")
assert.Contains(t, raw, "1.2.3")
assert.Contains(t, raw, "os_version")
}
func TestTrackShare_NilExtraStoresEmpty(t *testing.T) {
repo := newMockShareRepo()
repo.existingUsers[7] = true
svc := newTestService(repo, nil, newStubQRUploader("https://cdn/x.png"), "https://h5.example.com")
req := validTrackReq()
req.Extra = nil
resp, err := svc.TrackShare(context.Background(), req)
require.NoError(t, err)
require.NotNil(t, resp)
require.Len(t, repo.created, 1)
e := repo.created[0]
// 验证 extra 是有效 JSON 字符串(空对象 / null 都可以,只要非 nil)
assert.True(t, len(e.Extra) == 0 || strings.Contains(string(e.Extra), "{"),
"nil extra 应存为有效 JSON, got: %q", string(e.Extra))
}
// ---------------- Helpers (miniredis gating) ----------------
// skipWithoutMiniredis 当 miniredis 未在 go.mod 时直接 skip 子测试。
// 真实项目若有 miniredis,实现 newMiniredis 即可;否则视为环境不具备。
func skipWithoutMiniredis(t *testing.T) {
t.Helper()
// 通过 build tag / 检测 import 做不到(没有专用 build tag),
// 这里使用一个简单约定:环境变量 TEST_USE_MINIREDIS=1 才尝试连接。
// 默认情况下直接 skip,以免测试在缺 redis 时误失败。
if !miniredisEnabled() {
t.Skip("miniredis 未引入,跳过缓存命中实测;改用 cache_key 模板子测试覆盖")
}
}
func miniredisEnabled() bool {
// 简单约定:项目默认不引入 miniredis,如需启用在测试环境设该变量并提供 newMiniredis 实现。
return false
}
// newMiniredis 占位实现(默认不被调用),真实引入 miniredis 时替换为:
// mr, _ := miniredis.Run()
// rdb := redis.NewClient(&redis.Options{Addr: mr.Addr()})
// return mr, rdb
func newMiniredis(t *testing.T) (interface{ Close() }, *redis.Client) {
t.Skip("miniredis unavailable: see skipWithoutMiniredis")
return nil, nil
}