diff --git a/models/verification.go b/models/verification.go index d3292027..c836cd94 100644 --- a/models/verification.go +++ b/models/verification.go @@ -10,13 +10,12 @@ import ( type VerificationType string const ( - TypeEmailVerification VerificationType = "email_verification" - TypePasswordResetRequest VerificationType = "password_reset_request" - TypeEmailResetRequest VerificationType = "email_reset_request" - TypeMagicLinkSignInRequest VerificationType = "magic_link_sign_in_request" - TypeMagicLinkExchangeCode VerificationType = "magic_link_exchange_code" - TypeTOTPPendingAuth VerificationType = "totp_pending_auth" - TypeOrganizationInvitationVerify VerificationType = "organization_invitation_verify" + TypeEmailVerification VerificationType = "email_verification" + TypePasswordResetRequest VerificationType = "password_reset_request" + TypeEmailResetRequest VerificationType = "email_reset_request" + TypeMagicLinkSignInRequest VerificationType = "magic_link_sign_in_request" + TypeMagicLinkExchangeCode VerificationType = "magic_link_exchange_code" + TypeTOTPPendingAuth VerificationType = "totp_pending_auth" ) func (vt VerificationType) String() string { @@ -32,7 +31,6 @@ func (VerificationType) PrepareJSONSchema(schema *jsonschema.Schema) error { string(TypeMagicLinkSignInRequest), string(TypeMagicLinkExchangeCode), string(TypeTOTPPendingAuth), - string(TypeOrganizationInvitationVerify), } schema.WithDescription("The type of the verification") return nil diff --git a/plugins/organizations/api.go b/plugins/organizations/api.go index c5790769..f6fa70c2 100644 --- a/plugins/organizations/api.go +++ b/plugins/organizations/api.go @@ -50,12 +50,12 @@ func (a *API) CreateInvitation(ctx context.Context, actor *models.Actor, organiz return a.useCases.CreateOrganizationInvitation(ctx, actor, organizationID, request, redirectURL) } -func (a *API) GetInvitation(ctx context.Context, actor *models.Actor, organizationID string, invitationID string) (*types.OrganizationInvitation, error) { - return a.useCases.GetOrganizationInvitation(ctx, actor, organizationID, invitationID) +func (a *API) GetAllInvitations(ctx context.Context, actor *models.Actor, organizationID string) ([]types.GetOrganizationInvitationResponse, error) { + return a.useCases.GetAllOrganizationInvitations(ctx, actor, organizationID) } -func (a *API) GetAllInvitations(ctx context.Context, actor *models.Actor, organizationID string) ([]types.OrganizationInvitation, error) { - return a.useCases.GetAllOrganizationInvitations(ctx, actor, organizationID) +func (a *API) GetInvitation(ctx context.Context, actor *models.Actor, organizationID string, invitationID string) (*types.GetOrganizationInvitationResponse, error) { + return a.useCases.GetOrganizationInvitation(ctx, actor, organizationID, invitationID) } func (a *API) RevokeInvitation(ctx context.Context, actor *models.Actor, organizationID string, invitationID string) (*types.OrganizationInvitation, error) { @@ -70,10 +70,6 @@ func (a *API) RejectInvitation(ctx context.Context, actor *models.Actor, organiz return a.useCases.RejectOrganizationInvitation(ctx, actor, organizationID, invitationID) } -func (a *API) VerifyInvitation(ctx context.Context, actor *models.Actor, organizationID string, invitationID string, token string) (*types.VerifyOrganizationInvitationResponse, error) { - return a.useCases.VerifyOrganizationInvitation(ctx, actor, organizationID, invitationID, token) -} - // Members func (a *API) AddMember(ctx context.Context, actor *models.Actor, organizationID string, request types.AddOrganizationMemberRequest) (*types.OrganizationMember, error) { diff --git a/plugins/organizations/constants/errors.go b/plugins/organizations/constants/errors.go index e5be8d67..1290e07d 100644 --- a/plugins/organizations/constants/errors.go +++ b/plugins/organizations/constants/errors.go @@ -9,12 +9,10 @@ import ( ) var ( - ErrOrganizationsQuotaExceeded = errors.New("organizations quota exceeded") - ErrMembersQuotaExceeded = errors.New("members quota exceeded") - ErrInvitationsQuotaExceeded = errors.New("invitations quota exceeded") - ErrInvitationVerificationFailed = errors.New("verification token not found or already used") - ErrInvitationVerificationExpired = errors.New("verification token has expired") - ErrInvitationEmailMismatch = errors.New("this invitation was sent to a different email address") + ErrOrganizationsQuotaExceeded = errors.New("organizations quota exceeded") + ErrMembersQuotaExceeded = errors.New("members quota exceeded") + ErrInvitationsQuotaExceeded = errors.New("invitations quota exceeded") + ErrInvitationEmailMismatch = errors.New("this invitation was sent to a different email address") ) func HandleError(err error, reqCtx *models.RequestContext) { @@ -23,10 +21,6 @@ func HandleError(err error, reqCtx *models.RequestContext) { switch err { case ErrOrganizationsQuotaExceeded, ErrMembersQuotaExceeded, ErrInvitationsQuotaExceeded: status = http.StatusTooManyRequests - case ErrInvitationVerificationFailed: - status = http.StatusNotFound - case ErrInvitationVerificationExpired: - status = http.StatusGone case ErrInvitationEmailMismatch: status = http.StatusForbidden } diff --git a/plugins/organizations/handlers/handler_test_helpers.go b/plugins/organizations/handlers/handler_test_helpers.go index fec367f5..0b8fb8fc 100644 --- a/plugins/organizations/handlers/handler_test_helpers.go +++ b/plugins/organizations/handlers/handler_test_helpers.go @@ -29,12 +29,6 @@ func defaultMockUserService() *internaltests.MockUserService { return svc } -func defaultMockVerificationService() *internaltests.MockVerificationService { - svc := &internaltests.MockVerificationService{} - svc.On("DeleteByUserIDAndType", mock.Anything, mock.Anything, mock.Anything).Return(nil).Maybe() - return svc -} - func newOrgUseCases(svc orgservices.OrganizationService) *orgusecases.UseCases { return orgusecases.NewUseCases( svc, @@ -43,23 +37,19 @@ func newOrgUseCases(svc orgservices.OrganizationService) *orgusecases.UseCases { &orgtests.MockOrganizationTeamService{}, &orgtests.MockOrganizationTeamMemberService{}, defaultMockUserService(), - defaultMockVerificationService(), - &internaltests.MockTokenService{}, &models.Config{}, &noopAuthorizer{}, ) } -func newInvitationUseCases(svc orgservices.OrganizationInvitationService) *orgusecases.UseCases { +func newInvitationUseCases(orgSvc orgservices.OrganizationService, svc orgservices.OrganizationInvitationService) *orgusecases.UseCases { return orgusecases.NewUseCases( - &orgtests.MockOrganizationService{}, + orgSvc, svc, &orgtests.MockOrganizationMemberService{}, &orgtests.MockOrganizationTeamService{}, &orgtests.MockOrganizationTeamMemberService{}, defaultMockUserService(), - defaultMockVerificationService(), - &internaltests.MockTokenService{}, &models.Config{}, &noopAuthorizer{}, ) @@ -73,8 +63,6 @@ func newMemberUseCases(svc orgservices.OrganizationMemberService) *orgusecases.U &orgtests.MockOrganizationTeamService{}, &orgtests.MockOrganizationTeamMemberService{}, defaultMockUserService(), - defaultMockVerificationService(), - &internaltests.MockTokenService{}, &models.Config{}, &noopAuthorizer{}, ) @@ -88,8 +76,6 @@ func newTeamUseCases(svc orgservices.OrganizationTeamService) *orgusecases.UseCa svc, &orgtests.MockOrganizationTeamMemberService{}, defaultMockUserService(), - defaultMockVerificationService(), - &internaltests.MockTokenService{}, &models.Config{}, &noopAuthorizer{}, ) @@ -103,8 +89,6 @@ func newTeamMemberUseCases(svc orgservices.OrganizationTeamMemberService) *orgus &orgtests.MockOrganizationTeamService{}, svc, defaultMockUserService(), - defaultMockVerificationService(), - &internaltests.MockTokenService{}, &models.Config{}, &noopAuthorizer{}, ) diff --git a/plugins/organizations/handlers/organization_invitation_handlers.go b/plugins/organizations/handlers/organization_invitation_handlers.go index 6c1bed06..f7d23a9d 100644 --- a/plugins/organizations/handlers/organization_invitation_handlers.go +++ b/plugins/organizations/handlers/organization_invitation_handlers.go @@ -78,13 +78,13 @@ func (h *GetOrganizationInvitationHandler) Handle() http.HandlerFunc { organizationID := r.PathValue("organization_id") invitationID := r.PathValue("invitation_id") - invitation, err := h.UseCases.GetOrganizationInvitation(ctx, actor, organizationID, invitationID) + resp, err := h.UseCases.GetOrganizationInvitation(ctx, actor, organizationID, invitationID) if err != nil { orgconstants.HandleError(err, reqCtx) return } - reqCtx.SetJSONResponse(http.StatusOK, invitation) + reqCtx.SetJSONResponse(http.StatusOK, resp) } } @@ -151,53 +151,6 @@ func (h *AcceptOrganizationInvitationHandler) Handle() http.HandlerFunc { } } -type VerifyOrganizationInvitationHandler struct { - UseCases *orgusecases.UseCases - TrustedOrigins []string -} - -func (h *VerifyOrganizationInvitationHandler) Handle() http.HandlerFunc { - return func(w http.ResponseWriter, r *http.Request) { - ctx := r.Context() - reqCtx, _ := models.GetRequestContext(ctx) - actor := reqCtx.Actor - - organizationID := r.PathValue("organization_id") - invitationID := r.PathValue("invitation_id") - - token := r.URL.Query().Get("token") - if token == "" { - reqCtx.SetJSONResponse(http.StatusUnprocessableEntity, map[string]any{"message": "token is required"}) - reqCtx.Handled = true - return - } - - resp, err := h.UseCases.VerifyOrganizationInvitation(ctx, actor, organizationID, invitationID, token) - if err != nil { - orgconstants.HandleError(err, reqCtx) - return - } - - redirectURL := r.URL.Query().Get("redirect_url") - if redirectURL != "" { - validatedURL, err := util.IsTrustedCallbackURL(redirectURL, h.TrustedOrigins) - if err != nil { - reqCtx.SetJSONResponse(http.StatusBadRequest, map[string]any{"message": err.Error()}) - reqCtx.Handled = true - return - } - q := validatedURL.Query() - q.Set("invite_id", invitationID) - validatedURL.RawQuery = q.Encode() - reqCtx.RedirectURL = validatedURL.String() - reqCtx.ResponseStatus = http.StatusFound - return - } - - reqCtx.SetJSONResponse(http.StatusOK, resp) - } -} - type RejectOrganizationInvitationHandler struct { UseCases *orgusecases.UseCases } diff --git a/plugins/organizations/handlers/organization_invitation_handlers_test.go b/plugins/organizations/handlers/organization_invitation_handlers_test.go index 11dba41e..da53002d 100644 --- a/plugins/organizations/handlers/organization_invitation_handlers_test.go +++ b/plugins/organizations/handlers/organization_invitation_handlers_test.go @@ -20,6 +20,7 @@ import ( type organizationInvitationHandlerFixture struct { service *orgtests.MockOrganizationInvitationService + orgSvc *orgtests.MockOrganizationService } type organizationInvitationHandlerCase struct { @@ -35,7 +36,10 @@ type organizationInvitationHandlerCase struct { } func newOrganizationInvitationHandlerFixture() *organizationInvitationHandlerFixture { - return &organizationInvitationHandlerFixture{service: &orgtests.MockOrganizationInvitationService{}} + return &organizationInvitationHandlerFixture{ + service: &orgtests.MockOrganizationInvitationService{}, + orgSvc: &orgtests.MockOrganizationService{}, + } } func (f *organizationInvitationHandlerFixture) newRequest(t *testing.T, method, path string, body []byte, userID *string, organizationID, invitationID string) (*http.Request, *httptest.ResponseRecorder, *models.RequestContext) { @@ -84,6 +88,7 @@ func runOrganizationInvitationHandlerCases(t *testing.T, method, path string, bu tt.checkResponse(t, reqCtx) } fixture.service.AssertExpectations(t) + fixture.orgSvc.AssertExpectations(t) }) } } @@ -92,7 +97,7 @@ func TestCreateOrganizationInvitationHandler(t *testing.T) { t.Parallel() runOrganizationInvitationHandlerCases(t, http.MethodPost, "/organizations/org-1/invitations", func(fixture *organizationInvitationHandlerFixture) http.HandlerFunc { - return (&CreateOrganizationInvitationHandler{UseCases: newInvitationUseCases(fixture.service)}).Handle() + return (&CreateOrganizationInvitationHandler{UseCases: newInvitationUseCases(fixture.orgSvc, fixture.service)}).Handle() }, []organizationInvitationHandlerCase{ { name: "missing_user", @@ -155,7 +160,7 @@ func TestGetAllOrganizationInvitationsHandler(t *testing.T) { t.Parallel() runOrganizationInvitationHandlerCases(t, http.MethodGet, "/organizations/org-1/invitations", func(fixture *organizationInvitationHandlerFixture) http.HandlerFunc { - return (&GetAllOrganizationInvitationsHandler{UseCases: newInvitationUseCases(fixture.service)}).Handle() + return (&GetAllOrganizationInvitationsHandler{UseCases: newInvitationUseCases(fixture.orgSvc, fixture.service)}).Handle() }, []organizationInvitationHandlerCase{ { name: "missing_user", @@ -168,7 +173,7 @@ func TestGetAllOrganizationInvitationsHandler(t *testing.T) { userID: new("user-1"), organizationID: "org-1", prepare: func(fixture *organizationInvitationHandlerFixture) { - fixture.service.On("GetAllOrganizationInvitations", mock.Anything, "user-1", "org-1").Return(([]orgtypes.OrganizationInvitation)(nil), errors.New("some error")).Once() + fixture.service.On("GetAllOrganizationInvitationsByOrgIDWithOrg", mock.Anything, "org-1").Return(([]orgtypes.GetOrganizationInvitationResponse)(nil), errors.New("some error")).Once() }, expectedStatus: http.StatusBadRequest, expectedMessage: "some error", @@ -178,14 +183,20 @@ func TestGetAllOrganizationInvitationsHandler(t *testing.T) { userID: new("user-1"), organizationID: "org-1", prepare: func(fixture *organizationInvitationHandlerFixture) { - fixture.service.On("GetAllOrganizationInvitations", mock.Anything, "user-1", "org-1").Return([]orgtypes.OrganizationInvitation{{ID: "inv-1", OrganizationID: "org-1", Email: "user@example.com", Role: "member", Status: orgtypes.OrganizationInvitationStatusPending}}, nil).Once() + fixture.service.On("GetAllOrganizationInvitationsByOrgIDWithOrg", mock.Anything, "org-1").Return([]orgtypes.GetOrganizationInvitationResponse{ + { + Invitation: &orgtypes.OrganizationInvitation{ID: "inv-1", OrganizationID: "org-1", Email: "user@example.com", Role: "member", Status: orgtypes.OrganizationInvitationStatusPending}, + Organization: orgtypes.OrganizationSummary{ID: "org-1", Name: "Acme Corp", Slug: "acme"}, + }, + }, nil).Once() }, expectedStatus: http.StatusOK, checkResponse: func(t *testing.T, reqCtx *models.RequestContext) { - invitations := internaltests.DecodeResponseJSON[[]orgtypes.OrganizationInvitation](t, reqCtx) - require.Len(t, invitations, 1) - assert.Equal(t, "inv-1", invitations[0].ID) - assert.Equal(t, "org-1", invitations[0].OrganizationID) + resp := internaltests.DecodeResponseJSON[[]orgtypes.GetOrganizationInvitationResponse](t, reqCtx) + require.Len(t, resp, 1) + assert.Equal(t, "inv-1", resp[0].Invitation.ID) + assert.Equal(t, "org-1", resp[0].Invitation.OrganizationID) + assert.Equal(t, "Acme Corp", resp[0].Organization.Name) }, }, }) @@ -195,7 +206,7 @@ func TestGetOrganizationInvitationHandler(t *testing.T) { t.Parallel() runOrganizationInvitationHandlerCases(t, http.MethodGet, "/organizations/org-1/invitations/inv-1", func(fixture *organizationInvitationHandlerFixture) http.HandlerFunc { - return (&GetOrganizationInvitationHandler{UseCases: newInvitationUseCases(fixture.service)}).Handle() + return (&GetOrganizationInvitationHandler{UseCases: newInvitationUseCases(fixture.orgSvc, fixture.service)}).Handle() }, []organizationInvitationHandlerCase{ { name: "missing_user", @@ -210,7 +221,7 @@ func TestGetOrganizationInvitationHandler(t *testing.T) { organizationID: "org-1", invitationID: "inv-1", prepare: func(fixture *organizationInvitationHandlerFixture) { - fixture.service.On("GetOrganizationInvitationByID", mock.Anything, "inv-1").Return((*orgtypes.OrganizationInvitation)(nil), coreerrors.ErrNotFound).Once() + fixture.service.On("GetOrganizationInvitationByIDWithOrg", mock.Anything, "inv-1").Return((*orgtypes.GetOrganizationInvitationResponse)(nil), coreerrors.ErrNotFound).Once() }, expectedStatus: http.StatusNotFound, expectedMessage: "not found", @@ -221,14 +232,20 @@ func TestGetOrganizationInvitationHandler(t *testing.T) { organizationID: "org-1", invitationID: "inv-1", prepare: func(fixture *organizationInvitationHandlerFixture) { - fixture.service.On("GetOrganizationInvitationByID", mock.Anything, "inv-1").Return(&orgtypes.OrganizationInvitation{ID: "inv-1", OrganizationID: "org-1", Email: "user@example.com", Role: "member", Status: orgtypes.OrganizationInvitationStatusPending}, nil).Once() - fixture.service.On("GetOrganizationInvitation", mock.Anything, "user-1", "org-1", "inv-1").Return(&orgtypes.OrganizationInvitation{ID: "inv-1", OrganizationID: "org-1", Email: "user@example.com", Role: "member", Status: orgtypes.OrganizationInvitationStatusPending}, nil).Once() + fixture.service.On("GetOrganizationInvitationByIDWithOrg", mock.Anything, "inv-1").Return(&orgtypes.GetOrganizationInvitationResponse{ + Invitation: &orgtypes.OrganizationInvitation{ID: "inv-1", OrganizationID: "org-1", Email: "user@example.com", Role: "member", Status: orgtypes.OrganizationInvitationStatusPending}, + Organization: orgtypes.OrganizationSummary{ID: "org-1", Name: "Acme Corp", Slug: "acme"}, + }, nil).Once() }, expectedStatus: http.StatusOK, checkResponse: func(t *testing.T, reqCtx *models.RequestContext) { - invitation := internaltests.DecodeResponseJSON[orgtypes.OrganizationInvitation](t, reqCtx) - assert.Equal(t, "inv-1", invitation.ID) - assert.Equal(t, "org-1", invitation.OrganizationID) + resp := internaltests.DecodeResponseJSON[orgtypes.GetOrganizationInvitationResponse](t, reqCtx) + require.NotNil(t, resp) + assert.Equal(t, "inv-1", resp.Invitation.ID) + assert.Equal(t, "org-1", resp.Invitation.OrganizationID) + assert.Equal(t, "org-1", resp.Organization.ID) + assert.Equal(t, "Acme Corp", resp.Organization.Name) + assert.Equal(t, "acme", resp.Organization.Slug) }, }, }) @@ -238,7 +255,7 @@ func TestRevokeOrganizationInvitationHandler(t *testing.T) { t.Parallel() runOrganizationInvitationHandlerCases(t, http.MethodPatch, "/organizations/org-1/invitations/inv-1", func(fixture *organizationInvitationHandlerFixture) http.HandlerFunc { - return (&RevokeOrganizationInvitationHandler{UseCases: newInvitationUseCases(fixture.service)}).Handle() + return (&RevokeOrganizationInvitationHandler{UseCases: newInvitationUseCases(fixture.orgSvc, fixture.service)}).Handle() }, []organizationInvitationHandlerCase{ { name: "missing_user", @@ -280,7 +297,7 @@ func TestAcceptOrganizationInvitationHandler(t *testing.T) { t.Parallel() runOrganizationInvitationHandlerCases(t, http.MethodPost, "/organizations/org-1/invitations/inv-1/accept", func(fixture *organizationInvitationHandlerFixture) http.HandlerFunc { - return (&AcceptOrganizationInvitationHandler{UseCases: newInvitationUseCases(fixture.service)}).Handle() + return (&AcceptOrganizationInvitationHandler{UseCases: newInvitationUseCases(fixture.orgSvc, fixture.service)}).Handle() }, []organizationInvitationHandlerCase{ { name: "redirect_url", @@ -350,7 +367,7 @@ func TestRejectOrganizationInvitationHandler(t *testing.T) { t.Parallel() runOrganizationInvitationHandlerCases(t, http.MethodPost, "/organizations/org-1/invitations/inv-1/reject", func(fixture *organizationInvitationHandlerFixture) http.HandlerFunc { - return (&RejectOrganizationInvitationHandler{UseCases: newInvitationUseCases(fixture.service)}).Handle() + return (&RejectOrganizationInvitationHandler{UseCases: newInvitationUseCases(fixture.orgSvc, fixture.service)}).Handle() }, []organizationInvitationHandlerCase{ { name: "missing_user", diff --git a/plugins/organizations/openapi/openapi_docs.go b/plugins/organizations/openapi/openapi_docs.go index 3fd0822c..7e7a6359 100644 --- a/plugins/organizations/openapi/openapi_docs.go +++ b/plugins/organizations/openapi/openapi_docs.go @@ -83,27 +83,17 @@ func RegisterOpenAPIDocs(svc openapi.OpenAPIService) error { openapi.WithDescription("Lists all invitations for an organization."), openapi.WithTags("Organization Invitations"), openapi.WithRequest(&types.OrganizationID{}), - openapi.WithResponseStatus(http.StatusOK, &[]types.OrganizationInvitation{}), + openapi.WithResponseStatus(http.StatusOK, &[]types.GetOrganizationInvitationResponse{}), ), svc.AddOperation( http.MethodGet, "/organizations/{organization_id}/invitations/{invitation_id}", openapi.WithOperationID("getOrganizationInvitation"), openapi.WithSummary("Get invitation"), - openapi.WithDescription("Retrieves a single invitation by its ID."), + openapi.WithDescription("Retrieves a single invitation."), openapi.WithTags("Organization Invitations"), openapi.WithRequest(&types.InvitationID{}), - openapi.WithResponseStatus(http.StatusOK, &types.OrganizationInvitation{}), - ), - svc.AddOperation( - http.MethodGet, - "/organizations/{organization_id}/invitations/{invitation_id}/verify", - openapi.WithOperationID("verifyOrganizationInvitation"), - openapi.WithSummary("Verify invitation"), - openapi.WithDescription("Verifies an invitation token."), - openapi.WithTags("Organization Invitations"), - openapi.WithRequest(&types.VerifyOrganizationInvitationQuery{}), - openapi.WithResponseStatus(http.StatusOK, &types.VerifyOrganizationInvitationResponse{}), + openapi.WithResponseStatus(http.StatusOK, &types.GetOrganizationInvitationResponse{}), ), svc.AddOperation( http.MethodPatch, diff --git a/plugins/organizations/plugin.go b/plugins/organizations/plugin.go index 1522d797..be802c23 100644 --- a/plugins/organizations/plugin.go +++ b/plugins/organizations/plugin.go @@ -70,16 +70,6 @@ func (p *OrganizationsPlugin) Init(ctx *models.PluginContext) error { return fmt.Errorf("user service not available in service registry") } - verificationService, ok := ctx.ServiceRegistry.Get(models.ServiceVerification.String()).(rootservices.VerificationService) - if !ok { - return fmt.Errorf("verification service not available in service registry") - } - - tokenService, ok := ctx.ServiceRegistry.Get(models.ServiceToken.String()).(rootservices.TokenService) - if !ok { - return fmt.Errorf("token service not available in service registry") - } - mailerService, ok := ctx.ServiceRegistry.Get(models.ServiceMailer.String()).(rootservices.MailerService) if !ok { p.logger.Warn("mailer service not available in service registry: automatic email sending will be disabled for the organizations plugin") @@ -109,13 +99,13 @@ func (p *OrganizationsPlugin) Init(ctx *models.PluginContext) error { } p.emailTemplateManager = emailTemplateManager - p.invitationService = services.NewOrganizationInvitationService(ctx.DB, p.globalConfig, &p.pluginConfig, p.logger, ctx.EventBus, userService, mailerService, accessControlService, verificationService, tokenService, p.organizationRepo, p.invitationRepo, p.memberRepo, p.serviceUtils, p.emailTemplateManager, p.hooksExecutor) + p.invitationService = services.NewOrganizationInvitationService(ctx.DB, p.globalConfig, &p.pluginConfig, p.logger, ctx.EventBus, userService, mailerService, accessControlService, p.organizationRepo, p.invitationRepo, p.memberRepo, p.serviceUtils, p.emailTemplateManager, p.hooksExecutor) p.memberService = services.NewOrganizationMemberService(userService, accessControlService, p.organizationRepo, p.memberRepo, p.pluginConfig.MembersLimit, ctx.DB, p.serviceUtils, p.hooksExecutor) p.teamService = services.NewOrganizationTeamService(p.organizationRepo, p.memberRepo, p.teamRepo, p.teamMemberRepo, p.serviceUtils, ctx.DB, p.hooksExecutor) p.teamMemberService = services.NewOrganizationTeamMemberService(p.organizationRepo, p.memberRepo, p.teamRepo, p.teamMemberRepo, p.serviceUtils, p.hooksExecutor) authorizer := rootservices.NewDefaultAuthorizer() - p.useCases = usecases.NewUseCases(p.organizationService, p.invitationService, p.memberService, p.teamService, p.teamMemberService, userService, verificationService, tokenService, p.globalConfig, authorizer) + p.useCases = usecases.NewUseCases(p.organizationService, p.invitationService, p.memberService, p.teamService, p.teamMemberService, userService, p.globalConfig, authorizer) p.Api = BuildAPI(p) diff --git a/plugins/organizations/repositories/bun_organization_invitation_repository.go b/plugins/organizations/repositories/bun_organization_invitation_repository.go index 888d8d8a..96c6f7f6 100644 --- a/plugins/organizations/repositories/bun_organization_invitation_repository.go +++ b/plugins/organizations/repositories/bun_organization_invitation_repository.go @@ -119,3 +119,88 @@ func (r *BunOrganizationInvitationRepository) CountByOrganizationIDAndEmail(ctx func (r *BunOrganizationInvitationRepository) WithTx(tx bun.IDB) OrganizationInvitationRepository { return &BunOrganizationInvitationRepository{db: tx} } + +type invitationOrgRow struct { + ID string `bun:"column:id"` + Email string `bun:"column:email"` + InviterID string `bun:"column:inviter_id"` + OrganizationID string `bun:"column:organization_id"` + Role string `bun:"column:role"` + Status string `bun:"column:status"` + ExpiresAt time.Time `bun:"column:expires_at"` + CreatedAt time.Time `bun:"column:created_at"` + + OrgID string `bun:"column:org_id"` + OrgOwnerID string `bun:"column:org_owner_id"` + OrgName string `bun:"column:org_name"` + OrgSlug string `bun:"column:org_slug"` + OrgLogo *string `bun:"column:org_logo"` + OrgMetadata map[string]any `bun:"column:org_metadata,type:jsonb"` +} + +const invitationWithOrgColumns = `i.id, i.email, i.inviter_id, i.organization_id, i.role, i.status, i.expires_at, i.created_at,` + + ` o.id AS org_id, o.owner_id AS org_owner_id, o.name AS org_name,` + + ` o.slug AS org_slug, o.logo AS org_logo, o.metadata AS org_metadata` + +func mapToInvitationWithOrgResponse(row invitationOrgRow) types.GetOrganizationInvitationResponse { + return types.GetOrganizationInvitationResponse{ + Invitation: &types.OrganizationInvitation{ + ID: row.ID, + Email: row.Email, + InviterID: row.InviterID, + OrganizationID: row.OrganizationID, + Role: row.Role, + Status: types.OrganizationInvitationStatus(row.Status), + ExpiresAt: row.ExpiresAt, + CreatedAt: row.CreatedAt, + }, + Organization: types.OrganizationSummary{ + ID: row.OrgID, + OwnerID: row.OrgOwnerID, + Name: row.OrgName, + Slug: row.OrgSlug, + Logo: row.OrgLogo, + Metadata: row.OrgMetadata, + }, + } +} + +func (r *BunOrganizationInvitationRepository) GetByIDWithOrg(ctx context.Context, invitationID string) (*types.GetOrganizationInvitationResponse, error) { + var row invitationOrgRow + err := r.db.NewRaw(` + SELECT `+invitationWithOrgColumns+` + FROM organization_invitations i + INNER JOIN organizations o ON o.id = i.organization_id + WHERE i.id = ? + `, invitationID).Scan(ctx, &row) + if err == sql.ErrNoRows { + return nil, nil + } + if err != nil { + return nil, err + } + result := mapToInvitationWithOrgResponse(row) + return &result, nil +} + +func (r *BunOrganizationInvitationRepository) GetAllByOrganizationIDWithOrg(ctx context.Context, organizationID string) ([]types.GetOrganizationInvitationResponse, error) { + var rows []invitationOrgRow + err := r.db.NewRaw(` + SELECT `+invitationWithOrgColumns+` + FROM organization_invitations i + INNER JOIN organizations o ON o.id = i.organization_id + WHERE i.organization_id = ? + ORDER BY i.created_at DESC + `, organizationID).Scan(ctx, &rows) + if err == sql.ErrNoRows { + return []types.GetOrganizationInvitationResponse{}, nil + } + if err != nil { + return nil, err + } + results := make([]types.GetOrganizationInvitationResponse, len(rows)) + for i, row := range rows { + results[i] = mapToInvitationWithOrgResponse(row) + } + return results, nil +} diff --git a/plugins/organizations/repositories/interfaces.go b/plugins/organizations/repositories/interfaces.go index 708ebaf6..50e08c6e 100644 --- a/plugins/organizations/repositories/interfaces.go +++ b/plugins/organizations/repositories/interfaces.go @@ -21,8 +21,10 @@ type OrganizationRepository interface { type OrganizationInvitationRepository interface { Create(ctx context.Context, invitation *types.OrganizationInvitation) (*types.OrganizationInvitation, error) GetByID(ctx context.Context, invitationID string) (*types.OrganizationInvitation, error) + GetByIDWithOrg(ctx context.Context, invitationID string) (*types.GetOrganizationInvitationResponse, error) GetByOrganizationIDAndEmail(ctx context.Context, organizationID string, email string, status ...types.OrganizationInvitationStatus) (*types.OrganizationInvitation, error) GetAllByOrganizationID(ctx context.Context, organizationID string) ([]types.OrganizationInvitation, error) + GetAllByOrganizationIDWithOrg(ctx context.Context, organizationID string) ([]types.GetOrganizationInvitationResponse, error) GetAllPendingByEmail(ctx context.Context, email string) ([]types.OrganizationInvitation, error) Update(ctx context.Context, invitation *types.OrganizationInvitation) (*types.OrganizationInvitation, error) CountByOrganizationIDAndEmail(ctx context.Context, organizationID string, email string) (int, error) diff --git a/plugins/organizations/routes.go b/plugins/organizations/routes.go index 35933c14..4d8dbee5 100644 --- a/plugins/organizations/routes.go +++ b/plugins/organizations/routes.go @@ -18,7 +18,6 @@ func Routes(plugin *OrganizationsPlugin) []models.Route { createInvitationHandler := &handlers.CreateOrganizationInvitationHandler{UseCases: plugin.useCases} getInvitationHandler := &handlers.GetOrganizationInvitationHandler{UseCases: plugin.useCases} getAllInvitationsHandler := &handlers.GetAllOrganizationInvitationsHandler{UseCases: plugin.useCases} - verifyInvitationHandler := &handlers.VerifyOrganizationInvitationHandler{UseCases: plugin.useCases, TrustedOrigins: plugin.globalConfig.Security.TrustedOrigins} revokeInvitationHandler := &handlers.RevokeOrganizationInvitationHandler{UseCases: plugin.useCases} acceptInvitationHandler := &handlers.AcceptOrganizationInvitationHandler{UseCases: plugin.useCases, TrustedOrigins: plugin.globalConfig.Security.TrustedOrigins} rejectInvitationHandler := &handlers.RejectOrganizationInvitationHandler{UseCases: plugin.useCases} @@ -108,14 +107,6 @@ func Routes(plugin *OrganizationsPlugin) []models.Route { }, Handler: getInvitationHandler.Handle(), }, - { - Method: http.MethodGet, - Path: "/organizations/{organization_id}/invitations/{invitation_id}/verify", - Middleware: []func(http.Handler) http.Handler{ - middleware.RequireActor(models.ActorUser), - }, - Handler: verifyInvitationHandler.Handle(), - }, { Method: http.MethodPatch, Path: "/organizations/{organization_id}/invitations/{invitation_id}/revoke", diff --git a/plugins/organizations/services/interfaces.go b/plugins/organizations/services/interfaces.go index 16a9ac45..5900b632 100644 --- a/plugins/organizations/services/interfaces.go +++ b/plugins/organizations/services/interfaces.go @@ -14,14 +14,15 @@ type OrganizationService interface { UpdateOrganization(ctx context.Context, actor *models.Actor, organizationID string, request types.UpdateOrganizationRequest) (*types.Organization, error) DeleteOrganization(ctx context.Context, actor *models.Actor, organizationID string) error ExistsByID(ctx context.Context, organizationID string) (bool, error) - GetByIDNoAuth(ctx context.Context, organizationID string) (*types.Organization, error) } type OrganizationInvitationService interface { CreateOrganizationInvitation(ctx context.Context, actor *models.Actor, organizationID string, request types.CreateOrganizationInvitationRequest, redirectURL string) (*types.OrganizationInvitation, error) GetOrganizationInvitation(ctx context.Context, actor *models.Actor, organizationID string, invitationID string) (*types.OrganizationInvitation, error) GetOrganizationInvitationByID(ctx context.Context, invitationID string) (*types.OrganizationInvitation, error) + GetOrganizationInvitationByIDWithOrg(ctx context.Context, invitationID string) (*types.GetOrganizationInvitationResponse, error) GetAllOrganizationInvitations(ctx context.Context, actor *models.Actor, organizationID string) ([]types.OrganizationInvitation, error) + GetAllOrganizationInvitationsByOrgIDWithOrg(ctx context.Context, organizationID string) ([]types.GetOrganizationInvitationResponse, error) RevokeOrganizationInvitation(ctx context.Context, actor *models.Actor, organizationID string, invitationID string) (*types.OrganizationInvitation, error) AcceptOrganizationInvitation(ctx context.Context, actor *models.Actor, organizationID string, invitationID string) (*types.OrganizationInvitation, error) RejectOrganizationInvitation(ctx context.Context, actor *models.Actor, organizationID string, invitationID string) (*types.OrganizationInvitation, error) diff --git a/plugins/organizations/services/organization_invitation_service.go b/plugins/organizations/services/organization_invitation_service.go index 8ab6ee30..ca6f1e5f 100644 --- a/plugins/organizations/services/organization_invitation_service.go +++ b/plugins/organizations/services/organization_invitation_service.go @@ -3,7 +3,6 @@ package services import ( "context" "database/sql" - "fmt" "net/mail" "net/url" "strings" @@ -35,8 +34,6 @@ type organizationInvitationService struct { userService rootservices.UserService mailerService rootservices.MailerService accessControlService rootservices.AccessControlService - verificationService rootservices.VerificationService - tokenService rootservices.TokenService organizationRepo repositories.OrganizationRepository orgInvitationRepo repositories.OrganizationInvitationRepository orgMemberRepo repositories.OrganizationMemberRepository @@ -54,8 +51,6 @@ func NewOrganizationInvitationService( userService rootservices.UserService, mailerService rootservices.MailerService, accessControlService rootservices.AccessControlService, - verificationService rootservices.VerificationService, - tokenService rootservices.TokenService, organizationRepo repositories.OrganizationRepository, orgInvitationRepo repositories.OrganizationInvitationRepository, orgMemberRepo repositories.OrganizationMemberRepository, @@ -76,8 +71,6 @@ func NewOrganizationInvitationService( userService: userService, mailerService: mailerService, accessControlService: accessControlService, - verificationService: verificationService, - tokenService: tokenService, organizationRepo: organizationRepo, orgInvitationRepo: orgInvitationRepo, orgMemberRepo: orgMemberRepo, @@ -183,22 +176,7 @@ func (s *organizationInvitationService) CreateOrganizationInvitation(ctx context s.publishOrganizationInvitationCreatedEvent(created, organization) - rawToken, err := s.tokenService.Generate() - if err != nil { - return nil, err - } - hashedToken := s.tokenService.Hash(rawToken) - - inviteeUser, _ := s.userService.GetByEmail(ctx, request.Email) - userID := "" - if inviteeUser != nil { - userID = inviteeUser.ID - } - if _, err := s.verificationService.Create(ctx, userID, hashedToken, models.TypeOrganizationInvitationVerify, created.ID, s.pluginConfig.InvitationExpiresIn); err != nil { - return nil, err - } - - verifyURL := s.buildOrganizationInvitationVerifyURL(created, rawToken, redirectURL) + inviteURL := s.buildOrganizationInvitationURL(created, redirectURL) callbackHandled := false if s.pluginConfig.SendOrganizationInvitationEmail != nil { @@ -211,7 +189,7 @@ func (s *organizationInvitationService) CreateOrganizationInvitation(ctx context Organization: organization, Invitation: created, Inviter: inviter, - AcceptURL: verifyURL, + InviteURL: inviteURL, }, reqCtx) if err != nil { @@ -227,7 +205,7 @@ func (s *organizationInvitationService) CreateOrganizationInvitation(ctx context taskCtx, cancel := context.WithTimeout(detachedCtx, 15*time.Second) defer cancel() - if err := s.sendOrganizationInvitationEmail(taskCtx, created, organization, verifyURL); err != nil { + if err := s.sendOrganizationInvitationEmail(taskCtx, created, organization, inviteURL); err != nil { s.logger.Error("failed to send organization invitation email via built-in email service", "invitation_id", created.ID, "error", err) } }() @@ -236,13 +214,13 @@ func (s *organizationInvitationService) CreateOrganizationInvitation(ctx context return created, nil } -func (s *organizationInvitationService) sendOrganizationInvitationEmail(ctx context.Context, invitation *types.OrganizationInvitation, organization *types.Organization, acceptURL string) error { +func (s *organizationInvitationService) sendOrganizationInvitationEmail(ctx context.Context, invitation *types.OrganizationInvitation, organization *types.Organization, inviteURL string) error { subject, textBody, htmlBody, err := s.emailTemplateManager.Render(emailconstants.OrganizationInvitationEmailTemplateName, types.OrganizationInvitationContext{ CommonContext: emailtmpl.NewCommonContext(s.globalConfig.AppName, s.globalConfig.BaseURL), InvitationEmail: invitation.Email, OrganizationName: organization.Name, Role: invitation.Role, - AcceptLink: acceptURL, + InviteLink: inviteURL, Expiry: s.pluginConfig.InvitationExpiresIn, }) if err != nil { @@ -252,25 +230,23 @@ func (s *organizationInvitationService) sendOrganizationInvitationEmail(ctx cont return s.mailerService.SendEmail(ctx, invitation.Email, subject, textBody, htmlBody) } -func (s *organizationInvitationService) buildOrganizationInvitationVerifyURL(invitation *types.OrganizationInvitation, rawToken string, redirectURL string) string { - baseURL := s.globalConfig.BaseURL - basePath := s.globalConfig.BasePath - verifyPath := fmt.Sprintf("/organizations/%s/invitations/%s/verify", url.PathEscape(invitation.OrganizationID), url.PathEscape(invitation.ID)) +func (s *organizationInvitationService) buildOrganizationInvitationURL(invitation *types.OrganizationInvitation, redirectURL string) string { + base := redirectURL + if base == "" { + base = s.globalConfig.BaseURL + s.globalConfig.BasePath + } - fullURL := baseURL + basePath + verifyPath - parsedURL, err := url.Parse(fullURL) + parsed, err := url.Parse(base) if err != nil { - return fullURL + return base } - query := parsedURL.Query() - query.Set("token", rawToken) - if redirectURL != "" { - query.Set("redirect_url", redirectURL) - } - parsedURL.RawQuery = query.Encode() + q := parsed.Query() + q.Set("organization_id", invitation.OrganizationID) + q.Set("invitation_id", invitation.ID) + parsed.RawQuery = q.Encode() - return parsedURL.String() + return parsed.String() } func (s *organizationInvitationService) publishOrganizationInvitationCreatedEvent(invitation *types.OrganizationInvitation, organization *types.Organization) { @@ -343,6 +319,35 @@ func (s *organizationInvitationService) GetOrganizationInvitationByID(ctx contex return invitation, nil } +func (s *organizationInvitationService) GetOrganizationInvitationByIDWithOrg(ctx context.Context, invitationID string) (*types.GetOrganizationInvitationResponse, error) { + if invitationID == "" { + return nil, coreerrors.ErrNotFound + } + + resp, err := s.orgInvitationRepo.GetByIDWithOrg(ctx, invitationID) + if err != nil { + return nil, err + } + if resp == nil { + return nil, coreerrors.ErrNotFound + } + + return resp, nil +} + +func (s *organizationInvitationService) GetAllOrganizationInvitationsByOrgIDWithOrg(ctx context.Context, organizationID string) ([]types.GetOrganizationInvitationResponse, error) { + if organizationID == "" { + return nil, coreerrors.ErrNotFound + } + + resp, err := s.orgInvitationRepo.GetAllByOrganizationIDWithOrg(ctx, organizationID) + if err != nil { + return nil, err + } + + return resp, nil +} + func (s *organizationInvitationService) RevokeOrganizationInvitation(ctx context.Context, actor *models.Actor, organizationID string, invitationID string) (*types.OrganizationInvitation, error) { if actor == nil || actor.ID == "" || organizationID == "" || invitationID == "" { return nil, coreerrors.ErrUnauthorized diff --git a/plugins/organizations/services/organization_invitation_service_test.go b/plugins/organizations/services/organization_invitation_service_test.go index ecf44f70..1585cc35 100644 --- a/plugins/organizations/services/organization_invitation_service_test.go +++ b/plugins/organizations/services/organization_invitation_service_test.go @@ -49,15 +49,13 @@ func newTestOrganizationInvitationService( orgRepo repositories.OrganizationRepository, invRepo repositories.OrganizationInvitationRepository, memberRepo repositories.OrganizationMemberRepository, - verificationService rootservices.VerificationService, - tokenService rootservices.TokenService, ) *organizationInvitationService { serviceUtils := &ServiceUtils{orgRepo: orgRepo, orgMemberRepo: memberRepo} tmplMgr := emailtmpl.NewManager() _ = tmplMgr.Register(emailtmpl.Definition{ Name: "organization_invitation", Subject: "You're invited to join {{.OrganizationName}} on {{.AppName}}", - Text: "You have been invited to join {{.OrganizationName}} on {{.AppName}} as {{.Role}}. Open this link to accept: {{.AcceptLink}}", + Text: "You have been invited to join {{.OrganizationName}} on {{.AppName}} as {{.Role}}. Open this link to accept: {{.InviteLink}}", HTML: "

Invited to {{.OrganizationName}} as {{.Role}}

", }) return NewOrganizationInvitationService( @@ -69,8 +67,6 @@ func newTestOrganizationInvitationService( userService, nil, accessControlService, - verificationService, - tokenService, orgRepo, invRepo, memberRepo, @@ -459,8 +455,8 @@ func TestOrganizationInvitationService_CreateOrganizationInvitation(t *testing.T require.Equal(t, "org-1", params.Organization.ID) require.Equal(t, "inv-1", params.Invitation.ID) require.Equal(t, "user-1", params.Inviter.ID) - require.Contains(t, params.AcceptURL, "/organizations/org-1/invitations/inv-1/verify?") - require.Contains(t, params.AcceptURL, "token=test-raw-token") + require.Contains(t, params.InviteURL, "organization_id=org-1") + require.Contains(t, params.InviteURL, "invitation_id=inv-1") require.Nil(t, reqCtx) return nil } @@ -532,9 +528,7 @@ func TestOrganizationInvitationService_CreateOrganizationInvitation(t *testing.T pluginConfig: func(config *types.OrganizationsPluginConfig, callbackCalled *bool) { config.SendOrganizationInvitationEmail = func(params types.SendOrganizationInvitationEmailParams, reqCtx *models.RequestContext) error { *callbackCalled = true - require.Contains(t, params.AcceptURL, "/organizations/org-1/invitations/inv-1/verify") - require.Contains(t, params.AcceptURL, "token=test-raw-token") - require.Contains(t, params.AcceptURL, "redirect_url=https%3A%2F%2Fapp.example.com%2Fwelcome") + require.Equal(t, "https://app.example.com/welcome?invitation_id=inv-1&organization_id=org-1", params.InviteURL) return nil } }, @@ -553,8 +547,6 @@ func TestOrganizationInvitationService_CreateOrganizationInvitation(t *testing.T for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - t.Parallel() - logger := &testInvitationLogger{} pluginConfig := &types.OrganizationsPluginConfig{ Enabled: true, @@ -599,19 +591,12 @@ func TestOrganizationInvitationService_CreateOrganizationInvitation(t *testing.T accessControlService = orgtests.NewAccessControlServiceStub() } - tokenSvc := &internaltests.MockTokenService{} - tokenSvc.On("Generate").Return("test-raw-token", nil).Maybe() - tokenSvc.On("Hash", mock.Anything).Return("hashed-token").Maybe() - - verifSvc := &internaltests.MockVerificationService{} - verifSvc.On("Create", mock.Anything, mock.Anything, mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return(&models.Verification{}, nil).Maybe() - serviceUtils := &ServiceUtils{orgRepo: orgRepo, orgMemberRepo: memberRepo} tmplMgr := emailtmpl.NewManager() _ = tmplMgr.Register(emailtmpl.Definition{ Name: "organization_invitation", Subject: "You're invited to join {{.OrganizationName}} on {{.AppName}}", - Text: "You have been invited to join {{.OrganizationName}} on {{.AppName}} as {{.Role}}. Open this link to accept: {{.AcceptLink}}", + Text: "You have been invited to join {{.OrganizationName}} on {{.AppName}} as {{.Role}}. Open this link to accept: {{.InviteLink}}", HTML: "

Invited to {{.OrganizationName}} as {{.Role}}

", }) svc := NewOrganizationInvitationService( @@ -623,8 +608,6 @@ func TestOrganizationInvitationService_CreateOrganizationInvitation(t *testing.T userSvc, mailer, accessControlService, - verifSvc, - tokenSvc, orgRepo, orgInvitationRepo, memberRepo, @@ -740,7 +723,7 @@ func TestOrganizationInvitationService_GetOrganizationInvitation(t *testing.T) { tt.setup(orgRepo, invRepo, memberRepo) } - svc := newTestOrganizationInvitationService(&orgtests.MockOrganizationInvitationTxRunner{}, pluginConfig, &internaltests.MockUserService{}, orgtests.NewAccessControlServiceStub(), orgRepo, invRepo, memberRepo, &internaltests.MockVerificationService{}, &internaltests.MockTokenService{}) + svc := newTestOrganizationInvitationService(&orgtests.MockOrganizationInvitationTxRunner{}, pluginConfig, &internaltests.MockUserService{}, orgtests.NewAccessControlServiceStub(), orgRepo, invRepo, memberRepo) invitation, err := svc.GetOrganizationInvitation(context.Background(), orgtests.Actor(tt.actorUserID), tt.organizationID, tt.invitationID) if tt.expectErr != nil { require.Error(t, err) @@ -859,7 +842,7 @@ func TestOrganizationInvitationService_GetAllOrganizationInvitations(t *testing. tt.setup(orgRepo, invRepo, memberRepo) } - svc := newTestOrganizationInvitationService(&orgtests.MockOrganizationInvitationTxRunner{}, pluginConfig, &internaltests.MockUserService{}, orgtests.NewAccessControlServiceStub(), orgRepo, invRepo, memberRepo, &internaltests.MockVerificationService{}, &internaltests.MockTokenService{}) + svc := newTestOrganizationInvitationService(&orgtests.MockOrganizationInvitationTxRunner{}, pluginConfig, &internaltests.MockUserService{}, orgtests.NewAccessControlServiceStub(), orgRepo, invRepo, memberRepo) invitations, err := svc.GetAllOrganizationInvitations(context.Background(), orgtests.Actor(tt.actorUserID), tt.organizationID) if tt.expectErr != nil { require.Error(t, err) @@ -967,7 +950,7 @@ func TestOrganizationInvitationService_RevokeOrganizationInvitation(t *testing.T tt.setup(orgRepo, invRepo, memberRepo, hooks) } - svc := newTestOrganizationInvitationService(&orgtests.MockOrganizationInvitationTxRunner{}, pluginConfig, &internaltests.MockUserService{}, orgtests.NewAccessControlServiceStub(), orgRepo, invRepo, memberRepo, &internaltests.MockVerificationService{}, &internaltests.MockTokenService{}) + svc := newTestOrganizationInvitationService(&orgtests.MockOrganizationInvitationTxRunner{}, pluginConfig, &internaltests.MockUserService{}, orgtests.NewAccessControlServiceStub(), orgRepo, invRepo, memberRepo) invitation, err := svc.RevokeOrganizationInvitation(context.Background(), orgtests.Actor(tt.actorUserID), tt.organizationID, tt.invitationID) if tt.expectErr != nil { require.Error(t, err) @@ -1053,7 +1036,7 @@ func TestOrganizationInvitationService_AcceptPendingOrganizationInvitationsForEm } txRunner := &orgtests.MockOrganizationInvitationTxRunner{} - svc := newTestOrganizationInvitationService(txRunner, pluginConfig, userSvc, orgtests.NewAccessControlServiceStub(), orgRepo, invRepo, memberRepo, &internaltests.MockVerificationService{}, &internaltests.MockTokenService{}) + svc := newTestOrganizationInvitationService(txRunner, pluginConfig, userSvc, orgtests.NewAccessControlServiceStub(), orgRepo, invRepo, memberRepo) accepted, err := svc.AcceptPendingOrganizationInvitationsForEmail(context.Background(), tt.userID, tt.email) if tt.expectErr != nil { require.Error(t, err) @@ -1302,7 +1285,7 @@ func TestOrganizationInvitationService_AcceptOrganizationInvitation(t *testing.T tt.setup(userSvc, invRepo, memberRepo, hooks, memberHooks) } - svc := newTestOrganizationInvitationService(&orgtests.MockOrganizationInvitationTxRunner{}, pluginConfig, userSvc, orgtests.NewAccessControlServiceStub(), orgRepo, invRepo, memberRepo, &internaltests.MockVerificationService{}, &internaltests.MockTokenService{}) + svc := newTestOrganizationInvitationService(&orgtests.MockOrganizationInvitationTxRunner{}, pluginConfig, userSvc, orgtests.NewAccessControlServiceStub(), orgRepo, invRepo, memberRepo) invitation, err := svc.AcceptOrganizationInvitation(context.Background(), orgtests.Actor(tt.actorUserID), tt.organization, tt.invitationID) if tt.expectErr != nil { require.ErrorIs(t, err, tt.expectErr) @@ -1468,7 +1451,7 @@ func TestOrganizationInvitationService_RejectOrganizationInvitation(t *testing.T tt.setup(userSvc, invRepo) } - svc := newTestOrganizationInvitationService(&orgtests.MockOrganizationInvitationTxRunner{}, pluginConfig, userSvc, orgtests.NewAccessControlServiceStub(), &orgtests.MockOrganizationRepository{}, invRepo, &orgtests.MockOrganizationMemberRepository{}, &internaltests.MockVerificationService{}, &internaltests.MockTokenService{}) + svc := newTestOrganizationInvitationService(&orgtests.MockOrganizationInvitationTxRunner{}, pluginConfig, userSvc, orgtests.NewAccessControlServiceStub(), &orgtests.MockOrganizationRepository{}, invRepo, &orgtests.MockOrganizationMemberRepository{}) invitation, err := svc.RejectOrganizationInvitation(context.Background(), orgtests.Actor(tt.actorUserID), tt.organization, tt.invitationID) if tt.expectErr != nil { require.ErrorIs(t, err, tt.expectErr) diff --git a/plugins/organizations/services/organization_service.go b/plugins/organizations/services/organization_service.go index d8b213aa..54ad6039 100644 --- a/plugins/organizations/services/organization_service.go +++ b/plugins/organizations/services/organization_service.go @@ -333,20 +333,6 @@ func (s *organizationService) ExistsByID(ctx context.Context, organizationID str return org != nil, nil } -func (s *organizationService) GetByIDNoAuth(ctx context.Context, organizationID string) (*types.Organization, error) { - if organizationID == "" { - return nil, coreerrors.ErrNotFound - } - org, err := s.orgRepo.GetByID(ctx, organizationID) - if err != nil { - return nil, err - } - if org == nil { - return nil, coreerrors.ErrNotFound - } - return org, nil -} - func (s *organizationService) DeleteOrganization(ctx context.Context, actor *models.Actor, organizationID string) error { organization, err := s.serviceUtils.authorizeOwner(ctx, actor, organizationID) if err != nil { diff --git a/plugins/organizations/templates/organization_invitation/html.tmpl b/plugins/organizations/templates/organization_invitation/html.tmpl index d5d9b23d..c713d288 100644 --- a/plugins/organizations/templates/organization_invitation/html.tmpl +++ b/plugins/organizations/templates/organization_invitation/html.tmpl @@ -1,7 +1,7 @@

Hi {{.InvitationEmail}},

You have been invited to join {{.OrganizationName}} on {{.AppName}} as {{.Role}}.

-

Accept invitation

+

View Invitation

If the button does not work, copy this link:

-

{{.AcceptLink}}

+

{{.InviteLink}}

diff --git a/plugins/organizations/templates/organization_invitation/text.tmpl b/plugins/organizations/templates/organization_invitation/text.tmpl index 9817fb0a..d86465b5 100644 --- a/plugins/organizations/templates/organization_invitation/text.tmpl +++ b/plugins/organizations/templates/organization_invitation/text.tmpl @@ -1,3 +1,3 @@ You have been invited to join {{.OrganizationName}} as {{.Role}}. -Open this link to accept the invitation: {{.AcceptLink}} +Open this link to view the invitation: {{.InviteLink}} diff --git a/plugins/organizations/tests/repositories.go b/plugins/organizations/tests/repositories.go index 49648c98..859158b3 100644 --- a/plugins/organizations/tests/repositories.go +++ b/plugins/organizations/tests/repositories.go @@ -211,6 +211,22 @@ func (m *MockOrganizationInvitationRepository) GetAllByOrganizationID(ctx contex return args.Get(0).([]types.OrganizationInvitation), args.Error(1) } +func (m *MockOrganizationInvitationRepository) GetByIDWithOrg(ctx context.Context, invitationID string) (*types.GetOrganizationInvitationResponse, error) { + args := m.Called(ctx, invitationID) + if args.Get(0) == nil { + return nil, args.Error(1) + } + return args.Get(0).(*types.GetOrganizationInvitationResponse), args.Error(1) +} + +func (m *MockOrganizationInvitationRepository) GetAllByOrganizationIDWithOrg(ctx context.Context, organizationID string) ([]types.GetOrganizationInvitationResponse, error) { + args := m.Called(ctx, organizationID) + if args.Get(0) == nil { + return nil, args.Error(1) + } + return args.Get(0).([]types.GetOrganizationInvitationResponse), args.Error(1) +} + func (m *MockOrganizationInvitationRepository) GetAllPendingByEmail(ctx context.Context, email string) ([]types.OrganizationInvitation, error) { args := m.Called(ctx, email) if args.Get(0) == nil { diff --git a/plugins/organizations/tests/services.go b/plugins/organizations/tests/services.go index 5e9d6a53..e42ec080 100644 --- a/plugins/organizations/tests/services.go +++ b/plugins/organizations/tests/services.go @@ -63,14 +63,6 @@ func (m *MockOrganizationService) ExistsByID(ctx context.Context, organizationID return args.Bool(0), args.Error(1) } -func (m *MockOrganizationService) GetByIDNoAuth(ctx context.Context, organizationID string) (*types.Organization, error) { - args := m.Called(ctx, organizationID) - if args.Get(0) == nil { - return nil, args.Error(1) - } - return args.Get(0).(*types.Organization), args.Error(1) -} - type MockOrganizationInvitationService struct { mock.Mock } @@ -99,6 +91,22 @@ func (m *MockOrganizationInvitationService) GetOrganizationInvitationByID(ctx co return args.Get(0).(*types.OrganizationInvitation), args.Error(1) } +func (m *MockOrganizationInvitationService) GetOrganizationInvitationByIDWithOrg(ctx context.Context, invitationID string) (*types.GetOrganizationInvitationResponse, error) { + args := m.Called(ctx, invitationID) + if args.Get(0) == nil { + return nil, args.Error(1) + } + return args.Get(0).(*types.GetOrganizationInvitationResponse), args.Error(1) +} + +func (m *MockOrganizationInvitationService) GetAllOrganizationInvitationsByOrgIDWithOrg(ctx context.Context, organizationID string) ([]types.GetOrganizationInvitationResponse, error) { + args := m.Called(ctx, organizationID) + if args.Get(0) == nil { + return nil, args.Error(1) + } + return args.Get(0).([]types.GetOrganizationInvitationResponse), args.Error(1) +} + func (m *MockOrganizationInvitationService) GetAllOrganizationInvitations(ctx context.Context, actor *models.Actor, organizationID string) ([]types.OrganizationInvitation, error) { args := m.Called(ctx, actorID(actor), organizationID) if args.Get(0) == nil { diff --git a/plugins/organizations/types/api.go b/plugins/organizations/types/api.go index f0d42ad7..b62baf8b 100644 --- a/plugins/organizations/types/api.go +++ b/plugins/organizations/types/api.go @@ -58,23 +58,17 @@ type AcceptOrganizationInvitationQuery struct { } type OrganizationSummary struct { - ID string `json:"id"` - OwnerID string `json:"owner_id"` - Name string `json:"name"` - Slug string `json:"slug"` - Logo *string `json:"logo,omitempty"` -} - -type VerifyOrganizationInvitationResponse struct { - Invitation *OrganizationInvitation `json:"invitation"` - Organization OrganizationSummary `json:"organization"` + ID string `json:"id" required:"true" nullable:"false"` + OwnerID string `json:"owner_id" required:"true" nullable:"false"` + Name string `json:"name" required:"true" nullable:"false"` + Slug string `json:"slug" required:"true" nullable:"false"` + Logo *string `json:"logo,omitempty" nullable:"true"` + Metadata map[string]any `json:"metadata,omitempty" nullable:"true"` } -type VerifyOrganizationInvitationQuery struct { - OrganizationID string `path:"organization_id"` - InvitationID string `path:"invitation_id"` - Token string `query:"token" json:"token" nullable:"false"` - RedirectURL string `query:"redirect_url" json:"redirect_url,omitempty" nullable:"true"` +type GetOrganizationInvitationResponse struct { + Invitation *OrganizationInvitation `json:"invitation" required:"true" nullable:"false"` + Organization OrganizationSummary `json:"organization" required:"true" nullable:"false"` } type CreateOrganizationRequest struct { diff --git a/plugins/organizations/types/config.go b/plugins/organizations/types/config.go index a371550e..7ee93cf6 100644 --- a/plugins/organizations/types/config.go +++ b/plugins/organizations/types/config.go @@ -84,5 +84,5 @@ type SendOrganizationInvitationEmailParams struct { Organization *Organization Invitation *OrganizationInvitation Inviter *models.User - AcceptURL string + InviteURL string } diff --git a/plugins/organizations/types/template_contexts.go b/plugins/organizations/types/template_contexts.go index eaa3e054..209b5856 100644 --- a/plugins/organizations/types/template_contexts.go +++ b/plugins/organizations/types/template_contexts.go @@ -11,6 +11,6 @@ type OrganizationInvitationContext struct { InvitationEmail string OrganizationName string Role string - AcceptLink string + InviteLink string Expiry time.Duration } diff --git a/plugins/organizations/usecases/usecases.go b/plugins/organizations/usecases/usecases.go index 28096919..75f2d750 100644 --- a/plugins/organizations/usecases/usecases.go +++ b/plugins/organizations/usecases/usecases.go @@ -3,7 +3,6 @@ package usecases import ( "context" "strings" - "time" coreerrors "github.com/Authula/authula/core/errors" "github.com/Authula/authula/models" @@ -14,16 +13,14 @@ import ( ) type UseCases struct { - orgService orgservices.OrganizationService - invitationService orgservices.OrganizationInvitationService - memberService orgservices.OrganizationMemberService - teamService orgservices.OrganizationTeamService - teamMemberService orgservices.OrganizationTeamMemberService - userService rootservices.UserService - verificationService rootservices.VerificationService - tokenService rootservices.TokenService - globalConfig *models.Config - authorizer rootservices.Authorizer + orgService orgservices.OrganizationService + invitationService orgservices.OrganizationInvitationService + memberService orgservices.OrganizationMemberService + teamService orgservices.OrganizationTeamService + teamMemberService orgservices.OrganizationTeamMemberService + userService rootservices.UserService + globalConfig *models.Config + authorizer rootservices.Authorizer } func NewUseCases( @@ -33,22 +30,18 @@ func NewUseCases( teamService orgservices.OrganizationTeamService, teamMemberService orgservices.OrganizationTeamMemberService, userService rootservices.UserService, - verificationService rootservices.VerificationService, - tokenService rootservices.TokenService, globalConfig *models.Config, authorizer rootservices.Authorizer, ) *UseCases { return &UseCases{ - orgService: orgService, - invitationService: invitationService, - memberService: memberService, - teamService: teamService, - teamMemberService: teamMemberService, - userService: userService, - verificationService: verificationService, - tokenService: tokenService, - globalConfig: globalConfig, - authorizer: authorizer, + orgService: orgService, + invitationService: invitationService, + memberService: memberService, + teamService: teamService, + teamMemberService: teamMemberService, + userService: userService, + globalConfig: globalConfig, + authorizer: authorizer, } } @@ -108,104 +101,45 @@ func (u *UseCases) CreateOrganizationInvitation(ctx context.Context, actor *mode return u.invitationService.CreateOrganizationInvitation(ctx, actor, organizationID, request, redirectURL) } -func (u *UseCases) GetAllOrganizationInvitations(ctx context.Context, actor *models.Actor, organizationID string) ([]types.OrganizationInvitation, error) { +func (u *UseCases) GetAllOrganizationInvitations(ctx context.Context, actor *models.Actor, organizationID string) ([]types.GetOrganizationInvitationResponse, error) { if err := u.authorizer.AuthorizeOrganizationAccess(ctx, actor, organizationID); err != nil { return nil, err } if err := u.authorizer.AuthorizeScope(ctx, actor, orgconstants.OrganizationsInvitationsListPermission); err != nil { return nil, err } - return u.invitationService.GetAllOrganizationInvitations(ctx, actor, organizationID) -} - -func (u *UseCases) GetOrganizationInvitation(ctx context.Context, actor *models.Actor, organizationID string, invitationID string) (*types.OrganizationInvitation, error) { - invitation, err := u.invitationService.GetOrganizationInvitationByID(ctx, invitationID) - if err != nil { - return nil, err - } - if invitation.OrganizationID != organizationID { - return nil, coreerrors.ErrNotFound - } - user, err := u.userService.GetByID(ctx, actor.ID) + resp, err := u.invitationService.GetAllOrganizationInvitationsByOrgIDWithOrg(ctx, organizationID) if err != nil { return nil, err } - if user != nil && strings.EqualFold(invitation.Email, user.Email) { - return invitation, nil - } - if err := u.authorizer.AuthorizeOrganizationAccess(ctx, actor, organizationID); err != nil { - return nil, err - } - if err := u.authorizer.AuthorizeScope(ctx, actor, orgconstants.OrganizationsInvitationsReadPermission); err != nil { - return nil, err - } - return u.invitationService.GetOrganizationInvitation(ctx, actor, organizationID, invitationID) + return resp, nil } -func (u *UseCases) VerifyOrganizationInvitation(ctx context.Context, actor *models.Actor, organizationID string, invitationID string, rawToken string) (*types.VerifyOrganizationInvitationResponse, error) { - hashedToken := u.tokenService.Hash(rawToken) - - verification, err := u.verificationService.GetByToken(ctx, hashedToken) - if err != nil { - return nil, err - } - if verification == nil { - return nil, orgconstants.ErrInvitationVerificationFailed - } - if verification.Type != models.TypeOrganizationInvitationVerify { - return nil, orgconstants.ErrInvitationVerificationFailed - } - if u.verificationService.IsExpired(verification) { - return nil, orgconstants.ErrInvitationVerificationExpired - } - if verification.Identifier != invitationID { - return nil, orgconstants.ErrInvitationVerificationFailed - } - - invitation, err := u.invitationService.GetOrganizationInvitationByID(ctx, invitationID) +func (u *UseCases) GetOrganizationInvitation(ctx context.Context, actor *models.Actor, organizationID string, invitationID string) (*types.GetOrganizationInvitationResponse, error) { + resp, err := u.invitationService.GetOrganizationInvitationByIDWithOrg(ctx, invitationID) if err != nil { return nil, err } - if invitation.OrganizationID != organizationID { + if resp.Invitation.OrganizationID != organizationID { return nil, coreerrors.ErrNotFound } - if invitation.ExpiresAt.Before(time.Now().UTC()) { - return nil, orgconstants.ErrInvitationVerificationExpired - } - if invitation.Status != types.OrganizationInvitationStatusPending { - return nil, orgconstants.ErrInvitationVerificationFailed - } user, err := u.userService.GetByID(ctx, actor.ID) if err != nil { return nil, err } - if user == nil || !strings.EqualFold(invitation.Email, user.Email) { - return nil, orgconstants.ErrInvitationEmailMismatch + if user != nil && strings.EqualFold(resp.Invitation.Email, user.Email) { + return resp, nil } - if err := u.verificationService.Delete(ctx, verification.ID); err != nil { + if err := u.authorizer.AuthorizeOrganizationAccess(ctx, actor, organizationID); err != nil { return nil, err } - - org, err := u.orgService.GetByIDNoAuth(ctx, organizationID) - if err != nil { + if err := u.authorizer.AuthorizeScope(ctx, actor, orgconstants.OrganizationsInvitationsReadPermission); err != nil { return nil, err } - - resp := &types.VerifyOrganizationInvitationResponse{ - Invitation: invitation, - Organization: types.OrganizationSummary{ - ID: org.ID, - OwnerID: org.OwnerID, - Name: org.Name, - Slug: org.Slug, - Logo: org.Logo, - }, - } - return resp, nil } @@ -220,29 +154,11 @@ func (u *UseCases) RevokeOrganizationInvitation(ctx context.Context, actor *mode } func (u *UseCases) AcceptOrganizationInvitation(ctx context.Context, actor *models.Actor, organizationID string, invitationID string) (*types.OrganizationInvitation, error) { - invitation, err := u.invitationService.AcceptOrganizationInvitation(ctx, actor, organizationID, invitationID) - if err != nil { - return nil, err - } - - if err := u.verificationService.DeleteByUserIDAndType(ctx, actor.ID, models.TypeOrganizationInvitationVerify); err != nil { - return nil, err - } - - return invitation, nil + return u.invitationService.AcceptOrganizationInvitation(ctx, actor, organizationID, invitationID) } func (u *UseCases) RejectOrganizationInvitation(ctx context.Context, actor *models.Actor, organizationID string, invitationID string) (*types.OrganizationInvitation, error) { - invitation, err := u.invitationService.RejectOrganizationInvitation(ctx, actor, organizationID, invitationID) - if err != nil { - return nil, err - } - - if err := u.verificationService.DeleteByUserIDAndType(ctx, actor.ID, models.TypeOrganizationInvitationVerify); err != nil { - return nil, err - } - - return invitation, nil + return u.invitationService.RejectOrganizationInvitation(ctx, actor, organizationID, invitationID) } // ------------- OrganizationMemberService -------------