From 2e3e82e776dcea0adbc9f6b7b60dc972547cee48 Mon Sep 17 00:00:00 2001 From: Robin Dittmar Date: Tue, 1 Sep 2026 14:55:20 +0200 Subject: [PATCH 1/9] feat: added db transactions abstraction; cleaned up user registration; repos adapted new abstraction: all repos now support transactions, if present in ctx --- cmd/bootstrap/main.go | 4 +- internal/api/handler/list.go | 2 +- internal/api/handler/user.go | 44 +++++--------------- internal/api/request/list.go | 3 +- internal/api/router/router.go | 21 ++++------ internal/domain/list.go | 17 +++++--- internal/domain/registration.go | 49 ++++++++++++++++++++++ internal/domain/transactor.go | 7 ++++ internal/repository/auth.go | 18 ++++---- internal/repository/invite.go | 18 ++++---- internal/repository/list.go | 60 +++++++-------------------- internal/repository/repo.go | 73 +++++++++++++++++++++++++++++++++ internal/repository/store.go | 23 +++++++++++ internal/repository/user.go | 25 +++-------- 14 files changed, 219 insertions(+), 145 deletions(-) create mode 100644 internal/domain/registration.go create mode 100644 internal/domain/transactor.go create mode 100644 internal/repository/repo.go create mode 100644 internal/repository/store.go diff --git a/cmd/bootstrap/main.go b/cmd/bootstrap/main.go index 1d087c7..6b69b9e 100644 --- a/cmd/bootstrap/main.go +++ b/cmd/bootstrap/main.go @@ -70,8 +70,8 @@ func main() { } func seedAdminUser(db *sql.DB, email string, name string, password string) error { - userRepo := repository.NewUserRepo(db) - userService := domain.NewUserService(userRepo) + store := repository.NewStore(db) + userService := domain.NewUserService(store.User) _, err := userService.CreateUser(context.Background(), email, name, password) return err diff --git a/internal/api/handler/list.go b/internal/api/handler/list.go index 13f3445..9889809 100644 --- a/internal/api/handler/list.go +++ b/internal/api/handler/list.go @@ -47,7 +47,7 @@ func (h *ListHandler) CreateList(w http.ResponseWriter, r *http.Request) { return } - list, err := h.ListService.CreateList(ctx, authContext.UserID, payload.Name, payload.UserIDs) + list, err := h.ListService.CreateList(ctx, authContext.UserID, payload.Name) if err != nil { slog.ErrorContext(ctx, "failed to create list", slog.Any("error", err)) response.Error(ctx, w, http.StatusInternalServerError, "failed to create list") diff --git a/internal/api/handler/user.go b/internal/api/handler/user.go index d3be142..0db4a6e 100644 --- a/internal/api/handler/user.go +++ b/internal/api/handler/user.go @@ -10,13 +10,13 @@ import ( ) type UserHandler struct { - UserService *domain.UserService - AuthService *domain.AuthService - InviteService *domain.InviteService + UserService *domain.UserService + AuthService *domain.AuthService + RegistrationService *domain.RegistrationService } -func NewUserHandler(userService *domain.UserService, authService *domain.AuthService, inviteService *domain.InviteService) *UserHandler { - return &UserHandler{UserService: userService, AuthService: authService, InviteService: inviteService} +func NewUserHandler(userService *domain.UserService, authService *domain.AuthService, registrationService *domain.RegistrationService) *UserHandler { + return &UserHandler{UserService: userService, AuthService: authService, RegistrationService: registrationService} } // CreateUser handles the creation of a user @@ -42,41 +42,17 @@ func (h *UserHandler) CreateUser(w http.ResponseWriter, r *http.Request) { return } - invite, err := h.InviteService.GetInvite(ctx, payload.InviteCode) + user, err := h.RegistrationService.Register(ctx, payload.InviteCode, payload.Email, payload.Name, payload.Password) if err != nil { - slog.ErrorContext(ctx, "failed to get invite", slog.Any("error", err)) - response.Error(ctx, w, http.StatusBadRequest, "invite is invalid") - return - } - - // TODO: Creating user and consuming the invite must be in a transaction. - // The repositories must support tx in context, and the handler must be able to start a transaction - user, err := h.UserService.CreateUser(ctx, payload.Email, payload.Name, payload.Password) - if err != nil { - slog.ErrorContext(ctx, "failed to create user", slog.Any("error", err)) - response.Error(ctx, w, http.StatusInternalServerError, "failed to create user") - return - } - - err = h.InviteService.ConsumeInvite(ctx, invite.ID, user.ID) - if err != nil { - slog.ErrorContext(ctx, "failed to consume invite", slog.Any("error", err)) - - // TODO: This should be a transaction rollback, once we have db transactions in the handler - err = h.UserService.DeleteUser(ctx, user.ID) - if err != nil { - slog.ErrorContext(ctx, "failed to delete user again", - slog.Any("error", err), - slog.String("user_id", user.ID), - ) - } - response.Error(ctx, w, http.StatusInternalServerError, "failed to consume invite") + slog.ErrorContext(ctx, "failed to register user", + slog.Any("error", err), + slog.Any("payload", payload)) + response.Error(ctx, w, http.StatusInternalServerError, "failed to register") return } slog.InfoContext(ctx, "created user successfully", slog.String("user_id", user.ID), - slog.String("invite_id", invite.ID), ) response.JSON(ctx, w, http.StatusCreated, user) } diff --git a/internal/api/request/list.go b/internal/api/request/list.go index 00ef7ce..ca8bf1e 100644 --- a/internal/api/request/list.go +++ b/internal/api/request/list.go @@ -1,8 +1,7 @@ package request type CreateListPayload struct { - Name string `json:"name"` - UserIDs []string `json:"user_ids"` + Name string `json:"name"` } type AddUserToListPayload struct { diff --git a/internal/api/router/router.go b/internal/api/router/router.go index 2f79904..e027ab9 100644 --- a/internal/api/router/router.go +++ b/internal/api/router/router.go @@ -16,20 +16,17 @@ type Config struct { } func NewMux(cfg Config) http.Handler { - authRepo := repository.NewAuthRepo(cfg.Database) - authService := domain.NewAuthService(authRepo, []byte(cfg.JWTSecret)) + store := repository.NewStore(cfg.Database) + + authService := domain.NewAuthService(store.Auth, []byte(cfg.JWTSecret)) + inviteService := domain.NewInviteService(store.Invite) + userService := domain.NewUserService(store.User) + registrationService := domain.NewRegistrationService(store, userService, inviteService) + listService := domain.NewListService(store.List) + authHandler := handler.NewAuthHandler(authService) - - inviteRepo := repository.NewInviteRepo(cfg.Database) - inviteService := domain.NewInviteService(inviteRepo) inviteHandler := handler.NewInviteHandler(inviteService) - - userRepo := repository.NewUserRepo(cfg.Database) - userService := domain.NewUserService(userRepo) - userHandler := handler.NewUserHandler(userService, authService, inviteService) - - listRepo := repository.NewListRepo(cfg.Database) - listService := domain.NewListService(listRepo) + userHandler := handler.NewUserHandler(userService, authService, registrationService) listHandler := handler.NewListHandler(listService, userService) protected := middleware.WithJWT(authService) diff --git a/internal/domain/list.go b/internal/domain/list.go index 92a4665..b7a6aeb 100644 --- a/internal/domain/list.go +++ b/internal/domain/list.go @@ -4,7 +4,6 @@ import ( "context" "errors" "log/slog" - "slices" "time" ) @@ -36,7 +35,7 @@ type ListItem struct { } type ListRepository interface { - CreateList(ctx context.Context, name string, userIDs []string) (*List, error) + CreateList(ctx context.Context, name string) (*List, error) DeleteList(ctx context.Context, listID string) error GetLists(ctx context.Context, userID string) ([]List, error) AddUserToList(ctx context.Context, listID string, userID string) error @@ -58,16 +57,22 @@ func NewListService(r ListRepository) *ListService { return &ListService{repo: r} } -func (s *ListService) CreateList(ctx context.Context, authUserID string, name string, userIDs []string) (*List, error) { +func (s *ListService) CreateList(ctx context.Context, authUserID string, name string) (*List, error) { if name == "" { return nil, ErrListNameEmpty } - if !slices.Contains(userIDs, authUserID) { - userIDs = append(userIDs, authUserID) + list, err := s.repo.CreateList(ctx, name) + if err != nil { + return nil, err } - return s.repo.CreateList(ctx, name, userIDs) + err = s.repo.AddUserToList(ctx, list.ID, authUserID) + if err != nil { + return nil, err + } + + return list, nil } func (s *ListService) DeleteList(ctx context.Context, authUserID string, listID string) error { diff --git a/internal/domain/registration.go b/internal/domain/registration.go new file mode 100644 index 0000000..b136e93 --- /dev/null +++ b/internal/domain/registration.go @@ -0,0 +1,49 @@ +package domain + +import ( + "context" + "log/slog" +) + +type RegistrationService struct { + tx Transactor + UserService *UserService + InviteService *InviteService +} + +func NewRegistrationService(tx Transactor, u *UserService, i *InviteService) *RegistrationService { + return &RegistrationService{tx: tx, UserService: u, InviteService: i} +} + +func (s *RegistrationService) Register(ctx context.Context, inviteCode string, email string, username string, password string) (*User, error) { + invite, err := s.InviteService.GetInvite(ctx, inviteCode) + if err != nil { + slog.ErrorContext(ctx, "failed to get invite", slog.Any("error", err)) + return nil, err + } + + var user *User + err = s.tx.WithinTx(ctx, func(ctx context.Context) error { + user, err = s.UserService.CreateUser(ctx, email, username, password) + if err != nil { + slog.ErrorContext(ctx, "failed to create user", slog.Any("error", err)) + return err + } + + err = s.InviteService.ConsumeInvite(ctx, invite.ID, user.ID) + if err != nil { + slog.ErrorContext(ctx, "failed to consume invite", slog.Any("error", err)) + return err + } + + return nil + }) + if err != nil { + return nil, err + } + + slog.InfoContext(ctx, "user registration complete", + slog.String("user_id", user.ID), + slog.String("invite_id", invite.ID)) + return user, nil +} diff --git a/internal/domain/transactor.go b/internal/domain/transactor.go new file mode 100644 index 0000000..7007cb7 --- /dev/null +++ b/internal/domain/transactor.go @@ -0,0 +1,7 @@ +package domain + +import "context" + +type Transactor interface { + WithinTx(ctx context.Context, fn func(ctx context.Context) error) error +} diff --git a/internal/repository/auth.go b/internal/repository/auth.go index 37d50da..70f910a 100644 --- a/internal/repository/auth.go +++ b/internal/repository/auth.go @@ -11,17 +11,13 @@ import ( ) type AuthRepo struct { - db *sql.DB -} - -func NewAuthRepo(db *sql.DB) *AuthRepo { - return &AuthRepo{db: db} + Repo } func (r *AuthRepo) GetUserById(ctx context.Context, id string) (*domain.AuthUser, error) { user := &domain.AuthUser{} - err := r.db.QueryRowContext(ctx, + err := r.conn(ctx).QueryRowContext(ctx, "SELECT id, email, name, password_hash FROM users WHERE id = $1", id, ).Scan(&user.ID, &user.Email, &user.Name, &user.PasswordHash) @@ -35,7 +31,7 @@ func (r *AuthRepo) GetUserById(ctx context.Context, id string) (*domain.AuthUser func (r *AuthRepo) GetUserByEmail(ctx context.Context, email string) (*domain.AuthUser, error) { user := &domain.AuthUser{} - err := r.db.QueryRowContext(ctx, + err := r.conn(ctx).QueryRowContext(ctx, "SELECT id, email, name, password_hash FROM users WHERE email = $1", email, ).Scan(&user.ID, &user.Email, &user.Name, &user.PasswordHash) @@ -47,7 +43,7 @@ func (r *AuthRepo) GetUserByEmail(ctx context.Context, email string) (*domain.Au } func (r *AuthRepo) StoreRefreshToken(ctx context.Context, userID string, tokenHash string, expiresAt time.Time) error { - _, err := r.db.ExecContext(ctx, + _, err := r.conn(ctx).ExecContext(ctx, "INSERT INTO refresh_tokens (user_id, token_hash, expires_at) VALUES ($1, $2, $3)", userID, tokenHash, expiresAt, ) @@ -60,7 +56,7 @@ func (r *AuthRepo) StoreRefreshToken(ctx context.Context, userID string, tokenHa func (r *AuthRepo) ConsumeRefreshToken(ctx context.Context, tokenHash string) (string, error) { var userID string - err := r.db.QueryRowContext(ctx, + err := r.conn(ctx).QueryRowContext(ctx, "DELETE FROM refresh_tokens WHERE token_hash = $1 AND expires_at > NOW() RETURNING user_id", tokenHash, ).Scan(&userID) @@ -75,7 +71,7 @@ func (r *AuthRepo) ConsumeRefreshToken(ctx context.Context, tokenHash string) (s } func (r *AuthRepo) RevokeRefreshToken(ctx context.Context, tokenHash string) error { - _, err := r.db.ExecContext(ctx, "DELETE FROM refresh_tokens WHERE token_hash = $1", tokenHash) + _, err := r.conn(ctx).ExecContext(ctx, "DELETE FROM refresh_tokens WHERE token_hash = $1", tokenHash) if err != nil { return fmt.Errorf("failed to revoke refresh token: %w", err) } @@ -84,7 +80,7 @@ func (r *AuthRepo) RevokeRefreshToken(ctx context.Context, tokenHash string) err } func (r *AuthRepo) RevokeRefreshTokens(ctx context.Context, userID string) error { - _, err := r.db.ExecContext(ctx, "DELETE FROM refresh_tokens WHERE user_id = $1", userID) + _, err := r.conn(ctx).ExecContext(ctx, "DELETE FROM refresh_tokens WHERE user_id = $1", userID) if err != nil { return fmt.Errorf("failed to revoke refresh tokens: %w", err) } diff --git a/internal/repository/invite.go b/internal/repository/invite.go index 776b86c..d19d26e 100644 --- a/internal/repository/invite.go +++ b/internal/repository/invite.go @@ -11,16 +11,12 @@ import ( ) type InviteRepo struct { - db *sql.DB -} - -func NewInviteRepo(db *sql.DB) *InviteRepo { - return &InviteRepo{db: db} + Repo } func (r *InviteRepo) CreateInvite(ctx context.Context, inviterUserID string, code string, expiresAt time.Time) (*domain.Invite, error) { var id string - err := r.db.QueryRowContext(ctx, + err := r.conn(ctx).QueryRowContext(ctx, "INSERT INTO invites (inviter_user_id, code, expires_at) VALUES ($1, $2, $3) RETURNING id", inviterUserID, code, expiresAt, ).Scan(&id) @@ -37,7 +33,7 @@ func (r *InviteRepo) CreateInvite(ctx context.Context, inviterUserID string, cod } func (r *InviteRepo) DeleteInvite(ctx context.Context, userID string, inviteID string) error { - _, err := r.db.ExecContext(ctx, + _, err := r.conn(ctx).ExecContext(ctx, "DELETE FROM invites WHERE id = $1 AND inviter_user_id = $2 AND consumed_at IS NULL", inviteID, userID, ) @@ -49,7 +45,7 @@ func (r *InviteRepo) DeleteInvite(ctx context.Context, userID string, inviteID s } func (r *InviteRepo) ConsumeInvite(ctx context.Context, inviteID string, inviteeUserID string) error { - res, err := r.db.ExecContext(ctx, + res, err := r.conn(ctx).ExecContext(ctx, "UPDATE invites SET invitee_user_id=$1, consumed_at=NOW() WHERE id=$2 AND expires_at > NOW() AND consumed_at IS NULL", inviteeUserID, inviteID, ) @@ -70,7 +66,7 @@ func (r *InviteRepo) ConsumeInvite(ctx context.Context, inviteID string, invitee func (r *InviteRepo) GetInvite(ctx context.Context, code string) (*domain.Invite, error) { var invite domain.Invite - err := r.db.QueryRowContext(ctx, + err := r.conn(ctx).QueryRowContext(ctx, "SELECT id, code, expires_at, consumed_at FROM invites WHERE code=$1", code, ).Scan(&invite.ID, &invite.Code, &invite.ExpiresAt, &invite.ConsumedAt) @@ -82,7 +78,7 @@ func (r *InviteRepo) GetInvite(ctx context.Context, code string) (*domain.Invite } func (r *InviteRepo) GetInvites(ctx context.Context, userID string, offset int, count int) ([]domain.Invite, error) { - rows, err := r.db.QueryContext(ctx, + rows, err := r.conn(ctx).QueryContext(ctx, "SELECT id, code, expires_at, consumed_at FROM invites WHERE inviter_user_id=$1 ORDER BY created_at DESC OFFSET $2 LIMIT $3", userID, offset, count, ) @@ -110,7 +106,7 @@ func (r *InviteRepo) GetInvites(ctx context.Context, userID string, offset int, func (r *InviteRepo) CountInvites(ctx context.Context, userID string) (int, error) { var count int - err := r.db.QueryRowContext(ctx, + err := r.conn(ctx).QueryRowContext(ctx, "SELECT COUNT(*) FROM invites WHERE inviter_user_id=$1", userID, ).Scan(&count) diff --git a/internal/repository/list.go b/internal/repository/list.go index d58d90a..e588806 100644 --- a/internal/repository/list.go +++ b/internal/repository/list.go @@ -5,29 +5,18 @@ import ( "database/sql" "errors" "fmt" - "log/slog" "github.com/robindittmar/dttmr-api/internal/domain" ) type ListRepo struct { - db *sql.DB + Repo } -func NewListRepo(db *sql.DB) *ListRepo { - return &ListRepo{db: db} -} - -func (r *ListRepo) CreateList(ctx context.Context, name string, userIDs []string) (*domain.List, error) { - tx, err := r.db.BeginTx(ctx, nil) - if err != nil { - return nil, fmt.Errorf("begin transaction: %w", err) - } - defer tx.Rollback() - +func (r *ListRepo) CreateList(ctx context.Context, name string) (*domain.List, error) { list := &domain.List{Name: name} - err = tx.QueryRowContext(ctx, + err := r.conn(ctx).QueryRowContext(ctx, "INSERT INTO lists (name) VALUES ($1) RETURNING id, created_at, modified_at", name, ).Scan(&list.ID, &list.CreatedAt, &list.ModifiedAt) @@ -35,32 +24,11 @@ func (r *ListRepo) CreateList(ctx context.Context, name string, userIDs []string return nil, fmt.Errorf("failed to insert list: %w", err) } - stmt, err := tx.PrepareContext(ctx, "INSERT INTO list_users (list_id, user_id) VALUES ($1, $2)") - if err != nil { - return nil, fmt.Errorf("failed to prepare user/list association statement: %w", err) - } - defer func() { - err := stmt.Close() - if err != nil { - slog.Error("failed to close user/list association statement", slog.Any("error", err)) - } - }() - - for _, userID := range userIDs { - if _, err = stmt.ExecContext(ctx, list.ID, userID); err != nil { - return nil, fmt.Errorf("failed to insert user/list association: %w", err) - } - } - - if err := tx.Commit(); err != nil { - return nil, fmt.Errorf("commit transaction: %w", err) - } - return list, nil } func (r *ListRepo) DeleteList(ctx context.Context, listID string) error { - _, err := r.db.ExecContext(ctx, "DELETE FROM lists WHERE id = $1", listID) + _, err := r.conn(ctx).ExecContext(ctx, "DELETE FROM lists WHERE id = $1", listID) if err != nil { return fmt.Errorf("failed to delete list: %w", err) } @@ -69,7 +37,7 @@ func (r *ListRepo) DeleteList(ctx context.Context, listID string) error { } func (r *ListRepo) GetLists(ctx context.Context, userID string) ([]domain.List, error) { - rows, err := r.db.QueryContext(ctx, + rows, err := r.conn(ctx).QueryContext(ctx, "SELECT l.id, l.name, l.created_at, l.modified_at, (SELECT COUNT(*) FROM list_items WHERE list_id=l.id), (SELECT COUNT(*) FROM list_items WHERE list_id=l.id AND is_completed=true) FROM lists AS l INNER JOIN list_users ON l.id=list_users.list_id WHERE list_users.user_id = $1", userID, ) @@ -96,7 +64,7 @@ func (r *ListRepo) GetLists(ctx context.Context, userID string) ([]domain.List, } func (r *ListRepo) AddUserToList(ctx context.Context, listID string, userID string) error { - _, err := r.db.ExecContext(ctx, + _, err := r.conn(ctx).ExecContext(ctx, "INSERT INTO list_users (list_id, user_id) VALUES ($1, $2)", listID, userID, ) @@ -108,7 +76,7 @@ func (r *ListRepo) AddUserToList(ctx context.Context, listID string, userID stri } func (r *ListRepo) RemoveUserFromList(ctx context.Context, listID string, userID string) error { - _, err := r.db.ExecContext(ctx, + _, err := r.conn(ctx).ExecContext(ctx, "DELETE FROM list_users WHERE list_id = $1 AND user_id = $2", listID, userID, ) @@ -122,7 +90,7 @@ func (r *ListRepo) RemoveUserFromList(ctx context.Context, listID string, userID func (r *ListRepo) IsUserInList(ctx context.Context, listID string, userID string) (bool, error) { var cnt int - err := r.db.QueryRowContext(ctx, + err := r.conn(ctx).QueryRowContext(ctx, "SELECT COUNT(*) FROM list_users WHERE list_id = $1 AND user_id = $2", listID, userID, ).Scan(&cnt) @@ -136,7 +104,7 @@ func (r *ListRepo) IsUserInList(ctx context.Context, listID string, userID strin func (r *ListRepo) IsUserInListByItemID(ctx context.Context, listItemID string, userID string) (bool, error) { var cnt int - err := r.db.QueryRowContext(ctx, + err := r.conn(ctx).QueryRowContext(ctx, "SELECT COUNT(*) FROM list_users WHERE list_id = (SELECT list_id FROM list_items WHERE id = $1) AND user_id = $2", listItemID, userID, ).Scan(&cnt) @@ -150,7 +118,7 @@ func (r *ListRepo) IsUserInListByItemID(ctx context.Context, listItemID string, func (r *ListRepo) CreateListItem(ctx context.Context, listID string, title string) (*domain.ListItem, error) { l := &domain.ListItem{Title: title} - err := r.db.QueryRowContext(ctx, + err := r.conn(ctx).QueryRowContext(ctx, "INSERT INTO list_items (list_id, title) VALUES ($1, $2) RETURNING id, is_completed, created_at, modified_at", listID, title, ).Scan(&l.ID, &l.IsCompleted, &l.CreatedAt, &l.ModifiedAt) @@ -162,7 +130,7 @@ func (r *ListRepo) CreateListItem(ctx context.Context, listID string, title stri } func (r *ListRepo) DeleteListItem(ctx context.Context, listItemID string) error { - _, err := r.db.ExecContext(ctx, "DELETE FROM list_items WHERE id = $1", listItemID) + _, err := r.conn(ctx).ExecContext(ctx, "DELETE FROM list_items WHERE id = $1", listItemID) if err != nil { return fmt.Errorf("failed to delete list item: %w", err) } @@ -171,7 +139,7 @@ func (r *ListRepo) DeleteListItem(ctx context.Context, listItemID string) error } func (r *ListRepo) UpdateListItem(ctx context.Context, listItemID string, title string, isCompleted bool) error { - _, err := r.db.ExecContext(ctx, "UPDATE list_items SET title = $1, is_completed = $2, modified_at = NOW() WHERE id = $3", + _, err := r.conn(ctx).ExecContext(ctx, "UPDATE list_items SET title = $1, is_completed = $2, modified_at = NOW() WHERE id = $3", title, isCompleted, listItemID, ) if err != nil { @@ -182,7 +150,7 @@ func (r *ListRepo) UpdateListItem(ctx context.Context, listItemID string, title } func (r *ListRepo) SetListItemCompleted(ctx context.Context, listItemID string, isCompleted bool) error { - _, err := r.db.ExecContext(ctx, "UPDATE list_items SET is_completed = $1, modified_at = NOW() WHERE id = $2", + _, err := r.conn(ctx).ExecContext(ctx, "UPDATE list_items SET is_completed = $1, modified_at = NOW() WHERE id = $2", isCompleted, listItemID, ) if err != nil { @@ -193,7 +161,7 @@ func (r *ListRepo) SetListItemCompleted(ctx context.Context, listItemID string, } func (r *ListRepo) GetListItems(ctx context.Context, listID string) ([]domain.ListItem, error) { - rows, err := r.db.QueryContext(ctx, + rows, err := r.conn(ctx).QueryContext(ctx, "SELECT id, title, is_completed, created_at, modified_at FROM list_items WHERE list_id = $1 ORDER BY is_completed, modified_at DESC", listID, ) diff --git a/internal/repository/repo.go b/internal/repository/repo.go new file mode 100644 index 0000000..b5cefb9 --- /dev/null +++ b/internal/repository/repo.go @@ -0,0 +1,73 @@ +package repository + +import ( + "context" + "database/sql" + "fmt" + "sync/atomic" +) + +type Transaction interface { + ExecContext(ctx context.Context, query string, args ...any) (sql.Result, error) + QueryContext(ctx context.Context, query string, args ...any) (*sql.Rows, error) + QueryRowContext(ctx context.Context, query string, args ...any) *sql.Row + PrepareContext(ctx context.Context, query string) (*sql.Stmt, error) +} + +type Repo struct { + *Transactor +} + +func NewRepo(t *Transactor) Repo { + return Repo{t} +} + +func (r *Repo) conn(ctx context.Context) Transaction { + if tx, ok := ctx.Value(txKey{}).(*sql.Tx); ok { + return tx + } + return r.db +} + +type txKey struct{} + +type Transactor struct { + db *sql.DB + sp atomic.Uint64 +} + +func NewTransactor(db *sql.DB) *Transactor { + return &Transactor{db: db, sp: atomic.Uint64{}} +} + +func (t *Transactor) WithinTx(ctx context.Context, fn func(context.Context) error) error { + if tx, ok := ctx.Value(txKey{}).(*sql.Tx); ok { + return t.withinSavepoint(ctx, tx, fn) + } + + tx, err := t.db.BeginTx(ctx, nil) + if err != nil { + return err + } + defer tx.Rollback() + + if err = fn(ctx); err != nil { + return err + } + return tx.Commit() +} + +func (t *Transactor) withinSavepoint(ctx context.Context, tx *sql.Tx, fn func(context.Context) error) error { + name := fmt.Sprintf("sp_%d", t.sp.Add(1)) + if _, err := tx.ExecContext(ctx, "SAVEPOINT "+name); err != nil { + return err + } + + if err := fn(ctx); err != nil { + _, _ = tx.ExecContext(ctx, "ROLLBACK TO SAVEPOINT"+name) + return err + } + + _, err := tx.ExecContext(ctx, "RELEASE SAVEPOINT"+name) + return err +} diff --git a/internal/repository/store.go b/internal/repository/store.go new file mode 100644 index 0000000..27bade3 --- /dev/null +++ b/internal/repository/store.go @@ -0,0 +1,23 @@ +package repository + +import "database/sql" + +type Store struct { + *Transactor + Auth *AuthRepo + Invite *InviteRepo + List *ListRepo + User *UserRepo +} + +func NewStore(db *sql.DB) *Store { + t := NewTransactor(db) + r := NewRepo(t) + return &Store{ + Transactor: t, + Auth: &AuthRepo{r}, + Invite: &InviteRepo{r}, + List: &ListRepo{r}, + User: &UserRepo{r}, + } +} diff --git a/internal/repository/user.go b/internal/repository/user.go index e385fb1..798b5fb 100644 --- a/internal/repository/user.go +++ b/internal/repository/user.go @@ -2,30 +2,19 @@ package repository import ( "context" - "database/sql" "fmt" "github.com/robindittmar/dttmr-api/internal/domain" ) type UserRepo struct { - db *sql.DB -} - -func NewUserRepo(db *sql.DB) *UserRepo { - return &UserRepo{db: db} + Repo } func (r *UserRepo) CreateUser(ctx context.Context, email string, name string, passwordHash string) (*domain.User, error) { - tx, err := r.db.BeginTx(ctx, nil) - if err != nil { - return nil, fmt.Errorf("begin transaction: %w", err) - } - defer tx.Rollback() - user := &domain.User{Email: email, Name: name} - err = tx.QueryRowContext(ctx, + err := r.conn(ctx).QueryRowContext(ctx, "INSERT INTO users (email, name, password_hash) VALUES ($1, $2, $3) RETURNING id, created_at", email, name, passwordHash, ).Scan(&user.ID, &user.CreatedAt) @@ -33,15 +22,11 @@ func (r *UserRepo) CreateUser(ctx context.Context, email string, name string, pa return nil, fmt.Errorf("failed to insert user: %w", err) } - if err := tx.Commit(); err != nil { - return nil, fmt.Errorf("commit transaction: %w", err) - } - return user, nil } func (r *UserRepo) DeleteUser(ctx context.Context, userID string) error { - _, err := r.db.ExecContext(ctx, + _, err := r.conn(ctx).ExecContext(ctx, "DELETE FROM users WHERE id = $1", userID, ) @@ -53,7 +38,7 @@ func (r *UserRepo) DeleteUser(ctx context.Context, userID string) error { } func (r *UserRepo) ChangePassword(ctx context.Context, userID string, passwordHash string) error { - _, err := r.db.ExecContext(ctx, + _, err := r.conn(ctx).ExecContext(ctx, "UPDATE users SET password_hash = $1 WHERE id = $2", passwordHash, userID, ) @@ -67,7 +52,7 @@ func (r *UserRepo) ChangePassword(ctx context.Context, userID string, passwordHa func (r *UserRepo) GetUserByEmail(ctx context.Context, email string) (*domain.User, error) { user := &domain.User{} - err := r.db.QueryRowContext(ctx, + err := r.conn(ctx).QueryRowContext(ctx, "SELECT id, email, name FROM users WHERE email = $1", email, ).Scan(&user.ID, &user.Email, &user.Name) -- 2.54.0 From 4b703b077ce05fb8ed874e25732a1d9db2a0ec3d Mon Sep 17 00:00:00 2001 From: Robin Dittmar Date: Tue, 1 Sep 2026 15:06:30 +0200 Subject: [PATCH 2/9] fix: list service errors unification --- internal/domain/list.go | 37 ++++++++++++++++++------------------- 1 file changed, 18 insertions(+), 19 deletions(-) diff --git a/internal/domain/list.go b/internal/domain/list.go index b7a6aeb..a3f9240 100644 --- a/internal/domain/list.go +++ b/internal/domain/list.go @@ -8,12 +8,11 @@ import ( ) var ( - ErrListIDEmpty = errors.New("list id must not be empty") - ErrListNameEmpty = errors.New("list name must not be empty") - ErrUserIDEmpty = errors.New("user id must not be empty") - ErrListItemIDEmpty = errors.New("list item id must not be empty") - ErrListItemTitleEmpty = errors.New("list item title must not be empty") - ErrUserNotInList = errors.New("user not in list") + ErrListIDMissing = errors.New("list id is required") + ErrListNameMissing = errors.New("list name is required") + ErrListItemIDMissing = errors.New("list item id is required") + ErrListItemTitleMissing = errors.New("list item title is required") + ErrUserNotInList = errors.New("user not in list") ) type List struct { @@ -59,7 +58,7 @@ func NewListService(r ListRepository) *ListService { func (s *ListService) CreateList(ctx context.Context, authUserID string, name string) (*List, error) { if name == "" { - return nil, ErrListNameEmpty + return nil, ErrListNameMissing } list, err := s.repo.CreateList(ctx, name) @@ -77,7 +76,7 @@ func (s *ListService) CreateList(ctx context.Context, authUserID string, name st func (s *ListService) DeleteList(ctx context.Context, authUserID string, listID string) error { if listID == "" { - return ErrListIDEmpty + return ErrListIDMissing } if err := s.userAllowedToAccessList(ctx, authUserID, listID); err != nil { @@ -93,10 +92,10 @@ func (s *ListService) GetLists(ctx context.Context, authUserID string) ([]List, func (s *ListService) AddUserToList(ctx context.Context, authUserID string, listID string, userID string) error { if listID == "" { - return ErrListIDEmpty + return ErrListIDMissing } if userID == "" { - return ErrUserIDEmpty + return ErrUserIDMissing } if err := s.userAllowedToAccessList(ctx, authUserID, listID); err != nil { @@ -108,10 +107,10 @@ func (s *ListService) AddUserToList(ctx context.Context, authUserID string, list func (s *ListService) RemoveUserFromList(ctx context.Context, authUserID string, listID string, userID string) error { if listID == "" { - return ErrListIDEmpty + return ErrListIDMissing } if userID == "" { - return ErrUserIDEmpty + return ErrUserIDMissing } if err := s.userAllowedToAccessList(ctx, authUserID, listID); err != nil { @@ -123,10 +122,10 @@ func (s *ListService) RemoveUserFromList(ctx context.Context, authUserID string, func (s *ListService) CreateListItem(ctx context.Context, authUserID string, listID string, title string) (*ListItem, error) { if listID == "" { - return nil, ErrListIDEmpty + return nil, ErrListIDMissing } if title == "" { - return nil, ErrListItemTitleEmpty + return nil, ErrListItemTitleMissing } if err := s.userAllowedToAccessList(ctx, authUserID, listID); err != nil { @@ -138,7 +137,7 @@ func (s *ListService) CreateListItem(ctx context.Context, authUserID string, lis func (s *ListService) DeleteListItem(ctx context.Context, authUserID string, listItemID string) error { if listItemID == "" { - return ErrListItemIDEmpty + return ErrListItemIDMissing } if err := s.userAllowedToAccessListItem(ctx, authUserID, listItemID); err != nil { @@ -150,10 +149,10 @@ func (s *ListService) DeleteListItem(ctx context.Context, authUserID string, lis func (s *ListService) UpdateListItem(ctx context.Context, authUserID string, listItemID string, title string, isCompleted bool) error { if listItemID == "" { - return ErrListItemIDEmpty + return ErrListItemIDMissing } if title == "" { - return ErrListItemTitleEmpty + return ErrListItemTitleMissing } if err := s.userAllowedToAccessListItem(ctx, authUserID, listItemID); err != nil { @@ -165,7 +164,7 @@ func (s *ListService) UpdateListItem(ctx context.Context, authUserID string, lis func (s *ListService) SetListItemCompleted(ctx context.Context, authUserID string, listItemID string, isCompleted bool) error { if listItemID == "" { - return ErrListItemIDEmpty + return ErrListItemIDMissing } if err := s.userAllowedToAccessListItem(ctx, authUserID, listItemID); err != nil { @@ -177,7 +176,7 @@ func (s *ListService) SetListItemCompleted(ctx context.Context, authUserID strin func (s *ListService) GetListItems(ctx context.Context, authUserID string, listID string) ([]ListItem, error) { if listID == "" { - return nil, ErrListIDEmpty + return nil, ErrListIDMissing } if err := s.userAllowedToAccessList(ctx, authUserID, listID); err != nil { -- 2.54.0 From b35dc4841f1c3c16778d2c5a40dbace74a44eb55 Mon Sep 17 00:00:00 2001 From: Robin Dittmar Date: Tue, 1 Sep 2026 15:28:18 +0200 Subject: [PATCH 3/9] fix: /logout routes; improved error handling for invites --- internal/api/handler/auth.go | 1 + internal/api/handler/user.go | 14 +++++++++++++- internal/api/router/router.go | 2 +- internal/domain/invite.go | 1 + internal/domain/registration.go | 8 ++++++++ internal/repository/invite.go | 13 +++++++++++-- 6 files changed, 35 insertions(+), 4 deletions(-) diff --git a/internal/api/handler/auth.go b/internal/api/handler/auth.go index c4c5a5f..af924fb 100644 --- a/internal/api/handler/auth.go +++ b/internal/api/handler/auth.go @@ -89,6 +89,7 @@ func (h *AuthHandler) Refresh(w http.ResponseWriter, r *http.Request) { // @Tags Authorization // @Accept json // @Produce json +// @Param payload body request.LogoutPayload true "Logout payload" // @Success 200 {object} nil // @Error 400 {object} response.ErrorResponse "failed to decode request body" // @Error 500 {object} response.ErrorResponse "failed to logout" diff --git a/internal/api/handler/user.go b/internal/api/handler/user.go index 0db4a6e..1e40b66 100644 --- a/internal/api/handler/user.go +++ b/internal/api/handler/user.go @@ -1,6 +1,7 @@ package handler import ( + "errors" "log/slog" "net/http" @@ -29,6 +30,8 @@ func NewUserHandler(userService *domain.UserService, authService *domain.AuthSer // @Param payload body request.CreateUserPayload true "Create user payload" // @Success 201 {object} domain.User // @Error 400 {object} response.ErrorResponse "failed to decode request body" +// @Error 400 {object} response.ErrorResponse "invite is expired" +// @Error 409 {object} response.ErrorResponse "invite is already consumed" // @Error 400 {object} response.ErrorResponse "invite is invalid" // @Error 500 {object} response.ErrorResponse "failed to create user" // @Router /users [post] @@ -47,7 +50,16 @@ func (h *UserHandler) CreateUser(w http.ResponseWriter, r *http.Request) { slog.ErrorContext(ctx, "failed to register user", slog.Any("error", err), slog.Any("payload", payload)) - response.Error(ctx, w, http.StatusInternalServerError, "failed to register") + + if errors.Is(err, domain.ErrInviteExpired) { + response.Error(ctx, w, http.StatusBadRequest, "invite is expired") + } else if errors.Is(err, domain.ErrInviteConsumed) { + response.Error(ctx, w, http.StatusConflict, "invite is already consumed") + } else if errors.Is(err, domain.ErrInviteInvalid) { + response.Error(ctx, w, http.StatusBadRequest, "invite is invalid") + } else { + response.Error(ctx, w, http.StatusInternalServerError, "failed to register") + } return } diff --git a/internal/api/router/router.go b/internal/api/router/router.go index e027ab9..0296f3e 100644 --- a/internal/api/router/router.go +++ b/internal/api/router/router.go @@ -38,7 +38,7 @@ func NewMux(cfg Config) http.Handler { apiMux.HandleFunc("POST /login", authHandler.Login) apiMux.HandleFunc("POST /login/refresh", authHandler.Refresh) apiMux.HandleFunc("POST /logout", authHandler.Logout) - apiMux.HandleFunc("POST /logout/all", authHandler.LogoutAllDevices) + apiMux.HandleFunc("POST /logout/all", protected(authHandler.LogoutAllDevices)) // Users apiMux.HandleFunc("POST /users", userHandler.CreateUser) diff --git a/internal/domain/invite.go b/internal/domain/invite.go index c70119c..1ffb8ed 100644 --- a/internal/domain/invite.go +++ b/internal/domain/invite.go @@ -9,6 +9,7 @@ import ( var ( ErrInviteIDMissing = errors.New("invite id is required") ErrCodeMissing = errors.New("invite code is required") + ErrInviteInvalid = errors.New("invite is invalid") ErrInviteExpired = errors.New("invite is expired") ErrInviteConsumed = errors.New("invite is already consumed") ) diff --git a/internal/domain/registration.go b/internal/domain/registration.go index b136e93..678f9da 100644 --- a/internal/domain/registration.go +++ b/internal/domain/registration.go @@ -3,6 +3,7 @@ package domain import ( "context" "log/slog" + "time" ) type RegistrationService struct { @@ -22,6 +23,13 @@ func (s *RegistrationService) Register(ctx context.Context, inviteCode string, e return nil, err } + if invite.ConsumedAt != nil { + return nil, ErrInviteConsumed + } + if invite.ExpiresAt.Before(time.Now()) { + return nil, ErrInviteExpired + } + var user *User err = s.tx.WithinTx(ctx, func(ctx context.Context) error { user, err = s.UserService.CreateUser(ctx, email, username, password) diff --git a/internal/repository/invite.go b/internal/repository/invite.go index d19d26e..6932ed9 100644 --- a/internal/repository/invite.go +++ b/internal/repository/invite.go @@ -33,13 +33,22 @@ func (r *InviteRepo) CreateInvite(ctx context.Context, inviterUserID string, cod } func (r *InviteRepo) DeleteInvite(ctx context.Context, userID string, inviteID string) error { - _, err := r.conn(ctx).ExecContext(ctx, + res, err := r.conn(ctx).ExecContext(ctx, "DELETE FROM invites WHERE id = $1 AND inviter_user_id = $2 AND consumed_at IS NULL", inviteID, userID, ) if err != nil { return fmt.Errorf("failed to delete invite: %w", err) } + affected, err := res.RowsAffected() + if err != nil { + return fmt.Errorf("could not get rows affected: %w", err) + } + if affected < 1 { + // It's more of an assumption, + // but unless I encounter this being wrong, I'll keep it. + return domain.ErrInviteConsumed + } return nil } @@ -57,7 +66,7 @@ func (r *InviteRepo) ConsumeInvite(ctx context.Context, inviteID string, invitee return fmt.Errorf("could not get rows affected: %w", err) } if affected < 1 { - return fmt.Errorf("invite not found or expired") + return domain.ErrInviteInvalid } return nil -- 2.54.0 From 26aa0c2a9454a338160884f305a150fc98a9a4ad Mon Sep 17 00:00:00 2001 From: Robin Dittmar Date: Tue, 1 Sep 2026 15:38:17 +0200 Subject: [PATCH 4/9] fix: /login now has better error handling/reporting --- internal/api/handler/auth.go | 10 +++++++++- internal/domain/auth.go | 7 ++++++- internal/repository/auth.go | 3 +++ 3 files changed, 18 insertions(+), 2 deletions(-) diff --git a/internal/api/handler/auth.go b/internal/api/handler/auth.go index af924fb..fb9d03b 100644 --- a/internal/api/handler/auth.go +++ b/internal/api/handler/auth.go @@ -1,6 +1,7 @@ package handler import ( + "errors" "log/slog" "net/http" @@ -41,8 +42,15 @@ func (h *AuthHandler) Login(w http.ResponseWriter, r *http.Request) { tokens, err := h.AuthService.Login(ctx, payload.Email, payload.Password) if err != nil { + if errors.Is(err, domain.ErrEmailNotFound) { + response.Error(ctx, w, http.StatusUnauthorized, "email not found") + } else if errors.Is(err, domain.ErrPasswordWrong) { + response.Error(ctx, w, http.StatusUnauthorized, "password is wrong") + } else { + response.Error(ctx, w, http.StatusInternalServerError, "failed to login") + } + slog.ErrorContext(ctx, "failed to login", slog.Any("error", err)) - response.Error(ctx, w, http.StatusInternalServerError, "failed to login") return } diff --git a/internal/domain/auth.go b/internal/domain/auth.go index ba5d141..73070d9 100644 --- a/internal/domain/auth.go +++ b/internal/domain/auth.go @@ -14,6 +14,11 @@ import ( "golang.org/x/crypto/bcrypt" ) +var ( + ErrEmailNotFound = errors.New("email not found") + ErrPasswordWrong = errors.New("password is wrong") +) + type AuthRepository interface { GetUserById(ctx context.Context, id string) (*AuthUser, error) GetUserByEmail(ctx context.Context, email string) (*AuthUser, error) @@ -71,7 +76,7 @@ func (s *AuthService) Authenticate(ctx context.Context, email string, password s err = bcrypt.CompareHashAndPassword([]byte(user.PasswordHash), []byte(password)) if err != nil { if errors.Is(err, bcrypt.ErrMismatchedHashAndPassword) { - return user, errors.New("invalid email or password") + return user, ErrPasswordWrong } return user, err } diff --git a/internal/repository/auth.go b/internal/repository/auth.go index 70f910a..d4ed7c1 100644 --- a/internal/repository/auth.go +++ b/internal/repository/auth.go @@ -36,6 +36,9 @@ func (r *AuthRepo) GetUserByEmail(ctx context.Context, email string) (*domain.Au email, ).Scan(&user.ID, &user.Email, &user.Name, &user.PasswordHash) if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return nil, domain.ErrEmailNotFound + } return nil, fmt.Errorf("failed to get user: %w", err) } -- 2.54.0 From d3aa9523dc99977fd5c6dff35de4da250c68a915 Mon Sep 17 00:00:00 2001 From: Robin Dittmar Date: Tue, 1 Sep 2026 15:40:26 +0200 Subject: [PATCH 5/9] fix: updated doc strings --- internal/api/handler/auth.go | 2 ++ 1 file changed, 2 insertions(+) diff --git a/internal/api/handler/auth.go b/internal/api/handler/auth.go index fb9d03b..89e5e55 100644 --- a/internal/api/handler/auth.go +++ b/internal/api/handler/auth.go @@ -28,6 +28,8 @@ func NewAuthHandler(authService *domain.AuthService) *AuthHandler { // @Param payload body request.LoginPayload true "Login payload" // @Success 200 {object} domain.TokenPair // @Error 400 {object} response.ErrorResponse "failed to decode request body" +// @Error 401 {object} response.ErrorResponse "email not found" +// @Error 401 {object} response.ErrorResponse "password is wrong" // @Error 500 {object} response.ErrorResponse "failed to login" // @Router /login [post] func (h *AuthHandler) Login(w http.ResponseWriter, r *http.Request) { -- 2.54.0 From 173d56d399002e5cdc2a9a14cbca76482f5195e4 Mon Sep 17 00:00:00 2001 From: Robin Dittmar Date: Tue, 1 Sep 2026 15:44:27 +0200 Subject: [PATCH 6/9] fix: /lists/items now consistent in singular/plural usage --- internal/api/handler/list.go | 6 +++--- internal/api/router/router.go | 6 +++--- 2 files changed, 6 insertions(+), 6 deletions(-) diff --git a/internal/api/handler/list.go b/internal/api/handler/list.go index 9889809..ad2863f 100644 --- a/internal/api/handler/list.go +++ b/internal/api/handler/list.go @@ -243,7 +243,7 @@ func (h *ListHandler) RemoveUserFromList(w http.ResponseWriter, r *http.Request) // @Success 204 {object} nil // @Error 400 {object} response.ErrorResponse "failed to decode request body" // @Error 500 {object} response.ErrorResponse "failed to create list item" -// @Router /lists/item [post] +// @Router /lists/items [post] func (h *ListHandler) CreateListItem(w http.ResponseWriter, r *http.Request) { ctx := r.Context() @@ -283,7 +283,7 @@ func (h *ListHandler) CreateListItem(w http.ResponseWriter, r *http.Request) { // @Success 204 // @Error 400 {object} response.ErrorResponse "failed to decode request url" // @Error 500 {object} response.ErrorResponse "failed to delete list item" -// @Router /lists/item/{id} [delete] +// @Router /lists/items/{id} [delete] func (h *ListHandler) DeleteListItem(w http.ResponseWriter, r *http.Request) { ctx := r.Context() @@ -324,7 +324,7 @@ func (h *ListHandler) DeleteListItem(w http.ResponseWriter, r *http.Request) { // @Error 400 {object} response.ErrorResponse "failed to decode request body" // @Error 401 {object} response.ErrorResponse "not authorized" // @Error 500 {object} response.ErrorResponse "failed to update list item" -// @Router /lists/item [put] +// @Router /lists/items [put] func (h *ListHandler) UpdateListItem(w http.ResponseWriter, r *http.Request) { ctx := r.Context() diff --git a/internal/api/router/router.go b/internal/api/router/router.go index 0296f3e..16f8ac8 100644 --- a/internal/api/router/router.go +++ b/internal/api/router/router.go @@ -57,9 +57,9 @@ func NewMux(cfg Config) http.Handler { apiMux.Handle("GET /lists", protected(listHandler.GetLists)) apiMux.Handle("POST /lists/user", protected(listHandler.AddUserToList)) apiMux.Handle("DELETE /lists/user", protected(listHandler.RemoveUserFromList)) - apiMux.Handle("POST /lists/item", protected(listHandler.CreateListItem)) - apiMux.Handle("DELETE /lists/item/{id}", protected(listHandler.DeleteListItem)) - apiMux.Handle("PUT /lists/item", protected(listHandler.UpdateListItem)) + apiMux.Handle("POST /lists/items", protected(listHandler.CreateListItem)) + apiMux.Handle("DELETE /lists/items/{id}", protected(listHandler.DeleteListItem)) + apiMux.Handle("PUT /lists/items", protected(listHandler.UpdateListItem)) apiMux.Handle("POST /lists/items/{id}", protected(listHandler.SetListItemCompleted)) apiMux.Handle("GET /lists/{id}", protected(listHandler.GetListItems)) -- 2.54.0 From b4ab9eaba1849680c1e8995d3545c993b79973b1 Mon Sep 17 00:00:00 2001 From: Robin Dittmar Date: Tue, 1 Sep 2026 15:46:07 +0200 Subject: [PATCH 7/9] fix: POST /lists/items now returns "list_id" correctly --- internal/repository/list.go | 1 + 1 file changed, 1 insertion(+) diff --git a/internal/repository/list.go b/internal/repository/list.go index e588806..c8ae0fa 100644 --- a/internal/repository/list.go +++ b/internal/repository/list.go @@ -126,6 +126,7 @@ func (r *ListRepo) CreateListItem(ctx context.Context, listID string, title stri return nil, fmt.Errorf("failed to insert list item: %w", err) } + l.ListID = listID return l, nil } -- 2.54.0 From 24bb87fdca763bbd90286a87305acb7e143c9c6d Mon Sep 17 00:00:00 2001 From: Robin Dittmar Date: Tue, 1 Sep 2026 15:52:47 +0200 Subject: [PATCH 8/9] fix: removed double check for consumed/expired invitations --- internal/domain/registration.go | 8 -------- 1 file changed, 8 deletions(-) diff --git a/internal/domain/registration.go b/internal/domain/registration.go index 678f9da..b136e93 100644 --- a/internal/domain/registration.go +++ b/internal/domain/registration.go @@ -3,7 +3,6 @@ package domain import ( "context" "log/slog" - "time" ) type RegistrationService struct { @@ -23,13 +22,6 @@ func (s *RegistrationService) Register(ctx context.Context, inviteCode string, e return nil, err } - if invite.ConsumedAt != nil { - return nil, ErrInviteConsumed - } - if invite.ExpiresAt.Before(time.Now()) { - return nil, ErrInviteExpired - } - var user *User err = s.tx.WithinTx(ctx, func(ctx context.Context) error { user, err = s.UserService.CreateUser(ctx, email, username, password) -- 2.54.0 From 8e126c46fbc12031fe29198d95c8488ed4678cdc Mon Sep 17 00:00:00 2001 From: Robin Dittmar Date: Tue, 1 Sep 2026 15:53:08 +0200 Subject: [PATCH 9/9] fix: not finding an invite in the database now produced an invalid invite error --- internal/repository/invite.go | 3 +++ 1 file changed, 3 insertions(+) diff --git a/internal/repository/invite.go b/internal/repository/invite.go index 6932ed9..746651a 100644 --- a/internal/repository/invite.go +++ b/internal/repository/invite.go @@ -80,6 +80,9 @@ func (r *InviteRepo) GetInvite(ctx context.Context, code string) (*domain.Invite code, ).Scan(&invite.ID, &invite.Code, &invite.ExpiresAt, &invite.ConsumedAt) if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return nil, domain.ErrInviteInvalid + } return nil, fmt.Errorf("failed to get invite: %w", err) } -- 2.54.0