package user import ( "context" "fmt" "strings" "testing" "github.com/DATA-DOG/go-sqlmock" modeluser "github.com/perfect-panel/server/internal/model/user" "github.com/perfect-panel/server/internal/svc" "github.com/perfect-panel/server/internal/types" "github.com/perfect-panel/server/pkg/constant" "github.com/perfect-panel/server/pkg/hash" "gorm.io/driver/mysql" "gorm.io/gorm" ) func TestGetInviteRecordsInviter(t *testing.T) { svcCtx, mock, cleanup := newInviteRecordsTestSvc(t) defer cleanup() expectNoInviteRecordsFamily(t, mock, 100) mock.ExpectQuery("SELECT id, object_id, content"). WithArgs(34, int64(100), "邀请赠送"). WillReturnRows(sqlmock.NewRows([]string{"id", "object_id", "content", "created_at"}). AddRow(1, 100, `{"order_no":"order-1","amount":7,"remark":"邀请赠送"}`, 1779934580000)) mock.ExpectQuery("SELECT `order`.order_no, `order`.user_id, `order`.subscription_user_id, invitee.referer_id FROM `order`"). WithArgs("order-1"). WillReturnRows(sqlmock.NewRows([]string{"order_no", "user_id", "subscription_user_id", "referer_id"}).AddRow("order-1", 200, 200, 100)) resp, err := NewGetInviteRecordsLogic(inviteRecordsContext(100, 0), svcCtx).GetInviteRecords(&types.GetInviteRecordsRequest{Page: 1, Size: 10}) if err != nil { t.Fatalf("GetInviteRecords returned error: %v", err) } assertInviteRecordResponse(t, resp, types.InviteRecord{ Role: inviteRecordRoleInviter, PeerHash: hash.InvitePeerHash(200), GiftDays: 7, OrderNo: "order-1", CreatedAt: 1779934580000, }) assertInviteRecordsExpectations(t, mock) } func TestGetInviteRecordsInvitee(t *testing.T) { svcCtx, mock, cleanup := newInviteRecordsTestSvc(t) defer cleanup() expectNoInviteRecordsFamily(t, mock, 200) mock.ExpectQuery("SELECT id, object_id, content"). WithArgs(34, int64(200), "邀请赠送"). WillReturnRows(sqlmock.NewRows([]string{"id", "object_id", "content", "created_at"}). AddRow(2, 200, `{"order_no":"order-2","amount":7,"remark":"邀请赠送"}`, 1779934590000)) mock.ExpectQuery("SELECT `order`.order_no, `order`.user_id, `order`.subscription_user_id, invitee.referer_id FROM `order`"). WithArgs("order-2"). WillReturnRows(sqlmock.NewRows([]string{"order_no", "user_id", "subscription_user_id", "referer_id"}).AddRow("order-2", 200, 200, 100)) resp, err := NewGetInviteRecordsLogic(inviteRecordsContext(200, 100), svcCtx).GetInviteRecords(&types.GetInviteRecordsRequest{Page: 1, Size: 10}) if err != nil { t.Fatalf("GetInviteRecords returned error: %v", err) } assertInviteRecordResponse(t, resp, types.InviteRecord{ Role: inviteRecordRoleInvitee, PeerHash: hash.InvitePeerHash(100), GiftDays: 7, OrderNo: "order-2", CreatedAt: 1779934590000, }) assertInviteRecordsExpectations(t, mock) } func TestGetInviteRecordsMissingOrderReturnsDirtyRecord(t *testing.T) { svcCtx, mock, cleanup := newInviteRecordsTestSvc(t) defer cleanup() expectNoInviteRecordsFamily(t, mock, 100) mock.ExpectQuery("SELECT id, object_id, content"). WithArgs(34, int64(100), "邀请赠送"). WillReturnRows(sqlmock.NewRows([]string{"id", "object_id", "content", "created_at"}). AddRow(3, 100, `{"order_no":"missing-order","amount":7,"remark":"邀请赠送"}`, 1779934600000)) mock.ExpectQuery("SELECT `order`.order_no, `order`.user_id, `order`.subscription_user_id, invitee.referer_id FROM `order`"). WithArgs("missing-order"). WillReturnRows(sqlmock.NewRows([]string{"order_no", "user_id", "subscription_user_id", "referer_id"})) resp, err := NewGetInviteRecordsLogic(inviteRecordsContext(100, 0), svcCtx).GetInviteRecords(&types.GetInviteRecordsRequest{Page: 1, Size: 10}) if err != nil { t.Fatalf("GetInviteRecords returned error: %v", err) } assertInviteRecordResponse(t, resp, types.InviteRecord{ Role: inviteRecordRoleInviter, GiftDays: 7, OrderNo: "missing-order", CreatedAt: 1779934600000, }) assertInviteRecordsExpectations(t, mock) } func TestGetInviteRecordsFamilyMemberSeesOwnerGiftLog(t *testing.T) { svcCtx, mock, cleanup := newInviteRecordsTestSvc(t) defer cleanup() expectInviteRecordsFamilyMember(t, mock, 200, 900) mock.ExpectQuery("SELECT id, object_id, content"). WithArgs(34, int64(200), int64(900), "邀请赠送"). WillReturnRows(sqlmock.NewRows([]string{"id", "object_id", "content", "created_at"}). AddRow(4, 900, `{"order_no":"family-order","amount":7,"remark":"邀请赠送"}`, 1779934610000)) mock.ExpectQuery("SELECT `order`.order_no, `order`.user_id, `order`.subscription_user_id, invitee.referer_id FROM `order`"). WithArgs("family-order"). WillReturnRows(sqlmock.NewRows([]string{"order_no", "user_id", "subscription_user_id", "referer_id"}).AddRow("family-order", 200, 900, 100)) resp, err := NewGetInviteRecordsLogic(inviteRecordsContext(200, 100), svcCtx).GetInviteRecords(&types.GetInviteRecordsRequest{Page: 1, Size: 10}) if err != nil { t.Fatalf("GetInviteRecords returned error: %v", err) } assertInviteRecordResponse(t, resp, types.InviteRecord{ Role: inviteRecordRoleInvitee, PeerHash: hash.InvitePeerHash(100), GiftDays: 7, OrderNo: "family-order", CreatedAt: 1779934610000, }) assertInviteRecordsExpectations(t, mock) } func TestGetInviteRecordsFamilyMemberSeesAllOwnerGiftLogs(t *testing.T) { svcCtx, mock, cleanup := newInviteRecordsTestSvc(t) defer cleanup() expectInviteRecordsFamilyMember(t, mock, 51637, 510) mock.ExpectQuery("SELECT id, object_id, content"). WithArgs(34, int64(51637), int64(510), "邀请赠送"). WillReturnRows(sqlmock.NewRows([]string{"id", "object_id", "content", "created_at"}). AddRow(6, 510, `{"order_no":"owner-order","amount":7,"remark":"邀请赠送"}`, 1779934630000)) mock.ExpectQuery("SELECT `order`.order_no, `order`.user_id, `order`.subscription_user_id, invitee.referer_id FROM `order`"). WithArgs("owner-order"). WillReturnRows(sqlmock.NewRows([]string{"order_no", "user_id", "subscription_user_id", "referer_id"}).AddRow("owner-order", 571, 571, 510)) resp, err := NewGetInviteRecordsLogic(inviteRecordsContext(51637, 0), svcCtx).GetInviteRecords(&types.GetInviteRecordsRequest{Page: 1, Size: 10}) if err != nil { t.Fatalf("GetInviteRecords returned error: %v", err) } assertInviteRecordResponse(t, resp, types.InviteRecord{ Role: inviteRecordRoleInviter, PeerHash: hash.InvitePeerHash(571), GiftDays: 7, OrderNo: "owner-order", CreatedAt: 1779934630000, }) assertInviteRecordsExpectations(t, mock) } func TestGetInviteRecordsOwnerDoesNotSeeMemberGiftLog(t *testing.T) { svcCtx, mock, cleanup := newInviteRecordsTestSvc(t) defer cleanup() expectInviteRecordsFamilyOwner(t, mock, 900) mock.ExpectQuery("SELECT id, object_id, content"). WithArgs(34, int64(900), "邀请赠送"). WillReturnRows(sqlmock.NewRows([]string{"id", "object_id", "content", "created_at"}). AddRow(5, 900, `{"order_no":"member-order","amount":7,"remark":"邀请赠送"}`, 1779934620000)) mock.ExpectQuery("SELECT `order`.order_no, `order`.user_id, `order`.subscription_user_id, invitee.referer_id FROM `order`"). WithArgs("member-order"). WillReturnRows(sqlmock.NewRows([]string{"order_no", "user_id", "subscription_user_id", "referer_id"}).AddRow("member-order", 200, 900, 100)) resp, err := NewGetInviteRecordsLogic(inviteRecordsContext(900, 0), svcCtx).GetInviteRecords(&types.GetInviteRecordsRequest{Page: 1, Size: 10}) if err != nil { t.Fatalf("GetInviteRecords returned error: %v", err) } if resp == nil { t.Fatal("response is nil") } if resp.Total != 0 { t.Fatalf("Total = %d, want 0", resp.Total) } if len(resp.List) != 0 { t.Fatalf("len(List) = %d, want 0", len(resp.List)) } assertInviteRecordsExpectations(t, mock) } func newInviteRecordsTestSvc(t *testing.T) (*svc.ServiceContext, sqlmock.Sqlmock, func()) { t.Helper() sqlDB, mock, err := sqlmock.New(sqlmock.QueryMatcherOption(sqlmock.QueryMatcherFunc(func(expectedSQL, actualSQL string) error { if strings.Contains(actualSQL, expectedSQL) { return nil } return fmt.Errorf("actual sql %q does not contain %q", actualSQL, expectedSQL) }))) if err != nil { t.Fatalf("create sqlmock: %v", err) } db, err := gorm.Open(mysql.New(mysql.Config{Conn: sqlDB, SkipInitializeWithVersion: true}), &gorm.Config{}) if err != nil { _ = sqlDB.Close() t.Fatalf("open gorm db: %v", err) } return &svc.ServiceContext{DB: db}, mock, func() { _ = sqlDB.Close() } } func expectNoInviteRecordsFamily(t *testing.T, mock sqlmock.Sqlmock, userId int64) { t.Helper() mock.ExpectQuery("FROM `user_family_member` JOIN user_family"). WithArgs(1, userId, 1, 1). WillReturnError(gorm.ErrRecordNotFound) } func expectInviteRecordsFamilyMember(t *testing.T, mock sqlmock.Sqlmock, userId, ownerUserId int64) { t.Helper() mock.ExpectQuery("FROM `user_family_member` JOIN user_family"). WithArgs(1, userId, 1, 1). WillReturnRows(sqlmock.NewRows([]string{"family_id", "role", "owner_user_id"}).AddRow(800, 2, ownerUserId)) } func expectInviteRecordsFamilyOwner(t *testing.T, mock sqlmock.Sqlmock, userId int64) { t.Helper() mock.ExpectQuery("FROM `user_family_member` JOIN user_family"). WithArgs(1, userId, 1, 1). WillReturnRows(sqlmock.NewRows([]string{"family_id", "role", "owner_user_id"}).AddRow(800, 1, userId)) } func inviteRecordsContext(userId, refererId int64) context.Context { return context.WithValue(context.Background(), constant.CtxKeyUser, &modeluser.User{ Id: userId, RefererId: refererId, }) } func assertInviteRecordResponse(t *testing.T, resp *types.GetInviteRecordsResponse, want types.InviteRecord) { t.Helper() if resp == nil { t.Fatal("response is nil") } if resp.Total != 1 { t.Fatalf("Total = %d, want 1", resp.Total) } if len(resp.List) != 1 { t.Fatalf("len(List) = %d, want 1", len(resp.List)) } if got := resp.List[0]; got != want { t.Fatalf("record = %+v, want %+v", got, want) } } func assertInviteRecordsExpectations(t *testing.T, mock sqlmock.Sqlmock) { t.Helper() if err := mock.ExpectationsWereMet(); err != nil { t.Fatalf("unmet sql expectations: %v", err) } }