Files
zitadel/internal/command/user_human_recovery_codes.go
T
Gayathri Vijayan fb35fdd3c7 feat: add recovery codes to user repository (#11702)
# Which Problems Are Solved

Adds a repository implementation to add/remove recovery codes to the
users relational table.

# How the Problems Are Solved

This is achieved by:
* adding the following columns to the users table: `recovery_codes`,
`recovery_code_last_successful_check`, `recovery_code_failed_attempts`
* adding the `RecoveryCodes` field to the `HumanUser` domain with fields
to set recovery `codes`, `lastSuccessfullyCheckedAt` timestamp, and
`failedAttempts`
* setting `recoveryCodes` in the `Get` user query statement to return a
json object with details related to the recovery codes
* adding the repository-layer implementation to add/remove recovery
codes and set fields related to recovery code checks.
* adding projection reducers to handle the following events:
`HumanRecoveryCodesAddedEvent`, `HumanRecoveryCodesRemovedEvent`,
`HumanRecoveryCodeCheckSucceededEvent`, and
`HumanRecoveryCodeCheckFailedEvent`
* adding unit tests

# Additional Changes
* fix the error message when the recovery is empty in
`internal/command/user_human_recovery_codes.go`
* add a new error message for empty recovery code during checks in
`en.yaml`

# Additional Context
- Closes https://github.com/zitadel/zitadel/issues/11212
- Follow-up: integration tests for the reducers will be added in a
different PR after this
[PR](https://github.com/zitadel/zitadel/pull/11478) is merged
2026-02-26 19:22:37 +01:00

202 lines
6.9 KiB
Go

package command
import (
"context"
"github.com/zitadel/logging"
"github.com/zitadel/zitadel/internal/crypto"
"github.com/zitadel/zitadel/internal/domain"
"github.com/zitadel/zitadel/internal/eventstore"
"github.com/zitadel/zitadel/internal/repository/user"
"github.com/zitadel/zitadel/internal/telemetry/tracing"
"github.com/zitadel/zitadel/internal/zerrors"
)
func (c *Commands) ImportHumanRecoveryCodes(ctx context.Context, userID, resourceOwner string, codes []domain.ImportHumanRecoveryCode) (err error) {
ctx, span := tracing.NewSpan(ctx)
defer func() { span.EndWithError(err) }()
if len(codes) == 0 {
return zerrors.ThrowInvalidArgument(nil, "COMMAND-vee93", "Errors.User.MFA.RecoveryCodes.CountInvalid")
}
if _, err = c.checkUserExists(ctx, userID, resourceOwner); err != nil {
return err
}
recoveryCodeWriteModel := NewHumanRecoveryCodeWriteModel(userID, resourceOwner)
if err := c.eventstore.FilterToQueryReducer(ctx, recoveryCodeWriteModel); err != nil {
return err
}
if len(recoveryCodeWriteModel.Codes())+len(codes) > c.multifactors.RecoveryCodes.MaxCount {
return zerrors.ThrowPreconditionFailed(nil, "COMMAND-53cjw", "Errors.User.MFA.RecoveryCodes.MaxCountExceeded")
}
hashedCodes, err := domain.HashRecoveryCodesIfNeeded(ctx, codes, c.userPasswordHasher)
if err != nil {
return err
}
userAgg := UserAggregateFromWriteModelCtx(ctx, &recoveryCodeWriteModel.WriteModel)
_, err = c.eventstore.Push(ctx,
user.NewHumanRecoveryCodesAddedEvent(ctx, userAgg, hashedCodes, nil),
)
return err
}
type RecoveryCodesDetails struct {
*domain.ObjectDetails
RawCodes []string
}
func (c *Commands) GenerateRecoveryCodes(ctx context.Context, userID string, count int, resourceOwner string, authRequest *domain.AuthRequest) (*RecoveryCodesDetails, error) {
if userID == "" {
return nil, zerrors.ThrowInvalidArgument(nil, "COMMAND-4kje7", "Errors.User.UserIDMissing")
}
if count <= 0 {
return nil, zerrors.ThrowInvalidArgument(nil, "COMMAND-7c0nx", "Errors.User.RecoveryCodes.CountInvalid")
}
resourceOwner, err := c.checkUserExists(ctx, userID, resourceOwner)
if err != nil {
return nil, err
}
if err := c.checkPermissionUpdateUserCredentials(ctx, resourceOwner, userID); err != nil {
return nil, err
}
recoveryCodeWriteModel := NewHumanRecoveryCodeWriteModel(userID, resourceOwner)
if err := c.eventstore.FilterToQueryReducer(ctx, recoveryCodeWriteModel); err != nil {
return nil, err
}
if len(recoveryCodeWriteModel.Codes())+count > c.multifactors.RecoveryCodes.MaxCount {
return nil, zerrors.ThrowPreconditionFailed(nil, "COMMAND-8f2k9", "Errors.User.MFA.RecoveryCodes.MaxCountExceeded")
}
hashedCodes, rawCodes, err := domain.GenerateRecoveryCodes(ctx, count, c.multifactors.RecoveryCodes, c.userPasswordHasher)
if err != nil {
return nil, err
}
userAgg := UserAggregateFromWriteModelCtx(ctx, &recoveryCodeWriteModel.WriteModel)
_, err = c.eventstore.Push(ctx,
user.NewHumanRecoveryCodesAddedEvent(ctx, userAgg, hashedCodes, authRequestDomainToAuthRequestInfo(authRequest)),
)
if err != nil {
return nil, err
}
return &RecoveryCodesDetails{
ObjectDetails: writeModelToObjectDetails(&recoveryCodeWriteModel.WriteModel),
RawCodes: rawCodes,
}, nil
}
func (c *Commands) RemoveRecoveryCodes(ctx context.Context, userID, resourceOwner string, authRequest *domain.AuthRequest) (*domain.ObjectDetails, error) {
if userID == "" {
return nil, zerrors.ThrowInvalidArgument(nil, "COMMAND-l2n9r", "Errors.User.UserIDMissing")
}
writeModel := NewHumanRecoveryCodeWriteModel(userID, resourceOwner)
if err := c.eventstore.FilterToQueryReducer(ctx, writeModel); err != nil {
return nil, err
}
if err := c.checkPermissionUpdateUserCredentials(ctx, writeModel.ResourceOwner, userID); err != nil {
return nil, err
}
if writeModel.UserLocked() {
return nil, zerrors.ThrowPreconditionFailed(nil, "COMMAND-d9u8q", "Errors.User.Locked")
}
// if there aren't any recovery codes, we don't need to do anything
if writeModel.State != domain.MFAStateReady {
return writeModelToObjectDetails(&writeModel.WriteModel), nil
}
userAgg := UserAggregateFromWriteModelCtx(ctx, &writeModel.WriteModel)
_, err := c.eventstore.Push(ctx, user.NewHumanRecoveryCodeRemovedEvent(ctx, userAgg, authRequestDomainToAuthRequestInfo(authRequest)))
if err != nil {
return nil, err
}
return writeModelToObjectDetails(&writeModel.WriteModel), nil
}
func (c *Commands) HumanCheckRecoveryCode(ctx context.Context, userID, code, resourceOwner string, authRequest *domain.AuthRequest) error {
commands, err := checkRecoveryCode(ctx, userID, code, resourceOwner, authRequest, c.eventstore.FilterToQueryReducer, c.userPasswordHasher)
if len(commands) > 0 {
_, err = c.eventstore.Push(ctx, commands...)
logging.OnError(err).Error("failed to push recovery code check events")
}
return err
}
func checkRecoveryCode(
ctx context.Context,
userID, code, resourceOwner string,
authRequest *domain.AuthRequest,
queryReducer func(ctx context.Context, r eventstore.QueryReducer) error,
secretHasher *crypto.Hasher,
) ([]eventstore.Command, error) {
if code == "" {
return nil, zerrors.ThrowInvalidArgument(nil, "COMMAND-u0b6c", "Errors.User.MFA.RecoveryCodes.Empty")
}
if userID == "" {
return nil, zerrors.ThrowInvalidArgument(nil, "COMMAND-4m9s2", "Errors.User.UserIDMissing")
}
recoveryCodeWm := NewHumanRecoveryCodeWriteModel(userID, resourceOwner)
err := queryReducer(ctx, recoveryCodeWm)
if err != nil {
return nil, err
}
if recoveryCodeWm.UserLocked() {
return nil, zerrors.ThrowPreconditionFailed(nil, "COMMAND-2w6oa", "Errors.User.Locked")
}
if recoveryCodeWm.State != domain.MFAStateReady {
return nil, zerrors.ThrowPreconditionFailed(nil, "COMMAND-84rgg", "Errors.User.MFA.RecoveryCodes.NotReady")
}
hashedCode, err := domain.ValidateRecoveryCode(ctx, code, toHumanRecoveryCode(recoveryCodeWm), secretHasher)
// recheck for additional events (failed password checks or locks)
recheckErr := queryReducer(ctx, recoveryCodeWm)
if recheckErr != nil {
return nil, recheckErr
}
if recoveryCodeWm.UserLocked() {
return nil, zerrors.ThrowPreconditionFailed(nil, "COMMAND-ASV12", "Errors.User.Locked")
}
userAgg := UserAggregateFromWriteModelCtx(ctx, &recoveryCodeWm.WriteModel)
commands := make([]eventstore.Command, 0, 2)
authRequestInfo := authRequestDomainToAuthRequestInfo(authRequest)
if err == nil {
return append(commands, user.NewHumanRecoveryCodeCheckSucceededEvent(ctx, userAgg, hashedCode, authRequestInfo)), nil
}
commands = append(commands, user.NewHumanRecoveryCodeCheckFailedEvent(ctx, userAgg, authRequestInfo))
lockoutPolicy, lockoutErr := getLockoutPolicy(ctx, recoveryCodeWm.ResourceOwner, queryReducer)
logging.OnError(lockoutErr).Error("failed to get lockout policy")
if lockoutPolicy != nil && lockoutPolicy.MaxOTPAttempts > 0 && recoveryCodeWm.FailedAttempts+1 >= lockoutPolicy.MaxOTPAttempts {
commands = append(commands, user.NewUserLockedEvent(ctx, userAgg))
}
return commands, err
}