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 } // GetUserByOIDC returns the user with the given (issuer, subject) identity pair, // or domain.ErrNotFound. Both arguments must be non-empty. func (d *DB) GetUserByOIDC(ctx context.Context, issuer, subject string) (*domain.User, error) { u, err := scanUser(d.sql.QueryRowContext(ctx, `SELECT `+userColumns+` FROM users WHERE oidc_issuer = ? AND oidc_subject = ?`, issuer, subject)) if errors.Is(err, sql.ErrNoRows) { return nil, domain.ErrNotFound } if err != nil { return nil, fmt.Errorf("store: get user by oidc: %w", err) } return u, nil } // LinkOIDC stamps an OIDC identity onto an existing user (first OIDC login for a // pre-existing local account) and returns the updated row. // // The UPDATE only matches when the account carries no identity yet, or already // carries this exact one (idempotent) — it will NOT overwrite a different stored // identity, which would let anyone asserting the same email at a second IdP // hijack or lock out the account. A no-match (different identity, or the row is // gone) and a UNIQUE (issuer, subject) collision with another account both // surface as domain.ErrOIDCIdentityConflict. RETURNING folds the read-back into // the same statement. func (d *DB) LinkOIDC(ctx context.Context, userID int64, issuer, subject string) (*domain.User, error) { u, err := scanUser(d.sql.QueryRowContext(ctx, `UPDATE users SET oidc_issuer = ?, oidc_subject = ?, version = version + 1, updated_at = strftime('%Y-%m-%dT%H:%M:%SZ', 'now') WHERE id = ? AND (oidc_subject IS NULL OR (oidc_issuer = ? AND oidc_subject = ?)) RETURNING `+userColumns, issuer, subject, userID, issuer, subject, )) if errors.Is(err, sql.ErrNoRows) { // The account already has a different identity (or no longer exists). return nil, domain.ErrOIDCIdentityConflict } if err != nil { if isUniqueViolation(err) { return nil, domain.ErrOIDCIdentityConflict } return nil, fmt.Errorf("store: link oidc: %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") }