feat(07-01): add JWT mint, refresh, and blacklist primitives

Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
Jakub Zych
2026-09-22 13:34:58 +02:00
parent 251f3cc4a0
commit cad445a235
6 changed files with 348 additions and 11 deletions

126
bouncer/blacklist.go Normal file
View File

@@ -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
}

View File

@@ -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.

View File

@@ -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 == "" {

57
bouncer/mint.go Normal file
View File

@@ -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
}

69
bouncer/refresh.go Normal file
View File

@@ -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
}

View File

@@ -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")