feat(07-01): add JWT mint, refresh, and blacklist primitives
Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
126
bouncer/blacklist.go
Normal file
126
bouncer/blacklist.go
Normal 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
|
||||||
|
}
|
||||||
@@ -1,13 +1,19 @@
|
|||||||
package bouncer
|
package bouncer
|
||||||
|
|
||||||
import "context"
|
import (
|
||||||
|
"context"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
type userKey struct{}
|
type userKey struct{}
|
||||||
|
|
||||||
// Principal is the authenticated identity stored on the request context.
|
// Principal is the authenticated identity stored on the request context.
|
||||||
|
// PreferredLocale empty means no override. TokensValidAfter zero means no cutoff.
|
||||||
type Principal struct {
|
type Principal struct {
|
||||||
ID uint
|
ID uint
|
||||||
MustChangePassword bool
|
MustChangePassword bool
|
||||||
|
PreferredLocale string
|
||||||
|
TokensValidAfter time.Time
|
||||||
}
|
}
|
||||||
|
|
||||||
// WithUser stores the verified principal on ctx.
|
// WithUser stores the verified principal on ctx.
|
||||||
|
|||||||
@@ -9,6 +9,7 @@ import (
|
|||||||
"net/http"
|
"net/http"
|
||||||
"strconv"
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
|
"time"
|
||||||
|
|
||||||
"github.com/golang-jwt/jwt/v5"
|
"github.com/golang-jwt/jwt/v5"
|
||||||
)
|
)
|
||||||
@@ -67,6 +68,8 @@ func Middleware(secret string, users UserProvider) func(http.Handler) http.Handl
|
|||||||
type jwtGuard struct {
|
type jwtGuard struct {
|
||||||
secret string
|
secret string
|
||||||
users UserProvider
|
users UserProvider
|
||||||
|
bl BlacklistStore
|
||||||
|
cookieNames []string
|
||||||
}
|
}
|
||||||
|
|
||||||
var (
|
var (
|
||||||
@@ -74,19 +77,19 @@ var (
|
|||||||
_ UnauthorizedWriter = (*jwtGuard)(nil)
|
_ UnauthorizedWriter = (*jwtGuard)(nil)
|
||||||
)
|
)
|
||||||
|
|
||||||
// NewJWTGuard adapts the existing bearerToken -> Verify -> users.FindByID
|
// NewJWTGuard adapts bearer/cookie extraction, VerifyClaims, and users.FindByID
|
||||||
// chain (identical to Middleware's body) into a Guard + UnauthorizedWriter,
|
// into a Guard + UnauthorizedWriter. bl may be nil. An empty cookieNames list
|
||||||
// so Registry.Middleware("jwt") is byte-identical to bouncer.Middleware.
|
// is Bearer-only; otherwise each name is tried, in order, after the Authorization header.
|
||||||
func NewJWTGuard(secret string, users UserProvider) Guard {
|
func NewJWTGuard(secret string, users UserProvider, bl BlacklistStore, cookieNames ...string) Guard {
|
||||||
return &jwtGuard{secret: secret, users: users}
|
return &jwtGuard{secret: secret, users: users, bl: bl, cookieNames: cookieNames}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (g *jwtGuard) Authenticate(r *http.Request) (*Principal, error) {
|
func (g *jwtGuard) Authenticate(r *http.Request) (*Principal, error) {
|
||||||
raw, err := bearerToken(r)
|
raw, err := extractToken(r, g.cookieNames)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
sub, err := Verify(raw, g.secret)
|
sub, iat, _, jti, err := VerifyClaims(raw, g.secret)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
@@ -104,6 +107,18 @@ func (g *jwtGuard) Authenticate(r *http.Request) (*Principal, error) {
|
|||||||
if user == nil {
|
if user == nil {
|
||||||
return nil, errors.New(msgUserNotFound)
|
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
|
return user, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -131,6 +146,70 @@ func Verify(tokenString, secret string) (string, error) {
|
|||||||
return sub, nil
|
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) {
|
func bearerToken(r *http.Request) (string, error) {
|
||||||
h := strings.TrimSpace(r.Header.Get("Authorization"))
|
h := strings.TrimSpace(r.Header.Get("Authorization"))
|
||||||
if h == "" {
|
if h == "" {
|
||||||
|
|||||||
57
bouncer/mint.go
Normal file
57
bouncer/mint.go
Normal 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
69
bouncer/refresh.go
Normal 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
|
||||||
|
}
|
||||||
@@ -172,7 +172,7 @@ func TestJWTGuardRegistryMatchesMiddleware(t *testing.T) {
|
|||||||
w.WriteHeader(http.StatusNoContent)
|
w.WriteHeader(http.StatusNoContent)
|
||||||
}))
|
}))
|
||||||
reg := NewRegistry()
|
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)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
viaReg, err := reg.Middleware("jwt")
|
viaReg, err := reg.Middleware("jwt")
|
||||||
|
|||||||
Reference in New Issue
Block a user