Make create,update,delete in group use transactions

This commit is contained in:
Abin Simon
2022-03-21 12:10:18 +05:30
parent b7eede13c9
commit b088eaecef
2 changed files with 114 additions and 51 deletions
+96 -50
View File
@@ -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
}
+18 -1
View File
@@ -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},