package oidc import ( "context" "crypto/rand" "encoding/hex" "errors" "fmt" "time" "github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5/pgxpool" ) var ErrStateInvalid = errors.New("oidc: state ungueltig, bereits verwendet oder abgelaufen") // DefaultStateTTL begrenzt, wie lange ein ausgestellter State/Nonce fuer den // Redirect-Umweg zum Identity-Provider gueltig bleibt. const DefaultStateTTL = 10 * time.Minute type StateStore struct { pool *pgxpool.Pool } func NewStateStore(pool *pgxpool.Pool) *StateStore { return &StateStore{pool: pool} } // Generate stellt state+nonce fuer einen neuen Login-Redirect aus. state // wird als OAuth2-"state"-Parameter mitgeschickt (CSRF-Schutz), nonce // erscheint spaeter im ID-Token und muss exakt uebereinstimmen (Replay-Schutz). func (s *StateStore) Generate(ctx context.Context) (state, nonce string, err error) { state, err = randomValue() if err != nil { return "", "", err } nonce, err = randomValue() if err != nil { return "", "", err } _, err = s.pool.Exec(ctx, ` INSERT INTO oidc_states (state, nonce, expires_at) VALUES ($1, $2, $3) `, state, nonce, time.Now().Add(DefaultStateTTL)) if err != nil { return "", "", fmt.Errorf("state speichern: %w", err) } return state, nonce, nil } // Consume loest einen State EINMALIG ein (Akzeptanzkriterium/Pruefung 3: // Replay-Schutz) — atomar ueber WHERE used_at IS NULL, analog IAM-03/IAM-09. // Ein zweiter Callback mit demselben state (z.B. durch einen Angreifer, der // die Redirect-URL abgefangen hat) schlaegt fehl. func (s *StateStore) Consume(ctx context.Context, state string) (nonce string, err error) { err = s.pool.QueryRow(ctx, ` UPDATE oidc_states SET used_at = now() WHERE state = $1 AND used_at IS NULL AND expires_at > now() RETURNING nonce `, state).Scan(&nonce) if err != nil { if errors.Is(err, pgx.ErrNoRows) { return "", ErrStateInvalid } return "", fmt.Errorf("state einloesen: %w", err) } return nonce, nil } func randomValue() (string, error) { buf := make([]byte, 32) if _, err := rand.Read(buf); err != nil { return "", err } return hex.EncodeToString(buf), nil }