diff --git a/bouncer/phase07_coverage_test.go b/bouncer/phase07_coverage_test.go new file mode 100644 index 0000000..f005fe6 --- /dev/null +++ b/bouncer/phase07_coverage_test.go @@ -0,0 +1,115 @@ +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() +}