From e2c2165c358ca8c89616b2ca04908b08eab15e02 Mon Sep 17 00:00:00 2001 From: kubrickcode Date: Sat, 29 Nov 2025 05:49:53 +0000 Subject: [PATCH] refactor: standardize Go error handling with error chain support - Changed all error comparisons to errors.Is() for error chain support - Defined domain-specific custom error types (Configuration, Token, OAuth, Redis, Crypto) - Added error wrapping to preserve context and improve debugging - Added error chain validation tests fix #56 --- server/api/refresh/index.go | 4 +- server/api/verify/index.go | 3 +- server/pkg/auth/errors.go | 4 +- server/pkg/crypto/random.go | 5 +- server/pkg/errors/types.go | 51 +++++++ server/pkg/errors/types_test.go | 263 ++++++++++++++++++++++++++++++++ server/pkg/jwt/token.go | 12 +- server/pkg/oauth/client.go | 14 +- server/pkg/redis/client.go | 23 ++- 9 files changed, 351 insertions(+), 28 deletions(-) create mode 100644 server/pkg/errors/types.go create mode 100644 server/pkg/errors/types_test.go diff --git a/server/api/refresh/index.go b/server/api/refresh/index.go index 529b7c2..0a8634f 100644 --- a/server/api/refresh/index.go +++ b/server/api/refresh/index.go @@ -1,9 +1,11 @@ package handler import ( + "errors" "net/http" "github-project-status-viewer-server/pkg/crypto" + pkgerrors "github-project-status-viewer-server/pkg/errors" "github-project-status-viewer-server/pkg/httputil" "github-project-status-viewer-server/pkg/jwt" "github-project-status-viewer-server/pkg/oauth" @@ -43,7 +45,7 @@ func Handler(w http.ResponseWriter, r *http.Request) { storedSessionID, err := redisClient.Get(redis.RefreshTokenKeyPrefix + claims.RefreshTokenID) if err != nil { - if err == redis.ErrKeyNotFound { + if errors.Is(err, pkgerrors.ErrKeyNotFound) { httputil.WriteError(w, http.StatusUnauthorized, "refresh_token_revoked", "Refresh token has been revoked or expired") } else { httputil.WriteError(w, http.StatusInternalServerError, "server_error", "Failed to verify refresh token") diff --git a/server/api/verify/index.go b/server/api/verify/index.go index b4f1e89..405db11 100644 --- a/server/api/verify/index.go +++ b/server/api/verify/index.go @@ -4,6 +4,7 @@ import ( "errors" "net/http" + pkgerrors "github-project-status-viewer-server/pkg/errors" "github-project-status-viewer-server/pkg/httputil" "github-project-status-viewer-server/pkg/jwt" "github-project-status-viewer-server/pkg/oauth" @@ -42,7 +43,7 @@ func Handler(w http.ResponseWriter, r *http.Request) { githubAccessToken, err := redisClient.Get(redis.SessionKeyPrefix + claims.SessionID) if err != nil { - if errors.Is(err, redis.ErrKeyNotFound) { + if errors.Is(err, pkgerrors.ErrKeyNotFound) { httputil.WriteError(w, http.StatusUnauthorized, "session_not_found", "Session expired or invalid") } else { httputil.WriteError(w, http.StatusInternalServerError, "redis_error", "Failed to retrieve session") diff --git a/server/pkg/auth/errors.go b/server/pkg/auth/errors.go index 42a6239..972bf76 100644 --- a/server/pkg/auth/errors.go +++ b/server/pkg/auth/errors.go @@ -1,7 +1,7 @@ package auth -import "errors" +import pkgerrors "github-project-status-viewer-server/pkg/errors" var ( - ErrInvalidAuthHeader = errors.New("authorization header must be 'Bearer '") + ErrInvalidAuthHeader = pkgerrors.ErrInvalidAuthHeader ) diff --git a/server/pkg/crypto/random.go b/server/pkg/crypto/random.go index 56b26a9..1fdfaa7 100644 --- a/server/pkg/crypto/random.go +++ b/server/pkg/crypto/random.go @@ -3,6 +3,9 @@ package crypto import ( "crypto/rand" "encoding/hex" + "fmt" + + pkgerrors "github-project-status-viewer-server/pkg/errors" ) const ( @@ -13,7 +16,7 @@ const ( func generateRandomHex(byteLength int) (string, error) { bytes := make([]byte, byteLength) if _, err := rand.Read(bytes); err != nil { - return "", err + return "", fmt.Errorf("%w: %w", pkgerrors.ErrRandomGeneration, err) } return hex.EncodeToString(bytes), nil } diff --git a/server/pkg/errors/types.go b/server/pkg/errors/types.go new file mode 100644 index 0000000..e9b9f4f --- /dev/null +++ b/server/pkg/errors/types.go @@ -0,0 +1,51 @@ +package errors + +import "errors" + +// Configuration errors +var ( + ErrJWTSecretMissing = errors.New("JWT_SECRET not configured") + ErrOAuthConfigMissing = errors.New("OAuth configuration missing") + ErrRedisConfigMissing = errors.New("upstash redis configuration missing") + ErrInvalidAuthHeader = errors.New("authorization header must be 'Bearer '") + ErrMissingAuthCode = errors.New("authorization code is required") + ErrMissingStateParam = errors.New("state parameter is required for CSRF protection") + ErrBearerTokenRequired = errors.New("bearer token required") +) + +// Token errors +var ( + ErrInvalidTokenFormat = errors.New("invalid token format") + ErrTokenExpired = errors.New("token expired") + ErrInvalidSigningMethod = errors.New("unexpected signing method") + ErrInvalidAccessTokenClaims = errors.New("invalid access token claims type") + ErrInvalidRefreshTokenClaims = errors.New("invalid refresh token claims type") + ErrSessionNotFound = errors.New("session not found") + ErrSessionExpired = errors.New("session expired or invalid") + ErrSessionMismatch = errors.New("session mismatch detected") + ErrRefreshTokenRevoked = errors.New("refresh token has been revoked or expired") +) + +// OAuth errors +var ( + ErrOAuthExchangeFailed = errors.New("failed to exchange authorization code") + ErrOAuthRequestFailed = errors.New("OAuth request failed") + ErrAuthenticationFailed = errors.New("authentication failed") +) + +// Redis errors +var ( + ErrKeyNotFound = errors.New("key not found") + ErrRedisRequestFailed = errors.New("redis request failed") + ErrUnexpectedResponse = errors.New("unexpected response type") +) + +// Crypto errors +var ( + ErrRandomGeneration = errors.New("failed to generate random bytes") +) + +// HTTP errors +var ( + ErrMethodNotAllowed = errors.New("method not allowed") +) diff --git a/server/pkg/errors/types_test.go b/server/pkg/errors/types_test.go new file mode 100644 index 0000000..a7e7ef7 --- /dev/null +++ b/server/pkg/errors/types_test.go @@ -0,0 +1,263 @@ +package errors + +import ( + "errors" + "fmt" + "testing" +) + +func TestErrorWrapping(t *testing.T) { + tests := []struct { + baseErr error + name string + wrapContext string + }{ + { + name: "wrap JWT secret error", + baseErr: ErrJWTSecretMissing, + wrapContext: "manager initialization", + }, + { + name: "wrap OAuth config error", + baseErr: ErrOAuthConfigMissing, + wrapContext: "client initialization", + }, + { + name: "wrap session not found error", + baseErr: ErrSessionNotFound, + wrapContext: "verify access token", + }, + { + name: "wrap token expired error", + baseErr: ErrTokenExpired, + wrapContext: "validate refresh token", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + wrappedErr := fmt.Errorf("%s: %w", tt.wrapContext, tt.baseErr) + + if !errors.Is(wrappedErr, tt.baseErr) { + t.Errorf("errors.Is() failed: wrapped error does not match base error") + } + + if wrappedErr.Error() == tt.baseErr.Error() { + t.Errorf("wrapped error message should include context, got %q", wrappedErr.Error()) + } + }) + } +} + +func TestErrorChains(t *testing.T) { + tests := []struct { + buildChain func() error + name string + targetError error + }{ + { + name: "multi-level JWT error chain", + targetError: ErrJWTSecretMissing, + buildChain: func() error { + err := ErrJWTSecretMissing + err = fmt.Errorf("NewManager failed: %w", err) + err = fmt.Errorf("GetManager failed: %w", err) + return err + }, + }, + { + name: "multi-level session error chain", + targetError: ErrSessionNotFound, + buildChain: func() error { + err := ErrSessionNotFound + err = fmt.Errorf("redis.Get failed: %w", err) + err = fmt.Errorf("verify handler failed: %w", err) + return err + }, + }, + { + name: "OAuth error chain", + targetError: ErrOAuthExchangeFailed, + buildChain: func() error { + err := ErrOAuthExchangeFailed + err = fmt.Errorf("requestToken failed: %w", err) + err = fmt.Errorf("callback handler failed: %w", err) + return err + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + chainedErr := tt.buildChain() + + if !errors.Is(chainedErr, tt.targetError) { + t.Errorf("errors.Is() failed: error chain does not contain target error %v", tt.targetError) + } + }) + } +} + +func TestErrorUnwrap(t *testing.T) { + baseErr := ErrSessionNotFound + wrappedErr := fmt.Errorf("context: %w", baseErr) + + unwrapped := errors.Unwrap(wrappedErr) + if unwrapped != baseErr { + t.Errorf("errors.Unwrap() = %v, want %v", unwrapped, baseErr) + } +} + +func TestErrorEquality(t *testing.T) { + tests := []struct { + err1 error + err2 error + name string + wantEq bool + }{ + { + name: "same error instance", + err1: ErrSessionNotFound, + err2: ErrSessionNotFound, + wantEq: true, + }, + { + name: "different error instances", + err1: ErrSessionNotFound, + err2: ErrSessionExpired, + wantEq: false, + }, + { + name: "wrapped vs unwrapped same error", + err1: fmt.Errorf("context: %w", ErrSessionNotFound), + err2: ErrSessionNotFound, + wantEq: false, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + equal := (tt.err1 == tt.err2) + if equal != tt.wantEq { + t.Errorf("error equality = %v, want %v", equal, tt.wantEq) + } + }) + } +} + +func TestErrorIsComparison(t *testing.T) { + tests := []struct { + err error + name string + target error + want bool + }{ + { + name: "exact match", + err: ErrSessionNotFound, + target: ErrSessionNotFound, + want: true, + }, + { + name: "wrapped error matches", + err: fmt.Errorf("context: %w", ErrSessionNotFound), + target: ErrSessionNotFound, + want: true, + }, + { + name: "different errors", + err: ErrSessionNotFound, + target: ErrSessionExpired, + want: false, + }, + { + name: "double wrapped error matches", + err: fmt.Errorf("outer: %w", fmt.Errorf("inner: %w", ErrSessionNotFound)), + target: ErrSessionNotFound, + want: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := errors.Is(tt.err, tt.target) + if got != tt.want { + t.Errorf("errors.Is() = %v, want %v", got, tt.want) + } + }) + } +} + +func TestDefinedErrors(t *testing.T) { + errorGroups := []struct { + name string + errors []error + }{ + { + name: "Configuration Errors", + errors: []error{ + ErrJWTSecretMissing, + ErrOAuthConfigMissing, + ErrRedisConfigMissing, + }, + }, + { + name: "Token Errors", + errors: []error{ + ErrInvalidTokenFormat, + ErrTokenExpired, + ErrInvalidSigningMethod, + ErrInvalidAccessTokenClaims, + ErrInvalidRefreshTokenClaims, + ErrSessionNotFound, + ErrSessionExpired, + ErrSessionMismatch, + ErrRefreshTokenRevoked, + }, + }, + { + name: "Redis Errors", + errors: []error{ + ErrKeyNotFound, + ErrRedisRequestFailed, + ErrUnexpectedResponse, + }, + }, + { + name: "OAuth Errors", + errors: []error{ + ErrOAuthExchangeFailed, + ErrOAuthRequestFailed, + ErrAuthenticationFailed, + }, + }, + { + name: "Crypto Errors", + errors: []error{ + ErrRandomGeneration, + }, + }, + { + name: "HTTP Errors", + errors: []error{ + ErrInvalidAuthHeader, + ErrMissingAuthCode, + ErrMissingStateParam, + ErrBearerTokenRequired, + }, + }, + } + + for _, group := range errorGroups { + t.Run(group.name, func(t *testing.T) { + for _, err := range group.errors { + if err == nil { + t.Errorf("error in group %q should not be nil", group.name) + } + if err.Error() == "" { + t.Errorf("error message for %v in group %q should not be empty", err, group.name) + } + } + }) + } +} diff --git a/server/pkg/jwt/token.go b/server/pkg/jwt/token.go index 91b2011..117ce65 100644 --- a/server/pkg/jwt/token.go +++ b/server/pkg/jwt/token.go @@ -8,6 +8,8 @@ import ( "time" "github.com/golang-jwt/jwt/v5" + + pkgerrors "github-project-status-viewer-server/pkg/errors" ) const ( @@ -46,7 +48,7 @@ func GetManager() (*Manager, error) { func NewManager() (*Manager, error) { secret := os.Getenv("JWT_SECRET") if secret == "" { - return nil, fmt.Errorf("JWT_SECRET not configured") + return nil, pkgerrors.ErrJWTSecretMissing } return &Manager{secret: []byte(secret)}, nil } @@ -83,7 +85,7 @@ func (m *Manager) GenerateRefreshToken(refreshTokenID, sessionID string) (string func (m *Manager) ValidateAccessToken(tokenString string) (*AccessTokenClaims, error) { token, err := jwt.ParseWithClaims(tokenString, &AccessTokenClaims{}, func(token *jwt.Token) (any, error) { if _, ok := token.Method.(*jwt.SigningMethodHMAC); !ok { - return nil, fmt.Errorf("unexpected signing method: %v", token.Header["alg"]) + return nil, fmt.Errorf("%w: %v", pkgerrors.ErrInvalidSigningMethod, token.Header["alg"]) } return m.secret, nil }) @@ -94,7 +96,7 @@ func (m *Manager) ValidateAccessToken(tokenString string) (*AccessTokenClaims, e claims, ok := token.Claims.(*AccessTokenClaims) if !ok { - return nil, fmt.Errorf("invalid access token claims type") + return nil, pkgerrors.ErrInvalidAccessTokenClaims } return claims, nil @@ -103,7 +105,7 @@ func (m *Manager) ValidateAccessToken(tokenString string) (*AccessTokenClaims, e func (m *Manager) ValidateRefreshToken(tokenString string) (*RefreshTokenClaims, error) { token, err := jwt.ParseWithClaims(tokenString, &RefreshTokenClaims{}, func(token *jwt.Token) (any, error) { if _, ok := token.Method.(*jwt.SigningMethodHMAC); !ok { - return nil, fmt.Errorf("unexpected signing method: %v", token.Header["alg"]) + return nil, fmt.Errorf("%w: %v", pkgerrors.ErrInvalidSigningMethod, token.Header["alg"]) } return m.secret, nil }) @@ -114,7 +116,7 @@ func (m *Manager) ValidateRefreshToken(tokenString string) (*RefreshTokenClaims, claims, ok := token.Claims.(*RefreshTokenClaims) if !ok { - return nil, fmt.Errorf("invalid refresh token claims type") + return nil, pkgerrors.ErrInvalidRefreshTokenClaims } return claims, nil diff --git a/server/pkg/oauth/client.go b/server/pkg/oauth/client.go index 696200a..1d7a1af 100644 --- a/server/pkg/oauth/client.go +++ b/server/pkg/oauth/client.go @@ -11,6 +11,8 @@ import ( "strings" "sync" "time" + + pkgerrors "github-project-status-viewer-server/pkg/errors" ) const githubTokenURL = "https://github.com/login/oauth/access_token" @@ -39,7 +41,7 @@ func NewClient() (*Client, error) { clientSecret := os.Getenv("GITHUB_CLIENT_SECRET") if clientID == "" || clientSecret == "" { - return nil, fmt.Errorf("OAuth configuration missing") + return nil, pkgerrors.ErrOAuthConfigMissing } return &Client{ @@ -84,26 +86,26 @@ func (c *Client) requestToken(data url.Values) (*TokenResponse, error) { resp, err := c.HTTPClient.Do(req) if err != nil { - return nil, fmt.Errorf("request failed: %w", err) + return nil, fmt.Errorf("%w: %w", pkgerrors.ErrOAuthRequestFailed, err) } defer resp.Body.Close() body, err := io.ReadAll(resp.Body) if err != nil { - return nil, fmt.Errorf("failed to read response: %w", err) + return nil, fmt.Errorf("failed to read OAuth response: %w", err) } var tokenResp GitHubTokenResponse if err := json.Unmarshal(body, &tokenResp); err != nil { - return nil, fmt.Errorf("failed to parse response: %w", err) + return nil, fmt.Errorf("failed to parse OAuth response: %w", err) } if tokenResp.AccessToken == "" { var errResp GitHubErrorResponse if err := json.Unmarshal(body, &errResp); err == nil && errResp.Error != "" { - return nil, fmt.Errorf("authentication failed: %s - %s", errResp.Error, errResp.ErrorDescription) + return nil, fmt.Errorf("%w: %s - %s", pkgerrors.ErrAuthenticationFailed, errResp.Error, errResp.ErrorDescription) } - return nil, fmt.Errorf("authentication failed and could not parse error response from GitHub") + return nil, fmt.Errorf("%w: could not parse error response from GitHub", pkgerrors.ErrAuthenticationFailed) } return &TokenResponse{ diff --git a/server/pkg/redis/client.go b/server/pkg/redis/client.go index 3db67ce..c5646aa 100644 --- a/server/pkg/redis/client.go +++ b/server/pkg/redis/client.go @@ -3,7 +3,6 @@ package redis import ( "bytes" "encoding/json" - "errors" "fmt" "io" "log/slog" @@ -11,6 +10,8 @@ import ( "os" "sync" "time" + + pkgerrors "github-project-status-viewer-server/pkg/errors" ) const ( @@ -21,8 +22,6 @@ const ( defaultTimeout = 10 * time.Second ) -var ErrKeyNotFound = errors.New("key not found") - type Client struct { baseURL string token string @@ -51,7 +50,7 @@ func NewClient() (*Client, error) { token := os.Getenv("KV_REST_API_TOKEN") if baseURL == "" || token == "" { - return nil, fmt.Errorf("upstash redis configuration missing") + return nil, pkgerrors.ErrRedisConfigMissing } return &Client{ @@ -74,16 +73,16 @@ func (c *Client) Set(key string, value string, expiration time.Duration) error { func (c *Client) Get(key string) (string, error) { result, err := c.execute([]any{"GET", key}) if err != nil { - return "", err + return "", fmt.Errorf("redis get operation failed: %w", err) } if result == nil { - return "", ErrKeyNotFound + return "", pkgerrors.ErrKeyNotFound } str, ok := result.(string) if !ok { - return "", fmt.Errorf("unexpected response type") + return "", fmt.Errorf("%w: expected string, got %T", pkgerrors.ErrUnexpectedResponse, result) } return str, nil @@ -97,12 +96,12 @@ func (c *Client) Delete(key string) error { func (c *Client) Exists(key string) (bool, error) { result, err := c.execute([]any{"EXISTS", key}) if err != nil { - return false, err + return false, fmt.Errorf("redis exists operation failed: %w", err) } count, ok := result.(float64) if !ok { - return false, fmt.Errorf("unexpected response type") + return false, fmt.Errorf("%w: expected float64, got %T", pkgerrors.ErrUnexpectedResponse, result) } return count > 0, nil @@ -131,14 +130,14 @@ func (c *Client) execute(cmd []any) (any, error) { if resp.StatusCode != http.StatusOK { bodyBytes, err := io.ReadAll(resp.Body) if err != nil { - return nil, fmt.Errorf("redis request failed with status %d (failed to read error body: %w)", resp.StatusCode, err) + return nil, fmt.Errorf("%w: status %d (failed to read error body: %w)", pkgerrors.ErrRedisRequestFailed, resp.StatusCode, err) } - return nil, fmt.Errorf("redis request failed with status %d: %s", resp.StatusCode, string(bodyBytes)) + return nil, fmt.Errorf("%w: status %d: %s", pkgerrors.ErrRedisRequestFailed, resp.StatusCode, string(bodyBytes)) } var response upstashResponse if err := json.NewDecoder(resp.Body).Decode(&response); err != nil { - return nil, fmt.Errorf("failed to decode response: %w", err) + return nil, fmt.Errorf("failed to decode redis response: %w", err) } if response.Error != "" {