package bouncer import ( "context" "encoding/json" "errors" "fmt" "net/http" "strconv" "strings" "github.com/golang-jwt/jwt/v5" ) const ( msgTokenNotProvided = "Token not provided" msgTokenExpired = "Token has expired" msgUserNotFound = "User not found" msgBadSignature = "Token Signature could not be verified." msgMalformed = "Wrong number of segments" msgRequiredClaims = "JWT payload does not contain the required claims" ) // UserProvider loads a persisted user by JWT subject. type UserProvider interface { FindByID(ctx context.Context, id uint) (*Principal, error) } // Middleware validates a pinned HS256 bearer token and loads the user. func Middleware(secret string, users UserProvider) func(http.Handler) http.Handler { return func(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { raw, err := bearerToken(r) if err != nil { write401(w, err.Error()) return } sub, err := Verify(raw, secret) if err != nil { write401(w, err.Error()) return } id, err := strconv.ParseUint(sub, 10, 64) if err != nil || id == 0 { write401(w, msgUserNotFound) return } if users == nil { write401(w, msgUserNotFound) return } user, err := users.FindByID(r.Context(), uint(id)) if err != nil { write401(w, "Authentication error") return } if user == nil { write401(w, msgUserNotFound) return } next.ServeHTTP(w, r.WithContext(WithUser(r.Context(), user))) }) } } type jwtGuard struct { secret string users UserProvider } var ( _ Guard = (*jwtGuard)(nil) _ UnauthorizedWriter = (*jwtGuard)(nil) ) // NewJWTGuard adapts the existing bearerToken -> Verify -> users.FindByID // chain (identical to Middleware's body) into a Guard + UnauthorizedWriter, // so Registry.Middleware("jwt") is byte-identical to bouncer.Middleware. func NewJWTGuard(secret string, users UserProvider) Guard { return &jwtGuard{secret: secret, users: users} } func (g *jwtGuard) Authenticate(r *http.Request) (*Principal, error) { raw, err := bearerToken(r) if err != nil { return nil, err } sub, err := Verify(raw, g.secret) if err != nil { return nil, err } id, err := strconv.ParseUint(sub, 10, 64) if err != nil || id == 0 { return nil, errors.New(msgUserNotFound) } if g.users == nil { return nil, errors.New(msgUserNotFound) } user, err := g.users.FindByID(r.Context(), uint(id)) if err != nil { return nil, errors.New("Authentication error") } if user == nil { return nil, errors.New(msgUserNotFound) } return user, nil } func (g *jwtGuard) WriteUnauthorized(w http.ResponseWriter, err error) { write401(w, err.Error()) } // Verify parses a token with HS256 pinned and a required exp and sub. func Verify(tokenString, secret string) (string, error) { if strings.TrimSpace(secret) == "" { return "", fmt.Errorf("bouncer: jwt secret is empty") } parser := jwt.NewParser(jwt.WithValidMethods([]string{"HS256"}), jwt.WithExpirationRequired()) claims := jwt.MapClaims{} _, err := parser.ParseWithClaims(tokenString, claims, func(t *jwt.Token) (any, error) { return []byte(secret), nil }) if err != nil { return "", mapJWTError(err) } sub := subject(claims) if sub == "" { return "", errors.New(msgRequiredClaims) } return sub, nil } func bearerToken(r *http.Request) (string, error) { h := strings.TrimSpace(r.Header.Get("Authorization")) if h == "" { return "", errors.New(msgTokenNotProvided) } token, ok := strings.CutPrefix(h, "Bearer ") token = strings.TrimSpace(token) if !ok || token == "" { return "", errors.New(msgTokenNotProvided) } return token, nil } func subject(claims jwt.MapClaims) string { switch v := claims["sub"].(type) { case string: return strings.TrimSpace(v) case float64: if v <= 0 { return "" } return strconv.FormatInt(int64(v), 10) case json.Number: return strings.TrimSpace(v.String()) default: return "" } } func mapJWTError(err error) error { switch { case errors.Is(err, jwt.ErrTokenExpired): return errors.New(msgTokenExpired) case errors.Is(err, jwt.ErrTokenSignatureInvalid), errors.Is(err, jwt.ErrTokenUnverifiable): return errors.New(msgBadSignature) case errors.Is(err, jwt.ErrTokenMalformed): return errors.New(msgMalformed) default: msg := err.Error() if strings.Contains(strings.ToLower(msg), "expired") { return errors.New(msgTokenExpired) } if strings.Contains(strings.ToLower(msg), "malformed") || strings.Contains(msg, "segment") { return errors.New(msgMalformed) } return errors.New(msgBadSignature) } } func write401(w http.ResponseWriter, message string) { w.Header().Set("Content-Type", "application/json") w.WriteHeader(http.StatusUnauthorized) _ = json.NewEncoder(w).Encode(map[string]any{"error": true, "message": message}) }