package bouncer import ( "encoding/json" "errors" "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) } // 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) 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) } }