refactor(10.2-01): nest framework packages under modules
- Move remaining beach packages and embedded admin assets\n- Rewrite framework, example, build, and gate paths
This commit is contained in:
126
modules/bouncer/blacklist.go
Normal file
126
modules/bouncer/blacklist.go
Normal file
@@ -0,0 +1,126 @@
|
||||
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
|
||||
}
|
||||
Reference in New Issue
Block a user