Files
LK_API_Temp/internal/store/users.go
T

634 lines
14 KiB
Go

package store
import (
"context"
"crypto/sha256"
"database/sql"
"encoding/hex"
"time"
"golang.org/x/crypto/bcrypt"
sq "github.com/Masterminds/squirrel"
)
const (
adminLevel = 3
modLevel = 2
userLevel = 1
)
type User struct {
ID int64 `json:"id"`
FirstName string `json:"first_name"`
LastName string `json:"last_name"`
Username string `json:"username"`
Email string `json:"email"`
Password password `json:"-"`
IsActive bool `json:"is_active"`
RoleID int64 `json:"role_id"`
Role Role `json:"role"`
CreatedAt string `json:"created_at"`
UpdatedAt string `json:"updated_at"`
}
type password struct {
text *string `validate:"required,password"`
hash []byte
}
var psql = sq.StatementBuilder.PlaceholderFormat(sq.Dollar)
func (p *password) Set(password string) error {
hash, err := bcrypt.GenerateFromPassword([]byte(password), bcrypt.DefaultCost)
if err != nil {
return err
}
p.text = &password
p.hash = hash
return nil
}
type UserStore struct {
db *sql.DB
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
}
func NewUserStore(db *sql.DB) (*UserStore, error) {
store := &UserStore{
db: db,
}
var err error
store.getByIDStmt, 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, roles.*
FROM users u
JOIN roles ON u.role_id = roles.id
WHERE u.id = $1 AND u.is_active = true
`)
if err != nil {
return nil, err
}
store.getByEmailStmt, 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, roles.*
FROM users u
JOIN roles ON u.role_id = roles.id
WHERE u.email = $1 AND u.is_active = true
`)
if err != nil {
return nil, err
}
store.getByInvitationStmt, err = db.Prepare(`
SELECT u.id, u.first_name, u.last_name, u.username, u.email, u.created_at, u.updated_at, is_active
FROM users u
JOIN invitations i ON u.id = i.user_id
WHERE i.token = $1 AND i.expiry > $2
`)
if err != nil {
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, is_active = $5, updated_at = $6
WHERE id = $7
`)
if err != nil {
return nil, err
}
store.deleteStmt, err = db.Prepare(`
DELETE FROM users WHERE id = $1
`)
if err != nil {
return nil, err
}
store.invitationDeleteStmt, err = db.Prepare(`
DELETE FROM invitations WHERE user_id = $1
`)
if err != nil {
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)
`)
if err != nil {
return nil, err
}
store.passwordResetRequestCreateStmt, err = db.Prepare(`
INSERT INTO password_reset_requests (user_id, token, expiry) VALUES ($1, $2, $3)
`)
if err != nil {
return nil, err
}
return store, nil
}
func (s *UserStore) Close() error {
var errs []error
if err := s.getByIDStmt.Close(); err != nil {
errs = append(errs, err)
}
if err := s.getByEmailStmt.Close(); err != nil {
errs = append(errs, err)
}
if err := s.getByInvitationStmt.Close(); err != nil {
errs = append(errs, err)
}
if err := s.updateStmt.Close(); err != nil {
errs = append(errs, err)
}
if err := s.deleteStmt.Close(); err != nil {
errs = append(errs, err)
}
if err := s.invitationDeleteStmt.Close(); err != nil {
errs = append(errs, err)
}
if len(errs) > 0 {
return errs[0]
}
return nil
}
func (s *UserStore) Create(ctx context.Context, tx *sql.Tx, user *User) error {
role := user.Role.Name
if role == "" {
role = "user"
}
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()
if err != nil {
return err
}
ctx, cancel := context.WithTimeout(ctx, QueryTimeout)
defer cancel()
err = tx.QueryRowContext(
ctx,
query,
args...,
).Scan(&user.ID, &user.CreatedAt, &user.UpdatedAt)
if err != nil {
switch {
case err.Error() == `pq: duplicate key value violates unique constraint "users_username_key"`:
return ErrDuplicateUsername
case err.Error() == `pq: duplicate key value violates unique constraint "users_email_key"`:
return ErrDuplicateEmail
default:
return err
}
}
return nil
}
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) {
ctx, cancel := context.WithTimeout(ctx, QueryTimeout)
defer cancel()
user := &User{}
err := s.getByIDStmt.QueryRowContext(ctx, id).Scan(
&user.ID,
&user.FirstName,
&user.LastName,
&user.Username,
&user.Email,
&user.Password.hash,
&user.IsActive,
&user.CreatedAt,
&user.UpdatedAt,
&user.RoleID,
&user.Role.Name,
&user.Role.Level,
&user.Role.Description,
)
if err != nil {
switch err {
case sql.ErrNoRows:
return nil, ErrNotFound
default:
return nil, err
}
}
return user, nil
}
func (s *UserStore) GetByEmail(ctx context.Context, email string) (*User, error) {
ctx, cancel := context.WithTimeout(ctx, QueryTimeout)
defer cancel()
user := &User{}
err := s.getByEmailStmt.QueryRowContext(ctx, email).Scan(
&user.ID,
&user.FirstName,
&user.LastName,
&user.Username,
&user.Email,
&user.Password.hash,
&user.IsActive,
&user.CreatedAt,
&user.UpdatedAt,
&user.RoleID,
&user.Role.Name,
&user.Role.Level,
&user.Role.Description,
)
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, user *User) error {
ctx, cancel := context.WithTimeout(ctx, QueryTimeout)
defer cancel()
_, err := s.updateStmt.ExecContext(
ctx,
user.FirstName,
user.LastName,
user.Username,
user.Email,
user.Password.hash,
user.ID,
)
if err != nil {
return err
}
return nil
}
func (s *UserStore) CreateAndInvite(ctx context.Context, user *User, token string, invitationExp time.Duration) error {
return withTx(s.db, ctx, func(tx *sql.Tx) error {
if err := s.Create(ctx, tx, user); err != nil {
return err
}
if err := s.createUserInvitation(ctx, tx, user.ID, invitationExp, token); err != nil {
return err
}
return nil
})
}
func (s *UserStore) PasswordResetRequest(ctx context.Context, userID int64, token string, expiration time.Duration) error {
return withTx(s.db, ctx, func(tx *sql.Tx) error {
if err := s.createPasswordResetRequest(ctx, tx, userID, token, expiration); err != nil {
return err
}
return nil
})
}
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
user, err := s.getUserFromInvitation(ctx, tx, token)
if err != nil {
return err
}
// Update the user's status to active
user.IsActive = true
if err := s.update(ctx, tx, user); err != nil {
return err
}
// Delete the token from the invitations table
if err := s.deleteUserInvitation(ctx, tx, user.ID); err != nil {
return err
}
return nil
})
}
func (s *UserStore) Delete(ctx context.Context, id int64) error {
return withTx(s.db, ctx, func(tx *sql.Tx) error {
if err := s.delete(ctx, tx, id); err != nil {
return err
}
if err := s.deleteUserInvitation(ctx, tx, id); err != nil {
return err
}
return nil
})
}
// * Soft delete user
func (s *UserStore) DeleteUser(ctx context.Context, id int64) error {
query, _, err := psql.Update("users").
Set("is_active", false).
Where(sq.Eq{"id": id}).
ToSql()
if err != nil {
return err
}
ctx, cancel := context.WithTimeout(ctx, QueryTimeout)
defer cancel()
_, err = s.db.ExecContext(ctx, query, id)
if err != nil {
return err
}
return nil
}
func (p *password) Compare(password string) error {
return bcrypt.CompareHashAndPassword(p.hash, []byte(password))
}
func (s *UserStore) getUserFromInvitation(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.getByInvitationStmt)
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) createUserInvitation(ctx context.Context, tx *sql.Tx, userID int64, exp time.Duration, token string) error {
ctx, cancel := context.WithTimeout(ctx, QueryTimeout)
defer cancel()
stmt := tx.Stmt(s.inviteCreateStmt)
defer stmt.Close()
_, err := stmt.ExecContext(ctx, userID, token, time.Now().Add(exp))
if err != nil {
return err
}
return nil
}
func (s *UserStore) createPasswordResetRequest(ctx context.Context, tx *sql.Tx, userID int64, token string, expiry time.Duration) error {
ctx, cancel := context.WithTimeout(ctx, QueryTimeout)
defer cancel()
stmt := tx.Stmt(s.passwordResetRequestCreateStmt)
defer stmt.Close()
_, err := stmt.ExecContext(ctx, userID, token, time.Now().Add(expiry))
if err != nil {
return err
}
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()
stmt := tx.Stmt(s.updateStmt)
defer stmt.Close()
_, err := stmt.ExecContext(
ctx,
user.FirstName,
user.LastName,
user.Username,
user.Email,
user.IsActive,
user.ID,
)
if err != nil {
return err
}
return nil
}
func (s *UserStore) deleteUserInvitation(ctx context.Context, tx *sql.Tx, userID int64) error {
ctx, cancel := context.WithTimeout(ctx, QueryTimeout)
defer cancel()
stmt := tx.Stmt(s.invitationDeleteStmt)
defer stmt.Close()
_, err := stmt.ExecContext(ctx, userID)
if err != nil {
return err
}
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()
stmt := tx.Stmt(s.deleteStmt)
defer stmt.Close()
_, err := stmt.ExecContext(ctx, id)
if err != nil {
return err
}
return nil
}