diff options
Diffstat (limited to 'cmd/mirum-server/database.go')
| -rw-r--r-- | cmd/mirum-server/database.go | 130 | +36 −94 |
1 files changed, 36 insertions, 94 deletions
diff --git a/cmd/mirum-server/database.go b/cmd/mirum-server/database.go index 68356c3..27db982 100644 --- a/cmd/mirum-server/database.go +++ b/cmd/mirum-server/database.go @@ -1,4 +1,4 @@ -// SPDX-FileCopyrightText: 2026 Nikolay Govorov +// Copyright (c) 2026 Nikolay Govorov // SPDX-License-Identifier: AGPL-3.0-or-later package main @@ -32,27 +32,25 @@ import ( ) var ( - ErrAcquire = errors.New("database: failed to acquire connection") - ErrAlreadyMember = errors.New("database: already a member") - ErrEmailTaken = errors.New("database: email already taken") - ErrInvalidCreds = errors.New("database: invalid credentials") - ErrInvalidEmail = errors.New("database: invalid email") - ErrInvalidRole = errors.New("database: invalid role") - ErrInvalidSlug = errors.New("database: invalid slug") - ErrInvalidDateFormat = errors.New("database: invalid date format") - ErrInvalidTimezone = errors.New("database: invalid timezone") - ErrLastOwner = errors.New("database: last owner") - ErrMigrate = errors.New("database: failed to create migrator") - ErrNotImplemented = errors.New("database: filter not implemented") - ErrNotMember = errors.New("database: not a member") - ErrOpen = errors.New("database: failed to open") - ErrOrgNotFound = errors.New("database: organization not found") - ErrPing = errors.New("database: failed to ping") - ErrReservedEmail = errors.New("database: email uses a reserved domain") - ErrSlugTaken = errors.New("database: slug already taken") - ErrSoleOwner = errors.New("database: sole owner of an organization") - ErrUserNotFound = errors.New("database: user not found") - ErrWorkerNotFound = errors.New("database: worker not found") + ErrAcquire = errors.New("database: failed to acquire connection") + ErrAlreadyMember = errors.New("database: already a member") + ErrEmailTaken = errors.New("database: email already taken") + ErrInvalidCreds = errors.New("database: invalid credentials") + ErrInvalidEmail = errors.New("database: invalid email") + ErrInvalidRole = errors.New("database: invalid role") + ErrInvalidSlug = errors.New("database: invalid slug") + ErrLastOwner = errors.New("database: last owner") + ErrMigrate = errors.New("database: failed to create migrator") + ErrNotImplemented = errors.New("database: filter not implemented") + ErrNotMember = errors.New("database: not a member") + ErrOpen = errors.New("database: failed to open") + ErrOrgNotFound = errors.New("database: organization not found") + ErrPing = errors.New("database: failed to ping") + ErrReservedEmail = errors.New("database: email uses a reserved domain") + ErrSlugTaken = errors.New("database: slug already taken") + ErrSoleOwner = errors.New("database: sole owner of an organization") + ErrUserNotFound = errors.New("database: user not found") + ErrWorkerNotFound = errors.New("database: worker not found") ) // reservedEmailSuffix is the domain carved out for synthetic actors @@ -111,26 +109,11 @@ type DB struct { Pool *pgxpool.Pool } -type DateFormat int - -const ( - DateFormatDMY DateFormat = 1 - DateFormatMDY DateFormat = 2 - DateFormatYMD DateFormat = 3 -) - -type Locale struct { - Language *string - DateFormat *DateFormat -} - // User holds info about a user. type User struct { ID UserID Email string CreatedAt time.Time - Locale *Locale - Timezone *string } // Organization holds info about an organization. @@ -231,21 +214,13 @@ func (db *DB) Migrate(ctx context.Context) error { } migrator.AppendMigration("create_users", ` - CREATE TYPE locale_settings AS ( - language TEXT, - date_format INTEGER - ); - CREATE TABLE users ( id UUID PRIMARY KEY DEFAULT uuidv7(), email TEXT NOT NULL UNIQUE, password TEXT NOT NULL, superuser BOOLEAN NOT NULL DEFAULT false, created_at TIMESTAMPTZ NOT NULL DEFAULT now(), - deleted_at TIMESTAMPTZ, - - locale locale_settings, - timezone TEXT + deleted_at TIMESTAMPTZ ); CREATE FUNCTION app_user_id() RETURNS uuid STABLE AS $$ @@ -269,8 +244,6 @@ func (db *DB) Migrate(ctx context.Context) error { ); $$ LANGUAGE sql; `, ` - DROP TYPE locale_settings; - DROP TYPE date_format_t; DROP FUNCTION app_issuper; DROP FUNCTION app_user_id; DROP TABLE users; @@ -507,16 +480,12 @@ func (db *DB) UserGet(ctx context.Context, actor Actor, ref UserRef) (*User, err col, val := ref.where() q := sb.PostgreSQL.NewSelectBuilder() - sql, args := q.Select("id", "email", "created_at", - "(locale).language", "(locale).date_format", "timezone"). + sql, args := q.Select("id", "email", "created_at"). From("users"). Where(q.Equal(col, val), q.IsNull("deleted_at")). Build() - u.Locale = &Locale{} - if err := tx.QueryRow(ctx, sql, args...).Scan( - &u.ID, &u.Email, &u.CreatedAt, &u.Locale.Language, &u.Locale.DateFormat, &u.Timezone, - ); err != nil { + if err := tx.QueryRow(ctx, sql, args...).Scan(&u.ID, &u.Email, &u.CreatedAt); err != nil { if errors.Is(err, pgx.ErrNoRows) { return ErrUserNotFound } @@ -549,8 +518,7 @@ func (db *DB) UserList(ctx context.Context, actor Actor, cursor UserID, limit in } q := sb.PostgreSQL.NewSelectBuilder() - q.Select("id", "email", "created_at", - "(locale).language", "(locale).date_format", "timezone"). + q.Select("id", "email", "created_at"). From("users"). Where(q.IsNull("deleted_at")). OrderBy("id"). @@ -568,8 +536,7 @@ func (db *DB) UserList(ctx context.Context, actor Actor, cursor UserID, limit in for rows.Next() { var u User - u.Locale = &Locale{} - if err := rows.Scan(&u.ID, &u.Email, &u.CreatedAt, &u.Locale.Language, &u.Locale.DateFormat, &u.Timezone); err != nil { + if err := rows.Scan(&u.ID, &u.Email, &u.CreatedAt); err != nil { return err } users = append(users, u) @@ -580,16 +547,13 @@ func (db *DB) UserList(ctx context.Context, actor Actor, cursor UserID, limit in return users, total, err } -type UserUpdateParams struct { - Email *string - Password *string - Locale *Locale - Timezone *string -} - // UserUpdate updates a user's email and/or password. // Invalidates all sessions when password changes. -func (db *DB) UserUpdate(ctx context.Context, actor Actor, ref UserRef, p UserUpdateParams, pepper []byte) error { +func (db *DB) UserUpdate(ctx context.Context, actor Actor, ref UserRef, email *string, password *string, pepper []byte) error { + if email == nil && password == nil { + return nil + } + var id UserID return db.apicall(ctx, actor, func(tx pgx.Tx) error { @@ -601,46 +565,24 @@ func (db *DB) UserUpdate(ctx context.Context, actor Actor, ref UserRef, p UserUp return checkSelf(actor, id) }, func(tx pgx.Tx) error { - if p.Email != nil && strings.HasSuffix(strings.ToLower(*p.Email), reservedEmailSuffix) { + if email != nil && strings.HasSuffix(strings.ToLower(*email), reservedEmailSuffix) { return ErrReservedEmail } - if p.Timezone != nil { - if _, err := time.LoadLocation(*p.Timezone); err != nil { - return ErrInvalidTimezone - } - } return nil }, func(tx pgx.Tx) error { ub := sb.PostgreSQL.NewUpdateBuilder() ub.Update("users") - var hasSet bool - if p.Email != nil { - ub.SetMore(ub.Assign("email", *p.Email)) - hasSet = true + if email != nil { + ub.SetMore(ub.Assign("email", *email)) } - if p.Password != nil { - hash, err := hashPassword(*p.Password, pepper) + if password != nil { + hash, err := hashPassword(*password, pepper) if err != nil { return err } ub.SetMore(ub.Assign("password", hash)) - hasSet = true - } - if p.Locale != nil { - ub.SetMore(fmt.Sprintf( - "locale = ROW(COALESCE(%s, (locale).language), COALESCE(%s, (locale).date_format))::locale_settings", - ub.Var(p.Locale.Language), ub.Var(p.Locale.DateFormat), - )) - hasSet = true - } - if p.Timezone != nil { - ub.SetMore(ub.Assign("timezone", *p.Timezone)) - hasSet = true - } - if !hasSet { - return nil } ub.Where(ub.Equal("id", id)) @@ -653,7 +595,7 @@ func (db *DB) UserUpdate(ctx context.Context, actor Actor, ref UserRef, p UserUp return err } - if p.Password != nil { + if password != nil { if _, err := tx.Exec(ctx, `DELETE FROM sessions WHERE user_id = $1`, id); err != nil { return err } |
