diff --git a/internal/store/postgres/user_repository.go b/internal/store/postgres/user_repository.go index 70dacb512..2b402ac54 100644 --- a/internal/store/postgres/user_repository.go +++ b/internal/store/postgres/user_repository.go @@ -7,6 +7,7 @@ import ( "fmt" "slices" "strings" + "time" "github.com/raystack/frontier/pkg/utils" "github.com/raystack/salt/rql" @@ -14,10 +15,14 @@ import ( "github.com/pkg/errors" "github.com/doug-martin/goqu/v9" + "github.com/google/uuid" "github.com/jmoiron/sqlx" "github.com/raystack/frontier/core/consent" "github.com/raystack/frontier/core/user" + "github.com/raystack/frontier/internal/bootstrap/schema" + "github.com/raystack/frontier/pkg/auditrecord" "github.com/raystack/frontier/pkg/db" + "github.com/raystack/frontier/pkg/metadata" ) type UserRepository struct { @@ -186,6 +191,21 @@ func (r UserRepository) createWithTx(ctx context.Context, tx *sqlx.Tx, usr user. } } + record := buildUserAuditRecord(ctx, auditrecord.UserCreatedEvent, userModel, userModel.CreatedAt) + // signup runs unauthenticated, so with no caller the user created themselves + if record.ActorType == auditrecord.SystemActor { + record.ActorID, _ = uuid.Parse(userModel.ID) + record.ActorType = schema.UserPrincipal + record.ActorName = userModel.Title.String + if record.ActorName == "" { + record.ActorName = userModel.Email + } + record.ActorTitle = userModel.Title.String + } + if err := InsertAuditRecordInTx(ctx, tx, record); err != nil { + return user.User{}, err + } + transformedUser, err := userModel.transformToUser() if err != nil { return user.User{}, fmt.Errorf("%w: %w", errParse, err) @@ -193,47 +213,28 @@ func (r UserRepository) createWithTx(ctx context.Context, tx *sqlx.Tx, usr user. return transformedUser, nil } +func buildUserAuditRecord(ctx context.Context, event auditrecord.Event, u User, occurredAt time.Time) AuditRecord { + return BuildAuditRecord(ctx, event, + AuditResource{ID: schema.PlatformID, Type: auditrecord.PlatformType, Name: schema.PlatformID}, + &AuditTarget{ID: u.ID, Type: auditrecord.UserType, Name: u.Name, Metadata: metadata.Metadata{"email": u.Email}}, + schema.PlatformOrgID.String(), nil, occurredAt) +} + func (r UserRepository) Create(ctx context.Context, usr user.User) (user.User, error) { - if strings.TrimSpace(usr.Email) == "" || strings.TrimSpace(usr.Name) == "" { + var created user.User + err := r.dbc.WithTxn(ctx, sql.TxOptions{}, func(tx *sqlx.Tx) (err error) { + created, err = r.createWithTx(ctx, tx, usr) + return err + }) + switch { + case errors.Is(err, user.ErrConflict): + return user.User{}, user.ErrConflict + case errors.Is(err, user.ErrInvalidDetails): return user.User{}, user.ErrInvalidDetails - } - - createQuery, params, err := buildUserInsertQuery(usr) - if err != nil { - return user.User{}, fmt.Errorf("%w: %w", errQuery, err) - } - - tx, err := r.dbc.BeginTxx(ctx, nil) - if err != nil { + case err != nil: return user.User{}, err } - - var userModel User - if err = r.dbc.WithTimeout(ctx, TABLE_USERS, "Create", func(ctx context.Context) error { - return tx.QueryRowxContext(ctx, createQuery, params...). - StructScan(&userModel) - }); err != nil { - err = checkPostgresError(err) - switch { - case errors.Is(err, ErrDuplicateKey): - return user.User{}, user.ErrConflict - default: - if err := tx.Rollback(); err != nil { - return user.User{}, err - } - return user.User{}, err - } - } - - if err = tx.Commit(); err != nil { - return user.User{}, err - } - - transformedUser, err := userModel.transformToUser() - if err != nil { - return user.User{}, fmt.Errorf("%w: %w", errParse, err) - } - return transformedUser, nil + return created, nil } func (r UserRepository) List(ctx context.Context, flt user.Filter) ([]user.User, error) { @@ -564,9 +565,14 @@ func (r UserRepository) Delete(ctx context.Context, id string) error { return fmt.Errorf("%w: %s", errQuery, err) } - var userModel User - if err = r.dbc.WithTimeout(ctx, TABLE_USERS, "Delete", func(ctx context.Context) error { - return r.dbc.QueryRowxContext(ctx, query, params...).StructScan(&userModel) + if err = r.dbc.WithTxn(ctx, sql.TxOptions{}, func(tx *sqlx.Tx) error { + var userModel User + if err := r.dbc.WithTimeout(ctx, TABLE_USERS, "Delete", func(ctx context.Context) error { + return tx.QueryRowxContext(ctx, query, params...).StructScan(&userModel) + }); err != nil { + return err + } + return InsertAuditRecordInTx(ctx, tx, buildUserAuditRecord(ctx, auditrecord.UserDeletedEvent, userModel, time.Now().UTC())) }); err != nil { err = checkPostgresError(err) switch { diff --git a/internal/store/postgres/user_repository_test.go b/internal/store/postgres/user_repository_test.go index f072668d2..26bec102c 100644 --- a/internal/store/postgres/user_repository_test.go +++ b/internal/store/postgres/user_repository_test.go @@ -11,6 +11,7 @@ import ( "github.com/google/uuid" "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/raystack/frontier/pkg/metadata" "github.com/raystack/salt/rql" @@ -62,10 +63,21 @@ func (s *UserRepositoryTestSuite) TearDownTest() { func (s *UserRepositoryTestSuite) cleanup() error { queries := []string{ fmt.Sprintf("TRUNCATE TABLE %s RESTART IDENTITY CASCADE", postgres.TABLE_USERS), + fmt.Sprintf("TRUNCATE TABLE %s RESTART IDENTITY CASCADE", postgres.TABLE_AUDITRECORDS), } return execQueries(context.TODO(), s.client, queries) } +// assertAudited checks exactly one audit record of event exists for the user, written by actorID +func (s *UserRepositoryTestSuite) assertAudited(event, userID, actorID string) { + var n int + err := s.client.QueryRowxContext(s.ctx, fmt.Sprintf( + "SELECT count(*) FROM %s WHERE event = $1 AND target_id = $2 AND actor_id = $3", postgres.TABLE_AUDITRECORDS), + event, userID, actorID).Scan(&n) + s.Require().NoError(err) + s.Equal(1, n, "audit records for %s on %s", event, userID) +} + func (s *UserRepositoryTestSuite) TestGetByID() { type testCase struct { Description string @@ -211,6 +223,10 @@ func (s *UserRepositoryTestSuite) TestCreate() { if tc.ExpectedEmail != "" && (got.Email != tc.ExpectedEmail) { s.T().Fatalf("got result %+v, expected was %+v", got.ID, tc.ExpectedEmail) } + if tc.ErrString == "" { + // no caller in the context, so the user is their own actor + s.assertAudited(pkgAuditRecord.UserCreatedEvent.String(), got.ID, got.ID) + } }) } } @@ -483,6 +499,9 @@ func (s *UserRepositoryTestSuite) TestDelete() { s.T().Fatalf("got error %s, expected was %s", err.Error(), tc.Err) } } + if tc.Err == nil { + s.assertAudited(pkgAuditRecord.UserDeletedEvent.String(), tc.User, uuid.Nil.String()) + } }) } } diff --git a/pkg/auditrecord/consts.go b/pkg/auditrecord/consts.go index db630dabb..61e103d2e 100644 --- a/pkg/auditrecord/consts.go +++ b/pkg/auditrecord/consts.go @@ -77,6 +77,8 @@ const ( ResourceCreatedEvent Event = "resource.created" // User Events + UserCreatedEvent Event = "user.created" + UserDeletedEvent Event = "user.deleted" UserConsentGrantedEvent Event = "user.consent_granted" // PAT Events