Refactor invitation system
Signed-off-by: Sacha Al Himdani <sacha@getprobo.com>
This commit is contained in:
@@ -891,25 +891,45 @@ func (r *frameworkConnectionResolver) TotalCount(ctx context.Context, obj *types
|
||||
panic(fmt.Errorf("unsupported resolver: %T", obj.Resolver))
|
||||
}
|
||||
|
||||
// Organization is the resolver for the organization field.
|
||||
func (r *invitationResolver) Organization(ctx context.Context, obj *types.Invitation) (*types.Organization, error) {
|
||||
organization, err := r.authzSvc.GetOrganizationByInvitationID(ctx, obj.ID)
|
||||
if err != nil {
|
||||
panic(fmt.Errorf("cannot load organization: %w", err))
|
||||
}
|
||||
|
||||
return types.NewOrganization(organization), nil
|
||||
}
|
||||
|
||||
// TotalCount is the resolver for the totalCount field.
|
||||
func (r *invitationConnectionResolver) TotalCount(ctx context.Context, obj *types.InvitationConnection) (int, error) {
|
||||
currentUser := UserFromContext(ctx)
|
||||
if currentUser == nil {
|
||||
return 0, fmt.Errorf("no authenticated user")
|
||||
switch obj.Resolver.(type) {
|
||||
case *organizationResolver:
|
||||
authzSvc := r.AuthzService(ctx, obj.ParentID.TenantID())
|
||||
count, err := authzSvc.CountOrganizationInvitations(ctx, obj.ParentID)
|
||||
if err != nil {
|
||||
panic(fmt.Errorf("failed to count organization invitations: %w", err))
|
||||
}
|
||||
return count, nil
|
||||
case *viewerResolver:
|
||||
user := UserFromContext(ctx)
|
||||
if user == nil {
|
||||
panic(fmt.Errorf("no authenticated user"))
|
||||
}
|
||||
|
||||
invitationFilter := coredata.NewInvitationFilter(nil)
|
||||
if obj.Filter != nil {
|
||||
invitationFilter = coredata.NewInvitationFilter(obj.Filter.OnlyPending)
|
||||
}
|
||||
|
||||
count, err := r.authzSvc.CountUserInvitations(ctx, user.EmailAddress, invitationFilter)
|
||||
if err != nil {
|
||||
panic(fmt.Errorf("failed to count user invitations: %w", err))
|
||||
}
|
||||
return count, nil
|
||||
}
|
||||
|
||||
memberships, err := r.authzSvc.GetAllUserOrganizations(ctx, currentUser.ID)
|
||||
if err != nil || len(memberships) == 0 {
|
||||
return 0, fmt.Errorf("user has no organization memberships")
|
||||
}
|
||||
|
||||
orgID := memberships[0].ID
|
||||
count, err := r.authzSvc.CountOrganizationInvitations(ctx, orgID)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("failed to count invitations: %w", err)
|
||||
}
|
||||
|
||||
return count, nil
|
||||
panic(fmt.Errorf("unsupported resolver: %T", obj.Resolver))
|
||||
}
|
||||
|
||||
// Evidences is the resolver for the evidences field.
|
||||
@@ -1052,23 +1072,17 @@ func (r *measureConnectionResolver) TotalCount(ctx context.Context, obj *types.M
|
||||
|
||||
// TotalCount is the resolver for the totalCount field.
|
||||
func (r *membershipConnectionResolver) TotalCount(ctx context.Context, obj *types.MembershipConnection) (int, error) {
|
||||
currentUser := UserFromContext(ctx)
|
||||
if currentUser == nil {
|
||||
return 0, fmt.Errorf("no authenticated user")
|
||||
switch obj.Resolver.(type) {
|
||||
case *organizationResolver:
|
||||
authzSvc := r.AuthzService(ctx, obj.ParentID.TenantID())
|
||||
count, err := authzSvc.CountOrganizationMemberships(ctx, obj.ParentID)
|
||||
if err != nil {
|
||||
panic(fmt.Errorf("failed to count organization memberships: %w", err))
|
||||
}
|
||||
return count, nil
|
||||
default:
|
||||
panic(fmt.Errorf("unknown resolver type for membership connection"))
|
||||
}
|
||||
|
||||
memberships, err := r.authzSvc.GetAllUserOrganizations(ctx, currentUser.ID)
|
||||
if err != nil || len(memberships) == 0 {
|
||||
return 0, fmt.Errorf("user has no organization memberships")
|
||||
}
|
||||
|
||||
orgID := memberships[0].ID
|
||||
count, err := r.authzSvc.CountOrganizationMemberships(ctx, orgID)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("failed to count memberships: %w", err)
|
||||
}
|
||||
|
||||
return count, nil
|
||||
}
|
||||
|
||||
// CreateOrganization is the resolver for the createOrganization field.
|
||||
@@ -1104,7 +1118,6 @@ func (r *mutationResolver) CreateOrganization(ctx context.Context, input types.C
|
||||
ctx,
|
||||
probo.CreatePeopleRequest{
|
||||
OrganizationID: organization.ID,
|
||||
UserID: &UserFromContext(ctx).ID,
|
||||
FullName: UserFromContext(ctx).FullName,
|
||||
PrimaryEmailAddress: UserFromContext(ctx).EmailAddress,
|
||||
AdditionalEmailAddresses: []string{},
|
||||
@@ -1379,48 +1392,49 @@ func (r *mutationResolver) ConfirmEmail(ctx context.Context, input types.Confirm
|
||||
|
||||
// InviteUser is the resolver for the inviteUser field.
|
||||
func (r *mutationResolver) InviteUser(ctx context.Context, input types.InviteUserInput) (*types.InviteUserPayload, error) {
|
||||
user := UserFromContext(ctx)
|
||||
|
||||
organizations, err := r.authzSvc.GetAllUserOrganizations(ctx, user.ID)
|
||||
authzSvc := r.AuthzService(ctx, input.OrganizationID.TenantID())
|
||||
invitation, err := authzSvc.InviteUserToOrganization(ctx, input.OrganizationID, input.Email, input.FullName, string(authz.RoleMember))
|
||||
if err != nil {
|
||||
panic(fmt.Errorf("failed to list organizations for user: %w", err))
|
||||
panic(fmt.Errorf("failed to invite user to organization: %w", err))
|
||||
}
|
||||
|
||||
for _, organization := range organizations {
|
||||
if organization.ID == input.OrganizationID {
|
||||
invitation, err := r.authzSvc.InviteUserToOrganization(ctx, input.OrganizationID, input.Email, input.FullName, string(authz.RoleMember))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if input.CreatePeople {
|
||||
prb := r.ProboService(ctx, input.OrganizationID.TenantID())
|
||||
_, err := prb.Peoples.Create(ctx, probo.CreatePeopleRequest{
|
||||
OrganizationID: input.OrganizationID,
|
||||
FullName: input.FullName,
|
||||
PrimaryEmailAddress: input.Email,
|
||||
AdditionalEmailAddresses: []string{},
|
||||
Kind: coredata.PeopleKindEmployee,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create people record: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
return &types.InviteUserPayload{
|
||||
InvitationEdge: types.NewInvitationEdge(invitation, coredata.InvitationOrderFieldCreatedAt),
|
||||
}, nil
|
||||
if input.CreatePeople {
|
||||
prb := r.ProboService(ctx, input.OrganizationID.TenantID())
|
||||
_, err := prb.Peoples.Create(ctx, probo.CreatePeopleRequest{
|
||||
OrganizationID: input.OrganizationID,
|
||||
FullName: input.FullName,
|
||||
PrimaryEmailAddress: input.Email,
|
||||
AdditionalEmailAddresses: []string{},
|
||||
Kind: coredata.PeopleKindEmployee,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create people record: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
return nil, fmt.Errorf("organization not found")
|
||||
return &types.InviteUserPayload{
|
||||
InvitationEdge: types.NewInvitationEdge(invitation, coredata.InvitationOrderFieldCreatedAt),
|
||||
}, nil
|
||||
}
|
||||
|
||||
// AcceptInvitation is the resolver for the acceptInvitation field.
|
||||
func (r *mutationResolver) AcceptInvitation(ctx context.Context, input types.AcceptInvitationInput) (*types.AcceptInvitationPayload, error) {
|
||||
user := UserFromContext(ctx)
|
||||
|
||||
invitation, err := r.authzSvc.AcceptInvitationByID(ctx, input.InvitationID, user.ID)
|
||||
if err != nil {
|
||||
panic(fmt.Errorf("failed to accept invitation: %w", err))
|
||||
}
|
||||
|
||||
return &types.AcceptInvitationPayload{Invitation: types.NewInvitation(invitation)}, nil
|
||||
}
|
||||
|
||||
// DeleteInvitation is the resolver for the deleteInvitation field.
|
||||
func (r *mutationResolver) DeleteInvitation(ctx context.Context, input types.DeleteInvitationInput) (*types.DeleteInvitationPayload, error) {
|
||||
err := r.authzSvc.DeleteInvitation(ctx, input.InvitationID)
|
||||
authzSvc := r.AuthzService(ctx, input.InvitationID.TenantID())
|
||||
err := authzSvc.DeleteInvitation(ctx, input.InvitationID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
panic(fmt.Errorf("failed to delete invitation: %w", err))
|
||||
}
|
||||
|
||||
return &types.DeleteInvitationPayload{
|
||||
@@ -1430,25 +1444,13 @@ func (r *mutationResolver) DeleteInvitation(ctx context.Context, input types.Del
|
||||
|
||||
// RemoveMember is the resolver for the removeMember field.
|
||||
func (r *mutationResolver) RemoveMember(ctx context.Context, input types.RemoveMemberInput) (*types.RemoveMemberPayload, error) {
|
||||
user := UserFromContext(ctx)
|
||||
|
||||
organizations, err := r.authzSvc.GetAllUserOrganizations(ctx, user.ID)
|
||||
authzSvc := r.AuthzService(ctx, input.OrganizationID.TenantID())
|
||||
err := authzSvc.RemoveMemberFromOrganization(ctx, input.OrganizationID, input.MemberID)
|
||||
if err != nil {
|
||||
panic(fmt.Errorf("failed to list organizations for user: %w", err))
|
||||
return nil, err
|
||||
}
|
||||
|
||||
for _, organization := range organizations {
|
||||
if organization.ID == input.OrganizationID {
|
||||
err := r.authzSvc.RemoveMemberFromOrganization(ctx, input.OrganizationID, input.MemberID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return &types.RemoveMemberPayload{Success: true}, nil
|
||||
}
|
||||
}
|
||||
|
||||
return nil, fmt.Errorf("organization not found")
|
||||
return &types.RemoveMemberPayload{DeletedMemberID: input.MemberID}, nil
|
||||
}
|
||||
|
||||
// CreatePeople is the resolver for the createPeople field.
|
||||
@@ -3611,16 +3613,17 @@ func (r *organizationResolver) Memberships(ctx context.Context, obj *types.Organ
|
||||
|
||||
cursor := types.NewCursor(first, after, last, before, pageOrderBy)
|
||||
|
||||
page, err := r.authzSvc.GetAllOrganizationMemberships(ctx, obj.ID, cursor)
|
||||
authzSvc := r.AuthzService(ctx, obj.ID.TenantID())
|
||||
page, err := authzSvc.GetMembershipsByOrganizationID(ctx, obj.ID, cursor)
|
||||
if err != nil {
|
||||
panic(fmt.Errorf("cannot list memberships: %w", err))
|
||||
}
|
||||
|
||||
return types.NewMembershipConnection(page), nil
|
||||
return types.NewMembershipConnection(page, r, obj.ID), nil
|
||||
}
|
||||
|
||||
// Invitations is the resolver for the invitations field.
|
||||
func (r *organizationResolver) Invitations(ctx context.Context, obj *types.Organization, first *int, after *page.CursorKey, last *int, before *page.CursorKey, orderBy *types.InvitationOrder) (*types.InvitationConnection, error) {
|
||||
func (r *organizationResolver) Invitations(ctx context.Context, obj *types.Organization, first *int, after *page.CursorKey, last *int, before *page.CursorKey, orderBy *types.InvitationOrder, filter *types.InvitationFilter) (*types.InvitationConnection, error) {
|
||||
pageOrderBy := page.OrderBy[coredata.InvitationOrderField]{
|
||||
Field: coredata.InvitationOrderFieldCreatedAt,
|
||||
Direction: page.OrderDirectionDesc,
|
||||
@@ -3634,12 +3637,13 @@ func (r *organizationResolver) Invitations(ctx context.Context, obj *types.Organ
|
||||
|
||||
cursor := types.NewCursor(first, after, last, before, pageOrderBy)
|
||||
|
||||
page, err := r.authzSvc.GetAllOrganizationInvitations(ctx, obj.ID, cursor)
|
||||
authzSvc := r.AuthzService(ctx, obj.ID.TenantID())
|
||||
page, err := authzSvc.GetInvitationsByOrganizationID(ctx, obj.ID, cursor)
|
||||
if err != nil {
|
||||
panic(fmt.Errorf("cannot list invitations: %w", err))
|
||||
}
|
||||
|
||||
return types.NewInvitationConnection(page), nil
|
||||
return types.NewInvitationConnection(page, r, obj.ID, filter), nil
|
||||
}
|
||||
|
||||
// Connectors is the resolver for the connectors field.
|
||||
@@ -4943,41 +4947,19 @@ func (r *trustCenterReferenceConnectionResolver) TotalCount(ctx context.Context,
|
||||
return count, nil
|
||||
}
|
||||
|
||||
// People is the resolver for the people field.
|
||||
func (r *userResolver) People(ctx context.Context, obj *types.User, organizationID gid.GID) (*types.People, error) {
|
||||
prb := r.ProboService(ctx, organizationID.TenantID())
|
||||
|
||||
people, err := prb.Peoples.GetByUserID(ctx, obj.ID)
|
||||
if err != nil {
|
||||
var errPeopleNotFound *coredata.ErrPeopleNotFound
|
||||
if errors.As(err, &errPeopleNotFound) {
|
||||
return nil, nil
|
||||
}
|
||||
panic(fmt.Errorf("failed to get people: %w", err))
|
||||
}
|
||||
|
||||
return types.NewPeople(people), nil
|
||||
}
|
||||
|
||||
// TotalCount is the resolver for the totalCount field.
|
||||
func (r *userConnectionResolver) TotalCount(ctx context.Context, obj *types.UserConnection) (int, error) {
|
||||
currentUser := UserFromContext(ctx)
|
||||
if currentUser == nil {
|
||||
return 0, fmt.Errorf("no authenticated user")
|
||||
switch obj.Resolver.(type) {
|
||||
case *organizationResolver:
|
||||
authzSvc := r.AuthzService(ctx, obj.ParentID.TenantID())
|
||||
count, err := authzSvc.CountOrganizationUsers(ctx, obj.ParentID)
|
||||
if err != nil {
|
||||
panic(fmt.Errorf("failed to count organization users: %w", err))
|
||||
}
|
||||
return count, nil
|
||||
default:
|
||||
panic(fmt.Errorf("unknown resolver type for user connection"))
|
||||
}
|
||||
|
||||
memberships, err := r.authzSvc.GetAllUserOrganizations(ctx, currentUser.ID)
|
||||
if err != nil || len(memberships) == 0 {
|
||||
return 0, fmt.Errorf("user has no organization memberships")
|
||||
}
|
||||
|
||||
orgID := memberships[0].ID
|
||||
count, err := r.authzSvc.CountOrganizationMemberships(ctx, orgID)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("failed to count memberships: %w", err)
|
||||
}
|
||||
|
||||
return count, nil
|
||||
}
|
||||
|
||||
// Organization is the resolver for the organization field.
|
||||
@@ -5353,6 +5335,35 @@ func (r *viewerResolver) Organizations(ctx context.Context, obj *types.Viewer, f
|
||||
return types.NewOrganizationConnection(page), nil
|
||||
}
|
||||
|
||||
// Invitations is the resolver for the invitations field.
|
||||
func (r *viewerResolver) Invitations(ctx context.Context, obj *types.Viewer, first *int, after *page.CursorKey, last *int, before *page.CursorKey, orderBy *types.InvitationOrder, filter *types.InvitationFilter) (*types.InvitationConnection, error) {
|
||||
user := UserFromContext(ctx)
|
||||
|
||||
pageOrderBy := page.OrderBy[coredata.InvitationOrderField]{
|
||||
Field: coredata.InvitationOrderFieldCreatedAt,
|
||||
Direction: page.OrderDirectionDesc,
|
||||
}
|
||||
if orderBy != nil {
|
||||
pageOrderBy = page.OrderBy[coredata.InvitationOrderField]{
|
||||
Field: orderBy.Field,
|
||||
Direction: orderBy.Direction,
|
||||
}
|
||||
}
|
||||
cursor := types.NewCursor(first, after, last, before, pageOrderBy)
|
||||
|
||||
invitationFilter := coredata.NewInvitationFilter(nil)
|
||||
if filter != nil {
|
||||
invitationFilter = coredata.NewInvitationFilter(filter.OnlyPending)
|
||||
}
|
||||
|
||||
invitations, err := r.authzSvc.GetUserInvitations(ctx, user.EmailAddress, cursor, invitationFilter)
|
||||
if err != nil {
|
||||
panic(fmt.Errorf("failed to list invitations for user: %w", err))
|
||||
}
|
||||
|
||||
return types.NewInvitationConnection(invitations, r, gid.GID{}, filter), nil
|
||||
}
|
||||
|
||||
// Asset returns schema.AssetResolver implementation.
|
||||
func (r *Resolver) Asset() schema.AssetResolver { return &assetResolver{r} }
|
||||
|
||||
@@ -5432,6 +5443,9 @@ func (r *Resolver) FrameworkConnection() schema.FrameworkConnectionResolver {
|
||||
return &frameworkConnectionResolver{r}
|
||||
}
|
||||
|
||||
// Invitation returns schema.InvitationResolver implementation.
|
||||
func (r *Resolver) Invitation() schema.InvitationResolver { return &invitationResolver{r} }
|
||||
|
||||
// InvitationConnection returns schema.InvitationConnectionResolver implementation.
|
||||
func (r *Resolver) InvitationConnection() schema.InvitationConnectionResolver {
|
||||
return &invitationConnectionResolver{r}
|
||||
@@ -5541,9 +5555,6 @@ func (r *Resolver) TrustCenterReferenceConnection() schema.TrustCenterReferenceC
|
||||
return &trustCenterReferenceConnectionResolver{r}
|
||||
}
|
||||
|
||||
// User returns schema.UserResolver implementation.
|
||||
func (r *Resolver) User() schema.UserResolver { return &userResolver{r} }
|
||||
|
||||
// UserConnection returns schema.UserConnectionResolver implementation.
|
||||
func (r *Resolver) UserConnection() schema.UserConnectionResolver { return &userConnectionResolver{r} }
|
||||
|
||||
@@ -5603,6 +5614,7 @@ type evidenceConnectionResolver struct{ *Resolver }
|
||||
type fileResolver struct{ *Resolver }
|
||||
type frameworkResolver struct{ *Resolver }
|
||||
type frameworkConnectionResolver struct{ *Resolver }
|
||||
type invitationResolver struct{ *Resolver }
|
||||
type invitationConnectionResolver struct{ *Resolver }
|
||||
type measureResolver struct{ *Resolver }
|
||||
type measureConnectionResolver struct{ *Resolver }
|
||||
@@ -5630,7 +5642,6 @@ type trustCenterDocumentAccessResolver struct{ *Resolver }
|
||||
type trustCenterDocumentAccessConnectionResolver struct{ *Resolver }
|
||||
type trustCenterReferenceResolver struct{ *Resolver }
|
||||
type trustCenterReferenceConnectionResolver struct{ *Resolver }
|
||||
type userResolver struct{ *Resolver }
|
||||
type userConnectionResolver struct{ *Resolver }
|
||||
type vendorResolver struct{ *Resolver }
|
||||
type vendorBusinessAssociateAgreementResolver struct{ *Resolver }
|
||||
|
||||
Reference in New Issue
Block a user