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"` EmailSent bool `json:"email_sent"` EmailError string `json:"email_error,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 }