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