210 lines
5.1 KiB
Go
210 lines
5.1 KiB
Go
package users
|
|
|
|
import (
|
|
"context"
|
|
"crypto/rand"
|
|
"crypto/sha256"
|
|
"encoding/base64"
|
|
"encoding/hex"
|
|
"errors"
|
|
"time"
|
|
|
|
"github.com/jackc/pgx/v5"
|
|
)
|
|
|
|
var ErrInvitationNotFound = errors.New("invitation not found")
|
|
|
|
type Invitation struct {
|
|
ID string `json:"id"`
|
|
Email string `json:"email"`
|
|
Role string `json:"role"`
|
|
ClientID *string `json:"client_id"`
|
|
ClientName *string `json:"client_name"`
|
|
ExpiresAt time.Time `json:"expires_at"`
|
|
AcceptedAt *time.Time `json:"accepted_at"`
|
|
CreatedBy *string `json:"created_by"`
|
|
CreatedAt time.Time `json:"created_at"`
|
|
InvitationURL string `json:"invitation_url,omitempty"`
|
|
}
|
|
|
|
type CreateInvitationInput struct {
|
|
Email string
|
|
Role string
|
|
ClientID string
|
|
TokenHash string
|
|
ExpiresAt time.Time
|
|
CreatedBy string
|
|
}
|
|
|
|
func NewInvitationToken() (string, string, error) {
|
|
bytes := make([]byte, 32)
|
|
if _, err := rand.Read(bytes); err != nil {
|
|
return "", "", err
|
|
}
|
|
|
|
token := base64.RawURLEncoding.EncodeToString(bytes)
|
|
return token, HashInvitationToken(token), nil
|
|
}
|
|
|
|
func HashInvitationToken(token string) string {
|
|
sum := sha256.Sum256([]byte(token))
|
|
return hex.EncodeToString(sum[:])
|
|
}
|
|
|
|
func (s Store) CreateInvitation(ctx context.Context, input CreateInvitationInput) (Invitation, error) {
|
|
row := s.db.QueryRow(ctx, `
|
|
INSERT INTO invitations (email, role, client_id, token_hash, expires_at, created_by)
|
|
VALUES (lower($1), $2, NULLIF($3, '')::uuid, $4, $5, NULLIF($6, '')::uuid)
|
|
RETURNING
|
|
invitations.id::text,
|
|
invitations.email,
|
|
invitations.role,
|
|
invitations.client_id::text,
|
|
(SELECT name FROM clients WHERE clients.id = invitations.client_id),
|
|
invitations.expires_at,
|
|
invitations.accepted_at,
|
|
invitations.created_by::text,
|
|
invitations.created_at
|
|
`, input.Email, input.Role, input.ClientID, input.TokenHash, input.ExpiresAt, input.CreatedBy)
|
|
|
|
return scanInvitation(row)
|
|
}
|
|
|
|
func (s Store) ListInvitations(ctx context.Context) ([]Invitation, error) {
|
|
rows, err := s.db.Query(ctx, `
|
|
SELECT
|
|
invitations.id::text,
|
|
invitations.email,
|
|
invitations.role,
|
|
invitations.client_id::text,
|
|
clients.name,
|
|
invitations.expires_at,
|
|
invitations.accepted_at,
|
|
invitations.created_by::text,
|
|
invitations.created_at
|
|
FROM invitations
|
|
LEFT JOIN clients ON clients.id = invitations.client_id
|
|
ORDER BY invitations.created_at DESC
|
|
`)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
defer rows.Close()
|
|
|
|
invitations := []Invitation{}
|
|
for rows.Next() {
|
|
invitation, err := scanInvitation(rows)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
invitations = append(invitations, invitation)
|
|
}
|
|
|
|
return invitations, rows.Err()
|
|
}
|
|
|
|
func (s Store) AcceptInvitation(ctx context.Context, tokenHash string, name string, passwordHash string) (User, error) {
|
|
tx, err := s.db.Begin(ctx)
|
|
if err != nil {
|
|
return User{}, err
|
|
}
|
|
defer tx.Rollback(ctx)
|
|
|
|
var invitation Invitation
|
|
var rawTokenHash string
|
|
err = tx.QueryRow(ctx, `
|
|
SELECT
|
|
invitations.id::text,
|
|
invitations.email,
|
|
invitations.role,
|
|
invitations.client_id::text,
|
|
(SELECT name FROM clients WHERE clients.id = invitations.client_id),
|
|
invitations.expires_at,
|
|
invitations.accepted_at,
|
|
invitations.created_by::text,
|
|
invitations.created_at,
|
|
invitations.token_hash
|
|
FROM invitations
|
|
WHERE invitations.token_hash = $1
|
|
AND invitations.accepted_at IS NULL
|
|
AND invitations.expires_at > now()
|
|
FOR UPDATE
|
|
`, tokenHash).Scan(
|
|
&invitation.ID,
|
|
&invitation.Email,
|
|
&invitation.Role,
|
|
&invitation.ClientID,
|
|
&invitation.ClientName,
|
|
&invitation.ExpiresAt,
|
|
&invitation.AcceptedAt,
|
|
&invitation.CreatedBy,
|
|
&invitation.CreatedAt,
|
|
&rawTokenHash,
|
|
)
|
|
if errors.Is(err, pgx.ErrNoRows) {
|
|
return User{}, ErrInvitationNotFound
|
|
}
|
|
if err != nil {
|
|
return User{}, err
|
|
}
|
|
|
|
user, err := scanUser(tx.QueryRow(ctx, `
|
|
INSERT INTO users (email, name, password_hash, role, status)
|
|
VALUES ($1, $2, $3, $4, 'active')
|
|
RETURNING id::text, email, name, password_hash, role, status, created_at, updated_at
|
|
`, invitation.Email, name, passwordHash, invitation.Role))
|
|
if err != nil {
|
|
return User{}, err
|
|
}
|
|
|
|
if invitation.ClientID != nil {
|
|
if _, err := tx.Exec(ctx, `
|
|
INSERT INTO client_users (client_id, user_id)
|
|
VALUES ($1::uuid, $2::uuid)
|
|
ON CONFLICT DO NOTHING
|
|
`, *invitation.ClientID, user.ID); err != nil {
|
|
return User{}, err
|
|
}
|
|
}
|
|
|
|
if _, err := tx.Exec(ctx, `
|
|
UPDATE invitations
|
|
SET accepted_at = now()
|
|
WHERE token_hash = $1
|
|
`, rawTokenHash); err != nil {
|
|
return User{}, err
|
|
}
|
|
|
|
if err := tx.Commit(ctx); err != nil {
|
|
return User{}, err
|
|
}
|
|
|
|
return user, nil
|
|
}
|
|
|
|
type invitationScanner interface {
|
|
Scan(dest ...any) error
|
|
}
|
|
|
|
func scanInvitation(row invitationScanner) (Invitation, error) {
|
|
var invitation Invitation
|
|
err := row.Scan(
|
|
&invitation.ID,
|
|
&invitation.Email,
|
|
&invitation.Role,
|
|
&invitation.ClientID,
|
|
&invitation.ClientName,
|
|
&invitation.ExpiresAt,
|
|
&invitation.AcceptedAt,
|
|
&invitation.CreatedBy,
|
|
&invitation.CreatedAt,
|
|
)
|
|
if errors.Is(err, pgx.ErrNoRows) {
|
|
return Invitation{}, ErrInvitationNotFound
|
|
}
|
|
if err != nil {
|
|
return Invitation{}, err
|
|
}
|
|
return invitation, nil
|
|
}
|