package store import ( "context" "database/sql" "errors" "fmt" "strings" "gitea.stevedudenhoeffer.com/steve/pansy/internal/domain" ) // userColumns lists the users columns in the fixed order scanUser expects. const userColumns = `id, email, display_name, password_hash, oidc_issuer, oidc_subject, is_admin, version, created_at, updated_at` // scanner is satisfied by both *sql.Row and *sql.Rows, so scanUser works for // single-row and multi-row queries alike. type scanner interface { Scan(dest ...any) error } // scanUser reads one users row. is_admin is stored as INTEGER 0/1 (the driver // returns it as int64, which does not convert straight to bool), so it is read // into an int64 and mapped. func scanUser(s scanner) (*domain.User, error) { var ( u domain.User isAdmin int64 ) if err := s.Scan( &u.ID, &u.Email, &u.DisplayName, &u.PasswordHash, &u.OIDCIssuer, &u.OIDCSubject, &isAdmin, &u.Version, &u.CreatedAt, &u.UpdatedAt, ); err != nil { return nil, err } u.IsAdmin = isAdmin != 0 return &u, nil } // CreateUser inserts a new user and returns the stored row (with generated id // and timestamps). PasswordHash/OIDCIssuer/OIDCSubject may be nil. // // Two invariants are enforced inside the single INSERT statement so concurrent // registrations can't violate them (SQLite serializes writers, and the COUNT // subqueries see the latest committed state): // - is_admin is set iff this is the first user — no read-then-write window in // which two "first" registrations both become admin. // - the row is inserted only when the table is empty (bootstrap) or allowSignup // is true; otherwise zero rows are affected and ErrRegistrationClosed is // returned. This is the authoritative registration gate. // // A duplicate email trips the UNIQUE index and maps to ErrEmailTaken. func (d *DB) CreateUser(ctx context.Context, u *domain.User, allowSignup bool) (*domain.User, error) { res, err := d.sql.ExecContext(ctx, `INSERT INTO users (email, display_name, password_hash, oidc_issuer, oidc_subject, is_admin) SELECT ?, ?, ?, ?, ?, (SELECT count(*) FROM users) = 0 WHERE (SELECT count(*) FROM users) = 0 OR ?`, u.Email, u.DisplayName, u.PasswordHash, u.OIDCIssuer, u.OIDCSubject, boolToInt(allowSignup), ) if err != nil { if isUniqueViolation(err) { return nil, domain.ErrEmailTaken } return nil, fmt.Errorf("store: insert user: %w", err) } n, err := res.RowsAffected() if err != nil { return nil, fmt.Errorf("store: user insert rows: %w", err) } if n == 0 { return nil, domain.ErrRegistrationClosed } id, err := res.LastInsertId() if err != nil { return nil, fmt.Errorf("store: user insert id: %w", err) } return d.GetUserByID(ctx, id) } // GetUserByID returns the user with the given id, or domain.ErrNotFound. func (d *DB) GetUserByID(ctx context.Context, id int64) (*domain.User, error) { u, err := scanUser(d.sql.QueryRowContext(ctx, `SELECT `+userColumns+` FROM users WHERE id = ?`, id)) if errors.Is(err, sql.ErrNoRows) { return nil, domain.ErrNotFound } if err != nil { return nil, fmt.Errorf("store: get user by id: %w", err) } return u, nil } // GetUserByEmail returns the user with the given email (case-insensitive via the // column's NOCASE collation), or domain.ErrNotFound. func (d *DB) GetUserByEmail(ctx context.Context, email string) (*domain.User, error) { u, err := scanUser(d.sql.QueryRowContext(ctx, `SELECT `+userColumns+` FROM users WHERE email = ?`, email)) if errors.Is(err, sql.ErrNoRows) { return nil, domain.ErrNotFound } if err != nil { return nil, fmt.Errorf("store: get user by email: %w", err) } return u, nil } // CountUsers returns the number of user rows. Used to decide first-user-is-admin // and to allow bootstrap registration when signup is otherwise closed. func (d *DB) CountUsers(ctx context.Context) (int, error) { var n int if err := d.sql.QueryRowContext(ctx, `SELECT count(*) FROM users`).Scan(&n); err != nil { return 0, fmt.Errorf("store: count users: %w", err) } return n, nil } // boolToInt maps a Go bool to the 0/1 SQLite stores for INTEGER "boolean" columns. func boolToInt(b bool) int { if b { return 1 } return 0 } // isUniqueViolation reports whether err is a SQLite UNIQUE-constraint failure. // modernc surfaces these in the error text; the message is stable across SQLite // versions ("UNIQUE constraint failed: ."). func isUniqueViolation(err error) bool { return err != nil && strings.Contains(err.Error(), "UNIQUE constraint failed") }