Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion cmd/serve.go
Original file line number Diff line number Diff line change
Expand Up @@ -503,7 +503,7 @@ func buildAPIDependencies(
userProjectsService := userprojects.NewService(userProjectsRepository)

domainRepository := postgres.NewDomainRepository(logger, dbc)
domainService := domain.NewService(logger, domainRepository, userService, organizationService, membershipService)
domainService := domain.NewService(logger, domainRepository, userService, organizationService, membershipService, auditRecordRepository)

metaschemaRepository := postgres.NewMetaSchemaRepository(logger, dbc)
metaschemaService := metaschema.NewService(metaschemaRepository, logger, cfg.App.Metaschema.RefreshInterval)
Expand Down
94 changes: 94 additions & 0 deletions core/domain/mocks/audit_record_repository.go

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

60 changes: 58 additions & 2 deletions core/domain/mocks/org_service.go

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

74 changes: 59 additions & 15 deletions core/domain/service.go
Original file line number Diff line number Diff line change
Expand Up @@ -11,8 +11,10 @@ import (
"strings"
"time"

auditmodels "github.com/raystack/frontier/core/auditrecord/models"
"github.com/raystack/frontier/core/membership"
"github.com/raystack/frontier/core/organization"
pkgauditrecord "github.com/raystack/frontier/pkg/auditrecord"
"github.com/raystack/frontier/pkg/utils"

"log/slog"
Expand All @@ -29,20 +31,26 @@ type UserService interface {

type OrgService interface {
Get(ctx context.Context, id string) (organization.Organization, error)
GetRaw(ctx context.Context, id string) (organization.Organization, error)
}

type MembershipService interface {
AddOrganizationMember(ctx context.Context, orgID, principalID, principalType, roleID string) error
ListResourcesByPrincipal(ctx context.Context, principal authenticate.Principal, resourceType string, filter membership.ResourceFilter) ([]string, error)
}

type AuditRecordRepository interface {
Create(ctx context.Context, auditRecord auditmodels.AuditRecord) (auditmodels.AuditRecord, error)
}

type Service struct {
repository Repository
userService UserService
orgService OrgService
membershipService MembershipService
cron *cron.Cron
log *slog.Logger
repository Repository
userService UserService
orgService OrgService
membershipService MembershipService
auditRecordRepository AuditRecordRepository
cron *cron.Cron
log *slog.Logger
}

const (
Expand All @@ -52,14 +60,15 @@ const (
refreshTime = "0 0 * * *" // Once a day at midnight (UTC)
)

func NewService(logger *slog.Logger, repository Repository, userService UserService, orgService OrgService, membershipService MembershipService) *Service {
func NewService(logger *slog.Logger, repository Repository, userService UserService, orgService OrgService, membershipService MembershipService, auditRecordRepository AuditRecordRepository) *Service {
return &Service{
repository: repository,
userService: userService,
orgService: orgService,
membershipService: membershipService,
cron: cron.New(),
log: logger,
repository: repository,
userService: userService,
orgService: orgService,
membershipService: membershipService,
auditRecordRepository: auditRecordRepository,
cron: cron.New(),
log: logger,
}
}

Expand All @@ -73,9 +82,44 @@ func (s Service) List(ctx context.Context, flt Filter) ([]Domain, error) {
return s.repository.List(ctx, flt)
}

// Remove an organization's whitelisted domain from the database
// Delete marks an organization's whitelisted domain as deleted and writes an audit record
func (s Service) Delete(ctx context.Context, id string) error {
return s.repository.Delete(ctx, id)
dmn, err := s.repository.Get(ctx, id)
if err != nil {
return err
}
org, err := s.orgService.GetRaw(ctx, dmn.OrgID)
if err != nil {
return err
}
if err = s.repository.Delete(ctx, id); err != nil {
return err
}

s.createAuditRecord(ctx, pkgauditrecord.DomainDeletedEvent, dmn, org)
return nil
}

func (s Service) createAuditRecord(ctx context.Context, event pkgauditrecord.Event, dmn Domain, org organization.Organization) {
if _, err := s.auditRecordRepository.Create(ctx, auditmodels.AuditRecord{
Event: event,
Resource: auditmodels.Resource{
ID: org.ID,
Type: pkgauditrecord.OrganizationType,
Name: org.Title,
},
Target: &auditmodels.Target{
ID: dmn.ID,
Type: pkgauditrecord.DomainType,
Name: dmn.Name,
},
OrgID: org.ID,
OrgName: org.Title,
OccurredAt: time.Now(),
}); err != nil {
s.log.WarnContext(ctx, "failed to create domain audit record",
"event", event, "org_id", org.ID, "domain_id", dmn.ID, "domain_name", dmn.Name, "err", err)
}
}

// Creates a record for the domain in the database and returns the TXT record that needs to be added to the DNS for the domain verification
Expand Down
81 changes: 80 additions & 1 deletion core/domain/service_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,14 +2,19 @@ package domain_test

import (
"context"
"errors"
"log/slog"
"testing"
"time"

auditmodels "github.com/raystack/frontier/core/auditrecord/models"
"github.com/raystack/frontier/core/authenticate"
"github.com/raystack/frontier/core/domain"
"github.com/raystack/frontier/core/domain/mocks"
"github.com/raystack/frontier/core/organization"
"github.com/raystack/frontier/core/user"
"github.com/raystack/frontier/internal/bootstrap/schema"
pkgauditrecord "github.com/raystack/frontier/pkg/auditrecord"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/mock"
)
Expand All @@ -25,7 +30,7 @@ func TestService_ListJoinableOrgsByDomain(t *testing.T) {
userSvc := mocks.NewUserService(t)
orgSvc := mocks.NewOrgService(t)
memberSvc := mocks.NewMembershipService(t)
svc := domain.NewService(slog.Default(), repo, userSvc, orgSvc, memberSvc)
svc := domain.NewService(slog.Default(), repo, userSvc, orgSvc, memberSvc, mocks.NewAuditRecordRepository(t))
return svc, repo, userSvc, memberSvc
}

Expand Down Expand Up @@ -85,3 +90,77 @@ func TestService_ListJoinableOrgsByDomain(t *testing.T) {
assert.Equal(t, []string{"org-1", "org-2"}, got)
})
}

func TestService_Delete(t *testing.T) {
ctx := context.Background()
dmn := domain.Domain{ID: "3f1c9a2e-8b4d-4e6f-9a1b-2c3d4e5f6a7b", Name: "acme.com", OrgID: "org-1"}
org := organization.Organization{ID: "org-1", Name: "acme", Title: "Acme Inc"}
errDB := errors.New("connection reset")

newService := func(t *testing.T) (*domain.Service, *mocks.Repository, *mocks.OrgService, *mocks.AuditRecordRepository) {
t.Helper()
repo := mocks.NewRepository(t)
orgSvc := mocks.NewOrgService(t)
auditRepo := mocks.NewAuditRecordRepository(t)
svc := domain.NewService(slog.Default(), repo, mocks.NewUserService(t), orgSvc, mocks.NewMembershipService(t), auditRepo)
return svc, repo, orgSvc, auditRepo
}

t.Run("marks the domain deleted and writes a domain.deleted audit record", func(t *testing.T) {
svc, repo, orgSvc, auditRepo := newService(t)
repo.EXPECT().Get(ctx, dmn.ID).Return(dmn, nil)
orgSvc.EXPECT().GetRaw(ctx, org.ID).Return(org, nil)
repo.EXPECT().Delete(ctx, dmn.ID).Return(nil)

var got auditmodels.AuditRecord
auditRepo.EXPECT().Create(ctx, mock.Anything).
Run(func(_ context.Context, record auditmodels.AuditRecord) { got = record }).
Return(auditmodels.AuditRecord{}, nil)

assert.NoError(t, svc.Delete(ctx, dmn.ID))

assert.False(t, got.OccurredAt.IsZero())
got.OccurredAt = time.Time{}
assert.Equal(t, auditmodels.AuditRecord{
Event: pkgauditrecord.DomainDeletedEvent,
Resource: auditmodels.Resource{ID: org.ID, Type: pkgauditrecord.OrganizationType, Name: org.Title},
Target: &auditmodels.Target{ID: dmn.ID, Type: pkgauditrecord.DomainType, Name: dmn.Name},
OrgID: org.ID,
OrgName: org.Title,
}, got)
})

t.Run("an unknown domain writes no audit record", func(t *testing.T) {
svc, repo, _, _ := newService(t)
repo.EXPECT().Get(ctx, dmn.ID).Return(domain.Domain{}, domain.ErrNotExist)

assert.ErrorIs(t, svc.Delete(ctx, dmn.ID), domain.ErrNotExist)
})

t.Run("a failed org lookup leaves the domain alone", func(t *testing.T) {
svc, repo, orgSvc, _ := newService(t)
repo.EXPECT().Get(ctx, dmn.ID).Return(dmn, nil)
orgSvc.EXPECT().GetRaw(ctx, org.ID).Return(organization.Organization{}, errDB)

assert.ErrorIs(t, svc.Delete(ctx, dmn.ID), errDB)
})

t.Run("a failed delete writes no audit record", func(t *testing.T) {
svc, repo, orgSvc, _ := newService(t)
repo.EXPECT().Get(ctx, dmn.ID).Return(dmn, nil)
orgSvc.EXPECT().GetRaw(ctx, org.ID).Return(org, nil)
repo.EXPECT().Delete(ctx, dmn.ID).Return(errDB)

assert.ErrorIs(t, svc.Delete(ctx, dmn.ID), errDB)
})

t.Run("a failed audit write does not fail the delete", func(t *testing.T) {
svc, repo, orgSvc, auditRepo := newService(t)
repo.EXPECT().Get(ctx, dmn.ID).Return(dmn, nil)
orgSvc.EXPECT().GetRaw(ctx, org.ID).Return(org, nil)
repo.EXPECT().Delete(ctx, dmn.ID).Return(nil)
auditRepo.EXPECT().Create(ctx, mock.Anything).Return(auditmodels.AuditRecord{}, errDB)

assert.NoError(t, svc.Delete(ctx, dmn.ID))
})
}
Loading
Loading