From b088eaeceff20692e847627c835d8654d6f4c91e Mon Sep 17 00:00:00 2001 From: Abin Simon Date: Wed, 16 Mar 2022 15:16:06 +0530 Subject: [PATCH] Make create,update,delete in group use transactions --- pkg/service/group.go | 146 +++++++++++++++++++++++++------------- pkg/service/group_test.go | 19 ++++- 2 files changed, 114 insertions(+), 51 deletions(-) diff --git a/pkg/service/group.go b/pkg/service/group.go index 13d8d48..157a7ad 100644 --- a/pkg/service/group.go +++ b/pkg/service/group.go @@ -2,6 +2,7 @@ package service import ( "context" + "database/sql" "fmt" "strconv" "time" @@ -49,18 +50,18 @@ func NewGroupService(db *bun.DB, azc AuthzService) GroupService { return &groupService{db: db, azc: azc} } -func (s *groupService) deleteGroupRoleRelaitons(ctx context.Context, groupId uuid.UUID, group *userv3.Group) (*userv3.Group, error) { +func (s *groupService) deleteGroupRoleRelaitons(ctx context.Context, db bun.IDB, groupId uuid.UUID, group *userv3.Group) (*userv3.Group, error) { // delete previous entries // TODO: single delete command - err := pg.DeleteX(ctx, s.db, "group_id", groupId, &models.GroupRole{}) + err := pg.DeleteX(ctx, db, "group_id", groupId, &models.GroupRole{}) if err != nil { return &userv3.Group{}, err } - err = pg.DeleteX(ctx, s.db, "group_id", groupId, &models.ProjectGroupRole{}) + err = pg.DeleteX(ctx, db, "group_id", groupId, &models.ProjectGroupRole{}) if err != nil { return &userv3.Group{}, err } - err = pg.DeleteX(ctx, s.db, "group_id", groupId, &models.ProjectGroupNamespaceRole{}) + err = pg.DeleteX(ctx, db, "group_id", groupId, &models.ProjectGroupNamespaceRole{}) if err != nil { return &userv3.Group{}, err } @@ -73,7 +74,7 @@ func (s *groupService) deleteGroupRoleRelaitons(ctx context.Context, groupId uui } // Map roles to groups -func (s *groupService) createGroupRoleRelations(ctx context.Context, group *userv3.Group, ids parsedIds) (*userv3.Group, error) { +func (s *groupService) createGroupRoleRelations(ctx context.Context, db bun.IDB, group *userv3.Group, ids parsedIds) (*userv3.Group, error) { // TODO: add transactions projectNamespaceRoles := group.GetSpec().GetProjectNamespaceRoles() @@ -83,7 +84,7 @@ func (s *groupService) createGroupRoleRelations(ctx context.Context, group *user var ps []*authzv1.Policy for _, pnr := range projectNamespaceRoles { role := pnr.GetRole() - entity, err := pg.GetIdByName(ctx, s.db, role, &models.Role{}) + entity, err := pg.GetIdByName(ctx, db, role, &models.Role{}) if err != nil { return &userv3.Group{}, fmt.Errorf("unable to find role '%v'", role) } @@ -99,7 +100,7 @@ func (s *groupService) createGroupRoleRelations(ctx context.Context, group *user namespaceId := pnr.GetNamespace() // TODO: lookup id from name switch { case namespaceId != 0: - projectId, err := pg.GetProjectId(ctx, s.db, project) + projectId, err := pg.GetProjectId(ctx, db, project) if err != nil { return &userv3.Group{}, fmt.Errorf("unable to find project '%v'", project) } @@ -123,7 +124,7 @@ func (s *groupService) createGroupRoleRelations(ctx context.Context, group *user Obj: role, }) case project != "": - projectId, err := pg.GetProjectId(ctx, s.db, project) + projectId, err := pg.GetProjectId(ctx, db, project) if err != nil { return &userv3.Group{}, fmt.Errorf("unable to find project '%v'", project) } @@ -165,19 +166,19 @@ func (s *groupService) createGroupRoleRelations(ctx context.Context, group *user } } if len(pgnrs) > 0 { - _, err := pg.Create(ctx, s.db, &pgnrs) + _, err := pg.Create(ctx, db, &pgnrs) if err != nil { return &userv3.Group{}, err } } if len(pgrs) > 0 { - _, err := pg.Create(ctx, s.db, &pgrs) + _, err := pg.Create(ctx, db, &pgrs) if err != nil { return &userv3.Group{}, err } } if len(grs) > 0 { - _, err := pg.Create(ctx, s.db, &grs) + _, err := pg.Create(ctx, db, &grs) if err != nil { return &userv3.Group{}, err } @@ -193,8 +194,8 @@ func (s *groupService) createGroupRoleRelations(ctx context.Context, group *user return group, nil } -func (s *groupService) deleteGroupAccountRelations(ctx context.Context, groupId uuid.UUID, group *userv3.Group) (*userv3.Group, error) { - err := pg.DeleteX(ctx, s.db, "group_id", groupId, &models.GroupAccount{}) +func (s *groupService) deleteGroupAccountRelations(ctx context.Context, db bun.IDB, groupId uuid.UUID, group *userv3.Group) (*userv3.Group, error) { + err := pg.DeleteX(ctx, db, "group_id", groupId, &models.GroupAccount{}) if err != nil { return &userv3.Group{}, fmt.Errorf("unable to delete user; %v", err) } @@ -207,13 +208,13 @@ func (s *groupService) deleteGroupAccountRelations(ctx context.Context, groupId } // Update the users(account) mapped to each group -func (s *groupService) createGroupAccountRelations(ctx context.Context, groupId uuid.UUID, group *userv3.Group) (*userv3.Group, error) { +func (s *groupService) createGroupAccountRelations(ctx context.Context, db bun.IDB, groupId uuid.UUID, group *userv3.Group) (*userv3.Group, error) { // TODO: add transactions var grpaccs []models.GroupAccount var ugs []*authzv1.UserGroup for _, account := range unique(group.GetSpec().GetUsers()) { // FIXME: do combined lookup - entity, err := pg.GetIdByTraits(ctx, s.db, account, &models.KratosIdentities{}) + entity, err := pg.GetIdByTraits(ctx, db, account, &models.KratosIdentities{}) if err != nil { return &userv3.Group{}, fmt.Errorf("unable to find user '%v'", account) } @@ -236,7 +237,7 @@ func (s *groupService) createGroupAccountRelations(ctx context.Context, groupId if len(grpaccs) == 0 { return group, nil } - _, err := pg.Create(ctx, s.db, &grpaccs) + _, err := pg.Create(ctx, db, &grpaccs) if err != nil { return &userv3.Group{}, err } @@ -251,14 +252,14 @@ func (s *groupService) createGroupAccountRelations(ctx context.Context, groupId return group, nil } -func (s *groupService) getPartnerOrganization(ctx context.Context, group *userv3.Group) (uuid.UUID, uuid.UUID, error) { +func (s *groupService) getPartnerOrganization(ctx context.Context, db bun.IDB, group *userv3.Group) (uuid.UUID, uuid.UUID, error) { partner := group.GetMetadata().GetPartner() org := group.GetMetadata().GetOrganization() - partnerId, err := pg.GetPartnerId(ctx, s.db, partner) + partnerId, err := pg.GetPartnerId(ctx, db, partner) if err != nil { return uuid.Nil, uuid.Nil, err } - organizationId, err := pg.GetOrganizationId(ctx, s.db, org) + organizationId, err := pg.GetOrganizationId(ctx, db, org) if err != nil { return partnerId, uuid.Nil, err } @@ -267,7 +268,7 @@ func (s *groupService) getPartnerOrganization(ctx context.Context, group *userv3 } func (s *groupService) Create(ctx context.Context, group *userv3.Group) (*userv3.Group, error) { - partnerId, organizationId, err := s.getPartnerOrganization(ctx, group) + partnerId, organizationId, err := s.getPartnerOrganization(ctx, s.db, group) if err != nil { return nil, fmt.Errorf("unable to get partner and org id") } @@ -286,29 +287,44 @@ func (s *groupService) Create(ctx context.Context, group *userv3.Group) (*userv3 PartnerId: partnerId, Type: group.GetSpec().GetType(), } - entity, err := pg.Create(ctx, s.db, &grp) + + tx, err := s.db.BeginTx(ctx, &sql.TxOptions{}) if err != nil { return &userv3.Group{}, err } + entity, err := pg.Create(ctx, tx, &grp) + if err != nil { + tx.Rollback() // TODO: check errors for rollback (and do what?) + return &userv3.Group{}, err + } + //update v3 spec if grp, ok := entity.(*models.Group); ok { // we can get previous group using the id, find users/roles from that and delete those - group, err = s.createGroupAccountRelations(ctx, grp.ID, group) + group, err = s.createGroupAccountRelations(ctx, tx, grp.ID, group) if err != nil { + tx.Rollback() return &userv3.Group{}, err } - group, err = s.createGroupRoleRelations(ctx, group, parsedIds{Id: grp.ID, Partner: partnerId, Organization: organizationId}) + group, err = s.createGroupRoleRelations(ctx, tx, group, parsedIds{Id: grp.ID, Partner: partnerId, Organization: organizationId}) if err != nil { + tx.Rollback() return &userv3.Group{}, err } + + err = tx.Commit() + if err != nil { + fmt.Println("unable to commit changes", err) + } return group, nil } + tx.Rollback() return &userv3.Group{}, fmt.Errorf("unable to create group") } -func (s *groupService) toV3Group(ctx context.Context, group *userv3.Group, grp *models.Group) (*userv3.Group, error) { +func (s *groupService) toV3Group(ctx context.Context, db bun.IDB, group *userv3.Group, grp *models.Group) (*userv3.Group, error) { labels := make(map[string]string) labels["organization"] = group.GetMetadata().GetOrganization() labels["partner"] = group.GetMetadata().GetPartner() @@ -323,7 +339,7 @@ func (s *groupService) toV3Group(ctx context.Context, group *userv3.Group, grp * Labels: labels, ModifiedAt: timestamppb.New(grp.ModifiedAt), } - users, err := dao.GetUsers(ctx, s.db, grp.ID) + users, err := dao.GetUsers(ctx, db, grp.ID) if err != nil { return &userv3.Group{}, err } @@ -332,7 +348,7 @@ func (s *groupService) toV3Group(ctx context.Context, group *userv3.Group, grp * userNames = append(userNames, u.Traits["email"].(string)) } - roles, err := dao.GetGroupRoles(ctx, s.db, grp.ID) + roles, err := dao.GetGroupRoles(ctx, db, grp.ID) if err != nil { return &userv3.Group{}, err } @@ -356,7 +372,7 @@ func (s *groupService) GetByID(ctx context.Context, group *userv3.Group) (*userv } if grp, ok := entity.(*models.Group); ok { - return s.toV3Group(ctx, group, grp) + return s.toV3Group(ctx, s.db, group, grp) } return group, nil @@ -364,7 +380,7 @@ func (s *groupService) GetByID(ctx context.Context, group *userv3.Group) (*userv func (s *groupService) GetByName(ctx context.Context, group *userv3.Group) (*userv3.Group, error) { name := group.GetMetadata().GetName() - partnerId, organizationId, err := s.getPartnerOrganization(ctx, group) + partnerId, organizationId, err := s.getPartnerOrganization(ctx, s.db, group) if err != nil { return nil, fmt.Errorf("unable to get partner and org id") } @@ -374,7 +390,7 @@ func (s *groupService) GetByName(ctx context.Context, group *userv3.Group) (*use } if grp, ok := entity.(*models.Group); ok { - return s.toV3Group(ctx, group, grp) + return s.toV3Group(ctx, s.db, group, grp) } return group, nil @@ -383,7 +399,7 @@ func (s *groupService) GetByName(ctx context.Context, group *userv3.Group) (*use func (s *groupService) Update(ctx context.Context, group *userv3.Group) (*userv3.Group, error) { // TODO: inform when unchanged name := group.GetMetadata().GetName() - partnerId, organizationId, err := s.getPartnerOrganization(ctx, group) + partnerId, organizationId, err := s.getPartnerOrganization(ctx, s.db, group) if err != nil { return nil, fmt.Errorf("unable to get partner and org id") } @@ -399,28 +415,43 @@ func (s *groupService) Update(ctx context.Context, group *userv3.Group) (*userv3 grp.Type = group.Spec.Type grp.ModifiedAt = time.Now() - // update account/role links - group, err = s.deleteGroupAccountRelations(ctx, grp.ID, group) - if err != nil { - return &userv3.Group{}, err - } - group, err = s.createGroupAccountRelations(ctx, grp.ID, group) - if err != nil { - return &userv3.Group{}, err - } - group, err = s.deleteGroupRoleRelaitons(ctx, grp.ID, group) - if err != nil { - return &userv3.Group{}, err - } - group, err = s.createGroupRoleRelations(ctx, group, parsedIds{Id: grp.ID, Partner: partnerId, Organization: organizationId}) + tx, err := s.db.BeginTx(ctx, &sql.TxOptions{}) if err != nil { return &userv3.Group{}, err } - _, err = pg.Update(ctx, s.db, grp.ID, grp) + // update account/role links + group, err = s.deleteGroupAccountRelations(ctx, tx, grp.ID, group) if err != nil { + tx.Rollback() return &userv3.Group{}, err } + group, err = s.createGroupAccountRelations(ctx, tx, grp.ID, group) + if err != nil { + tx.Rollback() + return &userv3.Group{}, err + } + group, err = s.deleteGroupRoleRelaitons(ctx, tx, grp.ID, group) + if err != nil { + tx.Rollback() + return &userv3.Group{}, err + } + group, err = s.createGroupRoleRelations(ctx, tx, group, parsedIds{Id: grp.ID, Partner: partnerId, Organization: organizationId}) + if err != nil { + tx.Rollback() + return &userv3.Group{}, err + } + + _, err = pg.Update(ctx, tx, grp.ID, grp) + if err != nil { + tx.Rollback() + return &userv3.Group{}, err + } + + err = tx.Commit() + if err != nil { + fmt.Println("unable to commit changes", err) + } // update spec and status group.Spec = &userv3.GroupSpec{ @@ -435,7 +466,7 @@ func (s *groupService) Update(ctx context.Context, group *userv3.Group) (*userv3 func (s *groupService) Delete(ctx context.Context, group *userv3.Group) (*userv3.Group, error) { name := group.GetMetadata().GetName() - partnerId, organizationId, err := s.getPartnerOrganization(ctx, group) + partnerId, organizationId, err := s.getPartnerOrganization(ctx, s.db, group) if err != nil { return &userv3.Group{}, fmt.Errorf("unable to get partner and org id") } @@ -444,21 +475,36 @@ func (s *groupService) Delete(ctx context.Context, group *userv3.Group) (*userv3 return &userv3.Group{}, err } if grp, ok := entity.(*models.Group); ok { - group, err = s.deleteGroupRoleRelaitons(ctx, grp.ID, group) + + tx, err := s.db.BeginTx(ctx, &sql.TxOptions{}) if err != nil { return &userv3.Group{}, err } - group, err = s.deleteGroupAccountRelations(ctx, grp.ID, group) + + group, err = s.deleteGroupRoleRelaitons(ctx, s.db, grp.ID, group) if err != nil { + tx.Rollback() + return &userv3.Group{}, err + } + group, err = s.deleteGroupAccountRelations(ctx, s.db, grp.ID, group) + if err != nil { + tx.Rollback() return &userv3.Group{}, err } err = pg.Delete(ctx, s.db, grp.ID, grp) if err != nil { + tx.Rollback() return &userv3.Group{}, err } + + err = tx.Commit() + if err != nil { + fmt.Println("unable to commit changes", err) + } + return group, nil } - return group, nil + return &userv3.Group{}, fmt.Errorf("unable to delete group") } func (s *groupService) List(ctx context.Context, group *userv3.Group) (*userv3.GroupList, error) { @@ -487,7 +533,7 @@ func (s *groupService) List(ctx context.Context, group *userv3.Group) (*userv3.G if grps, ok := entities.(*[]models.Group); ok { for _, grp := range *grps { entry := &userv3.Group{Metadata: group.GetMetadata()} - entry, err = s.toV3Group(ctx, entry, &grp) + entry, err = s.toV3Group(ctx, s.db, entry, &grp) if err != nil { return groupList, err } diff --git a/pkg/service/group_test.go b/pkg/service/group_test.go index c65c8c1..c90b556 100644 --- a/pkg/service/group_test.go +++ b/pkg/service/group_test.go @@ -94,8 +94,11 @@ func TestCreateGroupNoUsersNoRoles(t *testing.T) { WithArgs().WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(ouuid)) mock.ExpectQuery(`SELECT "group"."id" FROM "authsrv_group" AS "group" WHERE .organization_id = '` + ouuid + `'. AND .partner_id = '` + puuid + `'. AND .name = 'group-` + guuid + `'.`). WillReturnError(fmt.Errorf("no data available")) + + mock.ExpectBegin() mock.ExpectQuery(`INSERT INTO "authsrv_group"`). WithArgs().WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(guuid)) + mock.ExpectCommit() group := &userv3.Group{ Metadata: &v3.Metadata{Partner: "partner-" + puuid, Organization: "org-" + ouuid, Name: "group-" + guuid}, @@ -133,9 +136,12 @@ func TestCreateGroupDuplicate(t *testing.T) { WithArgs().WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(puuid)) mock.ExpectQuery(`SELECT "organization"."id" FROM "authsrv_organization" AS "organization"`). WithArgs().WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(ouuid)) + + mock.ExpectBegin() // TODO: more precise checks mock.ExpectQuery(`INSERT INTO "authsrv_group"`). WithArgs().WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(guuid)) + mock.ExpectCommit() _, err := gs.Create(context.Background(), group) if err == nil { t.Fatal("should not be able to recreate group with same name") @@ -171,6 +177,8 @@ func TestCreateGroupWithUsersNoRoles(t *testing.T) { WithArgs().WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(ouuid)) mock.ExpectQuery(`SELECT "group"."id" FROM "authsrv_group" AS "group" WHERE .organization_id = '` + ouuid + `'. AND .partner_id = '` + puuid + `'. AND .name = 'group-` + guuid + `'.`). WillReturnError(fmt.Errorf("no data available")) + + mock.ExpectBegin() mock.ExpectQuery(`INSERT INTO "authsrv_group"`). WithArgs().WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(guuid)) for _, u := range tc.users { @@ -179,6 +187,7 @@ func TestCreateGroupWithUsersNoRoles(t *testing.T) { } mock.ExpectQuery(`INSERT INTO "authsrv_groupaccount"`). WithArgs().WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(uuid.New().String())) + mock.ExpectCommit() group := &userv3.Group{ Metadata: &v3.Metadata{Partner: "partner-" + puuid, Organization: "org-" + ouuid, Name: "group-" + guuid}, @@ -236,6 +245,8 @@ func TestCreateGroupNoUsersWithRoles(t *testing.T) { WithArgs().WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(ouuid)) mock.ExpectQuery(`SELECT "group"."id" FROM "authsrv_group" AS "group" WHERE .organization_id = '` + ouuid + `'. AND .partner_id = '` + puuid + `'. AND .name = 'group-` + guuid + `'.`). WillReturnError(fmt.Errorf("no data available")) + + mock.ExpectBegin() mock.ExpectQuery(`INSERT INTO "authsrv_group"`). WithArgs().WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(guuid)) mock.ExpectQuery(`SELECT "resourcerole"."id" FROM "authsrv_resourcerole" AS "resourcerole"`). @@ -246,6 +257,7 @@ func TestCreateGroupNoUsersWithRoles(t *testing.T) { } mock.ExpectQuery(fmt.Sprintf(`INSERT INTO "%v"`, tc.dbname)). WithArgs().WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(uuid.New().String())) + mock.ExpectCommit() group := &userv3.Group{ Metadata: &v3.Metadata{Partner: "partner-" + puuid, Organization: "org-" + ouuid, Name: "group-" + guuid}, @@ -311,6 +323,7 @@ func TestCreateGroupWithUsersWithRoles(t *testing.T) { WithArgs().WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(ouuid)) mock.ExpectQuery(`SELECT "group"."id" FROM "authsrv_group" AS "group" WHERE .organization_id = '` + ouuid + `'. AND .partner_id = '` + puuid + `'. AND .name = 'group-` + guuid + `'.`).WithArgs() + mock.ExpectBegin() // TODO: more precise checks mock.ExpectQuery(`INSERT INTO "authsrv_group"`). WithArgs().WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(guuid)) @@ -329,6 +342,7 @@ func TestCreateGroupWithUsersWithRoles(t *testing.T) { } mock.ExpectQuery(fmt.Sprintf(`INSERT INTO "%v"`, tc.dbname)). WithArgs().WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(uuid.New().String())) + mock.ExpectCommit() group := &userv3.Group{ Metadata: &v3.Metadata{Partner: "partner-" + puuid, Organization: "org-" + ouuid, Name: "group-" + guuid}, @@ -393,7 +407,7 @@ func TestUpdateGroupWithUsersWithRoles(t *testing.T) { mock.ExpectQuery(`SELECT "group"."id", "group"."name",.* FROM "authsrv_group" AS "group" WHERE .*name = 'group-` + guuid + `'`). WithArgs().WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(guuid, "group-"+guuid)) - // TODO: more precise checks + mock.ExpectBegin() mock.ExpectExec(`UPDATE "authsrv_groupaccount" AS "groupaccount" SET trash = TRUE WHERE ."group_id" = '` + guuid). WillReturnResult(sqlmock.NewResult(1, 1)) for _, u := range tc.users { @@ -418,6 +432,7 @@ func TestUpdateGroupWithUsersWithRoles(t *testing.T) { WithArgs().WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(uuid.New().String())) mock.ExpectExec(`UPDATE "authsrv_group"`). WillReturnResult(sqlmock.NewResult(1, 1)) + mock.ExpectCommit() group := &userv3.Group{ Metadata: &v3.Metadata{Partner: "partner-" + puuid, Organization: "org-" + ouuid, Name: "group-" + guuid}, @@ -463,6 +478,7 @@ func TestGroupDelete(t *testing.T) { mock.ExpectQuery(`SELECT "group"."id", "group"."name", .* FROM "authsrv_group" AS "group" WHERE`). WithArgs().WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(guuid, "group-"+guuid)) + mock.ExpectBegin() mock.ExpectExec(`UPDATE "authsrv_grouprole" AS "grouprole" SET trash = TRUE WHERE ."group_id" = '` + guuid). WillReturnResult(sqlmock.NewResult(1, 1)) mock.ExpectExec(`UPDATE "authsrv_projectgrouprole" AS "projectgrouprole" SET trash = TRUE WHERE ."group_id" = '` + guuid). @@ -473,6 +489,7 @@ func TestGroupDelete(t *testing.T) { WillReturnResult(sqlmock.NewResult(1, 1)) mock.ExpectExec(`UPDATE "authsrv_group" AS "group" SET trash = TRUE WHERE .id = '` + guuid). WillReturnResult(sqlmock.NewResult(1, 1)) + mock.ExpectCommit() group := &userv3.Group{ Metadata: &v3.Metadata{Partner: "partner-" + puuid, Organization: "org-" + ouuid, Name: "group-" + guuid},