package bouncer import ( "context" "encoding/json" "errors" "fmt" "math" "net/http" "strconv" "strings" "time" "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 bl BlacklistStore cookieNames []string audience string requireAudience bool writeFn func(http.ResponseWriter, error) } var ( _ Guard = (*jwtGuard)(nil) _ UnauthorizedWriter = (*jwtGuard)(nil) ) // NewJWTGuard adapts bearer/cookie extraction, VerifyClaims, and users.FindByID // into a Guard + UnauthorizedWriter. bl may be nil. An empty cookieNames list // is Bearer-only; otherwise each name is tried, in order, after the Authorization header. func NewJWTGuard(secret string, users UserProvider, bl BlacklistStore, cookieNames ...string) Guard { return &jwtGuard{secret: secret, users: users, bl: bl, cookieNames: cookieNames} } // NewBackendJWTGuard requires AudienceBackend. write may replace the // PHP-shaped 401 body; nil keeps write401. With no cookieNames it is // Bearer-only; otherwise each cookie is tried, in order, after the // Authorization header, so a Bearer token still wins when both are sent. func NewBackendJWTGuard(secret string, users UserProvider, bl BlacklistStore, write func(http.ResponseWriter, error), cookieNames ...string) Guard { return &jwtGuard{ secret: secret, users: users, bl: bl, cookieNames: cookieNames, audience: AudienceBackend, requireAudience: true, writeFn: write, } } func (g *jwtGuard) Authenticate(r *http.Request) (*Principal, error) { raw, err := extractToken(r, g.cookieNames) if err != nil { return nil, err } var sub, jti string var iat time.Time if g.requireAudience { sub, iat, _, jti, err = VerifyClaimsAudience(raw, g.secret, g.audience) } else { sub, iat, _, jti, err = VerifyClaims(raw, g.secret) } if err != nil { return nil, err } user, err := subjectPrincipal(r.Context(), g.users, sub) if err != nil { return nil, err } if g.bl != nil { blocked, err := g.bl.IsBlacklisted(r.Context(), jti) if err != nil { return nil, errors.New("Authentication error") } if blocked { return nil, errors.New(msgBadSignature) } } if issuedBeforeCutoff(user, iat) { return nil, ErrSubjectRejected } return user, nil } // ErrSubjectRejected reports that a token's subject is not a loadable user // (a non-numeric or zero sub, a nil provider, or a provider that returned no // principal for a missing, deleted or not-activated user), or that the token // was issued before Principal.TokensValidAfter. Its message is "User not found". var ErrSubjectRejected = errors.New(msgUserNotFound) // errAuthentication reports a provider failure while loading the subject. var errAuthentication = errors.New("Authentication error") // subjectPrincipal loads the token subject through users, shared by the JWT // guard and RefreshAudienceFor. A provider error is errAuthentication; every // other refusal is ErrSubjectRejected. func subjectPrincipal(ctx context.Context, users UserProvider, sub string) (*Principal, error) { id, err := strconv.ParseUint(sub, 10, 64) if err != nil || id == 0 { return nil, ErrSubjectRejected } if users == nil { return nil, ErrSubjectRejected } user, err := users.FindByID(ctx, uint(id)) if err != nil { return nil, errAuthentication } if user == nil { return nil, ErrSubjectRejected } return user, nil } // issuedBeforeCutoff reports whether iat predates the user's // tokens_valid_after cutoff (set by a password reset). A zero cutoff never cuts. func issuedBeforeCutoff(user *Principal, iat time.Time) bool { return !user.TokensValidAfter.IsZero() && iat.Before(user.TokensValidAfter) } func (g *jwtGuard) WriteUnauthorized(w http.ResponseWriter, err error) { if g.writeFn != nil { g.writeFn(w, err) return } write401(w, err.Error()) } // Verify parses a token with HS256 pinned and a required exp and sub. // A missing audience is accepted for PHP-issued frontend tokens. An explicit // audience must be AudienceUser, so a backend token cannot pass this guard. func Verify(tokenString, secret string) (string, error) { sub, _, _, _, err := verifyClaims(tokenString, secret, "") return sub, err } // VerifyClaims parses a token the same way Verify does and also returns iat, exp, and jti. // A missing audience stays valid so PHP-issued frontend tokens keep working. func VerifyClaims(tokenString, secret string) (sub string, iat, exp time.Time, jti string, err error) { return verifyClaims(tokenString, secret, "") } // VerifyClaimsAudience is VerifyClaims plus a required audience claim. func VerifyClaimsAudience(tokenString, secret, audience string) (sub string, iat, exp time.Time, jti string, err error) { if strings.TrimSpace(audience) == "" { return "", time.Time{}, time.Time{}, "", fmt.Errorf("bouncer: jwt audience is empty") } return verifyClaims(tokenString, secret, audience) } func verifyClaims(tokenString, secret, audience string) (sub string, iat, exp time.Time, jti string, err error) { if strings.TrimSpace(secret) == "" { return "", time.Time{}, time.Time{}, "", 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 "", time.Time{}, time.Time{}, "", mapJWTError(err) } if !frontendAudienceOK(claims, audience) { return "", time.Time{}, time.Time{}, "", errors.New(msgBadSignature) } sub = subject(claims) if sub == "" { return "", time.Time{}, time.Time{}, "", errors.New(msgRequiredClaims) } iat, _ = claimTime(claims, "iat") exp, _ = claimTime(claims, "exp") jti, _ = claims["jti"].(string) return sub, iat, exp, jti, nil } // frontendAudienceOK accepts a missing audience when expected is empty // (legacy frontend tokens) and otherwise requires expected. func frontendAudienceOK(claims jwt.MapClaims, expected string) bool { auds := claimAudiences(claims) if expected == "" { return len(auds) == 0 || audienceMatches(claims, AudienceUser) } return audienceMatches(claims, expected) } func audienceMatches(claims jwt.MapClaims, expected string) bool { for _, aud := range claimAudiences(claims) { if aud == expected { return true } } return false } func claimAudiences(claims jwt.MapClaims) []string { switch v := claims["aud"].(type) { case string: if strings.TrimSpace(v) == "" { return nil } return []string{v} case []string: return v case []any: out := make([]string, 0, len(v)) for _, item := range v { s, ok := item.(string) if ok && s != "" { out = append(out, s) } } return out default: return nil } } func extractToken(r *http.Request, cookieNames []string) (string, error) { raw, err := bearerToken(r) if err == nil { return raw, nil } if len(cookieNames) == 0 { return "", err } for _, name := range cookieNames { c, cerr := r.Cookie(name) if cerr != nil { continue } v := strings.TrimSpace(c.Value) if v != "" { return v, nil } } return "", errors.New(msgTokenNotProvided) } func claimTime(claims jwt.MapClaims, key string) (time.Time, bool) { switch v := claims[key].(type) { case float64: if math.IsNaN(v) || math.IsInf(v, 0) { return time.Time{}, false } return time.Unix(int64(v), 0), true case json.Number: n, err := v.Int64() if err != nil { return time.Time{}, false } return time.Unix(n, 0), true case int64: return time.Unix(v, 0), true default: return time.Time{}, false } } 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 math.IsNaN(v) || math.IsInf(v, 0) || v <= 0 || v != math.Trunc(v) || v > 9007199254740992 { return "" } return strconv.FormatInt(int64(v), 10) case json.Number: n, err := strconv.ParseInt(strings.TrimSpace(v.String()), 10, 64) if err != nil || n <= 0 { return "" } return strconv.FormatInt(n, 10) 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") // The reference backend (Laravel) sends this on every response, and a // replayed 401 compares it. w.Header().Set("Cache-Control", "no-cache, private") w.WriteHeader(http.StatusUnauthorized) _ = json.NewEncoder(w).Encode(map[string]any{"error": true, "message": message}) }