package store import ( "context" "database/sql" "errors" "fmt" "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 int 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. func (d *DB) CreateUser(ctx context.Context, u *domain.User) (*domain.User, error) { res, err := d.sql.ExecContext(ctx, `INSERT INTO users (email, display_name, password_hash, oidc_issuer, oidc_subject, is_admin) VALUES (?, ?, ?, ?, ?, ?)`, u.Email, u.DisplayName, u.PasswordHash, u.OIDCIssuer, u.OIDCSubject, boolToInt(u.IsAdmin), ) if err != nil { return nil, fmt.Errorf("store: insert user: %w", err) } 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 }