From cad445a23570562a8fa32bd21f9cdacffce41bf0 Mon Sep 17 00:00:00 2001 From: Jakub Zych Date: Tue, 22 Sep 2026 13:34:58 +0200 Subject: [PATCH] feat(07-01): add JWT mint, refresh, and blacklist primitives Co-authored-by: Cursor --- bouncer/blacklist.go | 126 +++++++++++++++++++++++++++++++++++++++ bouncer/context.go | 8 ++- bouncer/jwt.go | 97 +++++++++++++++++++++++++++--- bouncer/mint.go | 57 ++++++++++++++++++ bouncer/refresh.go | 69 +++++++++++++++++++++ bouncer/registry_test.go | 2 +- 6 files changed, 348 insertions(+), 11 deletions(-) create mode 100644 bouncer/blacklist.go create mode 100644 bouncer/mint.go create mode 100644 bouncer/refresh.go diff --git a/bouncer/blacklist.go b/bouncer/blacklist.go new file mode 100644 index 0000000..ac0957d --- /dev/null +++ b/bouncer/blacklist.go @@ -0,0 +1,126 @@ +package bouncer + +import ( + "context" + "database/sql" + "errors" + "fmt" + "regexp" + "sync" + "time" +) + +// BlacklistStore records revoked jti values. IsBlacklisted is true only once +// validUntil has been reached, so a grace window can keep a just-rotated +// token usable. Sweep drops rows whose storage expiry has passed. +type BlacklistStore interface { + Add(ctx context.Context, jti string, expiresAt, validUntil time.Time) error + IsBlacklisted(ctx context.Context, jti string) (bool, error) + Sweep(ctx context.Context, now time.Time) error +} + +type blEntry struct { + expiresAt time.Time + validUntil time.Time +} + +// MemoryBlacklist is an in-process store for tests. Expired rows are dropped +// on read, matching surf.MemoryStore. +type MemoryBlacklist struct { + mu sync.Mutex + entries map[string]blEntry +} + +// NewMemoryBlacklist returns an empty in-process blacklist. +func NewMemoryBlacklist() *MemoryBlacklist { + return &MemoryBlacklist{entries: make(map[string]blEntry)} +} + +func (m *MemoryBlacklist) Add(_ context.Context, jti string, expiresAt, validUntil time.Time) error { + m.mu.Lock() + defer m.mu.Unlock() + m.entries[jti] = blEntry{expiresAt: expiresAt, validUntil: validUntil} + return nil +} + +func (m *MemoryBlacklist) IsBlacklisted(_ context.Context, jti string) (bool, error) { + m.mu.Lock() + defer m.mu.Unlock() + e, ok := m.entries[jti] + if !ok { + return false, nil + } + now := time.Now() + if now.After(e.expiresAt) { + delete(m.entries, jti) + return false, nil + } + return !now.Before(e.validUntil), nil +} + +func (m *MemoryBlacklist) Sweep(_ context.Context, now time.Time) error { + m.mu.Lock() + defer m.mu.Unlock() + for k, e := range m.entries { + if e.expiresAt.Before(now) { + delete(m.entries, k) + } + } + return nil +} + +var blacklistIdent = regexp.MustCompile(`^[a-zA-Z_][a-zA-Z0-9_]*$`) + +// PostgresBlacklist stores revoked jti rows in a caller-supplied table. +// The table name is a plugin constant, still checked so it cannot be interpolated raw. +type PostgresBlacklist struct { + db *sql.DB + table string +} + +// NewPostgresBlacklist binds db to table. table must be a plain SQL identifier. +func NewPostgresBlacklist(db *sql.DB, table string) *PostgresBlacklist { + return &PostgresBlacklist{db: db, table: table} +} + +func (p *PostgresBlacklist) checkTable() error { + if p == nil || !blacklistIdent.MatchString(p.table) { + return fmt.Errorf("bouncer: blacklist table %q is not a safe identifier", p.table) + } + return nil +} + +func (p *PostgresBlacklist) Add(ctx context.Context, jti string, expiresAt, validUntil time.Time) error { + if err := p.checkTable(); err != nil { + return err + } + q := fmt.Sprintf(`INSERT INTO %s (jti, expires_at, valid_until) VALUES ($1, $2, $3) +ON CONFLICT (jti) DO UPDATE SET expires_at = EXCLUDED.expires_at, valid_until = EXCLUDED.valid_until`, p.table) + _, err := p.db.ExecContext(ctx, q, jti, expiresAt, validUntil) + return err +} + +func (p *PostgresBlacklist) IsBlacklisted(ctx context.Context, jti string) (bool, error) { + if err := p.checkTable(); err != nil { + return false, err + } + q := fmt.Sprintf(`SELECT valid_until FROM %s WHERE jti = $1`, p.table) + var validUntil time.Time + err := p.db.QueryRowContext(ctx, q, jti).Scan(&validUntil) + if errors.Is(err, sql.ErrNoRows) { + return false, nil + } + if err != nil { + return false, err + } + return !time.Now().Before(validUntil), nil +} + +func (p *PostgresBlacklist) Sweep(ctx context.Context, now time.Time) error { + if err := p.checkTable(); err != nil { + return err + } + q := fmt.Sprintf(`DELETE FROM %s WHERE expires_at < $1`, p.table) + _, err := p.db.ExecContext(ctx, q, now) + return err +} diff --git a/bouncer/context.go b/bouncer/context.go index fc4d94c..b1b3958 100644 --- a/bouncer/context.go +++ b/bouncer/context.go @@ -1,13 +1,19 @@ package bouncer -import "context" +import ( + "context" + "time" +) type userKey struct{} // Principal is the authenticated identity stored on the request context. +// PreferredLocale empty means no override. TokensValidAfter zero means no cutoff. type Principal struct { ID uint MustChangePassword bool + PreferredLocale string + TokensValidAfter time.Time } // WithUser stores the verified principal on ctx. diff --git a/bouncer/jwt.go b/bouncer/jwt.go index 2feb517..395bbc7 100644 --- a/bouncer/jwt.go +++ b/bouncer/jwt.go @@ -9,6 +9,7 @@ import ( "net/http" "strconv" "strings" + "time" "github.com/golang-jwt/jwt/v5" ) @@ -65,8 +66,10 @@ func Middleware(secret string, users UserProvider) func(http.Handler) http.Handl } type jwtGuard struct { - secret string - users UserProvider + secret string + users UserProvider + bl BlacklistStore + cookieNames []string } var ( @@ -74,19 +77,19 @@ var ( _ 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} +// 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} } func (g *jwtGuard) Authenticate(r *http.Request) (*Principal, error) { - raw, err := bearerToken(r) + raw, err := extractToken(r, g.cookieNames) if err != nil { return nil, err } - sub, err := Verify(raw, g.secret) + sub, iat, _, jti, err := VerifyClaims(raw, g.secret) if err != nil { return nil, err } @@ -104,6 +107,18 @@ func (g *jwtGuard) Authenticate(r *http.Request) (*Principal, 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 } @@ -131,6 +146,70 @@ func Verify(tokenString, secret string) (string, error) { return sub, nil } +// VerifyClaims parses a token the same way Verify does and also returns iat, exp, and jti. +func VerifyClaims(tokenString, secret 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) + } + 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 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 == "" { diff --git a/bouncer/mint.go b/bouncer/mint.go new file mode 100644 index 0000000..8570556 --- /dev/null +++ b/bouncer/mint.go @@ -0,0 +1,57 @@ +package bouncer + +import ( + "crypto/rand" + "encoding/hex" + "fmt" + "strings" + "time" + + "github.com/golang-jwt/jwt/v5" +) + +// prvHash is sha1("Golem15\User\Models\User"), the lock-subject PHP jwt-auth +// stamps on every user token. +const prvHash = "a867434cbc213adfbe78a02bed7082a6bd99c883" + +type registeredClaims struct { + jwt.RegisteredClaims + Prv string `json:"prv,omitempty"` +} + +// Mint signs an HS256 token whose iss is the full URL of the minting endpoint. +// The returned jti is the token's own jti claim. +func Mint(secret, sub, issuerURL string, ttl time.Duration) (string, string, error) { + if strings.TrimSpace(secret) == "" { + return "", "", fmt.Errorf("bouncer: jwt secret is empty") + } + jti, err := newJTI() + if err != nil { + return "", "", err + } + now := time.Now() + claims := registeredClaims{ + RegisteredClaims: jwt.RegisteredClaims{ + Issuer: issuerURL, + Subject: sub, + ExpiresAt: jwt.NewNumericDate(now.Add(ttl)), + NotBefore: jwt.NewNumericDate(now), + IssuedAt: jwt.NewNumericDate(now), + ID: jti, + }, + Prv: prvHash, + } + signed, err := jwt.NewWithClaims(jwt.SigningMethodHS256, claims).SignedString([]byte(secret)) + if err != nil { + return "", "", err + } + return signed, jti, nil +} + +func newJTI() (string, error) { + buf := make([]byte, 16) + if _, err := rand.Read(buf); err != nil { + return "", err + } + return hex.EncodeToString(buf), nil +} diff --git a/bouncer/refresh.go b/bouncer/refresh.go new file mode 100644 index 0000000..c5d1743 --- /dev/null +++ b/bouncer/refresh.go @@ -0,0 +1,69 @@ +package bouncer + +import ( + "context" + "errors" + "strings" + "time" + + "github.com/golang-jwt/jwt/v5" +) + +// Refresh issues a new token for a still-refreshable subject. exp is not +// required to be in the future; iat must fall inside refreshTTL. The previous +// jti is blacklisted with validUntil = now+grace. Storage expiry follows PHP +// jwt-auth: the later of the old exp and iat+refreshTTL, plus one minute, so +// a logged-out token cannot be refreshed again for the rest of its refresh window. +func Refresh(secret, tokenString string, refreshTTL time.Duration, bl BlacklistStore, grace time.Duration, issuerURL string) (string, error) { + if strings.TrimSpace(secret) == "" { + return "", errors.New("bouncer: jwt secret is empty") + } + parser := jwt.NewParser(jwt.WithValidMethods([]string{"HS256"}), jwt.WithoutClaimsValidation()) + claims := jwt.MapClaims{} + if _, err := parser.ParseWithClaims(tokenString, claims, func(*jwt.Token) (any, error) { + return []byte(secret), nil + }); err != nil { + return "", mapJWTError(err) + } + sub := subject(claims) + if sub == "" { + return "", errors.New(msgRequiredClaims) + } + iat, ok := claimTime(claims, "iat") + if !ok || time.Now().After(iat.Add(refreshTTL)) { + return "", errors.New("Token has expired and can no longer be refreshed") + } + jti, _ := claims["jti"].(string) + if bl != nil && jti != "" { + blocked, err := bl.IsBlacklisted(context.Background(), jti) + if err != nil { + return "", err + } + if blocked { + return "", errors.New("The token has been blacklisted") + } + } + exp, expOK := claimTime(claims, "exp") + ttl := exp.Sub(iat) + if !expOK || ttl <= 0 { + return "", errors.New(msgRequiredClaims) + } + next, _, err := Mint(secret, sub, issuerURL, ttl) + if err != nil { + return "", err + } + if bl != nil && jti != "" { + // PHP Blacklist::getMinutesUntilExpired keeps the row until the later of + // exp and iat+refreshTTL, plus one minute. Using exp alone would drop a + // logged-out token whose access exp has passed but whose refresh window + // has not, and the next Refresh would succeed. + expiresAt := iat.Add(refreshTTL).Add(time.Minute) + if until := exp.Add(time.Minute); until.After(expiresAt) { + expiresAt = until + } + if err := bl.Add(context.Background(), jti, expiresAt, time.Now().Add(grace)); err != nil { + return "", err + } + } + return next, nil +} diff --git a/bouncer/registry_test.go b/bouncer/registry_test.go index 4f78790..0d9703f 100644 --- a/bouncer/registry_test.go +++ b/bouncer/registry_test.go @@ -172,7 +172,7 @@ func TestJWTGuardRegistryMatchesMiddleware(t *testing.T) { w.WriteHeader(http.StatusNoContent) })) reg := NewRegistry() - if err := reg.Register("golem15.user", "jwt", NewJWTGuard(secret, users)); err != nil { + if err := reg.Register("golem15.user", "jwt", NewJWTGuard(secret, users, nil)); err != nil { t.Fatal(err) } viaReg, err := reg.Middleware("jwt")