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:
42
modules/bouncer/audience_test.go
Normal file
42
modules/bouncer/audience_test.go
Normal 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
|
||||
}
|
||||
171
modules/bouncer/backend_guard_test.go
Normal file
171
modules/bouncer/backend_guard_test.go
Normal 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
|
||||
}
|
||||
126
modules/bouncer/blacklist.go
Normal file
126
modules/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
|
||||
}
|
||||
61
modules/bouncer/blacklist_test.go
Normal file
61
modules/bouncer/blacklist_test.go
Normal 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")
|
||||
}
|
||||
}
|
||||
61
modules/bouncer/context.go
Normal file
61
modules/bouncer/context.go
Normal 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
|
||||
}
|
||||
32
modules/bouncer/context_test.go
Normal file
32
modules/bouncer/context_test.go
Normal 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)
|
||||
}
|
||||
}
|
||||
132
modules/bouncer/cookie_guard_test.go
Normal file
132
modules/bouncer/cookie_guard_test.go
Normal 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
25
modules/bouncer/guard.go
Normal 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
371
modules/bouncer/jwt.go
Normal 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})
|
||||
}
|
||||
128
modules/bouncer/jwt_guard_test.go
Normal file
128
modules/bouncer/jwt_guard_test.go
Normal 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
294
modules/bouncer/jwt_test.go
Normal 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
78
modules/bouncer/mint.go
Normal 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
|
||||
}
|
||||
62
modules/bouncer/mint_test.go
Normal file
62
modules/bouncer/mint_test.go
Normal 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)
|
||||
}
|
||||
27
modules/bouncer/password.go
Normal file
27
modules/bouncer/password.go
Normal 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
|
||||
}
|
||||
46
modules/bouncer/password_test.go
Normal file
46
modules/bouncer/password_test.go
Normal 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")
|
||||
}
|
||||
}
|
||||
115
modules/bouncer/phase07_coverage_test.go
Normal file
115
modules/bouncer/phase07_coverage_test.go
Normal 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
121
modules/bouncer/refresh.go
Normal 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
|
||||
}
|
||||
261
modules/bouncer/refresh_test.go
Normal file
261
modules/bouncer/refresh_test.go
Normal 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
103
modules/bouncer/registry.go
Normal 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")
|
||||
}
|
||||
85
modules/bouncer/registry_coverage_test.go
Normal file
85
modules/bouncer/registry_coverage_test.go
Normal 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")
|
||||
}
|
||||
}
|
||||
291
modules/bouncer/registry_test.go
Normal file
291
modules/bouncer/registry_test.go
Normal 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)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user