package bouncer import ( "context" "database/sql" "sync" "testing" "time" _ "github.com/jackc/pgx/v5/stdlib" "github.com/testcontainers/testcontainers-go/modules/postgres" ) func TestMintRefreshBlacklistRoundTrip(t *testing.T) { const secret = "summercms-test-only-hs256-secret" bl := NewMemoryBlacklist() token, jti, err := Mint(secret, "7", "http://example.com/_user/api/v1/login", time.Hour) if err != nil { t.Fatal(err) } sub, err := Verify(token, secret) if err != nil || sub != "7" { t.Fatalf("verify %q %v", sub, err) } next, err := Refresh(secret, token, 2*time.Hour, bl, 0, "http://example.com/_user/api/v1/refresh") if err != nil { t.Fatal(err) } if _, err := Verify(next, secret); err != nil { t.Fatal(err) } blocked, err := bl.IsBlacklisted(context.Background(), jti) if err != nil || !blocked { t.Fatalf("old jti blacklisted=%t err=%v", blocked, err) } if _, err := Refresh(secret, token, 2*time.Hour, bl, 0, "http://example.com/_user/api/v1/refresh"); err == nil { t.Fatal("blacklisted token refreshed") } } func TestMemoryBlacklistConcurrent(t *testing.T) { bl := NewMemoryBlacklist() var wg sync.WaitGroup for i := 0; i < 8; i++ { wg.Add(1) go func() { defer wg.Done() now := time.Now() if err := bl.Add(context.Background(), "jti", now.Add(time.Hour), now.Add(-time.Second)); err != nil { t.Error(err) return } if _, err := bl.IsBlacklisted(context.Background(), "jti"); err != nil { t.Error(err) } }() } wg.Wait() } func TestPostgresBlacklistConcurrent(t *testing.T) { if testing.Short() { t.Skip("postgres blacklist concurrency needs testcontainers") } ctx, cancel := context.WithTimeout(context.Background(), 2*time.Minute) defer cancel() ctr, err := postgres.Run(ctx, "postgres:16-alpine", postgres.WithDatabase("bl"), postgres.WithUsername("bl"), postgres.WithPassword("bl"), postgres.BasicWaitStrategies(), ) if err != nil { t.Fatal(err) } t.Cleanup(func() { _ = ctr.Terminate(context.Background()) }) dsn, err := ctr.ConnectionString(ctx, "sslmode=disable") if err != nil { t.Fatal(err) } db, err := sql.Open("pgx", dsn) if err != nil { t.Fatal(err) } t.Cleanup(func() { _ = db.Close() }) if _, err := db.ExecContext(ctx, `CREATE TABLE jwt_blacklist ( jti text PRIMARY KEY, expires_at timestamptz NOT NULL, valid_until timestamptz NOT NULL )`); err != nil { t.Fatal(err) } bl := NewPostgresBlacklist(db, "jwt_blacklist") var wg sync.WaitGroup for i := 0; i < 8; i++ { wg.Add(1) go func() { defer wg.Done() now := time.Now() if err := bl.Add(ctx, "shared", now.Add(time.Hour), now.Add(-time.Second)); err != nil { t.Error(err) return } blocked, err := bl.IsBlacklisted(ctx, "shared") if err != nil || !blocked { t.Errorf("blocked=%t err=%v", blocked, err) } }() } wg.Wait() }