From b1cfabc81a09fa685cfee7c51f7ab73f7c563498 Mon Sep 17 00:00:00 2001 From: Hailey Date: Sat, 12 Jul 2025 19:21:48 +0000 Subject: [PATCH] cleanup some error returns --- internal/helpers/helpers.go | 13 +++++++++++++ server/handle_server_confirm_email.go | 4 ++-- server/handle_server_reset_password.go | 4 ++-- server/handle_server_update_email.go | 7 +++---- server/middleware.go | 23 +++++++++++------------ 5 file(s) changed, 31 insertion(s)(+), 20 deletion(s)(-) diff --git a/internal/helpers/helpers.go b/internal/helpers/helpers.go --- a/internal/helpers/helpers.go +++ b/internal/helpers/helpers.go @@ -7,6 +7,7 @@ "errors" "math/rand" "net/url" + "github.com/Azure/go-autorest/autorest/to" "github.com/labstack/echo/v4" "github.com/lestrrat-go/jwx/v2/jwk" ) @@ -29,6 +30,18 @@ if suffix != nil { msg += ". " + *suffix } return genericError(e, 400, msg) +} + +func InvalidTokenError(e echo.Context) error { + return InputError(e, to.StringPtr("InvalidToken")) +} + +func ExpiredTokenError(e echo.Context) error { + // WARN: See https://github.com/bluesky-social/atproto/discussions/3319 + return e.JSON(400, map[string]string{ + "error": "ExpiredToken", + "message": "*", + }) } func genericError(e echo.Context, code int, msg string) error { diff --git a/server/handle_server_confirm_email.go b/server/handle_server_confirm_email.go --- a/server/handle_server_confirm_email.go +++ b/server/handle_server_confirm_email.go @@ -28,7 +28,7 @@ return helpers.InputError(e, nil) } if urepo.EmailVerificationCode == nil || urepo.EmailVerificationCodeExpiresAt == nil { - return helpers.InputError(e, to.StringPtr("ExpiredToken")) + return helpers.ExpiredTokenError(e) } if *urepo.EmailVerificationCode != req.Token { @@ -36,7 +36,7 @@ return helpers.InputError(e, to.StringPtr("InvalidToken")) } if time.Now().UTC().After(*urepo.EmailVerificationCodeExpiresAt) { - return helpers.InputError(e, to.StringPtr("ExpiredToken")) + return helpers.ExpiredTokenError(e) } now := time.Now().UTC() diff --git a/server/handle_server_reset_password.go b/server/handle_server_reset_password.go --- a/server/handle_server_reset_password.go +++ b/server/handle_server_reset_password.go @@ -33,11 +33,11 @@ return helpers.InputError(e, to.StringPtr("InvalidToken")) } if *urepo.PasswordResetCode != req.Token { - return helpers.InputError(e, to.StringPtr("InvalidToken")) + return helpers.InvalidTokenError(e) } if time.Now().UTC().After(*urepo.PasswordResetCodeExpiresAt) { - return helpers.InputError(e, to.StringPtr("ExpiredToken")) + return helpers.ExpiredTokenError(e) } hash, err := bcrypt.GenerateFromPassword([]byte(req.Password), 10) diff --git a/server/handle_server_update_email.go b/server/handle_server_update_email.go --- a/server/handle_server_update_email.go +++ b/server/handle_server_update_email.go @@ -3,7 +3,6 @@ import ( "time" - "github.com/Azure/go-autorest/autorest/to" "github.com/haileyok/cocoon/internal/helpers" "github.com/haileyok/cocoon/models" "github.com/labstack/echo/v4" @@ -29,15 +28,15 @@ return helpers.InputError(e, nil) } if urepo.EmailUpdateCode == nil || urepo.EmailUpdateCodeExpiresAt == nil { - return helpers.InputError(e, to.StringPtr("InvalidToken")) + return helpers.InvalidTokenError(e) } if *urepo.EmailUpdateCode != req.Token { - return helpers.InputError(e, to.StringPtr("InvalidToken")) + return helpers.InvalidTokenError(e) } if time.Now().UTC().After(*urepo.EmailUpdateCodeExpiresAt) { - return helpers.InputError(e, to.StringPtr("ExpiredToken")) + return helpers.ExpiredTokenError(e) } if err := s.db.Exec("UPDATE repos SET email_update_code = NULL, email_update_code_expires_at = NULL, email_confirmed_at = NULL, email = ? WHERE did = ?", nil, req.Email, urepo.Repo.Did).Error; err != nil { diff --git a/server/middleware.go b/server/middleware.go --- a/server/middleware.go +++ b/server/middleware.go @@ -54,7 +54,7 @@ tokenstr := pts[1] token, _, err := new(jwt.Parser).ParseUnverified(tokenstr, jwt.MapClaims{}) claims, ok := token.Claims.(jwt.MapClaims) if !ok { - return helpers.InputError(e, to.StringPtr("InvalidToken")) + return helpers.InvalidTokenError(e) } var did string @@ -93,12 +93,11 @@ return s.privateKey.Public(), nil }) if err != nil { s.logger.Error("error parsing jwt", "error", err) - // NOTE: https://github.com/bluesky-social/atproto/discussions/3319 - return e.JSON(400, map[string]string{"error": "ExpiredToken", "message": "token has expired"}) + return helpers.ExpiredTokenError(e) } if !token.Valid { - return helpers.InputError(e, to.StringPtr("InvalidToken")) + return helpers.InvalidTokenError(e) } } else { kpts := strings.Split(tokenstr, ".") @@ -143,9 +142,9 @@ isRefresh := e.Request().URL.Path == "/xrpc/com.atproto.server.refreshSession" scope, _ := claims["scope"].(string) if isRefresh && scope != "com.atproto.refresh" { - return helpers.InputError(e, to.StringPtr("InvalidToken")) + return helpers.InvalidTokenError(e) } else if !hasLxm && !isRefresh && scope != "com.atproto.access" { - return helpers.InputError(e, to.StringPtr("InvalidToken")) + return helpers.InvalidTokenError(e) } table := "tokens" @@ -160,7 +159,7 @@ } var result Result if err := s.db.Raw("SELECT EXISTS(SELECT 1 FROM "+table+" WHERE token = ?) AS found", nil, tokenstr).Scan(&result).Error; err != nil { if err == gorm.ErrRecordNotFound { - return helpers.InputError(e, to.StringPtr("InvalidToken")) + return helpers.InvalidTokenError(e) } s.logger.Error("error getting token from db", "error", err) @@ -168,7 +167,7 @@ return helpers.ServerError(e, nil) } if !result.Found { - return helpers.InputError(e, to.StringPtr("InvalidToken")) + return helpers.InvalidTokenError(e) } } @@ -179,7 +178,7 @@ return helpers.ServerError(e, nil) } if exp < float64(time.Now().UTC().Unix()) { - return helpers.InputError(e, to.StringPtr("ExpiredToken")) + return helpers.ExpiredTokenError(e) } if repo == nil { @@ -197,7 +196,7 @@ e.Set("did", did) e.Set("token", tokenstr) if err := next(e); err != nil { - e.Error(err) + return helpers.InvalidTokenError(e) } return nil @@ -241,7 +240,7 @@ return helpers.InputError(e, nil) } if oauthToken.Token == "" { - return helpers.InputError(e, to.StringPtr("InvalidToken")) + return helpers.InvalidTokenError(e) } if *oauthToken.Parameters.DpopJkt != proof.JKT { @@ -250,7 +249,7 @@ return helpers.InputError(e, to.StringPtr("dpop jkt mismatch")) } if time.Now().After(oauthToken.ExpiresAt) { - return e.JSON(400, map[string]string{"error": "ExpiredToken", "message": "token has expired"}) + return helpers.ExpiredTokenError(e) } repo, err := s.getRepoActorByDid(oauthToken.Sub) -- tangled.sh