diff --git a/api4/cloud.go b/api4/cloud.go index b5785b31c9..e0847fc63e 100644 --- a/api4/cloud.go +++ b/api4/cloud.go @@ -240,27 +240,40 @@ func getCloudProducts(c *Context, w http.ResponseWriter, r *http.Request) { return } - if !c.App.SessionHasPermissionTo(*c.AppContext.Session(), model.PermissionSysconsoleReadBilling) { - c.SetPermissionError(model.PermissionSysconsoleReadBilling) - return - } - includeLegacyProducts := r.URL.Query().Get("include_legacy") == "true" products, err := c.App.Cloud().GetCloudProducts(c.AppContext.Session().UserId, includeLegacyProducts) - if err != nil { c.Err = model.NewAppError("Api4.getCloudProducts", "api.cloud.request_error", nil, err.Error(), http.StatusInternalServerError) return } - json, err := json.Marshal(products) + byteProductsData, err := json.Marshal(products) if err != nil { c.Err = model.NewAppError("Api4.getCloudProducts", "api.cloud.app_error", nil, err.Error(), http.StatusInternalServerError) return } - w.Write(json) + if !c.App.SessionHasPermissionTo(*c.AppContext.Session(), model.PermissionSysconsoleReadBilling) { + + sanitizedProducts := []model.UserFacingProduct{} + err = json.Unmarshal(byteProductsData, &sanitizedProducts) + if err != nil { + c.Err = model.NewAppError("Api4.getCloudProducts", "api.cloud.app_error", nil, err.Error(), http.StatusInternalServerError) + return + } + + byteSanitizedProductsData, err := json.Marshal(sanitizedProducts) + if err != nil { + c.Err = model.NewAppError("Api4.getCloudProducts", "api.cloud.app_error", nil, err.Error(), http.StatusInternalServerError) + return + } + + w.Write(byteSanitizedProductsData) + return + } + + w.Write(byteProductsData) } func getCloudLimits(c *Context, w http.ResponseWriter, r *http.Request) { diff --git a/api4/cloud_test.go b/api4/cloud_test.go index d2a00ffe79..dfdc5213ac 100644 --- a/api4/cloud_test.go +++ b/api4/cloud_test.go @@ -238,3 +238,138 @@ func Test_validateBusinessEmail(t *testing.T) { require.Error(t, err) }) } + +func TestGetCloudProducts(t *testing.T) { + cloudProducts := []*model.Product{ + { + ID: "prod_test1", + Name: "name", + Description: "description", + PricePerSeat: 10, + SKU: "sku", + PriceID: "price_id", + Family: "family", + RecurringInterval: "recurring_interval", + BillingScheme: "billing_scheme", + }, + { + ID: "prod_test2", + Name: "name2", + Description: "description2", + PricePerSeat: 100, + SKU: "sku2", + PriceID: "price_id2", + Family: "family2", + RecurringInterval: "recurring_interval2", + BillingScheme: "billing_scheme2", + }, + { + ID: "prod_test3", + Name: "name3", + Description: "description3", + PricePerSeat: 1000, + SKU: "sku3", + PriceID: "price_id3", + Family: "family3", + RecurringInterval: "recurring_interval3", + BillingScheme: "billing_scheme3", + }, + } + + sanitizedProducts := []*model.Product{ + { + ID: "prod_test1", + Name: "name", + PricePerSeat: 10, + SKU: "sku", + }, + { + ID: "prod_test2", + Name: "name2", + PricePerSeat: 100, + SKU: "sku2", + }, + { + ID: "prod_test3", + Name: "name3", + PricePerSeat: 1000, + SKU: "sku3", + }, + } + t.Run("get products for admins", func(t *testing.T) { + th := Setup(t).InitBasic() + defer th.TearDown() + + th.Client.Login(th.SystemAdminUser.Email, th.SystemAdminUser.Password) + + th.App.Srv().SetLicense(model.NewTestLicense("cloud")) + + cloud := mocks.CloudInterface{} + cloud.Mock.On("GetCloudProducts", mock.Anything, mock.Anything).Return(cloudProducts, nil) + cloudImpl := th.App.Srv().Cloud + defer func() { + th.App.Srv().Cloud = cloudImpl + }() + th.App.Srv().Cloud = &cloud + + returnedProducts, r, err := th.Client.GetCloudProducts() + require.NoError(t, err) + require.Equal(t, http.StatusOK, r.StatusCode, "Status OK") + require.Equal(t, returnedProducts, cloudProducts) + }) + + t.Run("get products for non admins", func(t *testing.T) { + th := Setup(t).InitBasic() + defer th.TearDown() + + th.Client.Login(th.BasicUser.Email, th.BasicUser.Password) + + th.App.Srv().SetLicense(model.NewTestLicense("cloud")) + + cloud := mocks.CloudInterface{} + + cloud.Mock.On("GetCloudProducts", mock.Anything, mock.Anything).Return(cloudProducts, nil) + + cloudImpl := th.App.Srv().Cloud + defer func() { + th.App.Srv().Cloud = cloudImpl + }() + th.App.Srv().Cloud = &cloud + + returnedProducts, r, err := th.Client.GetCloudProducts() + require.NoError(t, err) + require.Equal(t, http.StatusOK, r.StatusCode, "Status OK") + require.Equal(t, returnedProducts, sanitizedProducts) + + // make a more explicit check + require.Equal(t, returnedProducts[0].ID, "prod_test1") + require.Equal(t, returnedProducts[0].Name, "name") + require.Equal(t, returnedProducts[0].SKU, "sku") + require.Equal(t, returnedProducts[0].PricePerSeat, float64(10)) + require.Equal(t, returnedProducts[0].Description, "") + require.Equal(t, returnedProducts[0].PriceID, "") + require.Equal(t, returnedProducts[0].Family, model.SubscriptionFamily("")) + require.Equal(t, returnedProducts[0].RecurringInterval, model.RecurringInterval("")) + require.Equal(t, returnedProducts[0].BillingScheme, model.BillingScheme("")) + + require.Equal(t, returnedProducts[1].ID, "prod_test2") + require.Equal(t, returnedProducts[1].Name, "name2") + require.Equal(t, returnedProducts[1].SKU, "sku2") + require.Equal(t, returnedProducts[1].PricePerSeat, float64(100)) + require.Equal(t, returnedProducts[1].Description, "") + require.Equal(t, returnedProducts[1].PriceID, "") + require.Equal(t, returnedProducts[1].Family, model.SubscriptionFamily("")) + require.Equal(t, returnedProducts[1].RecurringInterval, model.RecurringInterval("")) + require.Equal(t, returnedProducts[1].BillingScheme, model.BillingScheme("")) + + require.Equal(t, returnedProducts[2].ID, "prod_test3") + require.Equal(t, returnedProducts[2].Name, "name3") + require.Equal(t, returnedProducts[2].SKU, "sku3") + require.Equal(t, returnedProducts[2].PricePerSeat, float64(1000)) + require.Equal(t, returnedProducts[2].Description, "") + require.Equal(t, returnedProducts[2].PriceID, "") + require.Equal(t, returnedProducts[2].Family, model.SubscriptionFamily("")) + require.Equal(t, returnedProducts[2].RecurringInterval, model.RecurringInterval("")) + require.Equal(t, returnedProducts[2].BillingScheme, model.BillingScheme("")) + }) +} diff --git a/model/cloud.go b/model/cloud.go index ed5f6f35f2..e8fb1e3c4b 100644 --- a/model/cloud.go +++ b/model/cloud.go @@ -53,6 +53,13 @@ type Product struct { BillingScheme BillingScheme `json:"billing_scheme"` } +type UserFacingProduct struct { + ID string `json:"id"` + Name string `json:"name"` + SKU string `json:"sku"` + PricePerSeat float64 `json:"price_per_seat"` +} + // AddOn represents an addon to a product. type AddOn struct { ID string `json:"id"`