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
86 changes: 46 additions & 40 deletions internal/store/postgres/user_repository.go
Original file line number Diff line number Diff line change
Expand Up @@ -7,17 +7,22 @@ import (
"fmt"
"slices"
"strings"
"time"

"github.com/raystack/frontier/pkg/utils"
"github.com/raystack/salt/rql"

"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 {
Expand Down Expand Up @@ -186,54 +191,50 @@ 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)
}
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) {
Expand Down Expand Up @@ -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 {
Expand Down
19 changes: 19 additions & 0 deletions internal/store/postgres/user_repository_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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)
}
})
}
}
Expand Down Expand Up @@ -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())
}
})
}
}
Expand Down
2 changes: 2 additions & 0 deletions pkg/auditrecord/consts.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
Loading