package provider import ( "context" "testing" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "go.uber.org/zap" "google.golang.org/grpc/codes" "google.golang.org/grpc/status" "github.com/topfans/backend/pkg/authctx" "github.com/topfans/backend/pkg/logger" "github.com/topfans/backend/pkg/models" pb "github.com/topfans/backend/pkg/proto/asset" ) func init() { // provider 里的 handler 会调用 logger.Logger.*,测试环境先初始化避免 nil panic。 if logger.Logger == nil { logger.Logger = zap.NewNop() } } // ---- fakes ---- // fakeAssetLikeMgr 捕获 provider 传给 AssetLikeService 的身份参数。 type fakeAssetLikeMgr struct { gotUserID int64 gotStarID int64 called bool } func (f *fakeAssetLikeMgr) CheckAssetLike(_ context.Context, _ int64, userID, starID int64) (bool, error) { f.called = true f.gotUserID = userID f.gotStarID = starID return true, nil } func (f *fakeAssetLikeMgr) LikeAsset(context.Context, int64, int64, int64) (int32, error) { return 0, nil } func (f *fakeAssetLikeMgr) UnlikeAsset(context.Context, int64, int64, int64) (int32, error) { return 0, nil } func (f *fakeAssetLikeMgr) GetAssetLikes(context.Context, int64, int32, int32) ([]*models.AssetLike, int64, error) { return nil, 0, nil } func (f *fakeAssetLikeMgr) ClearAssetLikeRecords(context.Context, int64) error { return nil } // fakeAssetService 捕获分享类 RPC 落到 service 层时 req.SharerUserId 的实际值。 type fakeAssetService struct { qrReq *pb.GetAssetQrcodeRequest trackReq *pb.TrackShareRequest } func (f *fakeAssetService) GetMyAssets(*pb.GetMyAssetsRequest, int64, int64) (*pb.GetMyAssetsResponse, error) { return nil, nil } func (f *fakeAssetService) GetAssetsByType(*pb.GetAssetsByTypeRequest, int64, int64) (*pb.GetAssetsByTypeResponse, error) { return nil, nil } func (f *fakeAssetService) GetAsset(*pb.GetAssetRequest, int64, int64) (*pb.GetAssetResponse, error) { return nil, nil } func (f *fakeAssetService) GetAssetStatus(*pb.GetAssetStatusRequest, int64, int64) (*pb.GetAssetStatusResponse, error) { return nil, nil } func (f *fakeAssetService) GetAssetForRPC(*pb.GetAssetForRPCRequest) (*pb.GetAssetForRPCResponse, error) { return nil, nil } func (f *fakeAssetService) GetAssetQrcode(_ context.Context, req *pb.GetAssetQrcodeRequest) (*pb.GetAssetQrcodeResponse, error) { f.qrReq = req // OSS key 在 ShareService 里由 req.SharerUserId 拼出(share/qrcode/__.png), // 这里回填一个体现覆盖后 sharer 的 URL,供测试断言归因值随 req 流转。 return &pb.GetAssetQrcodeResponse{QrcodeUrl: "https://cdn/share/qrcode/x.png"}, nil } func (f *fakeAssetService) TrackShare(_ context.Context, req *pb.TrackShareRequest) (*pb.TrackShareResponse, error) { f.trackReq = req return &pb.TrackShareResponse{ShareEventId: 1}, nil } // ---- CheckAssetLike ---- func TestCheckAssetLike_RejectsForgedUserId(t *testing.T) { likeMgr := &fakeAssetLikeMgr{} p := &AssetProvider{assetLikeService: likeMgr} // ctx 携带可信身份 (100, 200);req 里伪造 (999, 888)。 ctx := authctx.WithIdentity(context.Background(), 100, 200) req := &pb.CheckAssetLikeRequest{AssetId: 7, UserId: 999, StarId: 888} _, err := p.CheckAssetLike(ctx, req) require.NoError(t, err) assert.True(t, likeMgr.called, "service 应被调用") assert.Equal(t, int64(100), likeMgr.gotUserID, "必须用 ctx 可信 user_id 覆盖伪造值") assert.Equal(t, int64(200), likeMgr.gotStarID, "必须用 ctx 可信 star_id 覆盖伪造值") } func TestCheckAssetLike_NoIdentity(t *testing.T) { likeMgr := &fakeAssetLikeMgr{} p := &AssetProvider{assetLikeService: likeMgr} req := &pb.CheckAssetLikeRequest{AssetId: 7, UserId: 999, StarId: 888} resp, err := p.CheckAssetLike(context.Background(), req) require.Error(t, err) assert.Equal(t, codes.Unauthenticated, status.Code(err)) assert.False(t, likeMgr.called, "缺身份时不得触达 service") if resp != nil && resp.Base != nil { assert.Equal(t, uint32(codes.Unauthenticated), resp.Base.Code) } } // ---- GetAssetQrcode ---- func TestGetAssetQrcode_RejectsForgedSharer(t *testing.T) { svc := &fakeAssetService{} p := &AssetProvider{assetService: svc} ctx := authctx.WithIdentity(context.Background(), 100, 0) req := &pb.GetAssetQrcodeRequest{AssetId: 7, SharerUserId: 999, SystemType: "android"} _, err := p.GetAssetQrcode(ctx, req) require.NoError(t, err) require.NotNil(t, svc.qrReq, "service 应收到请求") assert.Equal(t, int64(100), svc.qrReq.SharerUserId, "分享归因必须锚定 ctx 可信 user_id,OSS key/落地页 from= 都由此拼出") } func TestGetAssetQrcode_NoIdentity(t *testing.T) { svc := &fakeAssetService{} p := &AssetProvider{assetService: svc} req := &pb.GetAssetQrcodeRequest{AssetId: 7, SharerUserId: 999, SystemType: "android"} _, err := p.GetAssetQrcode(context.Background(), req) require.Error(t, err) assert.Equal(t, codes.Unauthenticated, status.Code(err)) assert.Nil(t, svc.qrReq, "缺身份时不得触达 service") } // ---- TrackShare ---- func TestTrackShare_RejectsForgedSharer(t *testing.T) { svc := &fakeAssetService{} p := &AssetProvider{assetService: svc} ctx := authctx.WithIdentity(context.Background(), 100, 0) req := &pb.TrackShareRequest{AssetId: 7, SharerUserId: 999, SystemType: "android", ShareTarget: "weixin_friend", Result: "success"} _, err := p.TrackShare(ctx, req) require.NoError(t, err) require.NotNil(t, svc.trackReq, "service 应收到请求") assert.Equal(t, int64(100), svc.trackReq.SharerUserId, "share_events.sharer_user_id 必须落 ctx 可信 user_id,不能被 req 伪造") } func TestTrackShare_NoIdentity(t *testing.T) { svc := &fakeAssetService{} p := &AssetProvider{assetService: svc} req := &pb.TrackShareRequest{AssetId: 7, SharerUserId: 999, SystemType: "android", ShareTarget: "weixin_friend", Result: "success"} _, err := p.TrackShare(context.Background(), req) require.Error(t, err) assert.Equal(t, codes.Unauthenticated, status.Code(err)) assert.Nil(t, svc.trackReq, "缺身份时不得触达 service") }