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") } }