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