diff --git a/internal/store/postgres/project_repository.go b/internal/store/postgres/project_repository.go index be4720c73..d75614444 100644 --- a/internal/store/postgres/project_repository.go +++ b/internal/store/postgres/project_repository.go @@ -7,11 +7,15 @@ import ( "errors" "fmt" "strings" + "time" "github.com/doug-martin/goqu/v9" + "github.com/jmoiron/sqlx" "github.com/raystack/frontier/core/organization" "github.com/raystack/frontier/core/project" + "github.com/raystack/frontier/pkg/auditrecord" "github.com/raystack/frontier/pkg/db" + "github.com/raystack/frontier/pkg/metadata" ) type ProjectRepository struct { @@ -24,6 +28,28 @@ func NewProjectRepository(dbc *db.Client) *ProjectRepository { } } +// projectWithOrgName is a written project row plus its org's title, for the audit record +type projectWithOrgName struct { + Project + OrgName sql.NullString `db:"org_name"` +} + +// projectReturning returns the written project row with its org's title +func projectReturning() []any { + return []any{ + goqu.I(TABLE_PROJECTS + ".*"), + dialect.From(TABLE_ORGANIZATIONS).Select("title"). + Where(goqu.Ex{"id": goqu.I(TABLE_PROJECTS + ".org_id")}).As("org_name"), + } +} + +func buildProjectAuditRecord(ctx context.Context, event auditrecord.Event, p projectWithOrgName, occurredAt time.Time) AuditRecord { + return BuildAuditRecord(ctx, event, + AuditResource{ID: p.OrgID, Type: auditrecord.OrganizationType, Name: p.OrgName.String}, + &AuditTarget{ID: p.ID, Type: auditrecord.ProjectType, Name: p.Title.String, Metadata: metadata.Metadata{"name": p.Name}}, + p.OrgID, nil, occurredAt) +} + var notDisabledProjectExp = goqu.Or( goqu.Ex{ "state": nil, @@ -123,14 +149,19 @@ func (r ProjectRepository) Create(ctx context.Context, prj project.Project) (pro if prj.State != "" { insertRow["state"] = prj.State } - query, params, err := dialect.Insert(TABLE_PROJECTS).Rows(insertRow).Returning(&Project{}).ToSQL() + query, params, err := dialect.Insert(TABLE_PROJECTS).Rows(insertRow).Returning(projectReturning()...).ToSQL() if err != nil { return project.Project{}, fmt.Errorf("%w: %w", errQuery, err) } - var projectModel Project - if err = r.dbc.WithTimeout(ctx, TABLE_PROJECTS, "Upsert", func(ctx context.Context) error { - return r.dbc.QueryRowxContext(ctx, query, params...).StructScan(&projectModel) + var result projectWithOrgName + if err = r.dbc.WithTxn(ctx, sql.TxOptions{}, func(tx *sqlx.Tx) error { + if err := r.dbc.WithTimeout(ctx, TABLE_PROJECTS, "Upsert", func(ctx context.Context) error { + return tx.QueryRowxContext(ctx, query, params...).StructScan(&result) + }); err != nil { + return err + } + return InsertAuditRecordInTx(ctx, tx, buildProjectAuditRecord(ctx, auditrecord.ProjectCreatedEvent, result, result.CreatedAt)) }); err != nil { err = checkPostgresError(err) switch { @@ -145,7 +176,7 @@ func (r ProjectRepository) Create(ctx context.Context, prj project.Project) (pro } } - transformedProj, err := projectModel.transformToProject() + transformedProj, err := result.transformToProject() if err != nil { return project.Project{}, fmt.Errorf("%w: %w", errParse, err) } @@ -353,16 +384,24 @@ func (r ProjectRepository) Delete(ctx context.Context, id string) error { goqu.Ex{ "id": id, }, - ).ToSQL() + ).Returning(projectReturning()...).ToSQL() if err != nil { return fmt.Errorf("%w: %s", errQuery, err) } - if err = r.dbc.WithTimeout(ctx, TABLE_PROJECTS, "Delete", func(ctx context.Context) error { - if _, err = r.dbc.DB.ExecContext(ctx, query, params...); err != nil { + if err = r.dbc.WithTxn(ctx, sql.TxOptions{}, func(tx *sqlx.Tx) error { + var result projectWithOrgName + err := r.dbc.WithTimeout(ctx, TABLE_PROJECTS, "Delete", func(ctx context.Context) error { + return tx.QueryRowxContext(ctx, query, params...).StructScan(&result) + }) + if errors.Is(err, sql.ErrNoRows) { + // already gone: nothing deleted, nothing to audit + return nil + } + if err != nil { return err } - return nil + return InsertAuditRecordInTx(ctx, tx, buildProjectAuditRecord(ctx, auditrecord.ProjectDeletedEvent, result, time.Now().UTC())) }); err != nil { err = checkPostgresError(err) switch { diff --git a/internal/store/postgres/project_repository_test.go b/internal/store/postgres/project_repository_test.go index 87a4e2cc5..283aaeecb 100644 --- a/internal/store/postgres/project_repository_test.go +++ b/internal/store/postgres/project_repository_test.go @@ -15,6 +15,7 @@ import ( "github.com/raystack/frontier/core/relation" "github.com/raystack/frontier/core/user" "github.com/raystack/frontier/internal/store/postgres" + pkgAuditRecord "github.com/raystack/frontier/pkg/auditrecord" "github.com/raystack/frontier/pkg/db" "github.com/stretchr/testify/suite" ) @@ -109,10 +110,21 @@ func (s *ProjectRepositoryTestSuite) cleanup() error { fmt.Sprintf("TRUNCATE TABLE %s RESTART IDENTITY CASCADE", postgres.TABLE_RELATIONS), fmt.Sprintf("TRUNCATE TABLE %s RESTART IDENTITY CASCADE", postgres.TABLE_ROLES), fmt.Sprintf("TRUNCATE TABLE %s RESTART IDENTITY CASCADE", postgres.TABLE_NAMESPACES), + fmt.Sprintf("TRUNCATE TABLE %s RESTART IDENTITY CASCADE", postgres.TABLE_AUDITRECORDS), } return execQueries(context.TODO(), s.client, queries) } +// auditCount counts audit records of event on the project, filed under its org +func (s *ProjectRepositoryTestSuite) auditCount(event pkgAuditRecord.Event, projectID string) int { + var n int + err := s.client.QueryRowxContext(s.ctx, fmt.Sprintf( + "SELECT count(*) FROM %s WHERE event = $1 AND target_id = $2 AND resource_type = $3", postgres.TABLE_AUDITRECORDS), + event.String(), projectID, pkgAuditRecord.OrganizationType.String()).Scan(&n) + s.Require().NoError(err) + return n +} + func (s *ProjectRepositoryTestSuite) TestGetByID() { type testCase struct { Description string @@ -289,10 +301,25 @@ func (s *ProjectRepositoryTestSuite) TestCreate() { if !cmp.Equal(got, tc.ExpectedProject, cmpopts.IgnoreFields(project.Project{}, "ID", "Organization", "Metadata", "CreatedAt", "UpdatedAt")) { s.T().Fatalf("got result %+v, expected was %+v", got, tc.ExpectedProject) } + if tc.ErrString == "" { + s.Equal(1, s.auditCount(pkgAuditRecord.ProjectCreatedEvent, got.ID)) + } }) } } +func (s *ProjectRepositoryTestSuite) TestDelete() { + s.Run("should delete a project and audit it", func() { + s.Require().NoError(s.repository.Delete(s.ctx, s.projects[1].ID)) + s.Equal(1, s.auditCount(pkgAuditRecord.ProjectDeletedEvent, s.projects[1].ID)) + }) + s.Run("should skip a project that does not exist without auditing", func() { + id := uuid.NewString() + s.Require().NoError(s.repository.Delete(s.ctx, id)) + s.Equal(0, s.auditCount(pkgAuditRecord.ProjectDeletedEvent, id)) + }) +} + func (s *ProjectRepositoryTestSuite) TestList() { type testCase struct { Description string diff --git a/pkg/auditrecord/consts.go b/pkg/auditrecord/consts.go index 40bdebec1..1ca4c8456 100644 --- a/pkg/auditrecord/consts.go +++ b/pkg/auditrecord/consts.go @@ -45,6 +45,10 @@ const ( // Domain Events DomainDeletedEvent Event = "domain.deleted" + // Project Events + ProjectCreatedEvent Event = "project.created" + ProjectDeletedEvent Event = "project.deleted" + // Project Member Events ProjectMemberRoleChangedEvent Event = "project.member_role_changed" ProjectMemberRemovedEvent Event = "project.member_removed"