- Move remaining beach packages and embedded admin assets\n- Rewrite framework, example, build, and gate paths
127 lines
3.5 KiB
Go
127 lines
3.5 KiB
Go
package bouncer
|
|
|
|
import (
|
|
"context"
|
|
"database/sql"
|
|
"errors"
|
|
"fmt"
|
|
"regexp"
|
|
"sync"
|
|
"time"
|
|
)
|
|
|
|
// BlacklistStore records revoked jti values. IsBlacklisted is true only once
|
|
// validUntil has been reached, so a grace window can keep a just-rotated
|
|
// token usable. Sweep drops rows whose storage expiry has passed.
|
|
type BlacklistStore interface {
|
|
Add(ctx context.Context, jti string, expiresAt, validUntil time.Time) error
|
|
IsBlacklisted(ctx context.Context, jti string) (bool, error)
|
|
Sweep(ctx context.Context, now time.Time) error
|
|
}
|
|
|
|
type blEntry struct {
|
|
expiresAt time.Time
|
|
validUntil time.Time
|
|
}
|
|
|
|
// MemoryBlacklist is an in-process store for tests. Expired rows are dropped
|
|
// on read, matching surf.MemoryStore.
|
|
type MemoryBlacklist struct {
|
|
mu sync.Mutex
|
|
entries map[string]blEntry
|
|
}
|
|
|
|
// NewMemoryBlacklist returns an empty in-process blacklist.
|
|
func NewMemoryBlacklist() *MemoryBlacklist {
|
|
return &MemoryBlacklist{entries: make(map[string]blEntry)}
|
|
}
|
|
|
|
func (m *MemoryBlacklist) Add(_ context.Context, jti string, expiresAt, validUntil time.Time) error {
|
|
m.mu.Lock()
|
|
defer m.mu.Unlock()
|
|
m.entries[jti] = blEntry{expiresAt: expiresAt, validUntil: validUntil}
|
|
return nil
|
|
}
|
|
|
|
func (m *MemoryBlacklist) IsBlacklisted(_ context.Context, jti string) (bool, error) {
|
|
m.mu.Lock()
|
|
defer m.mu.Unlock()
|
|
e, ok := m.entries[jti]
|
|
if !ok {
|
|
return false, nil
|
|
}
|
|
now := time.Now()
|
|
if now.After(e.expiresAt) {
|
|
delete(m.entries, jti)
|
|
return false, nil
|
|
}
|
|
return !now.Before(e.validUntil), nil
|
|
}
|
|
|
|
func (m *MemoryBlacklist) Sweep(_ context.Context, now time.Time) error {
|
|
m.mu.Lock()
|
|
defer m.mu.Unlock()
|
|
for k, e := range m.entries {
|
|
if e.expiresAt.Before(now) {
|
|
delete(m.entries, k)
|
|
}
|
|
}
|
|
return nil
|
|
}
|
|
|
|
var blacklistIdent = regexp.MustCompile(`^[a-zA-Z_][a-zA-Z0-9_]*$`)
|
|
|
|
// PostgresBlacklist stores revoked jti rows in a caller-supplied table.
|
|
// The table name is a plugin constant, still checked so it cannot be interpolated raw.
|
|
type PostgresBlacklist struct {
|
|
db *sql.DB
|
|
table string
|
|
}
|
|
|
|
// NewPostgresBlacklist binds db to table. table must be a plain SQL identifier.
|
|
func NewPostgresBlacklist(db *sql.DB, table string) *PostgresBlacklist {
|
|
return &PostgresBlacklist{db: db, table: table}
|
|
}
|
|
|
|
func (p *PostgresBlacklist) checkTable() error {
|
|
if p == nil || !blacklistIdent.MatchString(p.table) {
|
|
return fmt.Errorf("bouncer: blacklist table %q is not a safe identifier", p.table)
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func (p *PostgresBlacklist) Add(ctx context.Context, jti string, expiresAt, validUntil time.Time) error {
|
|
if err := p.checkTable(); err != nil {
|
|
return err
|
|
}
|
|
q := fmt.Sprintf(`INSERT INTO %s (jti, expires_at, valid_until) VALUES ($1, $2, $3)
|
|
ON CONFLICT (jti) DO UPDATE SET expires_at = EXCLUDED.expires_at, valid_until = EXCLUDED.valid_until`, p.table)
|
|
_, err := p.db.ExecContext(ctx, q, jti, expiresAt, validUntil)
|
|
return err
|
|
}
|
|
|
|
func (p *PostgresBlacklist) IsBlacklisted(ctx context.Context, jti string) (bool, error) {
|
|
if err := p.checkTable(); err != nil {
|
|
return false, err
|
|
}
|
|
q := fmt.Sprintf(`SELECT valid_until FROM %s WHERE jti = $1`, p.table)
|
|
var validUntil time.Time
|
|
err := p.db.QueryRowContext(ctx, q, jti).Scan(&validUntil)
|
|
if errors.Is(err, sql.ErrNoRows) {
|
|
return false, nil
|
|
}
|
|
if err != nil {
|
|
return false, err
|
|
}
|
|
return !time.Now().Before(validUntil), nil
|
|
}
|
|
|
|
func (p *PostgresBlacklist) Sweep(ctx context.Context, now time.Time) error {
|
|
if err := p.checkTable(); err != nil {
|
|
return err
|
|
}
|
|
q := fmt.Sprintf(`DELETE FROM %s WHERE expires_at < $1`, p.table)
|
|
_, err := p.db.ExecContext(ctx, q, now)
|
|
return err
|
|
}
|