refactor(10.2-01): nest framework packages under modules

- Move remaining beach packages and embedded admin assets\n- Rewrite framework, example, build, and gate paths
This commit is contained in:
Jakub Zych
2026-09-28 02:21:02 +02:00
parent ac1f6d14f4
commit 5e50b166ef
277 changed files with 303 additions and 303 deletions

View File

@@ -0,0 +1,42 @@
package bouncer
import (
"net/http"
"net/http/httptest"
"testing"
"time"
)
func TestAudienceCrossover(t *testing.T) {
const secret = "same-secret-for-both-guards"
users := memUsers{byID: map[uint]*Principal{1: {ID: 1}}}
userTok, _, err := MintAudience(secret, "1", "https://app.test/login", time.Hour, AudienceUser)
if err != nil {
t.Fatal(err)
}
backendTok, _, err := MintAudience(secret, "1", "https://app.test/_admin/api/v1/auth/login", time.Hour, AudienceBackend)
if err != nil {
t.Fatal(err)
}
frontend := NewJWTGuard(secret, users, nil)
backend := NewBackendJWTGuard(secret, users, nil, nil)
if principal, err := backend.Authenticate(withBearer(backendTok)); err != nil || principal == nil || principal.ID != 1 {
t.Fatalf("backend guard rejected its own audience: %v", err)
}
if principal, err := frontend.Authenticate(withBearer(userTok)); err != nil || principal == nil || principal.ID != 1 {
t.Fatalf("frontend guard rejected its own audience: %v", err)
}
if principal, err := backend.Authenticate(withBearer(userTok)); err == nil || principal != nil {
t.Fatal("backend guard accepted a frontend-audience token signed with the same secret")
}
if principal, err := frontend.Authenticate(withBearer(backendTok)); err == nil || principal != nil {
t.Fatal("frontend guard accepted a backend-audience token signed with the same secret")
}
}
func withBearer(token string) *http.Request {
req := httptest.NewRequest(http.MethodGet, "/", nil)
req.Header.Set("Authorization", "Bearer "+token)
return req
}

View File

@@ -0,0 +1,171 @@
package bouncer
import (
"net/http"
"net/http/httptest"
"testing"
"time"
)
// TestPhase09GuardIsolation is the phase gate for token crossover. Each
// subtest fails closed when a frontend and backend credential can be used
// on the other guard, including same-secret audience checks, distinct
// secrets, refresh, blacklist, and password-reset cutoff.
func TestPhase09GuardIsolation(t *testing.T) {
const (
userSecret = "frontend-secret"
backendSecret = "backend-secret"
sharedSecret = "same-secret-for-both-guards"
issuer = "https://app.test"
)
frontendUsers := memUsers{byID: map[uint]*Principal{1: {ID: 1}}}
backendUsers := memUsers{byID: map[uint]*Principal{2: {ID: 2, Backend: true}}}
userTok, _, err := MintAudience(sharedSecret, "1", issuer, time.Hour, AudienceUser)
if err != nil {
t.Fatal(err)
}
backendTok, _, err := MintAudience(sharedSecret, "2", issuer, time.Hour, AudienceBackend)
if err != nil {
t.Fatal(err)
}
frontend := NewJWTGuard(sharedSecret, frontendUsers, nil)
backend := NewBackendJWTGuard(sharedSecret, backendUsers, nil, nil)
t.Run("own audience", func(t *testing.T) {
principal, err := frontend.Authenticate(withBearer(userTok))
if err != nil || principal == nil || principal.ID != 1 {
t.Fatalf("frontend guard rejected its own token: %v %+v", err, principal)
}
principal, err = backend.Authenticate(withBearer(backendTok))
if err != nil || principal == nil || principal.ID != 2 || !principal.Backend {
t.Fatalf("backend guard rejected its own token: %v %+v", err, principal)
}
})
t.Run("audience crossover same secret", func(t *testing.T) {
if principal, err := backend.Authenticate(withBearer(userTok)); err == nil || principal != nil {
t.Fatal("backend guard accepted a frontend-audience token signed with the same secret")
}
if principal, err := frontend.Authenticate(withBearer(backendTok)); err == nil || principal != nil {
t.Fatal("frontend guard accepted a backend-audience token signed with the same secret")
}
})
t.Run("distinct secrets and registries", func(t *testing.T) {
userOnly, _, err := MintAudience(userSecret, "1", issuer, time.Hour, AudienceUser)
if err != nil {
t.Fatal(err)
}
adminOnly, _, err := MintAudience(backendSecret, "2", issuer, time.Hour, AudienceBackend)
if err != nil {
t.Fatal(err)
}
front := NewJWTGuard(userSecret, frontendUsers, nil)
back := NewBackendJWTGuard(backendSecret, backendUsers, nil, nil)
if principal, err := back.Authenticate(withBearer(userOnly)); err == nil || principal != nil {
t.Fatal("backend guard accepted a frontend token signed with the frontend secret")
}
if principal, err := front.Authenticate(withBearer(adminOnly)); err == nil || principal != nil {
t.Fatal("frontend guard accepted a backend token signed with the backend secret")
}
if principal, err := back.Authenticate(withBearer(adminOnly)); err != nil || principal == nil || principal.ID != 2 {
t.Fatalf("backend registry principal = %+v err=%v", principal, err)
}
if principal, err := back.Authenticate(withBearer(mustMint(t, backendSecret, "1", AudienceBackend))); err == nil || principal != nil {
t.Fatal("backend guard resolved a subject from the frontend user registry")
}
reg := NewRegistry()
if err := reg.Register("golem15.user", "jwt", front); err != nil {
t.Fatal(err)
}
if err := reg.Register("summercms.cabana", "backend", back); err != nil {
t.Fatal(err)
}
mw, err := reg.Middleware("backend")
if err != nil {
t.Fatal(err)
}
rec := httptest.NewRecorder()
mw(http.HandlerFunc(func(http.ResponseWriter, *http.Request) {
t.Fatal("backend middleware ran the handler for a frontend token")
})).ServeHTTP(rec, withBearer(userOnly))
if rec.Code != http.StatusUnauthorized {
t.Fatalf("crossover status=%d body=%s", rec.Code, rec.Body.String())
}
})
t.Run("malformed and wrong audience", func(t *testing.T) {
if _, err := backend.Authenticate(withBearer("not-a-jwt")); err == nil {
t.Fatal("malformed token was accepted")
}
other, _, err := MintAudience(sharedSecret, "2", issuer, time.Hour, "frontend")
if err != nil {
t.Fatal(err)
}
if principal, err := backend.Authenticate(withBearer(other)); err == nil || principal != nil {
t.Fatal("backend guard accepted an unexpected audience")
}
})
t.Run("blacklist and reset cutoff", func(t *testing.T) {
bl := NewMemoryBlacklist()
revoked, jti, err := MintAudience(backendSecret, "2", issuer, time.Hour, AudienceBackend)
if err != nil {
t.Fatal(err)
}
if err := bl.Add(t.Context(), jti, time.Now().Add(time.Hour), time.Now().Add(-time.Second)); err != nil {
t.Fatal(err)
}
guard := NewBackendJWTGuard(backendSecret, backendUsers, bl, nil)
if principal, err := guard.Authenticate(withBearer(revoked)); err == nil || principal != nil {
t.Fatal("blacklisted backend token was accepted")
}
if _, err := RefreshAudience(backendSecret, revoked, AudienceBackend, 2*time.Hour, bl, time.Minute, issuer); err == nil {
t.Fatal("blacklisted backend token was refreshed")
}
cutoff := memUsers{byID: map[uint]*Principal{2: {
ID: 2, Backend: true, TokensValidAfter: time.Now().Add(time.Minute),
}}}
stale, _, err := MintAudience(backendSecret, "2", issuer, time.Hour, AudienceBackend)
if err != nil {
t.Fatal(err)
}
if principal, err := NewBackendJWTGuard(backendSecret, cutoff, nil, nil).Authenticate(withBearer(stale)); err == nil || principal != nil {
t.Fatal("token issued before the password-reset cutoff was accepted")
}
})
t.Run("refresh does not cross audiences", func(t *testing.T) {
access, _, err := MintAudience(userSecret, "1", issuer, time.Hour, AudienceUser)
if err != nil {
t.Fatal(err)
}
admin, _, err := MintAudience(backendSecret, "2", issuer, time.Hour, AudienceBackend)
if err != nil {
t.Fatal(err)
}
if _, err := RefreshAudience(userSecret, access, AudienceBackend, 2*time.Hour, nil, time.Minute, issuer); err == nil {
t.Fatal("frontend token refreshed into the backend audience")
}
if _, err := RefreshAudience(backendSecret, admin, AudienceUser, 2*time.Hour, nil, time.Minute, issuer); err == nil {
t.Fatal("backend token refreshed into the frontend audience")
}
if _, err := RefreshAudience(userSecret, admin, AudienceBackend, 2*time.Hour, nil, time.Minute, issuer); err == nil {
t.Fatal("backend token refreshed with the frontend secret")
}
next, err := RefreshAudience(backendSecret, admin, AudienceBackend, 2*time.Hour, nil, time.Minute, issuer)
if err != nil || next == "" {
t.Fatalf("backend refresh failed: %v", err)
}
})
}
func mustMint(t *testing.T, secret, sub, audience string) string {
t.Helper()
token, _, err := MintAudience(secret, sub, "https://app.test", time.Hour, audience)
if err != nil {
t.Fatal(err)
}
return token
}

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

@@ -0,0 +1,61 @@
package bouncer
import (
"testing"
"time"
)
func TestBlacklistGraceWindow(t *testing.T) {
bl := NewMemoryBlacklist()
now := time.Now()
if err := bl.Add(t.Context(), "jti", now.Add(time.Hour), now.Add(time.Minute)); err != nil {
t.Fatal(err)
}
blocked, err := bl.IsBlacklisted(t.Context(), "jti")
if err != nil || blocked {
t.Fatalf("before validUntil: blocked=%t err=%v", blocked, err)
}
if err := bl.Add(t.Context(), "due", now.Add(time.Hour), now.Add(-time.Second)); err != nil {
t.Fatal(err)
}
blocked, err = bl.IsBlacklisted(t.Context(), "due")
if err != nil || !blocked {
t.Fatalf("after validUntil: blocked=%t err=%v", blocked, err)
}
if err := bl.Add(t.Context(), "now", now.Add(time.Hour), time.Now()); err != nil {
t.Fatal(err)
}
blocked, err = bl.IsBlacklisted(t.Context(), "now")
if err != nil || !blocked {
t.Fatalf("at validUntil: blocked=%t err=%v", blocked, err)
}
}
func TestBlacklistSweep(t *testing.T) {
bl := NewMemoryBlacklist()
now := time.Now()
if err := bl.Add(t.Context(), "gone", now.Add(-time.Minute), now.Add(-time.Minute)); err != nil {
t.Fatal(err)
}
if err := bl.Add(t.Context(), "stay", now.Add(time.Hour), now.Add(-time.Second)); err != nil {
t.Fatal(err)
}
if err := bl.Sweep(t.Context(), now); err != nil {
t.Fatal(err)
}
gone, err := bl.IsBlacklisted(t.Context(), "gone")
if err != nil || gone {
t.Fatalf("swept row still present: %t %v", gone, err)
}
stay, err := bl.IsBlacklisted(t.Context(), "stay")
if err != nil || !stay {
t.Fatalf("live row = %t %v", stay, err)
}
}
func TestPostgresBlacklistRejectsUnsafeTable(t *testing.T) {
bl := NewPostgresBlacklist(nil, "user_jwt;drop")
if err := bl.Add(t.Context(), "j", time.Now(), time.Now()); err == nil {
t.Fatal("want identifier error")
}
}

View File

@@ -0,0 +1,61 @@
package bouncer
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.
// IsSuperuser and PermissionGrants are set only for backend-admin principals.
// A grant ending in ".*" matches permission codes by prefix.
type Principal struct {
ID uint
MustChangePassword bool
PreferredLocale string
TokensValidAfter time.Time
Backend bool `json:"-"`
IsSuperuser bool `json:"-"`
PermissionGrants map[string]bool `json:"-"`
}
// WithUser stores the verified principal on ctx.
func WithUser(ctx context.Context, user *Principal) context.Context {
if ctx == nil {
ctx = context.Background()
}
return context.WithValue(ctx, userKey{}, user)
}
// User returns the verified principal from ctx.
func User(ctx context.Context) (*Principal, bool) {
if ctx == nil {
return nil, false
}
u, ok := ctx.Value(userKey{}).(*Principal)
return u, ok && u != nil
}
type credentialKey struct{}
// WithCredential stores the resolved credential (e.g. *ApiToken) on ctx.
func WithCredential(ctx context.Context, cred any) context.Context {
if ctx == nil {
ctx = context.Background()
}
return context.WithValue(ctx, credentialKey{}, cred)
}
// Credential returns the resolved credential from ctx.
func Credential(ctx context.Context) (any, bool) {
if ctx == nil {
return nil, false
}
c := ctx.Value(credentialKey{})
if c == nil {
return nil, false
}
return c, true
}

View File

@@ -0,0 +1,32 @@
package bouncer
import (
"testing"
"time"
)
func TestPrincipalLocaleAndCutoff(t *testing.T) {
when := time.Unix(1_700_000_000, 0)
p := Principal{PreferredLocale: "pl", TokensValidAfter: when}
if p.PreferredLocale != "pl" || !p.TokensValidAfter.Equal(when) {
t.Fatalf("%+v", p)
}
}
func TestContextCredentialRoundTrip(t *testing.T) {
if _, ok := Credential(t.Context()); ok {
t.Fatal("empty context must have no credential")
}
if _, ok := Credential(nil); ok {
t.Fatal("nil context must have no credential")
}
token := &struct{ Name string }{Name: "parity"}
got, ok := Credential(WithCredential(t.Context(), token))
if !ok || got != token {
t.Fatalf("got %#v ok=%t", got, ok)
}
got, ok = Credential(WithCredential(nil, token))
if !ok || got != token {
t.Fatalf("nil ctx got %#v ok=%t", got, ok)
}
}

View File

@@ -0,0 +1,132 @@
package bouncer
import (
"context"
"net/http"
"net/http/httptest"
"testing"
"time"
)
// TestPhase10CookieGuard covers the backend guard's summer_admin cookie
// transport (D-19): the cookie is read only when no Bearer header is sent,
// Bearer wins when both are present, an empty cookie is unauthenticated, and
// the cookie carries no weaker token than the header (audience, blacklist).
func TestPhase10CookieGuard(t *testing.T) {
const (
cookie = "summer_admin"
issuer = "https://app.test/backend"
)
users := memUsers{byID: map[uint]*Principal{
2: {ID: 2, Backend: true},
3: {ID: 3, Backend: true},
}}
withCookie := func(value string) *http.Request {
r := httptest.NewRequest(http.MethodGet, "/backend/api/v1/auth/me", nil)
r.AddCookie(&http.Cookie{Name: cookie, Value: value})
return r
}
mint := func(sub, audience string) (string, string) {
t.Helper()
tok, jti, err := MintAudience(secret, sub, issuer, time.Hour, audience)
if err != nil {
t.Fatal(err)
}
return tok, jti
}
guard := NewBackendJWTGuard(secret, users, nil, nil, cookie)
t.Run("cookie without bearer", func(t *testing.T) {
tok, _ := mint("2", AudienceBackend)
principal, err := guard.Authenticate(withCookie(tok))
if err != nil || principal == nil || principal.ID != 2 {
t.Fatalf("cookie token rejected: %v %+v", err, principal)
}
})
t.Run("cookie value is trimmed", func(t *testing.T) {
tok, _ := mint("2", AudienceBackend)
r := httptest.NewRequest(http.MethodGet, "/", nil)
r.Header.Set("Cookie", cookie+"= "+tok+" ")
if principal, err := guard.Authenticate(r); err != nil || principal == nil {
t.Fatalf("padded cookie rejected: %v", err)
}
})
t.Run("bearer wins over cookie", func(t *testing.T) {
bearerTok, _ := mint("3", AudienceBackend)
cookieTok, _ := mint("2", AudienceBackend)
r := withCookie(cookieTok)
r.Header.Set("Authorization", "Bearer "+bearerTok)
principal, err := guard.Authenticate(r)
if err != nil || principal == nil || principal.ID != 3 {
t.Fatalf("bearer did not win: %v %+v", err, principal)
}
// A bad Bearer is not rescued by a good cookie.
bad := withCookie(cookieTok)
bad.Header.Set("Authorization", "Bearer not-a-token")
if principal, err := guard.Authenticate(bad); err == nil || principal != nil {
t.Fatal("an invalid bearer fell back to the cookie")
}
})
t.Run("empty or missing cookie is unauthenticated", func(t *testing.T) {
for name, r := range map[string]*http.Request{
"empty": withCookie(""),
"blank": withCookie(" "),
"missing": httptest.NewRequest(http.MethodGet, "/", nil),
} {
principal, err := guard.Authenticate(r)
if err == nil || principal != nil || err.Error() != msgTokenNotProvided {
t.Fatalf("%s cookie: principal=%+v err=%v, want %q", name, principal, err, msgTokenNotProvided)
}
}
other := httptest.NewRequest(http.MethodGet, "/", nil)
tok, _ := mint("2", AudienceBackend)
other.AddCookie(&http.Cookie{Name: "summer_other", Value: tok})
if _, err := guard.Authenticate(other); err == nil {
t.Fatal("a token under another cookie name was accepted")
}
})
t.Run("frontend audience in the cookie is rejected", func(t *testing.T) {
tok, _ := mint("2", AudienceUser)
if principal, err := guard.Authenticate(withCookie(tok)); err == nil || principal != nil {
t.Fatal("backend guard accepted a frontend-audience cookie")
}
})
t.Run("blacklisted jti in the cookie is rejected", func(t *testing.T) {
bl := NewMemoryBlacklist()
blocking := NewBackendJWTGuard(secret, users, bl, nil, cookie)
tok, jti := mint("2", AudienceBackend)
if _, err := blocking.Authenticate(withCookie(tok)); err != nil {
t.Fatalf("fresh cookie rejected: %v", err)
}
if err := bl.Add(context.Background(), jti, time.Now().Add(time.Hour), time.Now().Add(-time.Second)); err != nil {
t.Fatal(err)
}
principal, err := blocking.Authenticate(withCookie(tok))
if err == nil || principal != nil || err.Error() != msgBadSignature {
t.Fatalf("blacklisted cookie: principal=%+v err=%v", principal, err)
}
})
t.Run("bearer-only guard ignores the cookie", func(t *testing.T) {
tok, _ := mint("2", AudienceBackend)
bearerOnly := NewBackendJWTGuard(secret, users, nil, nil)
if principal, err := bearerOnly.Authenticate(withCookie(tok)); err == nil || principal != nil {
t.Fatal("a guard without cookie names read the cookie")
}
})
t.Run("second cookie name is tried after an empty first", func(t *testing.T) {
tok, _ := mint("2", AudienceBackend)
two := NewBackendJWTGuard(secret, users, nil, nil, "legacy_admin", cookie)
r := withCookie(tok)
r.AddCookie(&http.Cookie{Name: "legacy_admin", Value: ""})
if principal, err := two.Authenticate(r); err != nil || principal == nil {
t.Fatalf("second cookie name not tried: %v", err)
}
})
}

25
modules/bouncer/guard.go Normal file
View File

@@ -0,0 +1,25 @@
package bouncer
import "net/http"
// Guard resolves the caller's Principal for r, or an error describing why not.
type Guard interface {
Authenticate(r *http.Request) (*Principal, error)
}
// CredentialGuard resolves the Principal AND its underlying credential (e.g.
// *models.ApiToken) in one pass -- a DB-backed guard must never verify twice
// per request (RESEARCH.md Pitfall 5: last_used_at must stamp once).
type CredentialGuard interface {
AuthenticateCredential(r *http.Request) (*Principal, any, error)
}
// UnauthorizedWriter lets a guard write its own failure response. jwtGuard
// implements this (reusing write401's {"error":true,"message":...} shape).
// TokenGuard does NOT implement it: PHP's TokenScope, not ApiTokenGuard, owns
// the {"error":"Invalid token"} 401 body (D-08) -- Registry.Middleware must
// pass an unauthenticated request through untouched when a guard has no
// UnauthorizedWriter, leaving denial to downstream middleware.
type UnauthorizedWriter interface {
WriteUnauthorized(w http.ResponseWriter, err error)
}

371
modules/bouncer/jwt.go Normal file
View File

@@ -0,0 +1,371 @@
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 requires AudienceBackend. write may replace the
// PHP-shaped 401 body; nil keeps write401. With no cookieNames it is
// Bearer-only; otherwise each cookie is tried, in order, after the
// Authorization header, so a Bearer token still wins when both are sent.
func NewBackendJWTGuard(secret string, users UserProvider, bl BlacklistStore, write func(http.ResponseWriter, error), cookieNames ...string) Guard {
return &jwtGuard{
secret: secret,
users: users,
bl: bl,
cookieNames: cookieNames,
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
}
user, err := subjectPrincipal(r.Context(), g.users, sub)
if err != nil {
return nil, err
}
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 issuedBeforeCutoff(user, iat) {
return nil, ErrSubjectRejected
}
return user, nil
}
// ErrSubjectRejected reports that a token's subject is not a loadable user
// (a non-numeric or zero sub, a nil provider, or a provider that returned no
// principal for a missing, deleted or not-activated user), or that the token
// was issued before Principal.TokensValidAfter. Its message is "User not found".
var ErrSubjectRejected = errors.New(msgUserNotFound)
// errAuthentication reports a provider failure while loading the subject.
var errAuthentication = errors.New("Authentication error")
// subjectPrincipal loads the token subject through users, shared by the JWT
// guard and RefreshAudienceFor. A provider error is errAuthentication; every
// other refusal is ErrSubjectRejected.
func subjectPrincipal(ctx context.Context, users UserProvider, sub string) (*Principal, error) {
id, err := strconv.ParseUint(sub, 10, 64)
if err != nil || id == 0 {
return nil, ErrSubjectRejected
}
if users == nil {
return nil, ErrSubjectRejected
}
user, err := users.FindByID(ctx, uint(id))
if err != nil {
return nil, errAuthentication
}
if user == nil {
return nil, ErrSubjectRejected
}
return user, nil
}
// issuedBeforeCutoff reports whether iat predates the user's
// tokens_valid_after cutoff (set by a password reset). A zero cutoff never cuts.
func issuedBeforeCutoff(user *Principal, iat time.Time) bool {
return !user.TokensValidAfter.IsZero() && iat.Before(user.TokensValidAfter)
}
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.
// A missing audience is accepted for PHP-issued frontend tokens. An explicit
// audience must be AudienceUser, so a backend token cannot pass this guard.
func Verify(tokenString, secret string) (string, error) {
sub, _, _, _, err := verifyClaims(tokenString, secret, "")
return sub, err
}
// 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 !frontendAudienceOK(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
}
// frontendAudienceOK accepts a missing audience when expected is empty
// (legacy frontend tokens) and otherwise requires expected.
func frontendAudienceOK(claims jwt.MapClaims, expected string) bool {
auds := claimAudiences(claims)
if expected == "" {
return len(auds) == 0 || audienceMatches(claims, AudienceUser)
}
return audienceMatches(claims, expected)
}
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})
}

View File

@@ -0,0 +1,128 @@
package bouncer
import (
"encoding/json"
"errors"
"net/http"
"net/http/httptest"
"testing"
"time"
"github.com/golang-jwt/jwt/v5"
)
func TestJWTGuardCookieFallback(t *testing.T) {
users := memUsers{byID: map[uint]*Principal{1: {ID: 1}}}
tok := sign(t, jwt.SigningMethodHS256, jwt.MapClaims{
"sub": "1",
"jti": "cookie-jti",
"exp": time.Now().Add(time.Hour).Unix(),
"iat": time.Now().Unix(),
}, []byte(secret))
g := NewJWTGuard(secret, users, nil, "token", "auth_token")
req := httptest.NewRequest(http.MethodGet, "/", nil)
req.AddCookie(&http.Cookie{Name: "token", Value: tok})
p, err := g.Authenticate(req)
if err != nil || p == nil || p.ID != 1 {
t.Fatalf("token cookie: %+v %v", p, err)
}
req = httptest.NewRequest(http.MethodGet, "/", nil)
req.AddCookie(&http.Cookie{Name: "auth_token", Value: tok})
p, err = g.Authenticate(req)
if err != nil || p == nil || p.ID != 1 {
t.Fatalf("auth_token cookie: %+v %v", p, err)
}
}
func TestJWTGuardBearerOnlyIgnoresCookies(t *testing.T) {
users := memUsers{byID: map[uint]*Principal{1: {ID: 1}}}
tok := sign(t, jwt.SigningMethodHS256, jwt.MapClaims{
"sub": "1",
"exp": time.Now().Add(time.Hour).Unix(),
}, []byte(secret))
g := NewJWTGuard(secret, users, nil)
req := httptest.NewRequest(http.MethodGet, "/", nil)
req.AddCookie(&http.Cookie{Name: "token", Value: tok})
if _, err := g.Authenticate(req); err == nil || err.Error() != msgTokenNotProvided {
t.Fatalf("cookie on bearer-only guard: %v", err)
}
}
func TestJWTGuardBlacklist(t *testing.T) {
users := memUsers{byID: map[uint]*Principal{1: {ID: 1}}}
tok := sign(t, jwt.SigningMethodHS256, jwt.MapClaims{
"sub": "1",
"jti": "revoked",
"exp": time.Now().Add(time.Hour).Unix(),
"iat": time.Now().Unix(),
}, []byte(secret))
bl := NewMemoryBlacklist()
if err := bl.Add(t.Context(), "revoked", time.Now().Add(time.Hour), time.Now().Add(-time.Second)); err != nil {
t.Fatal(err)
}
g := NewJWTGuard(secret, users, bl)
req := httptest.NewRequest(http.MethodGet, "/", nil)
req.Header.Set("Authorization", "Bearer "+tok)
_, err := g.Authenticate(req)
if err == nil || err.Error() != msgBadSignature {
t.Fatalf("blacklisted: %v", err)
}
reg := NewRegistry()
if err := reg.Register("golem15.user", "jwt", g); err != nil {
t.Fatal(err)
}
mw, err := reg.Middleware("jwt")
if err != nil {
t.Fatal(err)
}
rec := httptest.NewRecorder()
mw(http.HandlerFunc(func(http.ResponseWriter, *http.Request) {
t.Fatal("handler ran")
})).ServeHTTP(rec, req)
if rec.Code != http.StatusUnauthorized {
t.Fatalf("status = %d", rec.Code)
}
var body map[string]any
if err := json.Unmarshal(rec.Body.Bytes(), &body); err != nil {
t.Fatal(err)
}
if body["error"] != true || body["message"] != msgBadSignature {
t.Fatalf("body = %v", body)
}
}
func TestJWTGuardTokensValidAfter(t *testing.T) {
tok := sign(t, jwt.SigningMethodHS256, jwt.MapClaims{
"sub": "1",
"jti": "live",
"exp": time.Now().Add(time.Hour).Unix(),
"iat": time.Now().Add(-time.Minute).Unix(),
}, []byte(secret))
req := httptest.NewRequest(http.MethodGet, "/", nil)
req.Header.Set("Authorization", "Bearer "+tok)
cutoff := memUsers{byID: map[uint]*Principal{1: {ID: 1, TokensValidAfter: time.Now().Add(time.Minute)}}}
if _, err := NewJWTGuard(secret, cutoff, nil).Authenticate(req); err == nil || err.Error() != msgUserNotFound {
t.Fatalf("after cutoff: %v", err)
}
// The helper extraction keeps the guard's message and makes the refusal
// matchable, the same sentinel RefreshAudienceFor returns.
if _, err := NewJWTGuard(secret, cutoff, nil).Authenticate(req); err == nil || err.Error() != "User not found" || !errors.Is(err, ErrSubjectRejected) {
t.Fatalf("cutoff refusal = %v, want \"User not found\" matching ErrSubjectRejected", err)
}
open := memUsers{byID: map[uint]*Principal{1: {ID: 1, TokensValidAfter: time.Now().Add(-time.Hour)}}}
p, err := NewJWTGuard(secret, open, nil).Authenticate(req)
if err != nil || p == nil || p.ID != 1 {
t.Fatalf("before cutoff: %+v %v", p, err)
}
zero := memUsers{byID: map[uint]*Principal{1: {ID: 1}}}
p, err = NewJWTGuard(secret, zero, nil).Authenticate(req)
if err != nil || p == nil || p.ID != 1 {
t.Fatalf("zero cutoff: %+v %v", p, err)
}
}

294
modules/bouncer/jwt_test.go Normal file
View File

@@ -0,0 +1,294 @@
package bouncer
import (
"context"
"encoding/base64"
"encoding/json"
"net/http"
"net/http/httptest"
"strings"
"testing"
"time"
"github.com/golang-jwt/jwt/v5"
)
const secret = "test-secret"
type memUsers struct {
byID map[uint]*Principal
err error
}
func (m memUsers) FindByID(ctx context.Context, id uint) (*Principal, error) {
if m.err != nil {
return nil, m.err
}
return m.byID[id], nil
}
func sign(t *testing.T, method jwt.SigningMethod, claims jwt.MapClaims, key []byte) string {
t.Helper()
tok := jwt.NewWithClaims(method, claims)
s, err := tok.SignedString(key)
if err != nil {
t.Fatal(err)
}
return s
}
func TestVerifyRejectsBadTokens(t *testing.T) {
valid := sign(t, jwt.SigningMethodHS256, jwt.MapClaims{
"sub": "1",
"exp": time.Now().Add(time.Hour).Unix(),
}, []byte(secret))
if sub, err := Verify(valid, secret); err != nil || sub != "1" {
t.Fatalf("valid token: %s %v", sub, err)
}
expired := sign(t, jwt.SigningMethodHS256, jwt.MapClaims{
"sub": "1",
"exp": time.Now().Add(-time.Hour).Unix(),
}, []byte(secret))
if _, err := Verify(expired, secret); err == nil || err.Error() != msgTokenExpired {
t.Fatalf("expired: %v", err)
}
none := sign(t, jwt.SigningMethodHS384, jwt.MapClaims{
"sub": "1",
"exp": time.Now().Add(time.Hour).Unix(),
}, []byte(secret))
if _, err := Verify(none, secret); err == nil || err.Error() != msgBadSignature {
t.Fatalf("wrong alg: %v", err)
}
badSig := sign(t, jwt.SigningMethodHS256, jwt.MapClaims{
"sub": "1",
"exp": time.Now().Add(time.Hour).Unix(),
}, []byte("other-secret"))
if _, err := Verify(badSig, secret); err == nil || err.Error() != msgBadSignature {
t.Fatalf("bad sig: %v", err)
}
if _, err := Verify("not-a-jwt", secret); err == nil || err.Error() != msgMalformed {
t.Fatalf("malformed: %v", err)
}
missingSub := sign(t, jwt.SigningMethodHS256, jwt.MapClaims{
"exp": time.Now().Add(time.Hour).Unix(),
}, []byte(secret))
if _, err := Verify(missingSub, secret); err == nil || err.Error() != msgRequiredClaims {
t.Fatalf("missing sub: %v", err)
}
}
func TestMiddlewareStatusBodies(t *testing.T) {
users := memUsers{byID: map[uint]*Principal{1: {ID: 1}}}
h := Middleware(secret, users)(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("X-Hit", "1")
w.WriteHeader(http.StatusNoContent)
}))
assert401 := func(t *testing.T, req *http.Request, msg string) {
t.Helper()
rec := httptest.NewRecorder()
h.ServeHTTP(rec, req)
if rec.Code != http.StatusUnauthorized {
t.Fatalf("status = %d body %s", rec.Code, rec.Body.String())
}
if rec.Header().Get("X-Hit") != "" {
t.Fatal("handler ran")
}
var body map[string]any
if err := json.Unmarshal(rec.Body.Bytes(), &body); err != nil {
t.Fatal(err)
}
if body["error"] != true || body["message"] != msg {
t.Fatalf("body = %v want %s", body, msg)
}
}
t.Run("missing", func(t *testing.T) {
assert401(t, httptest.NewRequest(http.MethodGet, "/", nil), msgTokenNotProvided)
})
t.Run("unknown-user", func(t *testing.T) {
tok := sign(t, jwt.SigningMethodHS256, jwt.MapClaims{
"sub": "99",
"exp": time.Now().Add(time.Hour).Unix(),
}, []byte(secret))
req := httptest.NewRequest(http.MethodGet, "/", nil)
req.Header.Set("Authorization", "Bearer "+tok)
assert401(t, req, msgUserNotFound)
})
t.Run("valid", func(t *testing.T) {
tok := sign(t, jwt.SigningMethodHS256, jwt.MapClaims{
"sub": "1",
"exp": time.Now().Add(time.Hour).Unix(),
}, []byte(secret))
req := httptest.NewRequest(http.MethodGet, "/", nil)
req.Header.Set("Authorization", "Bearer "+tok)
rec := httptest.NewRecorder()
h.ServeHTTP(rec, req)
if rec.Code != http.StatusNoContent || rec.Header().Get("X-Hit") != "1" {
t.Fatalf("status=%d hit=%s", rec.Code, rec.Header().Get("X-Hit"))
}
})
t.Run("malformed", func(t *testing.T) {
req := httptest.NewRequest(http.MethodGet, "/", nil)
req.Header.Set("Authorization", "Bearer not-a-jwt")
assert401(t, req, msgMalformed)
})
t.Run("alg-none", func(t *testing.T) {
req := httptest.NewRequest(http.MethodGet, "/", nil)
req.Header.Set("Authorization", "Bearer "+noneToken(t, jwt.MapClaims{
"sub": "1",
"exp": time.Now().Add(time.Hour).Unix(),
}))
assert401(t, req, msgBadSignature)
})
t.Run("absent-exp", func(t *testing.T) {
tok := sign(t, jwt.SigningMethodHS256, jwt.MapClaims{"sub": "1"}, []byte(secret))
req := httptest.NewRequest(http.MethodGet, "/", nil)
req.Header.Set("Authorization", "Bearer "+tok)
assert401(t, req, msgBadSignature)
})
}
func TestVerifyRejectsAlgNoneEmptySecretAndAbsentExp(t *testing.T) {
valid := sign(t, jwt.SigningMethodHS256, jwt.MapClaims{
"sub": "1",
"exp": time.Now().Add(time.Hour).Unix(),
}, []byte(secret))
none := noneToken(t, jwt.MapClaims{
"sub": "1",
"exp": time.Now().Add(time.Hour).Unix(),
})
if _, err := Verify(none, secret); err == nil || err.Error() != msgBadSignature {
t.Fatalf("alg none: %v", err)
}
if _, err := Verify(valid, ""); err == nil || !strings.Contains(err.Error(), "jwt secret is empty") {
t.Fatalf("empty secret: %v", err)
}
if _, err := Verify(valid, " "); err == nil || !strings.Contains(err.Error(), "jwt secret is empty") {
t.Fatalf("blank secret: %v", err)
}
noExp := sign(t, jwt.SigningMethodHS256, jwt.MapClaims{"sub": "1"}, []byte(secret))
if sub, err := Verify(noExp, secret); err == nil || sub != "" {
t.Fatalf("absent exp must fail, got %q %v", sub, err)
}
}
func TestVerifyAndMiddlewareOmitTokenAndSecret(t *testing.T) {
const leakSecret = "unique-hs256-secret-value-9f3a"
tok := sign(t, jwt.SigningMethodHS256, jwt.MapClaims{
"sub": "1",
"exp": time.Now().Add(time.Hour).Unix(),
}, []byte(leakSecret))
assertClean := func(t *testing.T, msg string) {
t.Helper()
if strings.Contains(msg, leakSecret) || strings.Contains(msg, tok) {
t.Fatalf("leaked secret or token: %s", msg)
}
}
if _, err := Verify("not-a-jwt", leakSecret); err == nil {
t.Fatal("want malformed")
} else {
assertClean(t, err.Error())
}
if _, err := Verify(tok, "other-"+leakSecret); err == nil {
t.Fatal("want bad signature")
} else {
assertClean(t, err.Error())
}
h := Middleware(leakSecret, memUsers{})(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
t.Fatal("handler must not run")
}))
req := httptest.NewRequest(http.MethodGet, "/", nil)
req.Header.Set("Authorization", "Bearer "+tok)
rec := httptest.NewRecorder()
h.ServeHTTP(rec, req)
if rec.Code != http.StatusUnauthorized {
t.Fatalf("status = %d", rec.Code)
}
assertClean(t, rec.Body.String())
}
func TestContextUserRoundTrip(t *testing.T) {
if _, ok := User(t.Context()); ok {
t.Fatal("empty context must have no user")
}
p := &Principal{ID: 7, MustChangePassword: true}
got, ok := User(WithUser(t.Context(), p))
if !ok || got != p || got.ID != 7 || !got.MustChangePassword {
t.Fatalf("got %+v ok=%t", got, ok)
}
}
func noneToken(t *testing.T, claims jwt.MapClaims) string {
t.Helper()
payload, err := json.Marshal(claims)
if err != nil {
t.Fatal(err)
}
header := base64.RawURLEncoding.EncodeToString([]byte(`{"alg":"none","typ":"JWT"}`))
body := base64.RawURLEncoding.EncodeToString(payload)
return header + "." + body + "."
}
func TestVerifyRejectsFractionalSubject(t *testing.T) {
exp := time.Now().Add(time.Hour).Unix()
tok := sign(t, jwt.SigningMethodHS256, jwt.MapClaims{"sub": 12.5, "exp": exp}, []byte(secret))
if sub, err := Verify(tok, secret); err == nil || sub != "" {
t.Fatalf("fractional sub accepted: %q %v", sub, err)
}
tok = sign(t, jwt.SigningMethodHS256, jwt.MapClaims{"sub": 12, "exp": exp}, []byte(secret))
if sub, err := Verify(tok, secret); err != nil || sub != "12" {
t.Fatalf("whole sub: %q %v", sub, err)
}
}
func TestVerifySubjectMatrix(t *testing.T) {
exp := time.Now().Add(time.Hour).Unix()
cases := []struct {
name string
sub any
want string
}{
{"fractional", 12.5, ""},
{"huge float", 1e300, ""},
{"2^60 float", float64(1 << 60), ""},
{"negative", -1, ""},
{"zero", 0, ""},
{"whole number", 12, "12"},
{"string", "12", "12"},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
tok := sign(t, jwt.SigningMethodHS256, jwt.MapClaims{"sub": tc.sub, "exp": exp}, []byte(secret))
sub, err := Verify(tok, secret)
if tc.want == "" {
if err == nil || sub != "" {
t.Fatalf("sub %v accepted: %q", tc.sub, sub)
}
return
}
if err != nil || sub != tc.want {
t.Fatalf("sub %v: %q %v", tc.sub, sub, err)
}
})
}
}
func TestSubjectJSONNumber(t *testing.T) {
for in, want := range map[string]string{"1.5": "", "-3": "", "0": "", "12": "12"} {
if got := subject(jwt.MapClaims{"sub": json.Number(in)}); got != want {
t.Errorf("json.Number(%q) = %q, want %q", in, got, want)
}
}
}

78
modules/bouncer/mint.go Normal file
View File

@@ -0,0 +1,78 @@
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 frontend user token. Backend tokens do not carry it.
const prvHash = "a867434cbc213adfbe78a02bed7082a6bd99c883"
const (
// AudienceUser is the frontend jwt-guard audience.
AudienceUser = "user"
// AudienceBackend is the admin jwt-guard audience.
AudienceBackend = "backend"
)
type registeredClaims struct {
jwt.RegisteredClaims
Prv string `json:"prv,omitempty"`
}
// Mint signs a frontend-audience HS256 token. iss is the minting endpoint URL.
// The returned jti is the token's own jti claim.
func Mint(secret, sub, issuerURL string, ttl time.Duration) (string, string, error) {
return MintAudience(secret, sub, issuerURL, ttl, AudienceUser)
}
// MintAudience signs an HS256 token for audience. Frontend tokens keep the PHP
// prv lock-subject; backend tokens omit it.
func MintAudience(secret, sub, issuerURL string, ttl time.Duration, audience string) (string, string, error) {
if strings.TrimSpace(secret) == "" {
return "", "", fmt.Errorf("bouncer: jwt secret is empty")
}
if strings.TrimSpace(audience) == "" {
return "", "", fmt.Errorf("bouncer: jwt audience is empty")
}
jti, err := newJTI()
if err != nil {
return "", "", err
}
now := time.Now()
prv := ""
if audience == AudienceUser {
prv = prvHash
}
claims := registeredClaims{
RegisteredClaims: jwt.RegisteredClaims{
Issuer: issuerURL,
Subject: sub,
Audience: jwt.ClaimStrings{audience},
ExpiresAt: jwt.NewNumericDate(now.Add(ttl)),
NotBefore: jwt.NewNumericDate(now),
IssuedAt: jwt.NewNumericDate(now),
ID: jti,
},
Prv: prv,
}
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
}

View File

@@ -0,0 +1,62 @@
package bouncer
import (
"testing"
"time"
"github.com/golang-jwt/jwt/v5"
)
func TestMintClaims(t *testing.T) {
issuer := "https://app.test/_user/api/v1/login"
before := time.Now().Add(-2 * time.Second)
token, jti, err := Mint(secret, "42", issuer, 60*time.Minute)
after := time.Now().Add(2 * time.Second)
if err != nil {
t.Fatal(err)
}
if jti == "" {
t.Fatal("empty jti")
}
claims := decodeClaims(t, token, secret)
if claims["iss"] != issuer || claims["sub"] != "42" || claims["prv"] != "a867434cbc213adfbe78a02bed7082a6bd99c883" {
t.Fatalf("claims = %#v", claims)
}
if claims["jti"] != jti {
t.Fatalf("jti claim %v != returned %s", claims["jti"], jti)
}
iat := claimUnix(t, claims, "iat")
nbf := claimUnix(t, claims, "nbf")
exp := claimUnix(t, claims, "exp")
if iat.Before(before) || iat.After(after) || nbf.Before(before) || nbf.After(after) {
t.Fatalf("iat=%s nbf=%s want near now", iat, nbf)
}
if exp.Before(iat.Add(59*time.Minute)) || exp.After(iat.Add(61*time.Minute)) {
t.Fatalf("exp=%s iat=%s", exp, iat)
}
sub, err := Verify(token, secret)
if err != nil || sub != "42" {
t.Fatalf("Verify = %q %v", sub, err)
}
}
func decodeClaims(t *testing.T, token, key string) jwt.MapClaims {
t.Helper()
parser := jwt.NewParser(jwt.WithValidMethods([]string{"HS256"}), jwt.WithoutClaimsValidation())
claims := jwt.MapClaims{}
if _, err := parser.ParseWithClaims(token, claims, func(*jwt.Token) (any, error) {
return []byte(key), nil
}); err != nil {
t.Fatal(err)
}
return claims
}
func claimUnix(t *testing.T, claims jwt.MapClaims, key string) time.Time {
t.Helper()
v, ok := claims[key].(float64)
if !ok {
t.Fatalf("%s = %#v", key, claims[key])
}
return time.Unix(int64(v), 0)
}

View File

@@ -0,0 +1,27 @@
package bouncer
import "golang.org/x/crypto/bcrypt"
// HashPassword returns a bcrypt hash of plain at the given cost.
func HashPassword(cost int, plain string) (string, error) {
b, err := bcrypt.GenerateFromPassword([]byte(plain), cost)
if err != nil {
return "", err
}
return string(b), nil
}
// CheckPassword reports whether plain matches hash. A malformed hash returns false.
func CheckPassword(hash, plain string) bool {
return bcrypt.CompareHashAndPassword([]byte(hash), []byte(plain)) == nil
}
// NeedsRehash reports whether hash was produced below configuredCost.
// A hash bcrypt cannot parse needs a rehash.
func NeedsRehash(hash string, configuredCost int) bool {
cost, err := bcrypt.Cost([]byte(hash))
if err != nil {
return true
}
return cost < configuredCost
}

View File

@@ -0,0 +1,46 @@
package bouncer
import "testing"
func TestPasswordHashAndCheck(t *testing.T) {
hash, err := HashPassword(10, "secret")
if err != nil {
t.Fatal(err)
}
if !CheckPassword(hash, "secret") {
t.Fatal("matching password rejected")
}
if CheckPassword(hash, "wrong") {
t.Fatal("wrong password accepted")
}
if CheckPassword("not-a-hash", "secret") {
t.Fatal("malformed hash must not panic or match")
}
}
func TestPasswordAcceptsPHPHash(t *testing.T) {
// php -r 'echo password_hash("golem15-a1-check", PASSWORD_BCRYPT, ["cost"=>10]);'
const phpHash = "$2y$10$vvjEAuqFJXs6lWVy1eo5FuTZZr84LrP8Oz2c6pzAw4f2pk6u2xV5W"
if !CheckPassword(phpHash, "golem15-a1-check") {
t.Fatal("PHP $2y$ hash rejected")
}
if CheckPassword(phpHash, "other") {
t.Fatal("PHP hash matched the wrong password")
}
}
func TestNeedsRehash(t *testing.T) {
hash, err := HashPassword(10, "secret")
if err != nil {
t.Fatal(err)
}
if !NeedsRehash(hash, 12) {
t.Fatal("lower cost must need rehash")
}
if NeedsRehash(hash, 10) || NeedsRehash(hash, 8) {
t.Fatal("equal or higher cost must not need rehash")
}
if !NeedsRehash("not-a-hash", 10) {
t.Fatal("unparseable hash must need rehash")
}
}

View File

@@ -0,0 +1,115 @@
package bouncer
import (
"context"
"database/sql"
"sync"
"testing"
"time"
_ "github.com/jackc/pgx/v5/stdlib"
"github.com/testcontainers/testcontainers-go"
"github.com/testcontainers/testcontainers-go/modules/postgres"
)
func TestMintRefreshBlacklistRoundTrip(t *testing.T) {
const secret = "summercms-test-only-hs256-secret"
bl := NewMemoryBlacklist()
token, jti, err := Mint(secret, "7", "http://example.com/_user/api/v1/login", time.Hour)
if err != nil {
t.Fatal(err)
}
sub, err := Verify(token, secret)
if err != nil || sub != "7" {
t.Fatalf("verify %q %v", sub, err)
}
next, err := Refresh(secret, token, 2*time.Hour, bl, 0, "http://example.com/_user/api/v1/refresh")
if err != nil {
t.Fatal(err)
}
if _, err := Verify(next, secret); err != nil {
t.Fatal(err)
}
blocked, err := bl.IsBlacklisted(context.Background(), jti)
if err != nil || !blocked {
t.Fatalf("old jti blacklisted=%t err=%v", blocked, err)
}
if _, err := Refresh(secret, token, 2*time.Hour, bl, 0, "http://example.com/_user/api/v1/refresh"); err == nil {
t.Fatal("blacklisted token refreshed")
}
}
func TestMemoryBlacklistConcurrent(t *testing.T) {
bl := NewMemoryBlacklist()
var wg sync.WaitGroup
for i := 0; i < 8; i++ {
wg.Add(1)
go func() {
defer wg.Done()
now := time.Now()
if err := bl.Add(context.Background(), "jti", now.Add(time.Hour), now.Add(-time.Second)); err != nil {
t.Error(err)
return
}
if _, err := bl.IsBlacklisted(context.Background(), "jti"); err != nil {
t.Error(err)
}
}()
}
wg.Wait()
}
func TestPostgresBlacklistConcurrent(t *testing.T) {
if testing.Short() {
t.Skip("postgres blacklist concurrency needs testcontainers")
}
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Minute)
defer cancel()
ctr, err := postgres.Run(ctx, "postgres:16-alpine",
postgres.WithDatabase("bl"),
postgres.WithUsername("bl"),
postgres.WithPassword("bl"),
postgres.BasicWaitStrategies(),
testcontainers.WithEnv(map[string]string{
"POSTGRES_INITDB_ARGS": "--locale-provider=icu --icu-locale=pl-PL --encoding=UTF8",
}),
)
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = ctr.Terminate(context.Background()) })
dsn, err := ctr.ConnectionString(ctx, "sslmode=disable")
if err != nil {
t.Fatal(err)
}
db, err := sql.Open("pgx", dsn)
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = db.Close() })
if _, err := db.ExecContext(ctx, `CREATE TABLE jwt_blacklist (
jti text PRIMARY KEY,
expires_at timestamptz NOT NULL,
valid_until timestamptz NOT NULL
)`); err != nil {
t.Fatal(err)
}
bl := NewPostgresBlacklist(db, "jwt_blacklist")
var wg sync.WaitGroup
for i := 0; i < 8; i++ {
wg.Add(1)
go func() {
defer wg.Done()
now := time.Now()
if err := bl.Add(ctx, "shared", now.Add(time.Hour), now.Add(-time.Second)); err != nil {
t.Error(err)
return
}
blocked, err := bl.IsBlacklisted(ctx, "shared")
if err != nil || !blocked {
t.Errorf("blocked=%t err=%v", blocked, err)
}
}()
}
wg.Wait()
}

121
modules/bouncer/refresh.go Normal file
View File

@@ -0,0 +1,121 @@
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) {
return refreshAudience(secret, tokenString, AudienceUser, true, refreshTTL, bl, grace, issuerURL, nil)
}
// RefreshAudience reissues a token that already carries audience. Missing aud is rejected.
func RefreshAudience(secret, tokenString, audience string, refreshTTL time.Duration, bl BlacklistStore, grace time.Duration, issuerURL string) (string, error) {
if strings.TrimSpace(audience) == "" {
return "", errors.New("bouncer: jwt audience is empty")
}
return refreshAudience(secret, tokenString, audience, false, refreshTTL, bl, grace, issuerURL, nil)
}
// RefreshAudienceFor is RefreshAudience plus the JWT guard's subject checks
// before minting: the subject is loaded through users, and a missing,
// deleted or not-activated user, or a token issued before the user's
// TokensValidAfter cutoff, is refused. The lookup runs only after every
// token-only check (signature, audience, refresh window, blacklist, exp)
// passed, and a refused subject neither mints a token nor blacklists the old
// jti. Only subject refusals match errors.Is(err, ErrSubjectRejected); a
// provider failure returns a different error.
func RefreshAudienceFor(ctx context.Context, users UserProvider, secret, tokenString, audience string, refreshTTL time.Duration, bl BlacklistStore, grace time.Duration, issuerURL string) (string, error) {
if strings.TrimSpace(audience) == "" {
return "", errors.New("bouncer: jwt audience is empty")
}
check := func(sub string, iat time.Time) error {
user, err := subjectPrincipal(ctx, users, sub)
if err != nil {
return err
}
if issuedBeforeCutoff(user, iat) {
return ErrSubjectRejected
}
return nil
}
return refreshAudience(secret, tokenString, audience, false, refreshTTL, bl, grace, issuerURL, check)
}
// refreshAudience holds the shared refresh flow. check, when non-nil, runs
// after every token-only check and immediately before minting.
func refreshAudience(secret, tokenString, audience string, allowMissing bool, refreshTTL time.Duration, bl BlacklistStore, grace time.Duration, issuerURL string, check func(sub string, iat time.Time) error) (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)
}
auds := claimAudiences(claims)
if len(auds) == 0 {
if !allowMissing {
return "", errors.New(msgBadSignature)
}
} else if !audienceMatches(claims, audience) {
return "", errors.New(msgBadSignature)
}
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)
}
if check != nil {
if err := check(sub, iat); err != nil {
return "", err
}
}
next, _, err := MintAudience(secret, sub, issuerURL, ttl, audience)
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

@@ -0,0 +1,261 @@
package bouncer
import (
"context"
"errors"
"sync/atomic"
"testing"
"time"
"github.com/golang-jwt/jwt/v5"
)
func TestRefreshWithinWindow(t *testing.T) {
old := sign(t, jwt.SigningMethodHS256, jwt.MapClaims{
"sub": "42",
"iss": "https://app.test/_user/api/v1/login",
"prv": "a867434cbc213adfbe78a02bed7082a6bd99c883",
"jti": "old-jti",
"iat": time.Now().Add(-10 * time.Minute).Unix(),
"nbf": time.Now().Add(-10 * time.Minute).Unix(),
"exp": time.Now().Add(-time.Minute).Unix(),
}, []byte(secret))
next, err := Refresh(secret, old, time.Hour, NewMemoryBlacklist(), 0, "https://app.test/_user/api/v1/refresh")
if err != nil {
t.Fatal(err)
}
claims := decodeClaims(t, next, secret)
if claims["sub"] != "42" || claims["prv"] != "a867434cbc213adfbe78a02bed7082a6bd99c883" {
t.Fatalf("claims = %#v", claims)
}
if claims["iss"] != "https://app.test/_user/api/v1/refresh" {
t.Fatalf("iss = %v", claims["iss"])
}
if claims["jti"] == "old-jti" || claims["jti"] == "" {
t.Fatalf("jti = %v", claims["jti"])
}
iat := claimUnix(t, claims, "iat")
if time.Since(iat) > 5*time.Second {
t.Fatalf("iat not fresh: %s", iat)
}
}
func TestRefreshPastTTL(t *testing.T) {
old := sign(t, jwt.SigningMethodHS256, jwt.MapClaims{
"sub": "42",
"jti": "stale",
"iat": time.Now().Add(-2 * time.Hour).Unix(),
"exp": time.Now().Add(-time.Hour).Unix(),
}, []byte(secret))
if _, err := Refresh(secret, old, time.Hour, NewMemoryBlacklist(), 0, "https://app.test/_user/api/v1/refresh"); err == nil {
t.Fatal("want error past refresh TTL")
}
}
func TestRefreshRejectsBadSignature(t *testing.T) {
old := sign(t, jwt.SigningMethodHS256, jwt.MapClaims{
"sub": "42",
"jti": "x",
"iat": time.Now().Unix(),
"exp": time.Now().Add(time.Hour).Unix(),
}, []byte("other-secret"))
if _, err := Refresh(secret, old, time.Hour, nil, 0, "https://app.test/_user/api/v1/refresh"); err == nil {
t.Fatal("want signature error")
}
}
func TestRefreshBlacklistsOldJTI(t *testing.T) {
old := sign(t, jwt.SigningMethodHS256, jwt.MapClaims{
"sub": "42",
"jti": "rotate-me",
"iat": time.Now().Add(-time.Minute).Unix(),
"exp": time.Now().Add(-time.Second).Unix(),
}, []byte(secret))
bl := NewMemoryBlacklist()
if _, err := Refresh(secret, old, time.Hour, bl, 0, "https://app.test/_user/api/v1/refresh"); err != nil {
t.Fatal(err)
}
blocked, err := bl.IsBlacklisted(t.Context(), "rotate-me")
if err != nil || !blocked {
t.Fatalf("grace 0 blacklisted=%t err=%v", blocked, err)
}
old2 := sign(t, jwt.SigningMethodHS256, jwt.MapClaims{
"sub": "42",
"jti": "grace-me",
"iat": time.Now().Add(-time.Minute).Unix(),
"exp": time.Now().Add(-time.Second).Unix(),
}, []byte(secret))
bl2 := NewMemoryBlacklist()
if _, err := Refresh(secret, old2, time.Hour, bl2, time.Hour, "https://app.test/_user/api/v1/refresh"); err != nil {
t.Fatal(err)
}
blocked, err = bl2.IsBlacklisted(t.Context(), "grace-me")
if err != nil || blocked {
t.Fatalf("inside grace blacklisted=%t err=%v", blocked, err)
}
forever := sign(t, jwt.SigningMethodHS256, jwt.MapClaims{
"sub": "42",
"jti": "logged-out",
"iat": time.Now().Unix(),
"exp": time.Now().Add(time.Hour).Unix(),
}, []byte(secret))
if err := bl.Add(t.Context(), "logged-out", time.Now().Add(2*time.Hour), time.Now().Add(-time.Second)); err != nil {
t.Fatal(err)
}
if _, err := Refresh(secret, forever, time.Hour, bl, 0, "https://app.test/_user/api/v1/refresh"); err == nil {
t.Fatal("blacklisted token must not refresh")
}
}
// countingUsers wraps a provider and counts FindByID calls, proving which
// refusals happen before the subject lookup.
type countingUsers struct {
inner UserProvider
calls *atomic.Int32
}
func (c countingUsers) FindByID(ctx context.Context, id uint) (*Principal, error) {
c.calls.Add(1)
return c.inner.FindByID(ctx, id)
}
func TestRefreshAudienceForSubject(t *testing.T) {
const issuer = "https://app.test/backend/api/v1/auth/refresh"
iat := time.Now().Add(-10 * time.Minute).Truncate(time.Second)
token := func(t *testing.T, sub, jti, aud string, iat time.Time, key string) string {
t.Helper()
claims := jwt.MapClaims{
"sub": sub,
"jti": jti,
"iat": iat.Unix(),
"nbf": iat.Unix(),
"exp": iat.Add(5 * time.Minute).Unix(),
}
if aud != "" {
claims["aud"] = aud
}
return sign(t, jwt.SigningMethodHS256, claims, []byte(key))
}
backend := func(t *testing.T, sub, jti string) string {
t.Helper()
return token(t, sub, jti, AudienceBackend, iat, secret)
}
users := func(p *Principal) memUsers {
return memUsers{byID: map[uint]*Principal{7: p}}
}
blacklisted := func(t *testing.T, bl BlacklistStore, jti string) bool {
t.Helper()
blocked, err := bl.IsBlacklisted(context.Background(), jti)
if err != nil {
t.Fatal(err)
}
return blocked
}
t.Run("active principal without a cutoff refreshes", func(t *testing.T) {
bl := NewMemoryBlacklist()
next, err := RefreshAudienceFor(context.Background(), users(&Principal{ID: 7, Backend: true}), secret, backend(t, "7", "active"), AudienceBackend, time.Hour, bl, 0, issuer)
if err != nil {
t.Fatal(err)
}
claims := decodeClaims(t, next, secret)
if claims["sub"] != "7" || !audienceMatches(claims, AudienceBackend) {
t.Fatalf("claims = %#v", claims)
}
if !blacklisted(t, bl, "active") {
t.Fatal("old jti was not blacklisted after a successful refresh")
}
})
t.Run("token issued before the cutoff is a subject refusal", func(t *testing.T) {
bl := NewMemoryBlacklist()
p := &Principal{ID: 7, TokensValidAfter: iat.Add(time.Minute)}
next, err := RefreshAudienceFor(context.Background(), users(p), secret, backend(t, "7", "cut"), AudienceBackend, time.Hour, bl, 0, issuer)
if !errors.Is(err, ErrSubjectRejected) || next != "" {
t.Fatalf("pre-cutoff refresh = %q, %v; want ErrSubjectRejected", next, err)
}
if err.Error() != msgUserNotFound {
t.Fatalf("message = %q", err.Error())
}
if blacklisted(t, bl, "cut") {
t.Fatal("a refused subject blacklisted the old jti")
}
})
t.Run("token issued after the cutoff refreshes", func(t *testing.T) {
p := &Principal{ID: 7, TokensValidAfter: iat.Add(-time.Second)}
if _, err := RefreshAudienceFor(context.Background(), users(p), secret, backend(t, "7", "after"), AudienceBackend, time.Hour, NewMemoryBlacklist(), 0, issuer); err != nil {
t.Fatalf("post-cutoff refresh: %v", err)
}
})
t.Run("missing or not-activated principal is a subject refusal", func(t *testing.T) {
bl := NewMemoryBlacklist()
_, err := RefreshAudienceFor(context.Background(), memUsers{byID: map[uint]*Principal{}}, secret, backend(t, "7", "gone"), AudienceBackend, time.Hour, bl, 0, issuer)
if !errors.Is(err, ErrSubjectRejected) {
t.Fatalf("missing principal: %v", err)
}
if blacklisted(t, bl, "gone") {
t.Fatal("a refused subject blacklisted the old jti")
}
})
t.Run("nil provider is a subject refusal", func(t *testing.T) {
_, err := RefreshAudienceFor(context.Background(), nil, secret, backend(t, "7", "nil-users"), AudienceBackend, time.Hour, NewMemoryBlacklist(), 0, issuer)
if !errors.Is(err, ErrSubjectRejected) {
t.Fatalf("nil users: %v", err)
}
})
t.Run("non-numeric subject is a subject refusal", func(t *testing.T) {
_, err := RefreshAudienceFor(context.Background(), users(&Principal{ID: 7}), secret, backend(t, "seven", "nan"), AudienceBackend, time.Hour, NewMemoryBlacklist(), 0, issuer)
if !errors.Is(err, ErrSubjectRejected) {
t.Fatalf("non-numeric sub: %v", err)
}
})
t.Run("provider error is not a subject refusal", func(t *testing.T) {
bl := NewMemoryBlacklist()
_, err := RefreshAudienceFor(context.Background(), memUsers{err: errors.New("db down")}, secret, backend(t, "7", "db-error"), AudienceBackend, time.Hour, bl, 0, issuer)
if err == nil || errors.Is(err, ErrSubjectRejected) {
t.Fatalf("provider error: %v, want a non-subject error", err)
}
if blacklisted(t, bl, "db-error") {
t.Fatal("a provider error blacklisted the old jti")
}
})
t.Run("token-only refusals never reach the provider", func(t *testing.T) {
preBlocked := NewMemoryBlacklist()
if err := preBlocked.Add(context.Background(), "blocked", time.Now().Add(time.Hour), time.Now()); err != nil {
t.Fatal(err)
}
cases := map[string]struct {
tok string
bl BlacklistStore
}{
"outside the refresh window": {tok: token(t, "7", "stale", AudienceBackend, time.Now().Add(-2*time.Hour), secret), bl: NewMemoryBlacklist()},
"frontend audience": {tok: token(t, "7", "front", AudienceUser, iat, secret), bl: NewMemoryBlacklist()},
"wrong secret": {tok: token(t, "7", "forged", AudienceBackend, iat, "other-secret"), bl: NewMemoryBlacklist()},
"already blacklisted": {tok: backend(t, "7", "blocked"), bl: preBlocked},
}
for name, tc := range cases {
calls := &atomic.Int32{}
provider := countingUsers{inner: users(&Principal{ID: 7}), calls: calls}
if _, err := RefreshAudienceFor(context.Background(), provider, secret, tc.tok, AudienceBackend, time.Hour, tc.bl, 0, issuer); err == nil {
t.Fatalf("%s: refresh succeeded", name)
}
if n := calls.Load(); n != 0 {
t.Fatalf("%s: provider called %d times before the token checks refused", name, n)
}
}
})
t.Run("empty audience is rejected", func(t *testing.T) {
if _, err := RefreshAudienceFor(context.Background(), users(&Principal{ID: 7}), secret, backend(t, "7", "no-aud"), " ", time.Hour, nil, 0, issuer); err == nil {
t.Fatal("empty audience accepted")
}
})
}

103
modules/bouncer/registry.go Normal file
View File

@@ -0,0 +1,103 @@
package bouncer
import (
"errors"
"fmt"
"net/http"
"reflect"
)
type namedGuard struct {
pluginID string
g any
}
// Registry stores named Guard / CredentialGuard implementations and derives
// auth middleware from them.
type Registry struct {
guards map[string]namedGuard
}
// NewRegistry returns an empty named-guard registry.
func NewRegistry() *Registry {
return &Registry{guards: make(map[string]namedGuard)}
}
// Register stores g under name. g must implement Guard or CredentialGuard.
// Empty name, nil g, a type implementing neither, or a duplicate name all
// fail with a "bouncer: ..." error naming pluginID and name.
func (reg *Registry) Register(pluginID, name string, g any) error {
if reg == nil {
return fmt.Errorf("bouncer: registry is nil")
}
if name == "" || g == nil {
return fmt.Errorf("bouncer: plugin %q registered empty guard %q", pluginID, name)
}
switch rv := reflect.ValueOf(g); rv.Kind() {
case reflect.Pointer, reflect.Map, reflect.Slice, reflect.Func, reflect.Chan, reflect.Interface:
if rv.IsNil() {
return fmt.Errorf("bouncer: plugin %q registered empty guard %q", pluginID, name)
}
}
_, isGuard := g.(Guard)
_, isCred := g.(CredentialGuard)
if !isGuard && !isCred {
return fmt.Errorf("bouncer: plugin %q registered guard %q that implements neither Guard nor CredentialGuard", pluginID, name)
}
if existing, ok := reg.guards[name]; ok {
return fmt.Errorf("bouncer: guard %q already registered by %s", name, existing.pluginID)
}
if reg.guards == nil {
reg.guards = make(map[string]namedGuard)
}
reg.guards[name] = namedGuard{pluginID: pluginID, g: g}
return nil
}
// Middleware derives an http middleware from a registered guard. Unknown
// names fail (fail boot, mirrors surf.RegisterMiddleware's contract).
// On Authenticate/AuthenticateCredential success: WithUser (+WithCredential
// if a credential was returned) then next.ServeHTTP.
// On failure: if the guard implements UnauthorizedWriter, it writes the
// response and the chain stops; otherwise next.ServeHTTP runs unauthenticated.
func (reg *Registry) Middleware(name string) (func(http.Handler) http.Handler, error) {
if reg == nil {
return nil, fmt.Errorf("bouncer: registry is nil")
}
ng, ok := reg.guards[name]
if !ok {
return nil, fmt.Errorf("bouncer: unknown guard %q", name)
}
return func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
principal, cred, err := authenticate(ng.g, r)
if err != nil || principal == nil {
if wtr, ok := ng.g.(UnauthorizedWriter); ok {
if err == nil {
err = errors.New("unauthenticated")
}
wtr.WriteUnauthorized(w, err)
return
}
next.ServeHTTP(w, r)
return
}
ctx := WithUser(r.Context(), principal)
if cred != nil {
ctx = WithCredential(ctx, cred)
}
next.ServeHTTP(w, r.WithContext(ctx))
})
}, nil
}
func authenticate(g any, r *http.Request) (*Principal, any, error) {
if cg, ok := g.(CredentialGuard); ok {
return cg.AuthenticateCredential(r)
}
if gd, ok := g.(Guard); ok {
p, err := gd.Authenticate(r)
return p, nil, err
}
return nil, nil, fmt.Errorf("bouncer: guard implements neither Guard nor CredentialGuard")
}

View File

@@ -0,0 +1,85 @@
package bouncer
import (
"net/http"
"net/http/httptest"
"strings"
"testing"
)
// Gap (a): Registry.Register of a type implementing neither Guard nor
// CredentialGuard is already asserted by TestRegisterNeitherInterfaceNamesPluginAndName.
// This file covers the remaining Register fail-loud branches and the
// authenticate() default that Register itself makes unreachable through
// the public API.
func TestRegisterNilRegistryAndEmptyGuard(t *testing.T) {
var nilReg *Registry
if err := nilReg.Register("golem15.demo", "jwt", writerGuard{}); err == nil || !strings.Contains(err.Error(), "nil") {
t.Fatalf("nil registry: %v", err)
}
reg := NewRegistry()
if err := reg.Register("golem15.demo", "", writerGuard{}); err == nil || !strings.Contains(err.Error(), "golem15.demo") {
t.Fatalf("empty name: %v", err)
}
if err := reg.Register("golem15.demo", "jwt", nil); err == nil || !strings.Contains(err.Error(), "jwt") {
t.Fatalf("nil guard: %v", err)
}
}
func TestAuthenticateDefaultNeitherInterface(t *testing.T) {
// Register rejects this type; authenticate's default is only reachable
// by calling it directly (same-package coverage of the fail-closed branch).
p, cred, err := authenticate(notAGuard{}, httptest.NewRequest(http.MethodGet, "/", nil))
if p != nil || cred != nil {
t.Fatalf("principal=%v cred=%v", p, cred)
}
if err == nil || !strings.Contains(err.Error(), "neither Guard nor CredentialGuard") {
t.Fatalf("want neither-interface error, got %v", err)
}
}
func TestNilRegistryMiddleware(t *testing.T) {
var nilReg *Registry
if _, err := nilReg.Middleware("jwt"); err == nil || !strings.Contains(err.Error(), "nil") {
t.Fatalf("nil registry Middleware: %v", err)
}
}
func TestGuardAuthenticateSuccessAttachesUser(t *testing.T) {
reg := NewRegistry()
principal := &Principal{ID: 7}
if err := reg.Register("golem15.user", "jwt", writerGuard{principal: principal}); err != nil {
t.Fatal(err)
}
mw, err := reg.Middleware("jwt")
if err != nil {
t.Fatal(err)
}
h := mw(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
got, ok := User(r.Context())
if !ok || got != principal {
t.Fatalf("user = %+v ok=%t", got, ok)
}
if _, ok := Credential(r.Context()); ok {
t.Fatal("plain Guard must not attach a credential")
}
w.WriteHeader(http.StatusNoContent)
}))
rec := httptest.NewRecorder()
h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/", nil))
if rec.Code != http.StatusNoContent {
t.Fatalf("status = %d", rec.Code)
}
}
func TestWithUserNilContext(t *testing.T) {
p := &Principal{ID: 1}
got, ok := User(WithUser(nil, p))
if !ok || got != p {
t.Fatalf("got %+v ok=%t", got, ok)
}
if _, ok := User(nil); ok {
t.Fatal("nil context must have no user")
}
}

View File

@@ -0,0 +1,291 @@
package bouncer
import (
"errors"
"net/http"
"net/http/httptest"
"strings"
"testing"
"time"
"github.com/golang-jwt/jwt/v5"
)
type writerGuard struct {
principal *Principal
err error
wrote *bool
}
func (g writerGuard) Authenticate(*http.Request) (*Principal, error) {
return g.principal, g.err
}
func (g writerGuard) WriteUnauthorized(w http.ResponseWriter, err error) {
if g.wrote != nil {
*g.wrote = true
}
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(http.StatusUnauthorized)
_, _ = w.Write([]byte(`{"from":"writer"}`))
}
type credOnlyGuard struct {
principal *Principal
cred any
err error
}
func (g credOnlyGuard) AuthenticateCredential(*http.Request) (*Principal, any, error) {
return g.principal, g.cred, g.err
}
type notAGuard struct{}
func TestDuplicateGuardNameFailsWithPluginAndName(t *testing.T) {
reg := NewRegistry()
if err := reg.Register("golem15.user", "jwt", writerGuard{}); err != nil {
t.Fatal(err)
}
err := reg.Register("golem15.acme", "jwt", writerGuard{})
if err == nil || !strings.Contains(err.Error(), "golem15.user") || !strings.Contains(err.Error(), "jwt") {
t.Fatalf("want plugin and guard name in error, got %v", err)
}
}
func TestUnknownGuardNameFails(t *testing.T) {
reg := NewRegistry()
_, err := reg.Middleware("jwt")
if err == nil || !strings.Contains(err.Error(), "jwt") {
t.Fatalf("want guard name in error, got %v", err)
}
}
func TestRegisterNeitherInterfaceNamesPluginAndName(t *testing.T) {
reg := NewRegistry()
err := reg.Register("golem15.demo", "oops", notAGuard{})
if err == nil || !strings.Contains(err.Error(), "golem15.demo") || !strings.Contains(err.Error(), "oops") {
t.Fatalf("want plugin and guard name in error, got %v", err)
}
}
func TestWriterGuardFailureWritesOwnResponse(t *testing.T) {
wrote := false
reg := NewRegistry()
if err := reg.Register("golem15.user", "jwt", writerGuard{
err: errors.New("nope"),
wrote: &wrote,
}); err != nil {
t.Fatal(err)
}
mw, err := reg.Middleware("jwt")
if err != nil {
t.Fatal(err)
}
called := false
h := mw(http.HandlerFunc(func(http.ResponseWriter, *http.Request) {
called = true
}))
rec := httptest.NewRecorder()
h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/", nil))
if !wrote {
t.Fatal("WriteUnauthorized was not called")
}
if called {
t.Fatal("next ran on guard failure")
}
if rec.Code != http.StatusUnauthorized || rec.Body.String() != `{"from":"writer"}` {
t.Fatalf("status=%d body=%q", rec.Code, rec.Body.String())
}
}
func TestCredentialGuardSoftFailAndSuccess(t *testing.T) {
reg := NewRegistry()
failing := credOnlyGuard{err: errors.New("bad token")}
if err := reg.Register("golem15.acme", "inv_token", failing); err != nil {
t.Fatal(err)
}
mw, err := reg.Middleware("inv_token")
if err != nil {
t.Fatal(err)
}
t.Run("failure-passes-through", func(t *testing.T) {
called := false
h := mw(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
called = true
if _, ok := User(r.Context()); ok {
t.Fatal("principal attached on failure")
}
if _, ok := Credential(r.Context()); ok {
t.Fatal("credential attached on failure")
}
w.WriteHeader(http.StatusNoContent)
}))
rec := httptest.NewRecorder()
h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/", nil))
if !called {
t.Fatal("next did not run on CredentialGuard failure")
}
if rec.Code != http.StatusNoContent {
t.Fatalf("status = %d", rec.Code)
}
})
okReg := NewRegistry()
token := &struct{ ID uint }{ID: 9}
principal := &Principal{ID: 3}
if err := okReg.Register("golem15.acme", "inv_token", credOnlyGuard{
principal: principal,
cred: token,
}); err != nil {
t.Fatal(err)
}
okMW, err := okReg.Middleware("inv_token")
if err != nil {
t.Fatal(err)
}
t.Run("success-attaches-both", func(t *testing.T) {
h := okMW(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
got, ok := User(r.Context())
if !ok || got != principal {
t.Fatalf("user = %+v ok=%t", got, ok)
}
cred, ok := Credential(r.Context())
if !ok || cred != token {
t.Fatalf("cred = %#v ok=%t", cred, ok)
}
w.WriteHeader(http.StatusNoContent)
}))
rec := httptest.NewRecorder()
h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/", nil))
if rec.Code != http.StatusNoContent {
t.Fatalf("status = %d", rec.Code)
}
})
}
func TestJWTGuardRegistryMatchesMiddleware(t *testing.T) {
users := memUsers{byID: map[uint]*Principal{1: {ID: 1}}}
direct := Middleware(secret, users)(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("X-Hit", "1")
w.WriteHeader(http.StatusNoContent)
}))
reg := NewRegistry()
if err := reg.Register("golem15.user", "jwt", NewJWTGuard(secret, users, nil)); err != nil {
t.Fatal(err)
}
viaReg, err := reg.Middleware("jwt")
if err != nil {
t.Fatal(err)
}
hReg := viaReg(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("X-Hit", "1")
w.WriteHeader(http.StatusNoContent)
}))
expired := sign(t, jwt.SigningMethodHS256, jwt.MapClaims{
"sub": "1",
"exp": time.Now().Add(-time.Hour).Unix(),
}, []byte(secret))
badSig := sign(t, jwt.SigningMethodHS256, jwt.MapClaims{
"sub": "1",
"exp": time.Now().Add(time.Hour).Unix(),
}, []byte("other-secret"))
unknown := sign(t, jwt.SigningMethodHS256, jwt.MapClaims{
"sub": "99",
"exp": time.Now().Add(time.Hour).Unix(),
}, []byte(secret))
cases := []struct {
name string
header string
}{
{name: "missing"},
{name: "expired", header: "Bearer " + expired},
{name: "bad-signature", header: "Bearer " + badSig},
{name: "unknown-user", header: "Bearer " + unknown},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
req1 := httptest.NewRequest(http.MethodGet, "/", nil)
req2 := httptest.NewRequest(http.MethodGet, "/", nil)
if tc.header != "" {
req1.Header.Set("Authorization", tc.header)
req2.Header.Set("Authorization", tc.header)
}
rec1 := httptest.NewRecorder()
rec2 := httptest.NewRecorder()
direct.ServeHTTP(rec1, req1)
hReg.ServeHTTP(rec2, req2)
if rec1.Code != rec2.Code {
t.Fatalf("status direct=%d registry=%d", rec1.Code, rec2.Code)
}
if rec1.Body.String() != rec2.Body.String() {
t.Fatalf("body direct=%q registry=%q", rec1.Body.String(), rec2.Body.String())
}
if rec1.Header().Get("Content-Type") != rec2.Header().Get("Content-Type") {
t.Fatalf("content-type direct=%q registry=%q", rec1.Header().Get("Content-Type"), rec2.Header().Get("Content-Type"))
}
if rec1.Header().Get("X-Hit") != "" || rec2.Header().Get("X-Hit") != "" {
t.Fatal("handler ran")
}
})
}
}
func TestRegisterRejectsTypedNilGuard(t *testing.T) {
var reg Registry
var g *credPtrGuard
if err := reg.Register("p", "n", g); err == nil {
t.Fatal("typed-nil guard accepted")
}
}
type credPtrGuard struct{}
func (*credPtrGuard) Authenticate(*http.Request) (*Principal, error) { return nil, nil }
type nilFuncGuard func(*http.Request) (*Principal, error)
func (f nilFuncGuard) Authenticate(r *http.Request) (*Principal, error) { return f(r) }
type nilMapGuard map[string]string
func (nilMapGuard) Authenticate(*http.Request) (*Principal, error) { return nil, nil }
func TestRegisterRejectsTypedNilPointerFuncMapGuards(t *testing.T) {
var (
ptr *credPtrGuard
fn nilFuncGuard
mp nilMapGuard
)
for name, g := range map[string]any{"pointer": ptr, "func": fn, "map": mp} {
t.Run(name, func(t *testing.T) {
var reg Registry
err := reg.Register("golem15.p", "g", g)
if err == nil || !strings.Contains(err.Error(), "golem15.p") || !strings.Contains(err.Error(), `"g"`) {
t.Fatalf("typed-nil %s guard: err = %v", name, err)
}
if _, err := reg.Middleware("g"); err == nil {
t.Fatal("rejected guard must not be resolvable")
}
})
}
}
func TestRegisterAcceptsValidGuards(t *testing.T) {
var reg Registry
if err := reg.Register("p", "ptr", &credPtrGuard{}); err != nil {
t.Fatal(err)
}
if err := reg.Register("p", "fn", nilFuncGuard(func(*http.Request) (*Principal, error) { return nil, nil })); err != nil {
t.Fatal(err)
}
if err := reg.Register("p", "map", nilMapGuard{}); err != nil {
t.Fatal(err)
}
if err := reg.Register("p", "w", writerGuard{}); err != nil {
t.Fatal(err)
}
}