diff --git a/bouncer/jwt_guard_test.go b/bouncer/jwt_guard_test.go index 7a5479f..d0f8f60 100644 --- a/bouncer/jwt_guard_test.go +++ b/bouncer/jwt_guard_test.go @@ -2,6 +2,7 @@ package bouncer import ( "encoding/json" + "errors" "net/http" "net/http/httptest" "testing" @@ -107,6 +108,11 @@ func TestJWTGuardTokensValidAfter(t *testing.T) { 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) diff --git a/bouncer/refresh_test.go b/bouncer/refresh_test.go index 5fdb3ad..04b3b8e 100644 --- a/bouncer/refresh_test.go +++ b/bouncer/refresh_test.go @@ -1,6 +1,9 @@ package bouncer import ( + "context" + "errors" + "sync/atomic" "testing" "time" @@ -105,3 +108,154 @@ func TestRefreshBlacklistsOldJTI(t *testing.T) { 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") + } + }) +} diff --git a/cabana/phase10_coverage_test.go b/cabana/phase10_coverage_test.go index d0ebc96..05bc5ab 100644 --- a/cabana/phase10_coverage_test.go +++ b/cabana/phase10_coverage_test.go @@ -4,6 +4,7 @@ import ( "bytes" "context" "encoding/json" + "errors" "fmt" "net/http" "net/http/httptest" @@ -379,4 +380,104 @@ func TestPhase10Coverage(t *testing.T) { t.Fatalf("cookie refresh body=%s", rec.Body.String()) } }) + + t.Run("refresh enforces the guard's subject checks", func(t *testing.T) { + const secret = "phase10-subject-secret" + now := time.Now() + cutoff := now.Add(-10 * time.Minute) + subjects := refreshSubjects{byID: map[uint]*bouncer.Principal{ + 5: {ID: 5, Backend: true, TokensValidAfter: cutoff}, + }} + newService := func(users bouncer.UserProvider) *service { + return &service{ + secret: secret, + ttl: 15 * time.Minute, + refreshTTL: 2 * time.Hour, + issuer: "https://app.test" + DefaultAdminPrefix, + bl: bouncer.NewMemoryBlacklist(), + users: users, + } + } + sign := func(sub string, iat time.Time, jti string) string { + t.Helper() + tok, err := jwt.NewWithClaims(jwt.SigningMethodHS256, jwt.MapClaims{ + "sub": sub, "aud": bouncer.AudienceBackend, "jti": jti, + "iat": iat.Unix(), "nbf": iat.Unix(), "exp": iat.Add(15 * time.Minute).Unix(), + }).SignedString([]byte(secret)) + if err != nil { + t.Fatal(err) + } + return tok + } + call := func(svc *service, token string, cookie bool) *httptest.ResponseRecorder { + req := httptest.NewRequest(http.MethodPost, adminAPI("/auth/refresh"), nil) + req.Header.Set("X-Requested-With", "XMLHttpRequest") + if cookie { + req.AddCookie(&http.Cookie{Name: AdminCookieName, Value: token}) + } else { + req.Header.Set("Authorization", "Bearer "+token) + } + rec := httptest.NewRecorder() + requireAjax(svc.refresh)(rec, req) + return rec + } + errorCode := func(rec *httptest.ResponseRecorder) string { + t.Helper() + var body struct { + Error struct { + Code string `json:"code"` + } `json:"error"` + } + if err := json.Unmarshal(rec.Body.Bytes(), &body); err != nil { + t.Fatalf("error json: %v body=%s", err, rec.Body.String()) + } + return body.Error.Code + } + assertExpired := func(name string, rec *httptest.ResponseRecorder) { + t.Helper() + if rec.Code != http.StatusUnauthorized || errorCode(rec) != "unauthenticated" { + t.Fatalf("%s: status=%d body=%s, want 401 unauthenticated", name, rec.Code, rec.Body.String()) + } + var expired *http.Cookie + for _, c := range rec.Result().Cookies() { + if c.Name == AdminCookieName { + expired = c + } + } + if expired == nil || expired.Value != "" || expired.MaxAge >= 0 || expired.Path != DefaultAdminPrefix { + t.Fatalf("%s: cookie = %+v, want an expiring %s with Path %s", name, expired, AdminCookieName, DefaultAdminPrefix) + } + } + assertNoCookies := func(name string, rec *httptest.ResponseRecorder) { + t.Helper() + if rec.Code != http.StatusUnauthorized || errorCode(rec) != "unauthenticated" { + t.Fatalf("%s: status=%d body=%s, want 401 unauthenticated", name, rec.Code, rec.Body.String()) + } + if got := rec.Header().Values("Set-Cookie"); len(got) != 0 { + t.Fatalf("%s: set cookies %q", name, got) + } + } + + svc := newService(subjects) + assertExpired("pre-cutoff cookie", call(svc, sign("5", cutoff.Add(-time.Minute), "pre-cutoff-cookie"), true)) + assertExpired("unknown subject cookie", call(svc, sign("6", now.Add(-time.Minute), "unknown-cookie"), true)) + assertNoCookies("pre-cutoff bearer", call(svc, sign("5", cutoff.Add(-time.Minute), "pre-cutoff-bearer"), false)) + + failing := newService(refreshSubjects{err: errors.New("lookup failed")}) + assertNoCookies("provider error cookie", call(failing, sign("5", now.Add(-time.Minute), "provider-error"), true)) + + ok := call(svc, sign("5", cutoff.Add(time.Minute), "post-cutoff-cookie"), true) + if ok.Code != http.StatusOK { + t.Fatalf("post-cutoff cookie refresh status=%d body=%s", ok.Code, ok.Body.String()) + } + var rotated *http.Cookie + for _, c := range ok.Result().Cookies() { + if c.Name == AdminCookieName { + rotated = c + } + } + if rotated == nil || rotated.Value == "" || rotated.MaxAge <= 0 || rotated.Path != DefaultAdminPrefix { + t.Fatalf("post-cutoff rotated cookie = %+v", rotated) + } + }) }