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
|
||||
|
||||
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.
|
||||
|
||||
@@ -9,6 +9,7 @@ import (
|
||||
"net/http"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/golang-jwt/jwt/v5"
|
||||
)
|
||||
@@ -67,6 +68,8 @@ func Middleware(secret string, users UserProvider) func(http.Handler) http.Handl
|
||||
type jwtGuard struct {
|
||||
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
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)
|
||||
}))
|
||||
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")
|
||||
|
||||
Reference in New Issue
Block a user