A minted token refreshes once, then the old jti is blacklisted. Memory and Postgres blacklists take concurrent Add and IsBlacklisted calls. Co-authored-by: Cursor <cursoragent@cursor.com>
116 lines
3.0 KiB
Go
116 lines
3.0 KiB
Go
package bouncer
|
|
|
|
import (
|
|
"context"
|
|
"database/sql"
|
|
"sync"
|
|
"testing"
|
|
"time"
|
|
|
|
_ "github.com/jackc/pgx/v5/stdlib"
|
|
"github.com/testcontainers/testcontainers-go"
|
|
"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(),
|
|
testcontainers.WithEnv(map[string]string{
|
|
"POSTGRES_INITDB_ARGS": "--locale-provider=icu --icu-locale=pl-PL --encoding=UTF8",
|
|
}),
|
|
)
|
|
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()
|
|
}
|