refactor: updated store to use squirrel for complex querries
This commit is contained in:
+149
-11
@@ -12,6 +12,12 @@ import (
|
||||
sq "github.com/Masterminds/squirrel"
|
||||
)
|
||||
|
||||
const (
|
||||
adminLevel = 3
|
||||
modLevel = 2
|
||||
userLevel = 1
|
||||
)
|
||||
|
||||
type User struct {
|
||||
ID int64 `json:"id"`
|
||||
FirstName string `json:"first_name"`
|
||||
@@ -50,10 +56,12 @@ type UserStore struct {
|
||||
getByIDStmt *sql.Stmt
|
||||
getByEmailStmt *sql.Stmt
|
||||
getByInvitationStmt *sql.Stmt
|
||||
getUserFromPasswordTokenStmt *sql.Stmt
|
||||
updateStmt *sql.Stmt
|
||||
deleteStmt *sql.Stmt
|
||||
inviteCreateStmt *sql.Stmt
|
||||
invitationDeleteStmt *sql.Stmt
|
||||
deletePasswordResetRequestStmt *sql.Stmt
|
||||
passwordResetRequestCreateStmt *sql.Stmt
|
||||
}
|
||||
|
||||
@@ -94,10 +102,20 @@ func NewUserStore(db *sql.DB) (*UserStore, error) {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
store.getUserFromPasswordTokenStmt, err = db.Prepare(`
|
||||
SELECT u.id, u.first_name, u.last_name, u.username, u.email, u.password, u.is_active, u.created_at, u.updated_at
|
||||
FROM users u
|
||||
JOIN password_reset_requests pr ON u.id = pr.user_id
|
||||
WHERE u.id = $1 AND pr.expiry > $2
|
||||
`)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
store.updateStmt, err = db.Prepare(`
|
||||
UPDATE users
|
||||
SET first_name = $1, last_name = $2, username = $3, email = $4, password = $5
|
||||
WHERE id = $6
|
||||
SET first_name = $1, last_name = $2, username = $3, email = $4, is_active = $5, updated_at = $6
|
||||
WHERE id = $7
|
||||
`)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -117,6 +135,13 @@ func NewUserStore(db *sql.DB) (*UserStore, error) {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
store.deletePasswordResetRequestStmt, err = db.Prepare(`
|
||||
DELETE FROM password_reset_requests WHERE user_id = $1
|
||||
`)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
store.inviteCreateStmt, err = db.Prepare(`
|
||||
INSERT INTO invitations (user_id, token, expiry) VALUES ($1, $2, $3)
|
||||
`)
|
||||
@@ -169,7 +194,7 @@ func (s *UserStore) Create(ctx context.Context, tx *sql.Tx, user *User) error {
|
||||
role = "user"
|
||||
}
|
||||
|
||||
query, _, err := psql.Insert("users").
|
||||
query, args, err := psql.Insert("users").
|
||||
Columns("first_name", "last_name", "username", "email", "password", "role_id").
|
||||
Values(user.FirstName, user.LastName, user.Username, user.Email, user.Password.hash, sq.Expr("(SELECT id FROM roles WHERE name = ?)", role)).
|
||||
Suffix("RETURNING id, created_at, updated_at").ToSql()
|
||||
@@ -183,12 +208,7 @@ func (s *UserStore) Create(ctx context.Context, tx *sql.Tx, user *User) error {
|
||||
err = tx.QueryRowContext(
|
||||
ctx,
|
||||
query,
|
||||
user.FirstName,
|
||||
user.LastName,
|
||||
user.Username,
|
||||
user.Email,
|
||||
user.Password.hash,
|
||||
role,
|
||||
args...,
|
||||
).Scan(&user.ID, &user.CreatedAt, &user.UpdatedAt)
|
||||
if err != nil {
|
||||
switch {
|
||||
@@ -204,8 +224,58 @@ func (s *UserStore) Create(ctx context.Context, tx *sql.Tx, user *User) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *UserStore) GetAll(ctx context.Context, isAdmin bool) ([]User, error) {
|
||||
func (s *UserStore) GetAll(ctx context.Context, userID int64, roleLevel int64) ([]*User, error) {
|
||||
|
||||
var whereConditions []sq.Sqlizer
|
||||
|
||||
whereConditions = append(whereConditions, sq.Eq{"u.is_active": true})
|
||||
|
||||
if roleLevel == modLevel {
|
||||
whereConditions = append(whereConditions, sq.Or{
|
||||
sq.Lt{"roles.level": roleLevel},
|
||||
sq.Eq{"u.id": userID},
|
||||
})
|
||||
}
|
||||
|
||||
query, args, err := psql.Select("u.id", "u.first_name", "u.last_name", "u.username", "u.email", "u.is_active", "u.created_at", "u.updated_at", "roles.*").
|
||||
From("users u").Join("roles ON u.role_id = roles.id").Where(sq.And(
|
||||
whereConditions,
|
||||
)).OrderBy("u.created_at DESC").
|
||||
ToSql()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
rows, err := s.db.QueryContext(ctx, query, args...)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
users := make([]*User, 0)
|
||||
for rows.Next() {
|
||||
user := &User{}
|
||||
err := rows.Scan(
|
||||
&user.ID,
|
||||
&user.FirstName,
|
||||
&user.LastName,
|
||||
&user.Username,
|
||||
&user.Email,
|
||||
&user.IsActive,
|
||||
&user.CreatedAt,
|
||||
&user.UpdatedAt,
|
||||
&user.RoleID,
|
||||
&user.Role.Name,
|
||||
&user.Role.Level,
|
||||
&user.Role.Description,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
users = append(users, user)
|
||||
}
|
||||
|
||||
return users, nil
|
||||
}
|
||||
|
||||
func (s *UserStore) GetByID(ctx context.Context, id int64) (*User, error) {
|
||||
@@ -252,9 +322,9 @@ func (s *UserStore) GetByEmail(ctx context.Context, email string) (*User, error)
|
||||
&user.Username,
|
||||
&user.Email,
|
||||
&user.Password.hash,
|
||||
&user.IsActive,
|
||||
&user.CreatedAt,
|
||||
&user.UpdatedAt,
|
||||
&user.IsActive,
|
||||
&user.RoleID,
|
||||
&user.Role.Name,
|
||||
&user.Role.Level,
|
||||
@@ -316,6 +386,26 @@ func (s *UserStore) PasswordResetRequest(ctx context.Context, userID int64, toke
|
||||
})
|
||||
}
|
||||
|
||||
func (s *UserStore) ResetPassword(ctx context.Context, token string, newPassword string) error {
|
||||
return withTx(s.db, ctx, func(tx *sql.Tx) error {
|
||||
user, err := s.getUserFromPaswordToken(ctx, tx, token)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
user.Password.Set(newPassword)
|
||||
if err := s.update(ctx, tx, user); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if err := s.deletePasswordResetRequest(ctx, tx, user.ID); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
return nil
|
||||
})
|
||||
}
|
||||
|
||||
func (s *UserStore) Activate(ctx context.Context, token string) error {
|
||||
return withTx(s.db, ctx, func(tx *sql.Tx) error {
|
||||
// Find the user that this token belongs to
|
||||
@@ -441,6 +531,39 @@ func (s *UserStore) createPasswordResetRequest(ctx context.Context, tx *sql.Tx,
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *UserStore) getUserFromPaswordToken(ctx context.Context, tx *sql.Tx, token string) (*User, error) {
|
||||
hashed := sha256.Sum256([]byte(token))
|
||||
hashedToken := hex.EncodeToString(hashed[:])
|
||||
|
||||
ctx, cancel := context.WithTimeout(ctx, QueryTimeout)
|
||||
defer cancel()
|
||||
|
||||
stmt := tx.Stmt(s.getUserFromPasswordTokenStmt)
|
||||
defer stmt.Close()
|
||||
|
||||
user := &User{}
|
||||
err := stmt.QueryRowContext(ctx, hashedToken, time.Now()).Scan(
|
||||
&user.ID,
|
||||
&user.FirstName,
|
||||
&user.LastName,
|
||||
&user.Username,
|
||||
&user.Email,
|
||||
&user.CreatedAt,
|
||||
&user.UpdatedAt,
|
||||
&user.IsActive,
|
||||
)
|
||||
if err != nil {
|
||||
switch err {
|
||||
case sql.ErrNoRows:
|
||||
return nil, ErrNotFound
|
||||
default:
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
|
||||
return user, nil
|
||||
}
|
||||
|
||||
func (s *UserStore) update(ctx context.Context, tx *sql.Tx, user *User) error {
|
||||
ctx, cancel := context.WithTimeout(ctx, QueryTimeout)
|
||||
defer cancel()
|
||||
@@ -479,6 +602,21 @@ func (s *UserStore) deleteUserInvitation(ctx context.Context, tx *sql.Tx, userID
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *UserStore) deletePasswordResetRequest(ctx context.Context, tx *sql.Tx, userID int64) error {
|
||||
ctx, cancel := context.WithTimeout(ctx, QueryTimeout)
|
||||
defer cancel()
|
||||
|
||||
stmt := tx.Stmt(s.deletePasswordResetRequestStmt)
|
||||
defer stmt.Close()
|
||||
|
||||
_, err := stmt.ExecContext(ctx, userID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *UserStore) delete(ctx context.Context, tx *sql.Tx, id int64) error {
|
||||
ctx, cancel := context.WithTimeout(ctx, QueryTimeout)
|
||||
defer cancel()
|
||||
|
||||
Reference in New Issue
Block a user