diff --git a/pkg/services/ngalert/api/api_provisioning.go b/pkg/services/ngalert/api/api_provisioning.go index 262c8a05753..6471a4a38fa 100644 --- a/pkg/services/ngalert/api/api_provisioning.go +++ b/pkg/services/ngalert/api/api_provisioning.go @@ -46,7 +46,7 @@ type ContactPointService interface { type TemplateService interface { GetTemplates(ctx context.Context, orgID int64) ([]definitions.NotificationTemplate, error) GetTemplate(ctx context.Context, orgID int64, name string) (definitions.NotificationTemplate, error) - SetTemplate(ctx context.Context, orgID int64, tmpl definitions.NotificationTemplate) (definitions.NotificationTemplate, error) + UpsertTemplate(ctx context.Context, orgID int64, tmpl definitions.NotificationTemplate) (definitions.NotificationTemplate, error) DeleteTemplate(ctx context.Context, orgID int64, name string, provenance definitions.Provenance, version string) error } @@ -207,17 +207,12 @@ func (srv *ProvisioningSrv) RouteGetTemplates(c *contextmodel.ReqContext) respon return response.JSON(http.StatusOK, templates) } -func (srv *ProvisioningSrv) RouteGetTemplate(c *contextmodel.ReqContext, name string) response.Response { - templates, err := srv.templates.GetTemplates(c.Req.Context(), c.SignedInUser.GetOrgID()) +func (srv *ProvisioningSrv) RouteGetTemplate(c *contextmodel.ReqContext, nameOrUid string) response.Response { + template, err := srv.templates.GetTemplate(c.Req.Context(), c.SignedInUser.GetOrgID(), nameOrUid) if err != nil { return response.ErrOrFallback(http.StatusInternalServerError, "", err) } - for _, tmpl := range templates { - if tmpl.Name == name { - return response.JSON(http.StatusOK, tmpl) - } - } - return response.Err(provisioning.ErrTemplateNotFound) + return response.JSON(http.StatusOK, template) } func (srv *ProvisioningSrv) RoutePutTemplate(c *contextmodel.ReqContext, body definitions.NotificationTemplateContent, name string) response.Response { @@ -227,7 +222,7 @@ func (srv *ProvisioningSrv) RoutePutTemplate(c *contextmodel.ReqContext, body de Provenance: determineProvenance(c), ResourceVersion: body.ResourceVersion, } - modified, err := srv.templates.SetTemplate(c.Req.Context(), c.SignedInUser.GetOrgID(), tmpl) + modified, err := srv.templates.UpsertTemplate(c.Req.Context(), c.SignedInUser.GetOrgID(), tmpl) if err != nil { return response.ErrOrFallback(http.StatusInternalServerError, "", err) } diff --git a/pkg/services/ngalert/provisioning/errors.go b/pkg/services/ngalert/provisioning/errors.go index 5c23280000c..febdbe49923 100644 --- a/pkg/services/ngalert/provisioning/errors.go +++ b/pkg/services/ngalert/provisioning/errors.go @@ -24,6 +24,7 @@ var ( ErrTemplateNotFound = errutil.NotFound("alerting.notifications.templates.notFound") ErrTemplateInvalid = errutil.BadRequest("alerting.notifications.templates.invalidFormat").MustTemplate("Invalid format of the submitted template", errutil.WithPublic("Template is in invalid format. Correct the payload and try again.")) + ErrTemplateExists = errutil.BadRequest("alerting.notifications.templates.nameExists", errutil.WithPublicMessage("Template file with this name already exists. Use a different name or update existing one.")) ErrContactPointReferenced = errutil.Conflict("alerting.notifications.contact-points.referenced", errutil.WithPublicMessage("Contact point is currently referenced by a notification policy.")) ErrContactPointUsedInRule = errutil.Conflict("alerting.notifications.contact-points.used-by-rule", errutil.WithPublicMessage("Contact point is currently used in the notification settings of one or many alert rules.")) diff --git a/pkg/services/ngalert/provisioning/templates.go b/pkg/services/ngalert/provisioning/templates.go index bed740aa763..457c4549869 100644 --- a/pkg/services/ngalert/provisioning/templates.go +++ b/pkg/services/ngalert/provisioning/templates.go @@ -2,6 +2,7 @@ package provisioning import ( "context" + "errors" "fmt" "hash/fnv" "unsafe" @@ -9,6 +10,7 @@ import ( "github.com/grafana/grafana/pkg/infra/log" "github.com/grafana/grafana/pkg/services/ngalert/api/tooling/definitions" "github.com/grafana/grafana/pkg/services/ngalert/models" + "github.com/grafana/grafana/pkg/services/ngalert/notifier/legacy_storage" "github.com/grafana/grafana/pkg/services/ngalert/provisioning/validation" ) @@ -89,7 +91,7 @@ func (t *TemplateService) GetTemplate(ctx context.Context, orgID int64, name str return definitions.NotificationTemplate{}, ErrTemplateNotFound.Errorf("") } -func (t *TemplateService) SetTemplate(ctx context.Context, orgID int64, tmpl definitions.NotificationTemplate) (definitions.NotificationTemplate, error) { +func (t *TemplateService) UpsertTemplate(ctx context.Context, orgID int64, tmpl definitions.NotificationTemplate) (definitions.NotificationTemplate, error) { err := tmpl.Validate() if err != nil { return definitions.NotificationTemplate{}, MakeErrTemplateInvalid(err) @@ -100,32 +102,99 @@ func (t *TemplateService) SetTemplate(ctx context.Context, orgID int64, tmpl def return definitions.NotificationTemplate{}, err } + d, err := t.updateTemplate(ctx, revision, orgID, tmpl) + if err != nil { + if !errors.Is(err, ErrTemplateNotFound) { + return d, err + } + if tmpl.ResourceVersion != "" { // if version is set then it's an update operation. Fail because resource does not exist anymore + return definitions.NotificationTemplate{}, ErrTemplateNotFound.Errorf("") + } + return t.createTemplate(ctx, revision, orgID, tmpl) + } + return d, err +} + +func (t *TemplateService) CreateTemplate(ctx context.Context, orgID int64, tmpl definitions.NotificationTemplate) (definitions.NotificationTemplate, error) { + err := tmpl.Validate() + if err != nil { + return definitions.NotificationTemplate{}, MakeErrTemplateInvalid(err) + } + revision, err := t.configStore.Get(ctx, orgID) + if err != nil { + return definitions.NotificationTemplate{}, err + } + return t.createTemplate(ctx, revision, orgID, tmpl) +} + +func (t *TemplateService) createTemplate(ctx context.Context, revision *legacy_storage.ConfigRevision, orgID int64, tmpl definitions.NotificationTemplate) (definitions.NotificationTemplate, error) { if revision.Config.TemplateFiles == nil { revision.Config.TemplateFiles = map[string]string{} } - _, ok := revision.Config.TemplateFiles[tmpl.Name] - if ok { - // check that provenance is not changed in an invalid way - storedProvenance, err := t.provenanceStore.GetProvenance(ctx, &tmpl, orgID) - if err != nil { - return definitions.NotificationTemplate{}, err - } - if err := t.validator(storedProvenance, models.Provenance(tmpl.Provenance)); err != nil { - return definitions.NotificationTemplate{}, err - } + _, found := revision.Config.TemplateFiles[tmpl.Name] + if found { + return definitions.NotificationTemplate{}, ErrTemplateExists.Errorf("") } - existing, ok := revision.Config.TemplateFiles[tmpl.Name] - if ok { - err = t.checkOptimisticConcurrency(tmpl.Name, existing, models.Provenance(tmpl.Provenance), tmpl.ResourceVersion, "update") - if err != nil { - return definitions.NotificationTemplate{}, err + revision.Config.TemplateFiles[tmpl.Name] = tmpl.Template + + err := t.xact.InTransaction(ctx, func(ctx context.Context) error { + if err := t.configStore.Save(ctx, revision, orgID); err != nil { + return err } - } else if tmpl.ResourceVersion != "" { // if version is set then it's an update operation. Fail because resource does not exist anymore + return t.provenanceStore.SetProvenance(ctx, &tmpl, orgID, models.Provenance(tmpl.Provenance)) + }) + if err != nil { + return definitions.NotificationTemplate{}, err + } + + return definitions.NotificationTemplate{ + Name: tmpl.Name, + Template: tmpl.Template, + Provenance: tmpl.Provenance, + ResourceVersion: calculateTemplateFingerprint(tmpl.Template), + }, nil +} + +func (t *TemplateService) UpdateTemplate(ctx context.Context, orgID int64, tmpl definitions.NotificationTemplate) (definitions.NotificationTemplate, error) { + err := tmpl.Validate() + if err != nil { + return definitions.NotificationTemplate{}, MakeErrTemplateInvalid(err) + } + + revision, err := t.configStore.Get(ctx, orgID) + if err != nil { + return definitions.NotificationTemplate{}, err + } + return t.updateTemplate(ctx, revision, orgID, tmpl) +} + +func (t *TemplateService) updateTemplate(ctx context.Context, revision *legacy_storage.ConfigRevision, orgID int64, tmpl definitions.NotificationTemplate) (definitions.NotificationTemplate, error) { + if revision.Config.TemplateFiles == nil { + revision.Config.TemplateFiles = map[string]string{} + } + + existingName := tmpl.Name + exisitingContent, found := revision.Config.TemplateFiles[existingName] + if !found { return definitions.NotificationTemplate{}, ErrTemplateNotFound.Errorf("") } + // check that provenance is not changed in an invalid way + storedProvenance, err := t.provenanceStore.GetProvenance(ctx, &tmpl, orgID) + if err != nil { + return definitions.NotificationTemplate{}, err + } + if err := t.validator(storedProvenance, models.Provenance(tmpl.Provenance)); err != nil { + return definitions.NotificationTemplate{}, err + } + + err = t.checkOptimisticConcurrency(tmpl.Name, exisitingContent, models.Provenance(tmpl.Provenance), tmpl.ResourceVersion, "update") + if err != nil { + return definitions.NotificationTemplate{}, err + } + revision.Config.TemplateFiles[tmpl.Name] = tmpl.Template err = t.xact.InTransaction(ctx, func(ctx context.Context) error { diff --git a/pkg/services/ngalert/provisioning/templates_test.go b/pkg/services/ngalert/provisioning/templates_test.go index 2921a8b1499..d7e316defcd 100644 --- a/pkg/services/ngalert/provisioning/templates_test.go +++ b/pkg/services/ngalert/provisioning/templates_test.go @@ -200,7 +200,7 @@ func TestGetTemplate(t *testing.T) { }) } -func TestSetTemplate(t *testing.T) { +func TestUpsertTemplate(t *testing.T) { orgID := int64(1) templateName := "template1" currentTemplateContent := "test1" @@ -240,7 +240,7 @@ func TestSetTemplate(t *testing.T) { ResourceVersion: "", } - result, err := sut.SetTemplate(context.Background(), orgID, tmpl) + result, err := sut.UpsertTemplate(context.Background(), orgID, tmpl) require.NoError(t, err) require.Equal(t, definitions.NotificationTemplate{ @@ -281,7 +281,7 @@ func TestSetTemplate(t *testing.T) { ResourceVersion: calculateTemplateFingerprint("test1"), } - result, err := sut.SetTemplate(context.Background(), orgID, tmpl) + result, err := sut.UpsertTemplate(context.Background(), orgID, tmpl) require.NoError(t, err) assert.Equal(t, definitions.NotificationTemplate{ @@ -317,7 +317,7 @@ func TestSetTemplate(t *testing.T) { ResourceVersion: "", } - result, err := sut.SetTemplate(context.Background(), orgID, tmpl) + result, err := sut.UpsertTemplate(context.Background(), orgID, tmpl) require.NoError(t, err) assert.Equal(t, definitions.NotificationTemplate{ @@ -351,7 +351,7 @@ func TestSetTemplate(t *testing.T) { ResourceVersion: calculateTemplateFingerprint(currentTemplateContent), } - result, _ := sut.SetTemplate(context.Background(), orgID, tmpl) + result, _ := sut.UpsertTemplate(context.Background(), orgID, tmpl) expectedContent := fmt.Sprintf("{{ define \"%s\" }}\n content\n{{ end }}", templateName) require.Equal(t, definitions.NotificationTemplate{ @@ -375,7 +375,7 @@ func TestSetTemplate(t *testing.T) { Name: "name", Template: "{{ .NotAField }}", } - _, err := sut.SetTemplate(context.Background(), 1, tmpl) + _, err := sut.UpsertTemplate(context.Background(), 1, tmpl) require.NoError(t, err) }) @@ -388,7 +388,7 @@ func TestSetTemplate(t *testing.T) { Name: "", Template: "", } - _, err := sut.SetTemplate(context.Background(), orgID, tmpl) + _, err := sut.UpsertTemplate(context.Background(), orgID, tmpl) require.ErrorIs(t, err, ErrTemplateInvalid) }) @@ -397,7 +397,7 @@ func TestSetTemplate(t *testing.T) { Name: "", Template: "{{ .MyField }", } - _, err := sut.SetTemplate(context.Background(), orgID, tmpl) + _, err := sut.UpsertTemplate(context.Background(), orgID, tmpl) require.ErrorIs(t, err, ErrTemplateInvalid) }) @@ -426,7 +426,7 @@ func TestSetTemplate(t *testing.T) { } template.Provenance = definitions.Provenance(models.ProvenanceNone) - _, err := sut.SetTemplate(context.Background(), orgID, template) + _, err := sut.UpsertTemplate(context.Background(), orgID, template) require.ErrorIs(t, err, expectedErr) }) @@ -445,7 +445,7 @@ func TestSetTemplate(t *testing.T) { Provenance: definitions.Provenance(models.ProvenanceNone), } - _, err := sut.SetTemplate(context.Background(), orgID, template) + _, err := sut.UpsertTemplate(context.Background(), orgID, template) require.ErrorIs(t, err, ErrVersionConflict) prov.AssertExpectations(t) @@ -462,7 +462,7 @@ func TestSetTemplate(t *testing.T) { ResourceVersion: "version", Provenance: definitions.Provenance(models.ProvenanceNone), } - _, err := sut.SetTemplate(context.Background(), orgID, template) + _, err := sut.UpsertTemplate(context.Background(), orgID, template) require.ErrorIs(t, err, ErrTemplateNotFound) }) t.Run("propagates errors", func(t *testing.T) { @@ -478,7 +478,7 @@ func TestSetTemplate(t *testing.T) { return nil, expectedErr } - _, err := sut.SetTemplate(context.Background(), orgID, tmpl) + _, err := sut.UpsertTemplate(context.Background(), orgID, tmpl) require.ErrorIs(t, err, expectedErr) }) @@ -490,7 +490,7 @@ func TestSetTemplate(t *testing.T) { expectedErr := errors.New("test") prov.EXPECT().GetProvenance(mock.Anything, mock.Anything, mock.Anything).Return(models.ProvenanceNone, expectedErr) - _, err := sut.SetTemplate(context.Background(), orgID, tmpl) + _, err := sut.UpsertTemplate(context.Background(), orgID, tmpl) require.ErrorIs(t, err, expectedErr) @@ -506,7 +506,7 @@ func TestSetTemplate(t *testing.T) { prov.EXPECT().GetProvenance(mock.Anything, mock.Anything, mock.Anything).Return(models.ProvenanceNone, nil) prov.EXPECT().SetProvenance(mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return(expectedErr) - _, err := sut.SetTemplate(context.Background(), orgID, tmpl) + _, err := sut.UpsertTemplate(context.Background(), orgID, tmpl) require.ErrorIs(t, err, expectedErr) prov.AssertExpectations(t) @@ -524,7 +524,378 @@ func TestSetTemplate(t *testing.T) { prov.EXPECT().SaveSucceeds() prov.EXPECT().GetProvenance(mock.Anything, mock.Anything, mock.Anything).Return(models.ProvenanceNone, nil) - _, err := sut.SetTemplate(context.Background(), 1, tmpl) + _, err := sut.UpsertTemplate(context.Background(), 1, tmpl) + require.ErrorIs(t, err, expectedErr) + }) + }) +} + +func TestCreateTemplate(t *testing.T) { + orgID := int64(1) + amConfigToken := util.GenerateShortUID() + + tmpl := definitions.NotificationTemplate{ + Name: "new-template", + Template: "{{ define \"test\"}} test {{ end }}", + Provenance: definitions.Provenance(models.ProvenanceAPI), + } + + revision := func() *legacy_storage.ConfigRevision { + return &legacy_storage.ConfigRevision{ + Config: &definitions.PostableUserConfig{}, + ConcurrencyToken: amConfigToken, + } + } + + t.Run("adds new template to config file", func(t *testing.T) { + sut, store, prov := createTemplateServiceSut() + store.GetFn = func(ctx context.Context, org int64) (*legacy_storage.ConfigRevision, error) { + assert.Equal(t, orgID, org) + return revision(), nil + } + store.SaveFn = func(ctx context.Context, revision *legacy_storage.ConfigRevision) error { + assertInTransaction(t, ctx) + return nil + } + prov.EXPECT().SetProvenance(mock.Anything, mock.Anything, mock.Anything, mock.Anything).Run(func(ctx context.Context, o models.Provisionable, org int64, p models.Provenance) { + assertInTransaction(t, ctx) + }).Return(nil) + + result, err := sut.CreateTemplate(context.Background(), orgID, tmpl) + + require.NoError(t, err) + require.Equal(t, definitions.NotificationTemplate{ + Name: tmpl.Name, + Template: tmpl.Template, + Provenance: tmpl.Provenance, + ResourceVersion: calculateTemplateFingerprint(tmpl.Template), + }, result) + + require.Len(t, store.Calls, 2) + + require.Equal(t, "Save", store.Calls[1].Method) + saved := store.Calls[1].Args[1].(*legacy_storage.ConfigRevision) + assert.Equal(t, amConfigToken, saved.ConcurrencyToken) + assert.Contains(t, saved.Config.TemplateFiles, tmpl.Name) + assert.Equal(t, tmpl.Template, saved.Config.TemplateFiles[tmpl.Name]) + + prov.AssertCalled(t, "SetProvenance", mock.Anything, mock.MatchedBy(func(t *definitions.NotificationTemplate) bool { + return t.Name == tmpl.Name + }), orgID, models.ProvenanceAPI) + }) + + t.Run("returns ErrTemplateExists if template exists", func(t *testing.T) { + sut, store, _ := createTemplateServiceSut() + store.GetFn = func(ctx context.Context, org int64) (*legacy_storage.ConfigRevision, error) { + assert.Equal(t, orgID, org) + return &legacy_storage.ConfigRevision{ + Config: &definitions.PostableUserConfig{ + TemplateFiles: map[string]string{ + tmpl.Name: "test", + }, + }, + ConcurrencyToken: amConfigToken, + }, nil + } + + _, err := sut.CreateTemplate(context.Background(), orgID, tmpl) + + require.ErrorIs(t, err, ErrTemplateExists) + }) + + t.Run("rejects templates that fail validation", func(t *testing.T) { + sut, store, prov := createTemplateServiceSut() + + t.Run("empty content", func(t *testing.T) { + tmpl := definitions.NotificationTemplate{ + Name: "", + Template: "", + } + _, err := sut.CreateTemplate(context.Background(), orgID, tmpl) + require.ErrorIs(t, err, ErrTemplateInvalid) + }) + + t.Run("invalid content", func(t *testing.T) { + tmpl := definitions.NotificationTemplate{ + Name: "", + Template: "{{ .MyField }", + } + _, err := sut.CreateTemplate(context.Background(), orgID, tmpl) + require.ErrorIs(t, err, ErrTemplateInvalid) + }) + + require.Empty(t, store.Calls) + prov.AssertExpectations(t) + }) + + t.Run("propagates errors", func(t *testing.T) { + t.Run("when unable to read config", func(t *testing.T) { + sut, store, _ := createTemplateServiceSut() + expectedErr := errors.New("test") + store.GetFn = func(ctx context.Context, orgID int64) (*legacy_storage.ConfigRevision, error) { + return nil, expectedErr + } + + _, err := sut.CreateTemplate(context.Background(), orgID, tmpl) + require.ErrorIs(t, err, expectedErr) + }) + + t.Run("when provenance fails to save", func(t *testing.T) { + sut, store, prov := createTemplateServiceSut() + expectedErr := errors.New("test") + store.GetFn = func(ctx context.Context, orgID int64) (*legacy_storage.ConfigRevision, error) { + return revision(), nil + } + prov.EXPECT().SetProvenance(mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return(expectedErr) + + _, err := sut.CreateTemplate(context.Background(), orgID, tmpl) + require.ErrorIs(t, err, expectedErr) + + prov.AssertExpectations(t) + }) + + t.Run("when AM config fails to save", func(t *testing.T) { + sut, store, prov := createTemplateServiceSut() + store.GetFn = func(ctx context.Context, orgID int64) (*legacy_storage.ConfigRevision, error) { + return revision(), nil + } + expectedErr := errors.New("test") + store.SaveFn = func(ctx context.Context, revision *legacy_storage.ConfigRevision) error { + return expectedErr + } + prov.EXPECT().SaveSucceeds() + prov.EXPECT().GetProvenance(mock.Anything, mock.Anything, mock.Anything).Return(models.ProvenanceNone, nil) + + _, err := sut.CreateTemplate(context.Background(), 1, tmpl) + require.ErrorIs(t, err, expectedErr) + }) + }) +} + +func TestUpdateTemplate(t *testing.T) { + orgID := int64(1) + currentTemplateContent := "test1" + + tmpl := definitions.NotificationTemplate{ + Name: "template1", + Template: "{{ define \"test\"}} test {{ end }}", + Provenance: definitions.Provenance(models.ProvenanceAPI), + ResourceVersion: "", + } + + amConfigToken := util.GenerateShortUID() + revision := func() *legacy_storage.ConfigRevision { + return &legacy_storage.ConfigRevision{ + Config: &definitions.PostableUserConfig{ + TemplateFiles: map[string]string{ + tmpl.Name: currentTemplateContent, + }, + }, + ConcurrencyToken: amConfigToken, + } + } + + t.Run("returns ErrTemplateNotFound if template does not exist", func(t *testing.T) { + sut, store, prov := createTemplateServiceSut() + store.GetFn = func(ctx context.Context, org int64) (*legacy_storage.ConfigRevision, error) { + assert.Equal(t, orgID, org) + return &legacy_storage.ConfigRevision{ + Config: &definitions.PostableUserConfig{}, + ConcurrencyToken: amConfigToken, + }, nil + } + _, err := sut.UpdateTemplate(context.Background(), orgID, tmpl) + + require.ErrorIs(t, err, ErrTemplateNotFound) + + require.Len(t, store.Calls, 1) + prov.AssertExpectations(t) + }) + + t.Run("updates current template", func(t *testing.T) { + t.Run("when version matches", func(t *testing.T) { + sut, store, prov := createTemplateServiceSut() + store.GetFn = func(ctx context.Context, org int64) (*legacy_storage.ConfigRevision, error) { + return revision(), nil + } + prov.EXPECT().GetProvenance(mock.Anything, mock.Anything, mock.Anything).Return(models.ProvenanceAPI, nil) + prov.EXPECT().SetProvenance(mock.Anything, mock.Anything, mock.Anything, mock.Anything).Run(func(ctx context.Context, o models.Provisionable, org int64, p models.Provenance) { + assertInTransaction(t, ctx) + }).Return(nil) + + result, err := sut.UpdateTemplate(context.Background(), orgID, tmpl) + + require.NoError(t, err) + assert.Equal(t, definitions.NotificationTemplate{ + Name: tmpl.Name, + Template: tmpl.Template, + Provenance: tmpl.Provenance, + ResourceVersion: calculateTemplateFingerprint(tmpl.Template), + }, result) + + require.Len(t, store.Calls, 2) + require.Equal(t, "Save", store.Calls[1].Method) + saved := store.Calls[1].Args[1].(*legacy_storage.ConfigRevision) + assert.Equal(t, amConfigToken, saved.ConcurrencyToken) + assert.Contains(t, saved.Config.TemplateFiles, tmpl.Name) + assert.Equal(t, tmpl.Template, saved.Config.TemplateFiles[tmpl.Name]) + + prov.AssertExpectations(t) + }) + t.Run("bypasses optimistic concurrency validation when version is empty", func(t *testing.T) { + sut, store, prov := createTemplateServiceSut() + store.GetFn = func(ctx context.Context, org int64) (*legacy_storage.ConfigRevision, error) { + return revision(), nil + } + prov.EXPECT().GetProvenance(mock.Anything, mock.Anything, mock.Anything).Return(models.ProvenanceAPI, nil) + prov.EXPECT().SetProvenance(mock.Anything, mock.Anything, mock.Anything, mock.Anything).Run(func(ctx context.Context, o models.Provisionable, org int64, p models.Provenance) { + assertInTransaction(t, ctx) + }).Return(nil) + + result, err := sut.UpdateTemplate(context.Background(), orgID, tmpl) + + require.NoError(t, err) + assert.Equal(t, definitions.NotificationTemplate{ + Name: tmpl.Name, + Template: tmpl.Template, + Provenance: tmpl.Provenance, + ResourceVersion: calculateTemplateFingerprint(tmpl.Template), + }, result) + + require.Equal(t, "Save", store.Calls[1].Method) + saved := store.Calls[1].Args[1].(*legacy_storage.ConfigRevision) + assert.Equal(t, amConfigToken, saved.ConcurrencyToken) + assert.Contains(t, saved.Config.TemplateFiles, tmpl.Name) + assert.Equal(t, tmpl.Template, saved.Config.TemplateFiles[tmpl.Name]) + }) + }) + + t.Run("rejects templates that fail validation", func(t *testing.T) { + sut, store, prov := createTemplateServiceSut() + + t.Run("empty content", func(t *testing.T) { + tmpl := definitions.NotificationTemplate{ + Name: "", + Template: "", + } + _, err := sut.UpdateTemplate(context.Background(), orgID, tmpl) + require.ErrorIs(t, err, ErrTemplateInvalid) + }) + + t.Run("invalid content", func(t *testing.T) { + tmpl := definitions.NotificationTemplate{ + Name: "", + Template: "{{ .MyField }", + } + _, err := sut.UpdateTemplate(context.Background(), orgID, tmpl) + require.ErrorIs(t, err, ErrTemplateInvalid) + }) + + require.Empty(t, store.Calls) + prov.AssertExpectations(t) + }) + + t.Run("rejects existing templates if provenance is not right", func(t *testing.T) { + sut, store, prov := createTemplateServiceSut() + store.GetFn = func(ctx context.Context, org int64) (*legacy_storage.ConfigRevision, error) { + return revision(), nil + } + + prov.EXPECT().GetProvenance(mock.Anything, mock.Anything, mock.Anything).Return(models.ProvenanceAPI, nil) + + expectedErr := errors.New("test") + sut.validator = func(from, to models.Provenance) error { + assert.Equal(t, models.ProvenanceAPI, from) + assert.Equal(t, models.ProvenanceNone, to) + return expectedErr + } + + template := definitions.NotificationTemplate{ + Name: "template1", + Template: "asdf-new", + } + template.Provenance = definitions.Provenance(models.ProvenanceNone) + + _, err := sut.UpdateTemplate(context.Background(), orgID, template) + + require.ErrorIs(t, err, expectedErr) + }) + + t.Run("rejects existing templates if version is not right", func(t *testing.T) { + sut, store, prov := createTemplateServiceSut() + store.GetFn = func(ctx context.Context, org int64) (*legacy_storage.ConfigRevision, error) { + return revision(), nil + } + prov.EXPECT().GetProvenance(mock.Anything, mock.Anything, mock.Anything).Return(models.ProvenanceNone, nil) + + template := definitions.NotificationTemplate{ + Name: "template1", + Template: "asdf-new", + ResourceVersion: "bad-version", + Provenance: definitions.Provenance(models.ProvenanceNone), + } + + _, err := sut.UpdateTemplate(context.Background(), orgID, template) + + require.ErrorIs(t, err, ErrVersionConflict) + prov.AssertExpectations(t) + }) + + t.Run("propagates errors", func(t *testing.T) { + t.Run("when unable to read config", func(t *testing.T) { + sut, store, _ := createTemplateServiceSut() + expectedErr := errors.New("test") + store.GetFn = func(ctx context.Context, orgID int64) (*legacy_storage.ConfigRevision, error) { + return nil, expectedErr + } + + _, err := sut.UpdateTemplate(context.Background(), orgID, tmpl) + require.ErrorIs(t, err, expectedErr) + }) + + t.Run("when reading provenance status fails", func(t *testing.T) { + sut, store, prov := createTemplateServiceSut() + store.GetFn = func(ctx context.Context, org int64) (*legacy_storage.ConfigRevision, error) { + return revision(), nil + } + expectedErr := errors.New("test") + prov.EXPECT().GetProvenance(mock.Anything, mock.Anything, mock.Anything).Return(models.ProvenanceNone, expectedErr) + + _, err := sut.UpdateTemplate(context.Background(), orgID, tmpl) + + require.ErrorIs(t, err, expectedErr) + + prov.AssertExpectations(t) + }) + + t.Run("when provenance fails to save", func(t *testing.T) { + sut, store, prov := createTemplateServiceSut() + expectedErr := errors.New("test") + store.GetFn = func(ctx context.Context, orgID int64) (*legacy_storage.ConfigRevision, error) { + return revision(), nil + } + prov.EXPECT().GetProvenance(mock.Anything, mock.Anything, mock.Anything).Return(models.ProvenanceNone, nil) + prov.EXPECT().SetProvenance(mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return(expectedErr) + + _, err := sut.UpdateTemplate(context.Background(), orgID, tmpl) + require.ErrorIs(t, err, expectedErr) + + prov.AssertExpectations(t) + }) + + t.Run("when AM config fails to save", func(t *testing.T) { + sut, store, prov := createTemplateServiceSut() + store.GetFn = func(ctx context.Context, orgID int64) (*legacy_storage.ConfigRevision, error) { + return revision(), nil + } + expectedErr := errors.New("test") + store.SaveFn = func(ctx context.Context, revision *legacy_storage.ConfigRevision) error { + return expectedErr + } + prov.EXPECT().SaveSucceeds() + prov.EXPECT().GetProvenance(mock.Anything, mock.Anything, mock.Anything).Return(models.ProvenanceNone, nil) + + _, err := sut.UpdateTemplate(context.Background(), 1, tmpl) require.ErrorIs(t, err, expectedErr) }) }) diff --git a/pkg/services/provisioning/alerting/text_templates.go b/pkg/services/provisioning/alerting/text_templates.go index e64ccdd5959..056b41f1b95 100644 --- a/pkg/services/provisioning/alerting/text_templates.go +++ b/pkg/services/provisioning/alerting/text_templates.go @@ -32,7 +32,7 @@ func (c *defaultTextTemplateProvisioner) Provision(ctx context.Context, for _, file := range files { for _, template := range file.Templates { template.Data.Provenance = definitions.Provenance(models.ProvenanceFile) - _, err := c.templateService.SetTemplate(ctx, template.OrgID, template.Data) + _, err := c.templateService.UpsertTemplate(ctx, template.OrgID, template.Data) if err != nil { return err }