diff --git a/models/verification.go b/models/verification.go index c836cd9..d329202 100644 --- a/models/verification.go +++ b/models/verification.go @@ -10,12 +10,13 @@ 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" + 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" ) func (vt VerificationType) String() string { @@ -31,6 +32,7 @@ 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/openapi.json b/openapi.json index a7d9de7..3cbb8bf 100644 --- a/openapi.json +++ b/openapi.json @@ -2608,13 +2608,20 @@ "schema": { "type": "string" } + }, + { + "name": "redirect_url", + "in": "query", + "schema": { + "type": "string" + } } ], "requestBody": { "content": { "application/json": { "schema": { - "$ref": "#/components/schemas/CreateOrganizationInvitationRequest" + "$ref": "#/components/schemas/CreateOrganizationInvitationQuery" } } } @@ -2717,7 +2724,7 @@ "Organization Invitations" ], "summary": "Accept invitation", - "description": "Accepts an invitation to join an organization. Optionally redirects if a redirect_url is provided.", + "description": "Accepts an invitation to join an organization.", "operationId": "acceptOrganizationInvitation", "parameters": [ { @@ -2807,6 +2814,60 @@ } } }, + "/organizations/{organization_id}/invitations/{invitation_id}/verify": { + "get": { + "tags": [ + "Organization Invitations" + ], + "summary": "Verify invitation", + "description": "Verifies an invitation token.", + "operationId": "verifyOrganizationInvitation", + "parameters": [ + { + "name": "token", + "in": "query", + "schema": { + "type": "string" + } + }, + { + "name": "redirect_url", + "in": "query", + "schema": { + "type": "string" + } + }, + { + "name": "organization_id", + "in": "path", + "required": true, + "schema": { + "type": "string" + } + }, + { + "name": "invitation_id", + "in": "path", + "required": true, + "schema": { + "type": "string" + } + } + ], + "responses": { + "200": { + "description": "OK", + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/VerifyOrganizationInvitationResponse" + } + } + } + } + } + } + }, "/organizations/{organization_id}/members": { "get": { "tags": [ @@ -4346,16 +4407,21 @@ ], "type": "object" }, - "CreateOrganizationInvitationRequest": { + "CreateOrganizationInvitationQuery": { "properties": { - "email": { - "type": "string" - }, "redirect_url": { "type": [ "string", "null" ] + } + }, + "type": "object" + }, + "CreateOrganizationInvitationRequest": { + "properties": { + "email": { + "type": "string" }, "role": { "type": "string" @@ -5239,6 +5305,29 @@ ], "type": "object" }, + "OrganizationSummary": { + "properties": { + "id": { + "type": "string" + }, + "logo": { + "type": [ + "null", + "string" + ] + }, + "name": { + "type": "string" + }, + "owner_id": { + "type": "string" + }, + "slug": { + "type": "string" + } + }, + "type": "object" + }, "OrganizationTeam": { "properties": { "created_at": { @@ -6451,6 +6540,17 @@ ], "type": "object" }, + "VerifyOrganizationInvitationResponse": { + "properties": { + "invitation": { + "$ref": "#/components/schemas/OrganizationInvitation" + }, + "organization": { + "$ref": "#/components/schemas/OrganizationSummary" + } + }, + "type": "object" + }, "VerifyTOTPRequest": { "properties": { "code": { diff --git a/plugins/organizations/api.go b/plugins/organizations/api.go index 007bbdd..c579076 100644 --- a/plugins/organizations/api.go +++ b/plugins/organizations/api.go @@ -46,8 +46,8 @@ func (a *API) DeleteOrganization(ctx context.Context, actor *models.Actor, organ // Invitations -func (a *API) CreateInvitation(ctx context.Context, actor *models.Actor, organizationID string, request types.CreateOrganizationInvitationRequest) (*types.OrganizationInvitation, error) { - return a.useCases.CreateOrganizationInvitation(ctx, actor, organizationID, request) +func (a *API) CreateInvitation(ctx context.Context, actor *models.Actor, organizationID string, request types.CreateOrganizationInvitationRequest, redirectURL string) (*types.OrganizationInvitation, error) { + 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) { @@ -70,6 +70,10 @@ 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 ab4ca9f..e5be8d6 100644 --- a/plugins/organizations/constants/errors.go +++ b/plugins/organizations/constants/errors.go @@ -9,9 +9,12 @@ import ( ) var ( - ErrOrganizationsQuotaExceeded = errors.New("organizations quota exceeded") - ErrMembersQuotaExceeded = errors.New("members quota exceeded") - ErrInvitationsQuotaExceeded = errors.New("invitations quota exceeded") + 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") ) func HandleError(err error, reqCtx *models.RequestContext) { @@ -20,6 +23,12 @@ 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 } if status != 0 { diff --git a/plugins/organizations/handlers/handler_test_helpers.go b/plugins/organizations/handlers/handler_test_helpers.go index fa1fd4c..fec367f 100644 --- a/plugins/organizations/handlers/handler_test_helpers.go +++ b/plugins/organizations/handlers/handler_test_helpers.go @@ -3,6 +3,9 @@ package handlers import ( "context" + "github.com/stretchr/testify/mock" + + internaltests "github.com/Authula/authula/internal/tests" "github.com/Authula/authula/models" orgservices "github.com/Authula/authula/plugins/organizations/services" orgtests "github.com/Authula/authula/plugins/organizations/tests" @@ -19,6 +22,19 @@ func (a *noopAuthorizer) AuthorizeOrganizationAccess(_ context.Context, _ *model return nil } +func defaultMockUserService() *internaltests.MockUserService { + svc := &internaltests.MockUserService{} + svc.On("GetByID", mock.Anything, mock.Anything).Return((*models.User)(nil), nil).Maybe() + svc.On("GetByEmail", mock.Anything, mock.Anything).Return((*models.User)(nil), nil).Maybe() + 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, @@ -26,6 +42,10 @@ func newOrgUseCases(svc orgservices.OrganizationService) *orgusecases.UseCases { &orgtests.MockOrganizationMemberService{}, &orgtests.MockOrganizationTeamService{}, &orgtests.MockOrganizationTeamMemberService{}, + defaultMockUserService(), + defaultMockVerificationService(), + &internaltests.MockTokenService{}, + &models.Config{}, &noopAuthorizer{}, ) } @@ -37,6 +57,10 @@ func newInvitationUseCases(svc orgservices.OrganizationInvitationService) *orgus &orgtests.MockOrganizationMemberService{}, &orgtests.MockOrganizationTeamService{}, &orgtests.MockOrganizationTeamMemberService{}, + defaultMockUserService(), + defaultMockVerificationService(), + &internaltests.MockTokenService{}, + &models.Config{}, &noopAuthorizer{}, ) } @@ -48,6 +72,10 @@ func newMemberUseCases(svc orgservices.OrganizationMemberService) *orgusecases.U svc, &orgtests.MockOrganizationTeamService{}, &orgtests.MockOrganizationTeamMemberService{}, + defaultMockUserService(), + defaultMockVerificationService(), + &internaltests.MockTokenService{}, + &models.Config{}, &noopAuthorizer{}, ) } @@ -59,6 +87,10 @@ func newTeamUseCases(svc orgservices.OrganizationTeamService) *orgusecases.UseCa &orgtests.MockOrganizationMemberService{}, svc, &orgtests.MockOrganizationTeamMemberService{}, + defaultMockUserService(), + defaultMockVerificationService(), + &internaltests.MockTokenService{}, + &models.Config{}, &noopAuthorizer{}, ) } @@ -70,6 +102,10 @@ func newTeamMemberUseCases(svc orgservices.OrganizationTeamMemberService) *orgus &orgtests.MockOrganizationMemberService{}, &orgtests.MockOrganizationTeamService{}, svc, + defaultMockUserService(), + defaultMockVerificationService(), + &internaltests.MockTokenService{}, + &models.Config{}, &noopAuthorizer{}, ) } diff --git a/plugins/organizations/handlers/organization_handlers.go b/plugins/organizations/handlers/organization_handlers.go index baeff93..2bdea9a 100644 --- a/plugins/organizations/handlers/organization_handlers.go +++ b/plugins/organizations/handlers/organization_handlers.go @@ -22,7 +22,7 @@ func (h *CreateOrganizationHandler) Handle() http.HandlerFunc { var request types.CreateOrganizationRequest if err := util.ParseJSON(r, &request); err != nil { - reqCtx.SetJSONResponse(http.StatusUnprocessableEntity, map[string]any{"message": "invalid request body"}) + reqCtx.SetJSONResponse(http.StatusUnprocessableEntity, map[string]any{"message": err.Error()}) reqCtx.Handled = true return } @@ -102,7 +102,7 @@ func (h *UpdateOrganizationHandler) Handle() http.HandlerFunc { var request types.UpdateOrganizationRequest if err := util.ParseJSON(r, &request); err != nil { - reqCtx.SetJSONResponse(http.StatusUnprocessableEntity, map[string]any{"message": "invalid request body"}) + reqCtx.SetJSONResponse(http.StatusUnprocessableEntity, map[string]any{"message": err.Error()}) reqCtx.Handled = true return } diff --git a/plugins/organizations/handlers/organization_handlers_test.go b/plugins/organizations/handlers/organization_handlers_test.go index 68d8ef4..6df18e2 100644 --- a/plugins/organizations/handlers/organization_handlers_test.go +++ b/plugins/organizations/handlers/organization_handlers_test.go @@ -74,7 +74,7 @@ func TestCreateOrganizationHandler(t *testing.T) { userID: new("user-1"), body: []byte("{invalid"), expectedStatus: http.StatusUnprocessableEntity, - expectedMessage: "invalid request body", + expectedMessage: "invalid character 'i' looking for beginning of object key string", }, { name: "unprocessable_entity", @@ -356,7 +356,7 @@ func TestUpdateOrganizationHandler(t *testing.T) { organizationID: "org-1", body: []byte("{invalid"), expectedStatus: http.StatusUnprocessableEntity, - expectedMessage: "invalid request body", + expectedMessage: "invalid character 'i' looking for beginning of object key string", }, { name: "unprocessable_entity", diff --git a/plugins/organizations/handlers/organization_invitation_handlers.go b/plugins/organizations/handlers/organization_invitation_handlers.go index ade412a..6c1bed0 100644 --- a/plugins/organizations/handlers/organization_invitation_handlers.go +++ b/plugins/organizations/handlers/organization_invitation_handlers.go @@ -24,7 +24,7 @@ func (h *CreateOrganizationInvitationHandler) Handle() http.HandlerFunc { var request types.CreateOrganizationInvitationRequest if err := util.ParseJSON(r, &request); err != nil { - reqCtx.SetJSONResponse(http.StatusUnprocessableEntity, map[string]any{"message": "invalid request body"}) + reqCtx.SetJSONResponse(http.StatusUnprocessableEntity, map[string]any{"message": err.Error()}) reqCtx.Handled = true return } @@ -34,7 +34,8 @@ func (h *CreateOrganizationInvitationHandler) Handle() http.HandlerFunc { return } - invitation, err := h.UseCases.CreateOrganizationInvitation(ctx, actor, organizationID, request) + redirectURL := r.URL.Query().Get("redirect_url") + invitation, err := h.UseCases.CreateOrganizationInvitation(ctx, actor, organizationID, request, redirectURL) if err != nil { orgconstants.HandleError(err, reqCtx) return @@ -110,7 +111,8 @@ func (h *RevokeOrganizationInvitationHandler) Handle() http.HandlerFunc { } type AcceptOrganizationInvitationHandler struct { - UseCases *orgusecases.UseCases + UseCases *orgusecases.UseCases + TrustedOrigins []string } func (h *AcceptOrganizationInvitationHandler) Handle() http.HandlerFunc { @@ -133,20 +135,66 @@ func (h *AcceptOrganizationInvitationHandler) Handle() http.HandlerFunc { AssignerUserID: &invitation.InviterID, } - var request types.AcceptOrganizationInvitationRequest redirectURL := r.URL.Query().Get("redirect_url") - request.RedirectURL = &redirectURL - if err := request.Validate(); err != nil { - reqCtx.SetJSONResponse(http.StatusUnprocessableEntity, map[string]any{"message": err.Error()}) + 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 + } + reqCtx.RedirectURL = validatedURL.String() + return + } + + reqCtx.SetJSONResponse(http.StatusOK, invitation) + } +} + +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 } - if request.RedirectURL != nil && *request.RedirectURL != "" { - reqCtx.RedirectURL = *request.RedirectURL + + 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, invitation) + reqCtx.SetJSONResponse(http.StatusOK, resp) } } @@ -159,11 +207,6 @@ func (h *RejectOrganizationInvitationHandler) Handle() http.HandlerFunc { ctx := r.Context() reqCtx, _ := models.GetRequestContext(ctx) actor := reqCtx.Actor - if actor == nil || actor.ID == "" { - reqCtx.SetJSONResponse(http.StatusUnauthorized, map[string]any{"message": "Unauthorized"}) - reqCtx.Handled = true - return - } organizationID := r.PathValue("organization_id") invitationID := r.PathValue("invitation_id") diff --git a/plugins/organizations/handlers/organization_invitation_handlers_test.go b/plugins/organizations/handlers/organization_invitation_handlers_test.go index 6e7f0a6..11dba41 100644 --- a/plugins/organizations/handlers/organization_invitation_handlers_test.go +++ b/plugins/organizations/handlers/organization_invitation_handlers_test.go @@ -107,7 +107,7 @@ func TestCreateOrganizationInvitationHandler(t *testing.T) { organizationID: "org-1", body: []byte("{"), expectedStatus: http.StatusUnprocessableEntity, - expectedMessage: "invalid request body", + expectedMessage: "unexpected EOF", }, { name: "service_error", @@ -115,7 +115,7 @@ func TestCreateOrganizationInvitationHandler(t *testing.T) { organizationID: "org-1", body: internaltests.MarshalToJSON(t, orgtypes.CreateOrganizationInvitationRequest{Email: "user@example.com", Role: "member"}), prepare: func(fixture *organizationInvitationHandlerFixture) { - fixture.service.On("CreateOrganizationInvitation", mock.Anything, "user-1", "org-1", mock.Anything).Return((*orgtypes.OrganizationInvitation)(nil), errors.New("create failed")).Once() + fixture.service.On("CreateOrganizationInvitation", mock.Anything, "user-1", "org-1", mock.Anything, mock.Anything).Return((*orgtypes.OrganizationInvitation)(nil), errors.New("create failed")).Once() }, expectedStatus: http.StatusBadRequest, expectedMessage: "create failed", @@ -126,7 +126,7 @@ func TestCreateOrganizationInvitationHandler(t *testing.T) { organizationID: "org-1", body: internaltests.MarshalToJSON(t, orgtypes.CreateOrganizationInvitationRequest{Email: "user@example.com", Role: "member"}), prepare: func(fixture *organizationInvitationHandlerFixture) { - fixture.service.On("CreateOrganizationInvitation", mock.Anything, "user-1", "org-1", mock.Anything).Return((*orgtypes.OrganizationInvitation)(nil), orgconstants.ErrInvitationsQuotaExceeded).Once() + fixture.service.On("CreateOrganizationInvitation", mock.Anything, "user-1", "org-1", mock.Anything, mock.Anything).Return((*orgtypes.OrganizationInvitation)(nil), orgconstants.ErrInvitationsQuotaExceeded).Once() }, expectedStatus: http.StatusTooManyRequests, expectedMessage: orgconstants.ErrInvitationsQuotaExceeded.Error(), @@ -137,7 +137,7 @@ func TestCreateOrganizationInvitationHandler(t *testing.T) { organizationID: "org-1", body: internaltests.MarshalToJSON(t, orgtypes.CreateOrganizationInvitationRequest{Email: "user@example.com", Role: "member"}), prepare: func(fixture *organizationInvitationHandlerFixture) { - fixture.service.On("CreateOrganizationInvitation", mock.Anything, "user-1", "org-1", mock.Anything).Return(&orgtypes.OrganizationInvitation{ID: "inv-1", OrganizationID: "org-1", Email: "user@example.com", Role: "member", Status: orgtypes.OrganizationInvitationStatusPending}, nil).Once() + fixture.service.On("CreateOrganizationInvitation", mock.Anything, "user-1", "org-1", mock.Anything, mock.Anything).Return(&orgtypes.OrganizationInvitation{ID: "inv-1", OrganizationID: "org-1", Email: "user@example.com", Role: "member", Status: orgtypes.OrganizationInvitationStatusPending}, nil).Once() }, expectedStatus: http.StatusCreated, checkResponse: func(t *testing.T, reqCtx *models.RequestContext) { @@ -210,7 +210,7 @@ func TestGetOrganizationInvitationHandler(t *testing.T) { organizationID: "org-1", invitationID: "inv-1", prepare: func(fixture *organizationInvitationHandlerFixture) { - fixture.service.On("GetOrganizationInvitation", mock.Anything, "user-1", "org-1", "inv-1").Return((*orgtypes.OrganizationInvitation)(nil), coreerrors.ErrNotFound).Once() + fixture.service.On("GetOrganizationInvitationByID", mock.Anything, "inv-1").Return((*orgtypes.OrganizationInvitation)(nil), coreerrors.ErrNotFound).Once() }, expectedStatus: http.StatusNotFound, expectedMessage: "not found", @@ -221,6 +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{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() }, expectedStatus: http.StatusOK, diff --git a/plugins/organizations/handlers/organization_member_handlers.go b/plugins/organizations/handlers/organization_member_handlers.go index 3dfa1c5..8ba7c4c 100644 --- a/plugins/organizations/handlers/organization_member_handlers.go +++ b/plugins/organizations/handlers/organization_member_handlers.go @@ -24,7 +24,7 @@ func (h *AddOrganizationMemberHandler) Handle() http.HandlerFunc { var request types.AddOrganizationMemberRequest if err := util.ParseJSON(r, &request); err != nil { - reqCtx.SetJSONResponse(http.StatusUnprocessableEntity, map[string]any{"message": "invalid request body"}) + reqCtx.SetJSONResponse(http.StatusUnprocessableEntity, map[string]any{"message": err.Error()}) reqCtx.Handled = true return } @@ -132,7 +132,7 @@ func (h *UpdateOrganizationMemberHandler) Handle() http.HandlerFunc { var request types.UpdateOrganizationMemberRequest if err := util.ParseJSON(r, &request); err != nil { - reqCtx.SetJSONResponse(http.StatusUnprocessableEntity, map[string]any{"message": "invalid request body"}) + reqCtx.SetJSONResponse(http.StatusUnprocessableEntity, map[string]any{"message": err.Error()}) reqCtx.Handled = true return } diff --git a/plugins/organizations/handlers/organization_member_handlers_test.go b/plugins/organizations/handlers/organization_member_handlers_test.go index 64fa08c..22e6fa9 100644 --- a/plugins/organizations/handlers/organization_member_handlers_test.go +++ b/plugins/organizations/handlers/organization_member_handlers_test.go @@ -110,7 +110,7 @@ func TestAddOrganizationMemberHandler(t *testing.T) { organizationID: "org-1", body: []byte("{"), expectedStatus: http.StatusUnprocessableEntity, - expectedMessage: "invalid request body", + expectedMessage: "unexpected EOF", }, { name: "service_error", @@ -302,7 +302,7 @@ func TestUpdateOrganizationMemberHandler(t *testing.T) { memberID: "mem-1", body: []byte("{"), expectedStatus: http.StatusUnprocessableEntity, - expectedMessage: "invalid request body", + expectedMessage: "unexpected EOF", }, { name: "forbidden", diff --git a/plugins/organizations/handlers/organization_team_handlers.go b/plugins/organizations/handlers/organization_team_handlers.go index 5cd16d7..8ca2d5d 100644 --- a/plugins/organizations/handlers/organization_team_handlers.go +++ b/plugins/organizations/handlers/organization_team_handlers.go @@ -24,7 +24,7 @@ func (h *CreateOrganizationTeamHandler) Handle() http.HandlerFunc { var request types.CreateOrganizationTeamRequest if err := util.ParseJSON(r, &request); err != nil { - reqCtx.SetJSONResponse(http.StatusUnprocessableEntity, map[string]any{"message": "invalid request body"}) + reqCtx.SetJSONResponse(http.StatusUnprocessableEntity, map[string]any{"message": err.Error()}) reqCtx.Handled = true return } @@ -102,7 +102,7 @@ func (h *UpdateOrganizationTeamHandler) Handle() http.HandlerFunc { var request types.UpdateOrganizationTeamRequest if err := util.ParseJSON(r, &request); err != nil { - reqCtx.SetJSONResponse(http.StatusUnprocessableEntity, map[string]any{"message": "invalid request body"}) + reqCtx.SetJSONResponse(http.StatusUnprocessableEntity, map[string]any{"message": err.Error()}) reqCtx.Handled = true return } diff --git a/plugins/organizations/handlers/organization_team_handlers_test.go b/plugins/organizations/handlers/organization_team_handlers_test.go index 3b3c424..91f09db 100644 --- a/plugins/organizations/handlers/organization_team_handlers_test.go +++ b/plugins/organizations/handlers/organization_team_handlers_test.go @@ -105,7 +105,7 @@ func TestCreateOrganizationTeamHandler(t *testing.T) { organizationID: "org-1", body: []byte("{"), expectedStatus: http.StatusUnprocessableEntity, - expectedMessage: "invalid request body", + expectedMessage: "unexpected EOF", }, { name: "service_error", @@ -238,7 +238,7 @@ func TestUpdateOrganizationTeamHandler(t *testing.T) { teamID: "team-1", body: []byte("{"), expectedStatus: http.StatusUnprocessableEntity, - expectedMessage: "invalid request body", + expectedMessage: "unexpected EOF", }, { name: "service_error", diff --git a/plugins/organizations/handlers/organization_team_member_handlers.go b/plugins/organizations/handlers/organization_team_member_handlers.go index a183f84..de5e493 100644 --- a/plugins/organizations/handlers/organization_team_member_handlers.go +++ b/plugins/organizations/handlers/organization_team_member_handlers.go @@ -25,7 +25,7 @@ func (h *AddOrganizationTeamMemberHandler) Handle() http.HandlerFunc { var request types.AddOrganizationTeamMemberRequest if err := util.ParseJSON(r, &request); err != nil { - reqCtx.SetJSONResponse(http.StatusUnprocessableEntity, map[string]any{"message": "invalid request body"}) + reqCtx.SetJSONResponse(http.StatusUnprocessableEntity, map[string]any{"message": err.Error()}) reqCtx.Handled = true return } diff --git a/plugins/organizations/handlers/organization_team_member_handlers_test.go b/plugins/organizations/handlers/organization_team_member_handlers_test.go index a653a44..69ed825 100644 --- a/plugins/organizations/handlers/organization_team_member_handlers_test.go +++ b/plugins/organizations/handlers/organization_team_member_handlers_test.go @@ -111,7 +111,7 @@ func TestAddOrganizationTeamMemberHandler(t *testing.T) { teamID: "team-1", body: []byte("{"), expectedStatus: http.StatusUnprocessableEntity, - expectedMessage: "invalid request body", + expectedMessage: "unexpected EOF", }, { name: "service_error", diff --git a/plugins/organizations/openapi/openapi_docs.go b/plugins/organizations/openapi/openapi_docs.go index c3b3907..3647260 100644 --- a/plugins/organizations/openapi/openapi_docs.go +++ b/plugins/organizations/openapi/openapi_docs.go @@ -72,6 +72,7 @@ func RegisterOpenAPIDocs(svc openapi.OpenAPIService) error { openapi.WithTags("Organization Invitations"), openapi.WithRequest(&types.OrganizationID{}), openapi.WithRequest(&types.CreateOrganizationInvitationRequest{}), + openapi.WithRequest(&types.CreateOrganizationInvitationQuery{}), openapi.WithResponseStatus(http.StatusCreated, &types.OrganizationInvitation{}), ), svc.AddOperation( @@ -109,7 +110,7 @@ func RegisterOpenAPIDocs(svc openapi.OpenAPIService) error { "/organizations/{organization_id}/invitations/{invitation_id}/accept", openapi.WithOperationID("acceptOrganizationInvitation"), openapi.WithSummary("Accept invitation"), - openapi.WithDescription("Accepts an invitation to join an organization. Optionally redirects if a redirect_url is provided."), + openapi.WithDescription("Accepts an invitation to join an organization."), openapi.WithTags("Organization Invitations"), openapi.WithRequest(&types.AcceptOrganizationInvitationQuery{}), openapi.WithResponseStatus(http.StatusOK, &types.OrganizationInvitation{}), @@ -124,6 +125,16 @@ func RegisterOpenAPIDocs(svc openapi.OpenAPIService) error { 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{}), + ), // Members svc.AddOperation( diff --git a/plugins/organizations/plugin.go b/plugins/organizations/plugin.go index 93f98b3..1522d79 100644 --- a/plugins/organizations/plugin.go +++ b/plugins/organizations/plugin.go @@ -70,6 +70,16 @@ 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") @@ -99,13 +109,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, 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, verificationService, tokenService, 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, authorizer) + p.useCases = usecases.NewUseCases(p.organizationService, p.invitationService, p.memberService, p.teamService, p.teamMemberService, userService, verificationService, tokenService, p.globalConfig, authorizer) p.Api = BuildAPI(p) diff --git a/plugins/organizations/routes.go b/plugins/organizations/routes.go index 34a4f7b..a3faa27 100644 --- a/plugins/organizations/routes.go +++ b/plugins/organizations/routes.go @@ -18,8 +18,9 @@ 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} + acceptInvitationHandler := &handlers.AcceptOrganizationInvitationHandler{UseCases: plugin.useCases, TrustedOrigins: plugin.globalConfig.Security.TrustedOrigins} rejectInvitationHandler := &handlers.RejectOrganizationInvitationHandler{UseCases: plugin.useCases} addMemberHandler := &handlers.AddOrganizationMemberHandler{UseCases: plugin.useCases} @@ -107,6 +108,14 @@ 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}", @@ -119,7 +128,7 @@ func Routes(plugin *OrganizationsPlugin) []models.Route { Method: http.MethodPost, Path: "/organizations/{organization_id}/invitations/{invitation_id}/accept", Middleware: []func(http.Handler) http.Handler{ - middleware.RequireAuthenticated(), + middleware.RequireActor(models.ActorUser), }, Handler: acceptInvitationHandler.Handle(), }, @@ -127,7 +136,7 @@ func Routes(plugin *OrganizationsPlugin) []models.Route { Method: http.MethodPost, Path: "/organizations/{organization_id}/invitations/{invitation_id}/reject", Middleware: []func(http.Handler) http.Handler{ - middleware.RequireAuthenticated(), + middleware.RequireActor(models.ActorUser), }, Handler: rejectInvitationHandler.Handle(), }, diff --git a/plugins/organizations/services/interfaces.go b/plugins/organizations/services/interfaces.go index 36b30c2..16a9ac4 100644 --- a/plugins/organizations/services/interfaces.go +++ b/plugins/organizations/services/interfaces.go @@ -14,11 +14,13 @@ 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) (*types.OrganizationInvitation, error) + 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) GetAllOrganizationInvitations(ctx context.Context, actor *models.Actor, organizationID string) ([]types.OrganizationInvitation, 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) diff --git a/plugins/organizations/services/organization_invitation_service.go b/plugins/organizations/services/organization_invitation_service.go index bd1d20c..8ab6ee3 100644 --- a/plugins/organizations/services/organization_invitation_service.go +++ b/plugins/organizations/services/organization_invitation_service.go @@ -35,6 +35,8 @@ 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 @@ -52,6 +54,8 @@ 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, @@ -72,6 +76,8 @@ func NewOrganizationInvitationService( userService: userService, mailerService: mailerService, accessControlService: accessControlService, + verificationService: verificationService, + tokenService: tokenService, organizationRepo: organizationRepo, orgInvitationRepo: orgInvitationRepo, orgMemberRepo: orgMemberRepo, @@ -81,7 +87,7 @@ func NewOrganizationInvitationService( } } -func (s *organizationInvitationService) CreateOrganizationInvitation(ctx context.Context, actor *models.Actor, organizationID string, request types.CreateOrganizationInvitationRequest) (*types.OrganizationInvitation, error) { +func (s *organizationInvitationService) CreateOrganizationInvitation(ctx context.Context, actor *models.Actor, organizationID string, request types.CreateOrganizationInvitationRequest, redirectURL string) (*types.OrganizationInvitation, error) { reqCtx, _ := models.GetRequestContext(ctx) if actor == nil || actor.ID == "" || organizationID == "" { @@ -177,7 +183,22 @@ func (s *organizationInvitationService) CreateOrganizationInvitation(ctx context s.publishOrganizationInvitationCreatedEvent(created, organization) - acceptURL := s.buildOrganizationInvitationAcceptURL(created, request.RedirectURL) + 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) callbackHandled := false if s.pluginConfig.SendOrganizationInvitationEmail != nil { @@ -190,7 +211,7 @@ func (s *organizationInvitationService) CreateOrganizationInvitation(ctx context Organization: organization, Invitation: created, Inviter: inviter, - AcceptURL: acceptURL, + AcceptURL: verifyURL, }, reqCtx) if err != nil { @@ -206,7 +227,7 @@ func (s *organizationInvitationService) CreateOrganizationInvitation(ctx context taskCtx, cancel := context.WithTimeout(detachedCtx, 15*time.Second) defer cancel() - if err := s.sendOrganizationInvitationEmail(taskCtx, created, organization, acceptURL); err != nil { + if err := s.sendOrganizationInvitationEmail(taskCtx, created, organization, verifyURL); err != nil { s.logger.Error("failed to send organization invitation email via built-in email service", "invitation_id", created.ID, "error", err) } }() @@ -231,22 +252,23 @@ func (s *organizationInvitationService) sendOrganizationInvitationEmail(ctx cont return s.mailerService.SendEmail(ctx, invitation.Email, subject, textBody, htmlBody) } -func (s *organizationInvitationService) buildOrganizationInvitationAcceptURL(invitation *types.OrganizationInvitation, redirectURL string) string { +func (s *organizationInvitationService) buildOrganizationInvitationVerifyURL(invitation *types.OrganizationInvitation, rawToken string, redirectURL string) string { baseURL := s.globalConfig.BaseURL basePath := s.globalConfig.BasePath - acceptPath := fmt.Sprintf("/organizations/%s/invitations/%s/accept", url.PathEscape(invitation.OrganizationID), url.PathEscape(invitation.ID)) + verifyPath := fmt.Sprintf("/organizations/%s/invitations/%s/verify", url.PathEscape(invitation.OrganizationID), url.PathEscape(invitation.ID)) - fullURL := baseURL + basePath + acceptPath + fullURL := baseURL + basePath + verifyPath parsedURL, err := url.Parse(fullURL) if err != nil { return fullURL } + query := parsedURL.Query() + query.Set("token", rawToken) if redirectURL != "" { - query := parsedURL.Query() query.Set("redirect_url", redirectURL) - parsedURL.RawQuery = query.Encode() } + parsedURL.RawQuery = query.Encode() return parsedURL.String() } @@ -305,6 +327,22 @@ func (s *organizationInvitationService) GetOrganizationInvitation(ctx context.Co return invitation, nil } +func (s *organizationInvitationService) GetOrganizationInvitationByID(ctx context.Context, invitationID string) (*types.OrganizationInvitation, error) { + if invitationID == "" { + return nil, coreerrors.ErrNotFound + } + + invitation, err := s.orgInvitationRepo.GetByID(ctx, invitationID) + if err != nil { + return nil, err + } + if invitation == nil { + return nil, coreerrors.ErrNotFound + } + + return invitation, 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 2d556db..ecf44f7 100644 --- a/plugins/organizations/services/organization_invitation_service_test.go +++ b/plugins/organizations/services/organization_invitation_service_test.go @@ -49,6 +49,8 @@ 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() @@ -67,6 +69,8 @@ func newTestOrganizationInvitationService( userService, nil, accessControlService, + verificationService, + tokenService, orgRepo, invRepo, memberRepo, @@ -125,6 +129,7 @@ func TestOrganizationInvitationService_CreateOrganizationInvitation(t *testing.T actorUserID string organizationID string request types.CreateOrganizationInvitationRequest + redirectURL string invitationExpiresIn time.Duration accessControlService rootservices.AccessControlService setup func(*orgtests.MockOrganizationRepository, *orgtests.MockOrganizationInvitationRepository, *orgtests.MockOrganizationMemberRepository, *orgtests.MockOrganizationInvitationHooks) @@ -377,7 +382,8 @@ func TestOrganizationInvitationService_CreateOrganizationInvitation(t *testing.T actorUserID: "user-1", organizationID: "org-1", invitationExpiresIn: 36 * time.Hour, - request: types.CreateOrganizationInvitationRequest{Email: "user@example.com", Role: "member", RedirectURL: "https://app.example.com/welcome"}, + request: types.CreateOrganizationInvitationRequest{Email: "user@example.com", Role: "member"}, + redirectURL: "https://app.example.com/welcome", setup: func(orgRepo *orgtests.MockOrganizationRepository, invRepo *orgtests.MockOrganizationInvitationRepository, memberRepo *orgtests.MockOrganizationMemberRepository, hooks *orgtests.MockOrganizationInvitationHooks) { orgRepo.On("GetByID", mock.Anything, "org-1").Return(&types.Organization{ID: "org-1", OwnerID: "user-1", Name: "Acme"}, nil).Once() invRepo.On("GetByOrganizationIDAndEmail", mock.Anything, "org-1", "user@example.com", types.OrganizationInvitationStatusPending).Return(nil, nil).Once() @@ -435,7 +441,8 @@ func TestOrganizationInvitationService_CreateOrganizationInvitation(t *testing.T actorUserID: "user-1", organizationID: "org-1", invitationExpiresIn: 36 * time.Hour, - request: types.CreateOrganizationInvitationRequest{Email: "user@example.com", Role: "member", RedirectURL: "https://app.example.com/welcome"}, + request: types.CreateOrganizationInvitationRequest{Email: "user@example.com", Role: "member"}, + redirectURL: "https://app.example.com/welcome", setup: func(orgRepo *orgtests.MockOrganizationRepository, invRepo *orgtests.MockOrganizationInvitationRepository, memberRepo *orgtests.MockOrganizationMemberRepository, hooks *orgtests.MockOrganizationInvitationHooks) { orgRepo.On("GetByID", mock.Anything, "org-1").Return(&types.Organization{ID: "org-1", OwnerID: "user-1", Name: "Acme"}, nil).Once() invRepo.On("GetByOrganizationIDAndEmail", mock.Anything, "org-1", "user@example.com", types.OrganizationInvitationStatusPending).Return(nil, nil).Once() @@ -452,7 +459,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/accept") + require.Contains(t, params.AcceptURL, "/organizations/org-1/invitations/inv-1/verify?") + require.Contains(t, params.AcceptURL, "token=test-raw-token") require.Nil(t, reqCtx) return nil } @@ -509,7 +517,8 @@ func TestOrganizationInvitationService_CreateOrganizationInvitation(t *testing.T actorUserID: "user-1", organizationID: "org-1", invitationExpiresIn: 36 * time.Hour, - request: types.CreateOrganizationInvitationRequest{Email: "user@example.com", Role: "member", RedirectURL: "https://app.example.com/welcome"}, + request: types.CreateOrganizationInvitationRequest{Email: "user@example.com", Role: "member"}, + redirectURL: "https://app.example.com/welcome", setup: func(orgRepo *orgtests.MockOrganizationRepository, invRepo *orgtests.MockOrganizationInvitationRepository, memberRepo *orgtests.MockOrganizationMemberRepository, hooks *orgtests.MockOrganizationInvitationHooks) { orgRepo.On("GetByID", mock.Anything, "org-1").Return(&types.Organization{ID: "org-1", OwnerID: "user-1", Name: "Acme"}, nil).Once() invRepo.On("GetByOrganizationIDAndEmail", mock.Anything, "org-1", "user@example.com", types.OrganizationInvitationStatusPending).Return(nil, nil).Once() @@ -523,7 +532,9 @@ 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.Equal(t, "https://example.com/auth/organizations/org-1/invitations/inv-1/accept?redirect_url=https%3A%2F%2Fapp.example.com%2Fwelcome", params.AcceptURL) + 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") return nil } }, @@ -567,6 +578,7 @@ func TestOrganizationInvitationService_CreateOrganizationInvitation(t *testing.T if tt.userSetup != nil { tt.userSetup(userSvc) } + userSvc.On("GetByEmail", mock.Anything, mock.Anything).Return(&models.User{}, nil).Maybe() var mailer rootservices.MailerService var mailerCalls chan invitationEmailCall if tt.mailerFactory != nil { @@ -587,6 +599,13 @@ 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{ @@ -604,13 +623,15 @@ func TestOrganizationInvitationService_CreateOrganizationInvitation(t *testing.T userSvc, mailer, accessControlService, + verifSvc, + tokenSvc, orgRepo, orgInvitationRepo, memberRepo, serviceUtils, tmplMgr, ) - inv, err := svc.CreateOrganizationInvitation(context.Background(), orgtests.Actor(tt.actorUserID), tt.organizationID, tt.request) + inv, err := svc.CreateOrganizationInvitation(context.Background(), orgtests.Actor(tt.actorUserID), tt.organizationID, tt.request, tt.redirectURL) if tt.expectErr != nil { require.Error(t, err) require.Nil(t, inv) @@ -719,7 +740,7 @@ func TestOrganizationInvitationService_GetOrganizationInvitation(t *testing.T) { tt.setup(orgRepo, invRepo, memberRepo) } - svc := newTestOrganizationInvitationService(&orgtests.MockOrganizationInvitationTxRunner{}, pluginConfig, &internaltests.MockUserService{}, orgtests.NewAccessControlServiceStub(), orgRepo, invRepo, memberRepo) + svc := newTestOrganizationInvitationService(&orgtests.MockOrganizationInvitationTxRunner{}, pluginConfig, &internaltests.MockUserService{}, orgtests.NewAccessControlServiceStub(), orgRepo, invRepo, memberRepo, &internaltests.MockVerificationService{}, &internaltests.MockTokenService{}) invitation, err := svc.GetOrganizationInvitation(context.Background(), orgtests.Actor(tt.actorUserID), tt.organizationID, tt.invitationID) if tt.expectErr != nil { require.Error(t, err) @@ -838,7 +859,7 @@ func TestOrganizationInvitationService_GetAllOrganizationInvitations(t *testing. tt.setup(orgRepo, invRepo, memberRepo) } - svc := newTestOrganizationInvitationService(&orgtests.MockOrganizationInvitationTxRunner{}, pluginConfig, &internaltests.MockUserService{}, orgtests.NewAccessControlServiceStub(), orgRepo, invRepo, memberRepo) + svc := newTestOrganizationInvitationService(&orgtests.MockOrganizationInvitationTxRunner{}, pluginConfig, &internaltests.MockUserService{}, orgtests.NewAccessControlServiceStub(), orgRepo, invRepo, memberRepo, &internaltests.MockVerificationService{}, &internaltests.MockTokenService{}) invitations, err := svc.GetAllOrganizationInvitations(context.Background(), orgtests.Actor(tt.actorUserID), tt.organizationID) if tt.expectErr != nil { require.Error(t, err) @@ -946,7 +967,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) + svc := newTestOrganizationInvitationService(&orgtests.MockOrganizationInvitationTxRunner{}, pluginConfig, &internaltests.MockUserService{}, orgtests.NewAccessControlServiceStub(), orgRepo, invRepo, memberRepo, &internaltests.MockVerificationService{}, &internaltests.MockTokenService{}) invitation, err := svc.RevokeOrganizationInvitation(context.Background(), orgtests.Actor(tt.actorUserID), tt.organizationID, tt.invitationID) if tt.expectErr != nil { require.Error(t, err) @@ -1032,7 +1053,7 @@ func TestOrganizationInvitationService_AcceptPendingOrganizationInvitationsForEm } txRunner := &orgtests.MockOrganizationInvitationTxRunner{} - svc := newTestOrganizationInvitationService(txRunner, pluginConfig, userSvc, orgtests.NewAccessControlServiceStub(), orgRepo, invRepo, memberRepo) + svc := newTestOrganizationInvitationService(txRunner, pluginConfig, userSvc, orgtests.NewAccessControlServiceStub(), orgRepo, invRepo, memberRepo, &internaltests.MockVerificationService{}, &internaltests.MockTokenService{}) accepted, err := svc.AcceptPendingOrganizationInvitationsForEmail(context.Background(), tt.userID, tt.email) if tt.expectErr != nil { require.Error(t, err) @@ -1281,7 +1302,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) + svc := newTestOrganizationInvitationService(&orgtests.MockOrganizationInvitationTxRunner{}, pluginConfig, userSvc, orgtests.NewAccessControlServiceStub(), orgRepo, invRepo, memberRepo, &internaltests.MockVerificationService{}, &internaltests.MockTokenService{}) invitation, err := svc.AcceptOrganizationInvitation(context.Background(), orgtests.Actor(tt.actorUserID), tt.organization, tt.invitationID) if tt.expectErr != nil { require.ErrorIs(t, err, tt.expectErr) @@ -1447,7 +1468,7 @@ func TestOrganizationInvitationService_RejectOrganizationInvitation(t *testing.T tt.setup(userSvc, invRepo) } - svc := newTestOrganizationInvitationService(&orgtests.MockOrganizationInvitationTxRunner{}, pluginConfig, userSvc, orgtests.NewAccessControlServiceStub(), &orgtests.MockOrganizationRepository{}, invRepo, &orgtests.MockOrganizationMemberRepository{}) + svc := newTestOrganizationInvitationService(&orgtests.MockOrganizationInvitationTxRunner{}, pluginConfig, userSvc, orgtests.NewAccessControlServiceStub(), &orgtests.MockOrganizationRepository{}, invRepo, &orgtests.MockOrganizationMemberRepository{}, &internaltests.MockVerificationService{}, &internaltests.MockTokenService{}) 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 54ad603..d8b213a 100644 --- a/plugins/organizations/services/organization_service.go +++ b/plugins/organizations/services/organization_service.go @@ -333,6 +333,20 @@ 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/tests/services.go b/plugins/organizations/tests/services.go index e86a1eb..5e9d6a5 100644 --- a/plugins/organizations/tests/services.go +++ b/plugins/organizations/tests/services.go @@ -63,12 +63,20 @@ 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 } -func (m *MockOrganizationInvitationService) CreateOrganizationInvitation(ctx context.Context, actor *models.Actor, organizationID string, request types.CreateOrganizationInvitationRequest) (*types.OrganizationInvitation, error) { - args := m.Called(ctx, actorID(actor), organizationID, request) +func (m *MockOrganizationInvitationService) CreateOrganizationInvitation(ctx context.Context, actor *models.Actor, organizationID string, request types.CreateOrganizationInvitationRequest, redirectURL string) (*types.OrganizationInvitation, error) { + args := m.Called(ctx, actorID(actor), organizationID, request, redirectURL) if args.Get(0) == nil { return nil, args.Error(1) } @@ -83,6 +91,14 @@ func (m *MockOrganizationInvitationService) GetOrganizationInvitation(ctx contex return args.Get(0).(*types.OrganizationInvitation), args.Error(1) } +func (m *MockOrganizationInvitationService) GetOrganizationInvitationByID(ctx context.Context, invitationID string) (*types.OrganizationInvitation, error) { + args := m.Called(ctx, invitationID) + if args.Get(0) == nil { + return nil, args.Error(1) + } + return args.Get(0).(*types.OrganizationInvitation), 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 2c5ebd9..f0d42ad 100644 --- a/plugins/organizations/types/api.go +++ b/plugins/organizations/types/api.go @@ -57,6 +57,26 @@ type AcceptOrganizationInvitationQuery struct { RedirectURL string `query:"redirect_url" json:"redirect_url,omitempty" nullable:"true"` } +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"` +} + +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 CreateOrganizationRequest struct { Name string `json:"name" required:"true" nullable:"false"` Role string `json:"role" required:"true" nullable:"false"` @@ -110,9 +130,12 @@ func (r *UpdateOrganizationRequest) Validate() error { } type CreateOrganizationInvitationRequest struct { - Email string `json:"email" required:"true" nullable:"false"` - Role string `json:"role" required:"true" nullable:"false"` - RedirectURL string `json:"redirect_url,omitempty" nullable:"true"` + Email string `json:"email" required:"true" nullable:"false"` + Role string `json:"role" required:"true" nullable:"false"` +} + +type CreateOrganizationInvitationQuery struct { + RedirectURL string `query:"redirect_url" json:"redirect_url,omitempty" nullable:"true"` } func (r *CreateOrganizationInvitationRequest) Validate() error { diff --git a/plugins/organizations/usecases/usecases.go b/plugins/organizations/usecases/usecases.go index 1e04096..2809691 100644 --- a/plugins/organizations/usecases/usecases.go +++ b/plugins/organizations/usecases/usecases.go @@ -2,7 +2,10 @@ package usecases import ( "context" + "strings" + "time" + coreerrors "github.com/Authula/authula/core/errors" "github.com/Authula/authula/models" orgconstants "github.com/Authula/authula/plugins/organizations/constants" orgservices "github.com/Authula/authula/plugins/organizations/services" @@ -11,12 +14,16 @@ import ( ) type UseCases struct { - orgService orgservices.OrganizationService - invitationService orgservices.OrganizationInvitationService - memberService orgservices.OrganizationMemberService - teamService orgservices.OrganizationTeamService - teamMemberService orgservices.OrganizationTeamMemberService - authorizer rootservices.Authorizer + 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 } func NewUseCases( @@ -25,15 +32,23 @@ func NewUseCases( memberService orgservices.OrganizationMemberService, 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, - authorizer: authorizer, + orgService: orgService, + invitationService: invitationService, + memberService: memberService, + teamService: teamService, + teamMemberService: teamMemberService, + userService: userService, + verificationService: verificationService, + tokenService: tokenService, + globalConfig: globalConfig, + authorizer: authorizer, } } @@ -83,14 +98,14 @@ func (u *UseCases) ExistsByID(ctx context.Context, organizationID string) (bool, // ------------- OrganizationInvitationService ------------- -func (u *UseCases) CreateOrganizationInvitation(ctx context.Context, actor *models.Actor, organizationID string, request types.CreateOrganizationInvitationRequest) (*types.OrganizationInvitation, error) { +func (u *UseCases) CreateOrganizationInvitation(ctx context.Context, actor *models.Actor, organizationID string, request types.CreateOrganizationInvitationRequest, redirectURL string) (*types.OrganizationInvitation, error) { if err := u.authorizer.AuthorizeOrganizationAccess(ctx, actor, organizationID); err != nil { return nil, err } if err := u.authorizer.AuthorizeScope(ctx, actor, orgconstants.OrganizationsInvitationsCreatePermission); err != nil { return nil, err } - return u.invitationService.CreateOrganizationInvitation(ctx, actor, organizationID, request) + return u.invitationService.CreateOrganizationInvitation(ctx, actor, organizationID, request, redirectURL) } func (u *UseCases) GetAllOrganizationInvitations(ctx context.Context, actor *models.Actor, organizationID string) ([]types.OrganizationInvitation, error) { @@ -104,6 +119,22 @@ func (u *UseCases) GetAllOrganizationInvitations(ctx context.Context, actor *mod } 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) + 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 } @@ -113,6 +144,71 @@ func (u *UseCases) GetOrganizationInvitation(ctx context.Context, actor *models. return u.invitationService.GetOrganizationInvitation(ctx, actor, organizationID, invitationID) } +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) + if err != nil { + return nil, err + } + if 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 err := u.verificationService.Delete(ctx, verification.ID); err != nil { + return nil, err + } + + org, err := u.orgService.GetByIDNoAuth(ctx, organizationID) + if 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 +} + func (u *UseCases) RevokeOrganizationInvitation(ctx context.Context, actor *models.Actor, organizationID string, invitationID string) (*types.OrganizationInvitation, error) { if err := u.authorizer.AuthorizeOrganizationAccess(ctx, actor, organizationID); err != nil { return nil, err @@ -124,11 +220,29 @@ 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) { - return u.invitationService.AcceptOrganizationInvitation(ctx, actor, organizationID, invitationID) + 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 } func (u *UseCases) RejectOrganizationInvitation(ctx context.Context, actor *models.Actor, organizationID string, invitationID string) (*types.OrganizationInvitation, error) { - return u.invitationService.RejectOrganizationInvitation(ctx, actor, organizationID, invitationID) + 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 } // ------------- OrganizationMemberService -------------