package auth import ( "context" "errors" "fmt" "time" "github.com/google/uuid" "github.com/jackc/pgx/v5" ) // createSession legt eine Sitzung an. // // Gespeichert werden ausschließlich die Hashes der Tokens; die Klartextwerte // verlassen den Server nur einmal in der Anmeldeantwort. func (repository *Repository) createSession(createContext context.Context, userID uuid.UUID, accessTokenHash string, refreshTokenHash string, accessExpiresAt time.Time, refreshExpiresAt time.Time, ipAddress string, userAgent string) (uuid.UUID, error) { const insertStatement = ` INSERT INTO sessions (user_id, access_token_hash, refresh_token_hash, access_expires_at, refresh_expires_at, ip_address, user_agent) VALUES ($1, $2, $3, $4, $5, $6, $7) RETURNING id` var createdSessionID uuid.UUID insertError := repository.connectionPool.QueryRow(createContext, insertStatement, userID, accessTokenHash, refreshTokenHash, accessExpiresAt, refreshExpiresAt, nullIfEmptyInet(ipAddress), nullIfEmpty(userAgent)).Scan(&createdSessionID) if insertError != nil { return uuid.Nil, fmt.Errorf("die sitzung konnte nicht angelegt werden: %w", insertError) } return createdSessionID, nil } // findSessionByAccessTokenHash sucht eine gültige Sitzung zum Zugriffstoken. // // Abgelaufene und widerrufene Sitzungen werden in der Abfrage ausgeschlossen, // damit ein vergessener Zustandsvergleich im Code nicht zu einem Zugriff führt. func (repository *Repository) findSessionByAccessTokenHash(queryContext context.Context, accessTokenHash string) (Session, error) { const selectStatement = ` SELECT id, user_id, access_expires_at, refresh_expires_at FROM sessions WHERE access_token_hash = $1 AND revoked_at IS NULL AND access_expires_at > now()` var foundSession Session scanError := repository.connectionPool.QueryRow(queryContext, selectStatement, accessTokenHash).Scan( &foundSession.ID, &foundSession.UserID, &foundSession.AccessExpiresAt, &foundSession.RefreshExpiresAt) if errors.Is(scanError, pgx.ErrNoRows) { return Session{}, ErrSessionInvalid } if scanError != nil { return Session{}, fmt.Errorf("die sitzung konnte nicht gelesen werden: %w", scanError) } return foundSession, nil } // findSessionByRefreshTokenHash sucht eine gültige Sitzung zum Erneuerungstoken. func (repository *Repository) findSessionByRefreshTokenHash(queryContext context.Context, refreshTokenHash string) (Session, error) { const selectStatement = ` SELECT id, user_id, access_expires_at, refresh_expires_at FROM sessions WHERE refresh_token_hash = $1 AND revoked_at IS NULL AND refresh_expires_at > now()` var foundSession Session scanError := repository.connectionPool.QueryRow(queryContext, selectStatement, refreshTokenHash).Scan( &foundSession.ID, &foundSession.UserID, &foundSession.AccessExpiresAt, &foundSession.RefreshExpiresAt) if errors.Is(scanError, pgx.ErrNoRows) { return Session{}, ErrSessionInvalid } if scanError != nil { return Session{}, fmt.Errorf("die sitzung konnte nicht gelesen werden: %w", scanError) } return foundSession, nil } // touchSession vermerkt die Verwendung einer Sitzung. func (repository *Repository) touchSession(updateContext context.Context, sessionID uuid.UUID) error { if _, updateError := repository.connectionPool.Exec(updateContext, "UPDATE sessions SET last_used_at = now() WHERE id = $1", sessionID); updateError != nil { return fmt.Errorf("die sitzung konnte nicht aktualisiert werden: %w", updateError) } return nil } // rotateSessionTokens tauscht die Tokens einer Sitzung aus. // // Bei jeder Erneuerung werden beide Tokens ersetzt. Ein abgefangenes // Erneuerungstoken wird damit spätestens beim nächsten regulären Gebrauch wertlos. func (repository *Repository) rotateSessionTokens(updateContext context.Context, sessionID uuid.UUID, newAccessTokenHash string, newRefreshTokenHash string, accessExpiresAt time.Time, refreshExpiresAt time.Time) error { const updateStatement = ` UPDATE sessions SET access_token_hash = $2, refresh_token_hash = $3, access_expires_at = $4, refresh_expires_at = $5, last_used_at = now() WHERE id = $1 AND revoked_at IS NULL` commandTag, updateError := repository.connectionPool.Exec(updateContext, updateStatement, sessionID, newAccessTokenHash, newRefreshTokenHash, accessExpiresAt, refreshExpiresAt) if updateError != nil { return fmt.Errorf("die sitzung konnte nicht erneuert werden: %w", updateError) } // Keine betroffene Zeile bedeutet: die Sitzung wurde zwischenzeitlich widerrufen. if commandTag.RowsAffected() == 0 { return ErrSessionInvalid } return nil } // revokeSession widerruft eine einzelne Sitzung. func (repository *Repository) revokeSession(updateContext context.Context, sessionID uuid.UUID) error { if _, updateError := repository.connectionPool.Exec(updateContext, "UPDATE sessions SET revoked_at = now() WHERE id = $1 AND revoked_at IS NULL", sessionID); updateError != nil { return fmt.Errorf("die sitzung konnte nicht widerrufen werden: %w", updateError) } return nil } // RevokeAllSessionsOfUser widerruft sämtliche Sitzungen eines Benutzers. // // Das ist die Reaktion auf eine Passwortänderung, eine Sperre oder einen // Sicherheitsvorfall. func (repository *Repository) RevokeAllSessionsOfUser(updateContext context.Context, userID uuid.UUID) error { if _, updateError := repository.connectionPool.Exec(updateContext, "UPDATE sessions SET revoked_at = now() WHERE user_id = $1 AND revoked_at IS NULL", userID); updateError != nil { return fmt.Errorf("die sitzungen konnten nicht widerrufen werden: %w", updateError) } return nil } // DeleteExpiredSessions entfernt abgelaufene Sitzungen. // // Die Bereinigung hält die Tabelle klein; sicherheitsrelevant ist sie nicht, // da abgelaufene Sitzungen bereits in den Abfragen ausgeschlossen werden. func (repository *Repository) DeleteExpiredSessions(deleteContext context.Context) (int64, error) { commandTag, deleteError := repository.connectionPool.Exec(deleteContext, "DELETE FROM sessions WHERE refresh_expires_at < now() - interval '7 days'") if deleteError != nil { return 0, fmt.Errorf("die abgelaufenen sitzungen konnten nicht entfernt werden: %w", deleteError) } return commandTag.RowsAffected(), nil } // --------------------------------------------------------------------------- // MFA-Herausforderungen // --------------------------------------------------------------------------- // createMFAChallenge legt eine offene MFA-Herausforderung an. func (repository *Repository) createMFAChallenge(createContext context.Context, userID uuid.UUID, expiresAt time.Time, ipAddress string) (uuid.UUID, error) { const insertStatement = ` INSERT INTO mfa_challenges (user_id, expires_at, ip_address) VALUES ($1, $2, $3) RETURNING id` var createdChallengeID uuid.UUID insertError := repository.connectionPool.QueryRow(createContext, insertStatement, userID, expiresAt, nullIfEmptyInet(ipAddress)).Scan(&createdChallengeID) if insertError != nil { return uuid.Nil, fmt.Errorf("die anmeldung konnte nicht fortgesetzt werden: %w", insertError) } return createdChallengeID, nil } // challengeRecord ist eine gelesene MFA-Herausforderung. type challengeRecord struct { // id ist der Bezeichner der Herausforderung. id uuid.UUID // userID ist der zugehörige Benutzer. userID uuid.UUID // attempts ist die Anzahl bisheriger Fehlversuche. attempts int } // findOpenMFAChallenge liest eine noch offene Herausforderung. func (repository *Repository) findOpenMFAChallenge(queryContext context.Context, challengeID uuid.UUID) (challengeRecord, error) { const selectStatement = ` SELECT id, user_id, attempts FROM mfa_challenges WHERE id = $1 AND consumed_at IS NULL AND expires_at > now()` var record challengeRecord scanError := repository.connectionPool.QueryRow(queryContext, selectStatement, challengeID).Scan( &record.id, &record.userID, &record.attempts) if errors.Is(scanError, pgx.ErrNoRows) { return challengeRecord{}, ErrChallengeNotFound } if scanError != nil { return challengeRecord{}, fmt.Errorf("die anmeldung konnte nicht geprüft werden: %w", scanError) } return record, nil } // incrementMFAChallengeAttempts erhöht den Fehlversuchszähler einer Herausforderung. func (repository *Repository) incrementMFAChallengeAttempts(updateContext context.Context, challengeID uuid.UUID) error { if _, updateError := repository.connectionPool.Exec(updateContext, "UPDATE mfa_challenges SET attempts = attempts + 1 WHERE id = $1", challengeID); updateError != nil { return fmt.Errorf("der fehlversuch konnte nicht vermerkt werden: %w", updateError) } return nil } // consumeMFAChallenge markiert eine Herausforderung als eingelöst. // // Die Bedingung auf consumed_at stellt sicher, dass zwei gleichzeitige Anfragen // nicht beide eine Sitzung erhalten. func (repository *Repository) consumeMFAChallenge(updateContext context.Context, challengeID uuid.UUID) error { commandTag, updateError := repository.connectionPool.Exec(updateContext, "UPDATE mfa_challenges SET consumed_at = now() WHERE id = $1 AND consumed_at IS NULL", challengeID) if updateError != nil { return fmt.Errorf("die anmeldung konnte nicht abgeschlossen werden: %w", updateError) } if commandTag.RowsAffected() == 0 { return ErrChallengeNotFound } return nil } // --------------------------------------------------------------------------- // Zweiter Faktor // --------------------------------------------------------------------------- // mfaMethodRecord ist eine gelesene MFA-Methode. type mfaMethodRecord struct { // id ist der Bezeichner der Methode. id uuid.UUID // secretCiphertext ist das verschlüsselte Secret. secretCiphertext []byte // keyVersion benennt den zur Entschlüsselung nötigen Schlüssel. keyVersion string // enabled meldet, ob die Methode bestätigt wurde. enabled bool // lastUsedTimeStep ist der zuletzt akzeptierte TOTP-Zeitschritt. lastUsedTimeStep *int64 } // upsertTOTPMethod legt eine TOTP-Methode an oder ersetzt eine unbestätigte. func (repository *Repository) upsertTOTPMethod(upsertContext context.Context, userID uuid.UUID, secretCiphertext []byte, keyVersion string) error { // Eine bereits bestätigte Methode wird nicht überschrieben: sonst könnte ein // Angreifer mit einer offenen Sitzung den zweiten Faktor stillschweigend ersetzen. const upsertStatement = ` INSERT INTO user_mfa_methods (user_id, type, secret_ciphertext, key_version, enabled) VALUES ($1, 'totp', $2, $3, FALSE) ON CONFLICT (user_id, type) DO UPDATE SET secret_ciphertext = EXCLUDED.secret_ciphertext, key_version = EXCLUDED.key_version, last_used_time_step = NULL WHERE user_mfa_methods.enabled = FALSE` commandTag, upsertError := repository.connectionPool.Exec(upsertContext, upsertStatement, userID, secretCiphertext, keyVersion) if upsertError != nil { return fmt.Errorf("der zweite faktor konnte nicht gespeichert werden: %w", upsertError) } if commandTag.RowsAffected() == 0 { return ErrMFAAlreadyEnabled } return nil } // findTOTPMethod liest die TOTP-Methode eines Benutzers. func (repository *Repository) findTOTPMethod(queryContext context.Context, userID uuid.UUID) (mfaMethodRecord, error) { const selectStatement = ` SELECT id, secret_ciphertext, key_version, enabled, last_used_time_step FROM user_mfa_methods WHERE user_id = $1 AND type = 'totp'` var record mfaMethodRecord scanError := repository.connectionPool.QueryRow(queryContext, selectStatement, userID).Scan( &record.id, &record.secretCiphertext, &record.keyVersion, &record.enabled, &record.lastUsedTimeStep) if errors.Is(scanError, pgx.ErrNoRows) { return mfaMethodRecord{}, ErrMFANotEnrolled } if scanError != nil { return mfaMethodRecord{}, fmt.Errorf("der zweite faktor konnte nicht gelesen werden: %w", scanError) } return record, nil } // confirmTOTPMethod schaltet eine Methode scharf und aktiviert MFA am Konto. func (repository *Repository) confirmTOTPMethod(confirmContext context.Context, userID uuid.UUID, usedTimeStep int64) error { databaseTransaction, transactionError := repository.connectionPool.Begin(confirmContext) if transactionError != nil { return fmt.Errorf("die transaktion konnte nicht begonnen werden: %w", transactionError) } defer func() { _ = databaseTransaction.Rollback(confirmContext) }() if _, updateError := databaseTransaction.Exec(confirmContext, ` UPDATE user_mfa_methods SET enabled = TRUE, confirmed_at = now(), last_used_time_step = $2 WHERE user_id = $1 AND type = 'totp'`, userID, usedTimeStep); updateError != nil { return fmt.Errorf("der zweite faktor konnte nicht bestätigt werden: %w", updateError) } if _, updateError := databaseTransaction.Exec(confirmContext, "UPDATE users SET mfa_enabled = TRUE, updated_at = now() WHERE id = $1", userID); updateError != nil { return fmt.Errorf("das konto konnte nicht aktualisiert werden: %w", updateError) } if commitError := databaseTransaction.Commit(confirmContext); commitError != nil { return fmt.Errorf("die bestätigung konnte nicht gespeichert werden: %w", commitError) } return nil } // updateTOTPTimeStep vermerkt den zuletzt verwendeten Zeitschritt. // // Die Bedingung stellt sicher, dass der Wert nur steigt: bei gleichzeitigen // Anfragen darf ein älterer Schritt den neueren nicht überschreiben. func (repository *Repository) updateTOTPTimeStep(updateContext context.Context, methodID uuid.UUID, usedTimeStep int64) error { const updateStatement = ` UPDATE user_mfa_methods SET last_used_time_step = $2 WHERE id = $1 AND (last_used_time_step IS NULL OR last_used_time_step < $2)` if _, updateError := repository.connectionPool.Exec(updateContext, updateStatement, methodID, usedTimeStep); updateError != nil { return fmt.Errorf("der verwendete code konnte nicht vermerkt werden: %w", updateError) } return nil } // disableMFA entfernt alle zweiten Faktoren und Wiederherstellungscodes. func (repository *Repository) disableMFA(disableContext context.Context, userID uuid.UUID) error { databaseTransaction, transactionError := repository.connectionPool.Begin(disableContext) if transactionError != nil { return fmt.Errorf("die transaktion konnte nicht begonnen werden: %w", transactionError) } defer func() { _ = databaseTransaction.Rollback(disableContext) }() if _, deleteError := databaseTransaction.Exec(disableContext, "DELETE FROM user_mfa_methods WHERE user_id = $1", userID); deleteError != nil { return fmt.Errorf("der zweite faktor konnte nicht entfernt werden: %w", deleteError) } if _, deleteError := databaseTransaction.Exec(disableContext, "DELETE FROM user_recovery_codes WHERE user_id = $1", userID); deleteError != nil { return fmt.Errorf("die wiederherstellungscodes konnten nicht entfernt werden: %w", deleteError) } if _, updateError := databaseTransaction.Exec(disableContext, "UPDATE users SET mfa_enabled = FALSE, updated_at = now() WHERE id = $1", userID); updateError != nil { return fmt.Errorf("das konto konnte nicht aktualisiert werden: %w", updateError) } if commitError := databaseTransaction.Commit(disableContext); commitError != nil { return fmt.Errorf("die änderung konnte nicht gespeichert werden: %w", commitError) } return nil } // --------------------------------------------------------------------------- // Wiederherstellungscodes // --------------------------------------------------------------------------- // replaceRecoveryCodes ersetzt sämtliche Wiederherstellungscodes eines Benutzers. func (repository *Repository) replaceRecoveryCodes(replaceContext context.Context, userID uuid.UUID, codeHashes []string) error { databaseTransaction, transactionError := repository.connectionPool.Begin(replaceContext) if transactionError != nil { return fmt.Errorf("die transaktion konnte nicht begonnen werden: %w", transactionError) } defer func() { _ = databaseTransaction.Rollback(replaceContext) }() if _, deleteError := databaseTransaction.Exec(replaceContext, "DELETE FROM user_recovery_codes WHERE user_id = $1", userID); deleteError != nil { return fmt.Errorf("die bisherigen codes konnten nicht entfernt werden: %w", deleteError) } for _, codeHash := range codeHashes { if _, insertError := databaseTransaction.Exec(replaceContext, "INSERT INTO user_recovery_codes (user_id, code_hash) VALUES ($1, $2)", userID, codeHash); insertError != nil { return fmt.Errorf("ein wiederherstellungscode konnte nicht gespeichert werden: %w", insertError) } } if commitError := databaseTransaction.Commit(replaceContext); commitError != nil { return fmt.Errorf("die wiederherstellungscodes konnten nicht gespeichert werden: %w", commitError) } return nil } // recoveryCodeRecord ist ein ungenutzter Wiederherstellungscode. type recoveryCodeRecord struct { // id ist der Bezeichner des Codes. id uuid.UUID // codeHash ist der gespeicherte Argon2id-Hash. codeHash string } // listUnusedRecoveryCodes liest alle noch nicht verwendeten Codes. func (repository *Repository) listUnusedRecoveryCodes(queryContext context.Context, userID uuid.UUID) ([]recoveryCodeRecord, error) { const selectStatement = ` SELECT id, code_hash FROM user_recovery_codes WHERE user_id = $1 AND used_at IS NULL` codeRows, queryError := repository.connectionPool.Query(queryContext, selectStatement, userID) if queryError != nil { return nil, fmt.Errorf("die wiederherstellungscodes konnten nicht gelesen werden: %w", queryError) } defer codeRows.Close() var unusedCodes []recoveryCodeRecord for codeRows.Next() { var record recoveryCodeRecord if scanError := codeRows.Scan(&record.id, &record.codeHash); scanError != nil { return nil, fmt.Errorf("ein wiederherstellungscode konnte nicht gelesen werden: %w", scanError) } unusedCodes = append(unusedCodes, record) } if rowsError := codeRows.Err(); rowsError != nil { return nil, fmt.Errorf("die wiederherstellungscodes konnten nicht vollständig gelesen werden: %w", rowsError) } return unusedCodes, nil } // consumeRecoveryCode entwertet einen Wiederherstellungscode. // // Die Bedingung auf used_at verhindert, dass derselbe Code bei gleichzeitigen // Anfragen zweimal gilt. func (repository *Repository) consumeRecoveryCode(updateContext context.Context, codeID uuid.UUID) error { commandTag, updateError := repository.connectionPool.Exec(updateContext, "UPDATE user_recovery_codes SET used_at = now() WHERE id = $1 AND used_at IS NULL", codeID) if updateError != nil { return fmt.Errorf("der wiederherstellungscode konnte nicht entwertet werden: %w", updateError) } if commandTag.RowsAffected() == 0 { return ErrInvalidMFACode } return nil } // nullIfEmptyInet wandelt eine leere Adresse in NULL. func nullIfEmptyInet(ipAddress string) *string { if ipAddress == "" { return nil } return &ipAddress }