diff --git a/server/channels/api4/cloud.go b/server/channels/api4/cloud.go index c455a82be3..3769403950 100644 --- a/server/channels/api4/cloud.go +++ b/server/channels/api4/cloud.go @@ -60,7 +60,21 @@ func (api *API) InitCloud() { api.BaseRoutes.Cloud.Handle("/delete-workspace", api.APISessionRequired(selfServeDeleteWorkspace)).Methods(http.MethodDelete) } +func ensureCloudInterface(c *Context, where string) bool { + cloud := c.App.Cloud() + if cloud == nil { + c.Err = model.NewAppError(where, "api.server.cws.needs_enterprise_edition", nil, "", http.StatusBadRequest) + return false + } + return true +} + func getSubscription(c *Context, w http.ResponseWriter, r *http.Request) { + ensured := ensureCloudInterface(c, "Api4.getSubscription") + if !ensured { + return + } + if !c.App.Channels().License().IsCloud() { c.Err = model.NewAppError("Api4.getSubscription", "api.cloud.license_error", nil, "", http.StatusForbidden) return @@ -102,6 +116,10 @@ func getSubscription(c *Context, w http.ResponseWriter, r *http.Request) { } func changeSubscription(c *Context, w http.ResponseWriter, r *http.Request) { + ensured := ensureCloudInterface(c, "Api4.changeSubscription") + if !ensured { + return + } userId := c.AppContext.Session().UserId if !c.App.Channels().License().IsCloud() { @@ -176,6 +194,11 @@ func changeSubscription(c *Context, w http.ResponseWriter, r *http.Request) { } func requestCloudTrial(c *Context, w http.ResponseWriter, r *http.Request) { + ensured := ensureCloudInterface(c, "Api4.requestCloudTrial") + if !ensured { + return + } + if !c.App.Channels().License().IsCloud() { c.Err = model.NewAppError("Api4.requestCloudTrial", "api.cloud.license_error", nil, "", http.StatusForbidden) return @@ -218,6 +241,11 @@ func requestCloudTrial(c *Context, w http.ResponseWriter, r *http.Request) { } func validateBusinessEmail(c *Context, w http.ResponseWriter, r *http.Request) { + ensured := ensureCloudInterface(c, "Api4.validateBusinessEmail") + if !ensured { + return + } + if !c.App.Channels().License().IsCloud() { c.Err = model.NewAppError("Api4.validateBusinessEmail", "api.cloud.license_error", nil, "", http.StatusForbidden) return @@ -263,6 +291,11 @@ func validateBusinessEmail(c *Context, w http.ResponseWriter, r *http.Request) { } func validateWorkspaceBusinessEmail(c *Context, w http.ResponseWriter, r *http.Request) { + ensured := ensureCloudInterface(c, "Api4.validateWorkspaceBusinessEmail") + if !ensured { + return + } + if !c.App.Channels().License().IsCloud() { c.Err = model.NewAppError("Api4.validateWorkspaceBusinessEmail", "api.cloud.license_error", nil, "", http.StatusForbidden) return @@ -309,6 +342,11 @@ func validateWorkspaceBusinessEmail(c *Context, w http.ResponseWriter, r *http.R } func getSelfHostedProducts(c *Context, w http.ResponseWriter, r *http.Request) { + ensured := ensureCloudInterface(c, "Api4.getSelfHostedProducts") + if !ensured { + return + } + products, err := c.App.Cloud().GetSelfHostedProducts(c.AppContext.Session().UserId) if err != nil { c.Err = model.NewAppError("Api4.getSelfHostedProducts", "api.cloud.request_error", nil, "", http.StatusInternalServerError).Wrap(err) @@ -343,6 +381,11 @@ func getSelfHostedProducts(c *Context, w http.ResponseWriter, r *http.Request) { } func getCloudProducts(c *Context, w http.ResponseWriter, r *http.Request) { + ensured := ensureCloudInterface(c, "Api4.getCloudProducts") + if !ensured { + return + } + if !c.App.Channels().License().IsCloud() { c.Err = model.NewAppError("Api4.getCloudProducts", "api.cloud.license_error", nil, "", http.StatusForbidden) return @@ -384,6 +427,11 @@ func getCloudProducts(c *Context, w http.ResponseWriter, r *http.Request) { } func getCloudLimits(c *Context, w http.ResponseWriter, r *http.Request) { + ensured := ensureCloudInterface(c, "Api4.getCloudLimits") + if !ensured { + return + } + if !c.App.Channels().License().IsCloud() { c.Err = model.NewAppError("Api4.getCloudLimits", "api.cloud.license_error", nil, "", http.StatusForbidden) return @@ -405,6 +453,11 @@ func getCloudLimits(c *Context, w http.ResponseWriter, r *http.Request) { } func getCloudCustomer(c *Context, w http.ResponseWriter, r *http.Request) { + ensured := ensureCloudInterface(c, "Api4.getCloudCustomer") + if !ensured { + return + } + if !c.App.Channels().License().IsCloud() { c.Err = model.NewAppError("Api4.getCloudCustomer", "api.cloud.license_error", nil, "", http.StatusForbidden) return @@ -432,6 +485,11 @@ func getCloudCustomer(c *Context, w http.ResponseWriter, r *http.Request) { // getLicenseSelfServeStatus makes check for the license in the CWS self-serve portal and establishes if the license is renewable, expandable etc. func getLicenseSelfServeStatus(c *Context, w http.ResponseWriter, r *http.Request) { + ensured := ensureCloudInterface(c, "Api4.getLicenseSelfServeStatus") + if !ensured { + return + } + if !c.App.SessionHasPermissionTo(*c.AppContext.Session(), model.PermissionManageLicenseInformation) { c.SetPermissionError(model.PermissionManageLicenseInformation) return @@ -460,6 +518,11 @@ func getLicenseSelfServeStatus(c *Context, w http.ResponseWriter, r *http.Reques } func updateCloudCustomer(c *Context, w http.ResponseWriter, r *http.Request) { + ensured := ensureCloudInterface(c, "Api4.updateCloudCustomer") + if !ensured { + return + } + if !c.App.Channels().License().IsCloud() { c.Err = model.NewAppError("Api4.updateCloudCustomer", "api.cloud.license_error", nil, "", http.StatusForbidden) return @@ -498,6 +561,11 @@ func updateCloudCustomer(c *Context, w http.ResponseWriter, r *http.Request) { } func updateCloudCustomerAddress(c *Context, w http.ResponseWriter, r *http.Request) { + ensured := ensureCloudInterface(c, "Api4.updateCloudCustomerAddress") + if !ensured { + return + } + if !c.App.Channels().License().IsCloud() { c.Err = model.NewAppError("Api4.updateCloudCustomerAddress", "api.cloud.license_error", nil, "", http.StatusForbidden) return @@ -536,6 +604,11 @@ func updateCloudCustomerAddress(c *Context, w http.ResponseWriter, r *http.Reque } func createCustomerPayment(c *Context, w http.ResponseWriter, r *http.Request) { + ensured := ensureCloudInterface(c, "Api4.createCustomerPayment") + if !ensured { + return + } + if !c.App.Channels().License().IsCloud() { c.Err = model.NewAppError("Api4.createCustomerPayment", "api.cloud.license_error", nil, "", http.StatusForbidden) return @@ -567,6 +640,11 @@ func createCustomerPayment(c *Context, w http.ResponseWriter, r *http.Request) { } func confirmCustomerPayment(c *Context, w http.ResponseWriter, r *http.Request) { + ensured := ensureCloudInterface(c, "Api4.confirmCustomerPayment") + if !ensured { + return + } + if !c.App.Channels().License().IsCloud() { c.Err = model.NewAppError("Api4.confirmCustomerPayment", "api.cloud.license_error", nil, "", http.StatusForbidden) return @@ -604,6 +682,11 @@ func confirmCustomerPayment(c *Context, w http.ResponseWriter, r *http.Request) } func getInvoicesForSubscription(c *Context, w http.ResponseWriter, r *http.Request) { + ensured := ensureCloudInterface(c, "Api4.getInvoicesForSubscription") + if !ensured { + return + } + if !c.App.Channels().License().IsCloud() { c.Err = model.NewAppError("Api4.getInvoicesForSubscription", "api.cloud.license_error", nil, "", http.StatusForbidden) return @@ -630,6 +713,11 @@ func getInvoicesForSubscription(c *Context, w http.ResponseWriter, r *http.Reque } func getSubscriptionInvoicePDF(c *Context, w http.ResponseWriter, r *http.Request) { + ensured := ensureCloudInterface(c, "Api4.getSubscriptionInvoicePDF") + if !ensured { + return + } + if !c.App.Channels().License().IsCloud() { c.Err = model.NewAppError("Api4.getSubscriptionInvoicePDF", "api.cloud.license_error", nil, "", http.StatusForbidden) return @@ -665,6 +753,11 @@ func getSubscriptionInvoicePDF(c *Context, w http.ResponseWriter, r *http.Reques } func handleCWSWebhook(c *Context, w http.ResponseWriter, r *http.Request) { + ensured := ensureCloudInterface(c, "Api4.handleCWSWebhook") + if !ensured { + return + } + if !c.App.Channels().License().IsCloud() { c.Err = model.NewAppError("Api4.handleCWSWebhook", "api.cloud.license_error", nil, "", http.StatusForbidden) return @@ -765,12 +858,12 @@ func handleCWSWebhook(c *Context, w http.ResponseWriter, r *http.Request) { } func handleCheckCWSConnection(c *Context, w http.ResponseWriter, r *http.Request) { - cloud := c.App.Cloud() - if cloud == nil { - c.Err = model.NewAppError("Api4.handleCWSHealthCheck", "api.server.cws.needs_enterprise_edition", nil, "", http.StatusBadRequest) + ensured := ensureCloudInterface(c, "Api4.handleCheckCWSConnection") + if !ensured { return } - if err := cloud.CheckCWSConnection(c.AppContext.Session().UserId); err != nil { + + if err := c.App.Cloud().CheckCWSConnection(c.AppContext.Session().UserId); err != nil { c.Err = model.NewAppError("Api4.handleCWSHealthCheck", "api.server.cws.health_check.app_error", nil, "CWS Server is not available.", http.StatusInternalServerError) return } @@ -779,6 +872,11 @@ func handleCheckCWSConnection(c *Context, w http.ResponseWriter, r *http.Request } func selfServeDeleteWorkspace(c *Context, w http.ResponseWriter, r *http.Request) { + ensured := ensureCloudInterface(c, "Api4.selfServeDeleteWorkspace") + if !ensured { + return + } + bodyBytes, err := io.ReadAll(r.Body) if err != nil { c.Err = model.NewAppError("Api4.selfServeDeleteWorkspace", "api.cloud.app_error", nil, err.Error(), http.StatusBadRequest) diff --git a/server/channels/api4/cloud_test.go b/server/channels/api4/cloud_test.go index 81d92dcbaa..1114ef4ba4 100644 --- a/server/channels/api4/cloud_test.go +++ b/server/channels/api4/cloud_test.go @@ -20,6 +20,15 @@ func Test_getCloudLimits(t *testing.T) { th := Setup(t).InitBasic() defer th.TearDown() + cloud := &mocks.CloudInterface{} + cloud.Mock.On("GetCloudLimits", mock.Anything).Return(nil, errors.New("Unable to get limits")) + + cloudImpl := th.App.Srv().Cloud + defer func() { + th.App.Srv().Cloud = cloudImpl + }() + th.App.Srv().Cloud = cloud + th.App.Srv().RemoveLicense() th.Client.Login(th.BasicUser.Email, th.BasicUser.Password) @@ -34,6 +43,15 @@ func Test_getCloudLimits(t *testing.T) { th := Setup(t).InitBasic() defer th.TearDown() + cloud := &mocks.CloudInterface{} + cloud.Mock.On("GetCloudLimits", mock.Anything).Return(nil, errors.New("Unable to get limits")) + + cloudImpl := th.App.Srv().Cloud + defer func() { + th.App.Srv().Cloud = cloudImpl + }() + th.App.Srv().Cloud = cloud + th.App.Srv().SetLicense(model.NewTestLicense()) th.Client.Login(th.BasicUser.Email, th.BasicUser.Password) diff --git a/server/channels/api4/hosted_customer.go b/server/channels/api4/hosted_customer.go index cead966f41..bbed311ea4 100644 --- a/server/channels/api4/hosted_customer.go +++ b/server/channels/api4/hosted_customer.go @@ -39,9 +39,8 @@ func (api *API) InitHostedCustomer() { } func ensureSelfHostedAdmin(c *Context, where string) { - cloud := c.App.Cloud() - if cloud == nil { - c.Err = model.NewAppError(where, "api.server.cws.needs_enterprise_edition", nil, "", http.StatusBadRequest) + ensured := ensureCloudInterface(c, where) + if !ensured { return }