Files
summercms/bouncer/jwt.go
Jakub Zych dfa00f7e3a feat(09-01): implement separate-admin genre list tracer
- Audience-aware mint, verify, refresh, and backend guard keep frontend tokens compatible
- Cabana mounts raw admin login, list schema, and record list behind admin.jwt.secret
- Framework migration seeds Winter backend users and developer/publisher roles
2026-09-24 17:17:19 +02:00

345 lines
9.2 KiB
Go

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 is bearer-only and requires AudienceBackend.
// write may replace the PHP-shaped 401 body; nil keeps write401.
func NewBackendJWTGuard(secret string, users UserProvider, bl BlacklistStore, write func(http.ResponseWriter, error)) Guard {
return &jwtGuard{
secret: secret,
users: users,
bl: bl,
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
}
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)
}
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 !user.TokensValidAfter.IsZero() && iat.Before(user.TokensValidAfter) {
return nil, errors.New(msgUserNotFound)
}
return user, nil
}
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.
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
}
// 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 audience != "" && !audienceMatches(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
}
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")
w.WriteHeader(http.StatusUnauthorized)
_ = json.NewEncoder(w).Encode(map[string]any{"error": true, "message": message})
}