From 251f3cc4a00692eb502a31ec75cc1f76af7d46a7 Mon Sep 17 00:00:00 2001 From: Jakub Zych Date: Tue, 22 Sep 2026 13:33:37 +0200 Subject: [PATCH] test(07-01): add failing tests for JWT lifecycle primitives Co-authored-by: Cursor --- bouncer/blacklist_test.go | 61 +++++++++++++++++++ bouncer/context_test.go | 13 +++- bouncer/jwt_guard_test.go | 122 ++++++++++++++++++++++++++++++++++++++ bouncer/mint_test.go | 62 +++++++++++++++++++ bouncer/refresh_test.go | 107 +++++++++++++++++++++++++++++++++ 5 files changed, 364 insertions(+), 1 deletion(-) create mode 100644 bouncer/blacklist_test.go create mode 100644 bouncer/jwt_guard_test.go create mode 100644 bouncer/mint_test.go create mode 100644 bouncer/refresh_test.go diff --git a/bouncer/blacklist_test.go b/bouncer/blacklist_test.go new file mode 100644 index 0000000..9a27ba6 --- /dev/null +++ b/bouncer/blacklist_test.go @@ -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") + } +} diff --git a/bouncer/context_test.go b/bouncer/context_test.go index af92aff..f7c62a7 100644 --- a/bouncer/context_test.go +++ b/bouncer/context_test.go @@ -1,6 +1,17 @@ package bouncer -import "testing" +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 { diff --git a/bouncer/jwt_guard_test.go b/bouncer/jwt_guard_test.go new file mode 100644 index 0000000..7a5479f --- /dev/null +++ b/bouncer/jwt_guard_test.go @@ -0,0 +1,122 @@ +package bouncer + +import ( + "encoding/json" + "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) + } + + 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) + } +} diff --git a/bouncer/mint_test.go b/bouncer/mint_test.go new file mode 100644 index 0000000..030c4fa --- /dev/null +++ b/bouncer/mint_test.go @@ -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) +} diff --git a/bouncer/refresh_test.go b/bouncer/refresh_test.go new file mode 100644 index 0000000..5fdb3ad --- /dev/null +++ b/bouncer/refresh_test.go @@ -0,0 +1,107 @@ +package bouncer + +import ( + "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") + } +}