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 }