Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
96 changes: 96 additions & 0 deletions backend/internal/application/auth/mfa_setup_reenrol_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,96 @@
// Copyright (c) 2026 OpenDefender Contributors
// SPDX-License-Identifier: AGPL-3.0-only
// This program is free software: you can redistribute it and/or modify it under
// the terms of the GNU Affero General Public License v3.0 (see LICENSE).

package auth

import (
"context"
"errors"
"testing"
"time"

"github.com/google/uuid"
"github.com/pquerna/otp/totp"

"github.com/opendefender/openrisk/internal/domain"
"github.com/opendefender/openrisk/pkg/crypto"
"github.com/opendefender/openrisk/pkg/otp"
)

// uniqueMFARepository enforces mfa_secrets.user_id UNIQUE, which the plain mock
// does not: it overwrote on Create, so the defect behind #889 never showed.
type uniqueMFARepository struct{ *MockMFARepository }

func (r uniqueMFARepository) CreateMFASecret(ctx context.Context, s *domain.MFASecret) error {
for _, existing := range r.secrets {
if existing.UserID == s.UserID {
return errors.New("duplicated key not allowed")
}
}
return r.MockMFARepository.CreateMFASecret(ctx, s)
}

func TestSetupMFA_ReplacesAnUnverifiedSecret(t *testing.T) {
ctx := context.Background()
key := make([]byte, 32)
repo := uniqueMFARepository{NewMockMFARepository()}
uc := NewSetupMFAUseCase(repo, key)
in := SetupMFAInput{UserID: uuid.New(), TenantID: uuid.New(), Email: "awa@example.test"}

first, err := uc.Execute(ctx, in)
if err != nil {
t.Fatalf("first setup: %v", err)
}
// The enrolment is abandoned: the secret stays unverified. Starting over
// must work, not fail on the unique user_id forever.
second, err := uc.Execute(ctx, in)
if err != nil {
t.Fatalf("setup after an unfinished one must succeed, got %v", err)
}
if second.Secret == first.Secret {
t.Fatal("a fresh enrolment must get a fresh secret")
}

stored, _ := repo.GetMFASecret(ctx, in.UserID, in.TenantID)
plain, err := crypto.DecryptAES256GCM(stored.SecretEncrypted, key)
if err != nil {
t.Fatal(err)
}
if plain != second.Secret {
t.Fatal("the stored secret must be the new one, so only the new QR code verifies")
}
if stored.IsVerified {
t.Fatal("the replaced secret must still await verification")
}
code, _ := totp.GenerateCode(second.Secret, time.Now())
if !otp.VerifyTOTP(plain, code) {
t.Fatal("a code from the new QR code must verify against the stored secret")
}
oldCode, _ := totp.GenerateCode(first.Secret, time.Now())
if oldCode != code && otp.VerifyTOTP(plain, oldCode) {
t.Fatal("a code from the abandoned QR code must no longer verify")
}
}

func TestSetupMFA_VerifiedSecretIsKept(t *testing.T) {
ctx := context.Background()
key := make([]byte, 32)
repo := uniqueMFARepository{NewMockMFARepository()}
userID, tenantID := uuid.New(), uuid.New()
enc, _ := crypto.EncryptAES256GCM("KEEPME", key)
_ = repo.MockMFARepository.CreateMFASecret(ctx, &domain.MFASecret{
ID: uuid.New(), UserID: userID, TenantID: tenantID, SecretEncrypted: enc, IsVerified: true,
})

_, err := NewSetupMFAUseCase(repo, key).Execute(ctx, SetupMFAInput{UserID: userID, TenantID: tenantID, Email: "a@example.test"})
var appErr *domain.AppError
if !errors.As(err, &appErr) || !errors.Is(appErr.Err, domain.ErrConflict) {
t.Fatalf("a verified secret must answer Conflict, got %v", err)
}
stored, _ := repo.GetMFASecret(ctx, userID, tenantID)
if stored.SecretEncrypted != enc {
t.Fatal("a verified secret must not be touched")
}
}
34 changes: 24 additions & 10 deletions backend/internal/application/auth/mfa_usecase.go
Original file line number Diff line number Diff line change
Expand Up @@ -83,16 +83,30 @@ func (uc *SetupMFAUseCase) Execute(ctx context.Context, input SetupMFAInput) (*S
return nil, fmt.Errorf("failed to encrypt secret: %w", err)
}

// Store encrypted secret (not yet verified)
mfaSecret := &domain.MFASecret{
UserID: input.UserID,
TenantID: input.TenantID,
SecretEncrypted: encryptedSecret,
IsVerified: false,
}

if err := uc.mfaRepo.CreateMFASecret(ctx, mfaSecret); err != nil {
return nil, fmt.Errorf("failed to store MFA secret: %w", err)
// Store encrypted secret (not yet verified). An unverified secret left by
// an enrolment that was never finished (tab closed, token expired) is
// replaced: inserting a second row hit the unique user_id, and the account
// could never enrol again — locked out for good if its role requires MFA
// (#889). A verified secret was refused above and is never touched.
if existingSecret != nil {
replaced, err := uc.mfaRepo.ReplaceUnverifiedMFASecret(ctx, input.UserID, input.TenantID, encryptedSecret)
if err != nil {
return nil, fmt.Errorf("failed to replace unverified MFA secret: %w", err)
}
if !replaced {
// Verified between the read above and this write.
return nil, domain.NewConflictError("MFA", "already_enabled")
}
} else {
mfaSecret := &domain.MFASecret{
UserID: input.UserID,
TenantID: input.TenantID,
SecretEncrypted: encryptedSecret,
IsVerified: false,
}
if err := uc.mfaRepo.CreateMFASecret(ctx, mfaSecret); err != nil {
return nil, fmt.Errorf("failed to store MFA secret: %w", err)
}
}

// Generate backup codes (CSPRNG, unique per user — see otp.GenerateBackupCodes).
Expand Down
12 changes: 12 additions & 0 deletions backend/internal/application/auth/mfa_usecase_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -56,6 +56,18 @@ func (m *MockMFARepository) ConsumeTOTPStep(ctx context.Context, userID, tenantI
return true, nil
}

func (m *MockMFARepository) ReplaceUnverifiedMFASecret(ctx context.Context, userID, tenantID uuid.UUID, secretEncrypted string) (bool, error) {
key := userID.String() + ":" + tenantID.String()
s, ok := m.secrets[key]
if !ok || s.IsVerified {
return false, nil
}
s.SecretEncrypted = secretEncrypted
s.LastTOTPStep = nil
s.LastUsedAt = nil
return true, nil
}

func (m *MockMFARepository) DisableMFA(ctx context.Context, userID, tenantID uuid.UUID) error {
key := userID.String() + ":" + tenantID.String()
delete(m.secrets, key)
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,88 @@
// Copyright (c) 2026 OpenDefender Contributors
// SPDX-License-Identifier: AGPL-3.0-only
// This program is free software: you can redistribute it and/or modify it under
// the terms of the GNU Affero General Public License v3.0 (see LICENSE).

package repository

import (
"context"
"errors"
"os"
"testing"
"time"

"github.com/google/uuid"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/driver/postgres"
"gorm.io/gorm"
"gorm.io/gorm/logger"

"github.com/opendefender/openrisk/internal/domain"
)

// #889 — an unfinished enrolment can be started over: the unverified secret's
// key is replaced in place, and a verified one is never touched.

func checkReplaceUnverified(t *testing.T, db *gorm.DB) {
t.Helper()
repo := NewGormMFARepository(db)
ctx := context.Background()
step := int64(42)
used := time.Now()

pending, tenant := uuid.New(), uuid.New()
require.NoError(t, db.Create(&domain.MFASecret{
ID: uuid.New(), UserID: pending, TenantID: tenant, SecretEncrypted: "old",
LastTOTPStep: &step, LastUsedAt: &used,
}).Error)
enrolled := uuid.New()
require.NoError(t, db.Create(&domain.MFASecret{
ID: uuid.New(), UserID: enrolled, TenantID: tenant, SecretEncrypted: "kept", IsVerified: true,
}).Error)

// Another tenant's request cannot reach the row.
ok, err := repo.ReplaceUnverifiedMFASecret(ctx, pending, uuid.New(), "intruder")
require.NoError(t, err)
assert.False(t, ok, "tenant-scoped: another tenant replaces nothing")

ok, err = repo.ReplaceUnverifiedMFASecret(ctx, pending, tenant, "new")
require.NoError(t, err)
assert.True(t, ok)
var got domain.MFASecret
require.NoError(t, db.Where("user_id = ?", pending).First(&got).Error)
assert.Equal(t, "new", got.SecretEncrypted)
assert.False(t, got.IsVerified)
assert.Nil(t, got.LastTOTPStep, "the abandoned key's replay step must not carry over")
assert.Nil(t, got.LastUsedAt)

ok, err = repo.ReplaceUnverifiedMFASecret(ctx, enrolled, tenant, "attack")
require.NoError(t, err)
assert.False(t, ok, "a verified secret is never replaced")
var kept domain.MFASecret
require.NoError(t, db.Where("user_id = ?", enrolled).First(&kept).Error)
assert.Equal(t, "kept", kept.SecretEncrypted)
}

func TestGormMFARepository_ReplaceUnverifiedMFASecret(t *testing.T) {
_, db := setupMFARepo(t)
checkReplaceUnverified(t, db)
}

func TestGormMFARepository_ReplaceUnverifiedMFASecret_Postgres(t *testing.T) {
dsn := os.Getenv("DATABASE_URL")
if dsn == "" {
t.Skip("DATABASE_URL not set")
}
db, err := gorm.Open(postgres.Open(dsn), &gorm.Config{Logger: logger.Default.LogMode(logger.Silent)})
require.NoError(t, err)

rollback := errors.New("rollback")
err = db.Transaction(func(tx *gorm.DB) error {
require.NoError(t, tx.Exec(`CREATE TEMP TABLE mfa_secrets (LIKE public.mfa_secrets INCLUDING ALL) ON COMMIT DROP`).Error)
checkReplaceUnverified(t, tx)
return rollback
})
require.ErrorIs(t, err, rollback)
}
21 changes: 21 additions & 0 deletions backend/internal/infrastructure/repository/gorm_mfa_repository.go
Original file line number Diff line number Diff line change
Expand Up @@ -71,6 +71,27 @@ func (r *GormMFARepository) ConsumeTOTPStep(ctx context.Context, userID, tenantI
return res.RowsAffected == 1, nil
}

// ReplaceUnverifiedMFASecret gives an unfinished enrolment a new key (#889).
//
// One conditional UPDATE, not a read and a Save: the is_verified = false guard
// makes it lose cleanly against a verification that lands at the same moment,
// and last_totp_step is reset here explicitly rather than by a Save, which
// ConsumeTOTPStep's contract forbids (#849). The step and last use belonged to
// the abandoned key, so they start over with the new one.
func (r *GormMFARepository) ReplaceUnverifiedMFASecret(ctx context.Context, userID, tenantID uuid.UUID, secretEncrypted string) (bool, error) {
res := r.db.WithContext(ctx).Model(&domain.MFASecret{}).
Where("user_id = ? AND tenant_id = ? AND is_verified = ?", userID, tenantID, false).
Updates(map[string]any{
"secret_encrypted": secretEncrypted,
"last_totp_step": nil,
"last_used_at": nil,
})
if res.Error != nil {
return false, res.Error
}
return res.RowsAffected == 1, nil
}

// DisableMFA removes the TOTP secret and every backup code of one user, in one
// transaction (#754).
//
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,10 @@ type MFARepository interface {
// ConsumeTOTPStep marks a TOTP step as used; false means it (or a later one)
// already was, i.e. a replay (#849).
ConsumeTOTPStep(ctx context.Context, userID, tenantID uuid.UUID, step int64) (bool, error)
// ReplaceUnverifiedMFASecret swaps the key material of a secret that was
// never verified, for an enrolment started over (#889). false means there
// is no unverified secret to replace, e.g. it was verified meanwhile.
ReplaceUnverifiedMFASecret(ctx context.Context, userID, tenantID uuid.UUID, secretEncrypted string) (bool, error)
// DisableMFA deletes the secret AND the backup codes, atomically.
DisableMFA(ctx context.Context, userID, tenantID uuid.UUID) error

Expand Down
Loading