307 lines
12 KiB
Go
307 lines
12 KiB
Go
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")
|
|
}
|
|
})
|
|
}
|
|
|
|
func TestVerifyRefreshableClaimsAudience(t *testing.T) {
|
|
now := time.Now()
|
|
token := func(method jwt.SigningMethod, key []byte, claims jwt.MapClaims) string {
|
|
return sign(t, method, claims, key)
|
|
}
|
|
base := func() jwt.MapClaims {
|
|
return jwt.MapClaims{"sub": "7", "jti": "rj", "aud": AudienceBackend, "iat": now.Add(-10 * time.Minute).Unix(), "exp": now.Add(-time.Minute).Unix()}
|
|
}
|
|
|
|
sub, iat, exp, jti, err := VerifyRefreshableClaimsAudience(token(jwt.SigningMethodHS256, []byte(secret), base()), secret, AudienceBackend, time.Hour)
|
|
if err != nil || sub != "7" || jti != "rj" || !exp.Before(now) || !iat.Before(exp) {
|
|
t.Fatalf("expired token inside the refresh window: %q %v %v %q %v", sub, iat, exp, jti, err)
|
|
}
|
|
// VerifyClaimsAudience rejects the same token: the difference is the point.
|
|
if _, _, _, _, err := VerifyClaimsAudience(token(jwt.SigningMethodHS256, []byte(secret), base()), secret, AudienceBackend); err == nil {
|
|
t.Fatal("VerifyClaimsAudience accepted an expired token")
|
|
}
|
|
|
|
tests := map[string]string{
|
|
"wrong secret": token(jwt.SigningMethodHS256, []byte("another-secret-value-for-the-test!"), base()),
|
|
"empty token": "",
|
|
"window closed": token(jwt.SigningMethodHS256, []byte(secret), func() jwt.MapClaims {
|
|
c := base()
|
|
c["iat"] = now.Add(-2 * time.Hour).Unix()
|
|
return c
|
|
}()),
|
|
"wrong audience": token(jwt.SigningMethodHS256, []byte(secret), func() jwt.MapClaims { c := base(); c["aud"] = AudienceUser; return c }()),
|
|
"missing audience": token(jwt.SigningMethodHS256, []byte(secret), func() jwt.MapClaims {
|
|
c := base()
|
|
delete(c, "aud")
|
|
return c
|
|
}()),
|
|
"missing jti": token(jwt.SigningMethodHS256, []byte(secret), func() jwt.MapClaims { c := base(); delete(c, "jti"); return c }()),
|
|
"missing iat": token(jwt.SigningMethodHS256, []byte(secret), func() jwt.MapClaims { c := base(); delete(c, "iat"); return c }()),
|
|
}
|
|
for name, raw := range tests {
|
|
if _, _, _, _, err := VerifyRefreshableClaimsAudience(raw, secret, AudienceBackend, time.Hour); err == nil {
|
|
t.Fatalf("%s: token was accepted", name)
|
|
}
|
|
}
|
|
if _, _, _, _, err := VerifyRefreshableClaimsAudience("x", secret, "", time.Hour); err == nil {
|
|
t.Fatal("empty audience was accepted")
|
|
}
|
|
}
|