fix: sync harden auth, privacy, and PTS state
This commit is contained in:
parent
c1597696af
commit
2512eab51d
24 changed files with 817 additions and 154 deletions
|
|
@ -3,7 +3,9 @@ package account
|
|||
import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"crypto/sha256"
|
||||
"crypto/subtle"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
|
@ -26,6 +28,10 @@ const (
|
|||
codeChannelEmailChange = "email_change"
|
||||
codeChannelEmailLogin = "email_login"
|
||||
codeChannelEmailSetupRequired = "email_setup_required"
|
||||
codeChannelPasswordRecovery = "password_recovery"
|
||||
passwordRecoveryCodePrefix = "password-recovery:"
|
||||
passwordRecoveryCodeTTL = 15 * time.Minute
|
||||
passwordRecoveryCASRetries = 32
|
||||
)
|
||||
|
||||
// Service 提供账号安全配置查询。
|
||||
|
|
@ -388,12 +394,48 @@ func (s *Service) RequestPasswordRecovery(ctx context.Context, userID int64) (st
|
|||
if !settings.HasPassword || settings.RecoveryEmail == "" {
|
||||
return "", domain.ErrPasswordRecoveryNA
|
||||
}
|
||||
settings.RecoveryCode = recoveryCode
|
||||
settings.RecoveryCodeExpiresAt = time.Now().Unix() + recoveryCodeTTL
|
||||
if s.passwords != nil {
|
||||
if err := s.passwords.Save(ctx, userID, settings); err != nil {
|
||||
return "", err
|
||||
// Recovery has no development-code fallback. If the server cannot both
|
||||
// persist and deliver a fresh code, it must report the flow unavailable.
|
||||
if s == nil || s.codes == nil || s.loginEmailSender == nil {
|
||||
return "", domain.ErrPasswordRecoveryNA
|
||||
}
|
||||
code, err := randomDigits(s.loginEmailCodeLength)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
deliveryID, err := otpdelivery.NewDeliveryID()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
expiresAt := time.Now().Add(passwordRecoveryCodeTTL)
|
||||
key := passwordRecoveryCodeKey(userID)
|
||||
rec := store.PhoneCode{
|
||||
Version: store.PhoneCodeVersionCurrent,
|
||||
UserID: userID,
|
||||
Email: normalizeLoginEmail(settings.RecoveryEmail),
|
||||
Code: code,
|
||||
DeliveryID: deliveryID,
|
||||
Channel: codeChannelPasswordRecovery,
|
||||
MaxAttempts: s.loginEmailCodeMaxAttempts,
|
||||
RecoveryBinding: passwordRecoveryBinding(settings),
|
||||
}
|
||||
if err := s.codes.Set(ctx, key, rec, passwordRecoveryCodeTTL); err != nil {
|
||||
return "", err
|
||||
}
|
||||
if err := deliverOTP(ctx, s.loginEmailSender, otpdelivery.Request{
|
||||
DeliveryID: deliveryID,
|
||||
Purpose: otpdelivery.PurposePasswordRecovery,
|
||||
Channel: otpdelivery.ChannelEmail,
|
||||
Recipient: settings.RecoveryEmail,
|
||||
Code: code,
|
||||
ExpiresAt: expiresAt,
|
||||
}); err != nil {
|
||||
cleanupCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), 2*time.Second)
|
||||
defer cancel()
|
||||
if cleanupErr := s.deletePasswordRecoveryCode(cleanupCtx, key, deliveryID); cleanupErr != nil {
|
||||
return "", fmt.Errorf("deliver password recovery code: %w (cleanup failed: %v)", err, cleanupErr)
|
||||
}
|
||||
return "", fmt.Errorf("deliver password recovery code: %w", err)
|
||||
}
|
||||
return emailPattern(settings.RecoveryEmail), nil
|
||||
}
|
||||
|
|
@ -403,23 +445,35 @@ func (s *Service) CheckRecoveryPassword(ctx context.Context, userID int64, code
|
|||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return checkRecoveryCode(settings, code)
|
||||
if !settings.HasPassword || settings.RecoveryEmail == "" {
|
||||
return domain.ErrPasswordRecoveryExpired
|
||||
}
|
||||
return s.verifyPasswordRecoveryCode(ctx, userID, passwordRecoveryBinding(settings), code, false)
|
||||
}
|
||||
|
||||
func (s *Service) RecoverPassword(ctx context.Context, userID int64, code string, input *domain.PasswordInputSettings) error {
|
||||
if s == nil || s.passwords == nil || userID == 0 {
|
||||
return nil
|
||||
return domain.ErrPasswordRecoveryExpired
|
||||
}
|
||||
settings, err := s.GetPasswordWithoutRefresh(ctx, userID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := checkRecoveryCode(settings, code); err != nil {
|
||||
return err
|
||||
if !settings.HasPassword || settings.RecoveryEmail == "" {
|
||||
return domain.ErrPasswordRecoveryExpired
|
||||
}
|
||||
binding := passwordRecoveryBinding(settings)
|
||||
if input == nil || len(input.NewPasswordHash) == 0 {
|
||||
settings = defaultPasswordSettings()
|
||||
return s.passwords.Save(ctx, userID, settings)
|
||||
if err := s.verifyPasswordRecoveryCode(ctx, userID, binding, code, true); err != nil {
|
||||
return err
|
||||
}
|
||||
return s.passwords.Save(ctx, userID, defaultPasswordSettings())
|
||||
}
|
||||
// Reject invalid proofs before doing the comparatively expensive SRP
|
||||
// verifier/challenge work. The final consuming check below still decides
|
||||
// the single winner if the code changes concurrently.
|
||||
if err := s.verifyPasswordRecoveryCode(ctx, userID, binding, code, false); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := validateNewPasswordSettings(*input); err != nil {
|
||||
return err
|
||||
|
|
@ -438,8 +492,9 @@ func (s *Service) RecoverPassword(ctx context.Context, userID int64, code string
|
|||
if input.HasHint {
|
||||
settings.Hint = input.Hint
|
||||
}
|
||||
settings.RecoveryCode = ""
|
||||
settings.RecoveryCodeExpiresAt = 0
|
||||
if err := s.verifyPasswordRecoveryCode(ctx, userID, binding, code, true); err != nil {
|
||||
return err
|
||||
}
|
||||
return s.passwords.Save(ctx, userID, normalizePasswordSettings(settings))
|
||||
}
|
||||
|
||||
|
|
@ -498,33 +553,123 @@ func (s *Service) ResendPasswordEmail(ctx context.Context, userID int64) error {
|
|||
}
|
||||
|
||||
func (s *Service) CancelPasswordEmail(ctx context.Context, userID int64) error {
|
||||
if s != nil && s.codes != nil && userID != 0 {
|
||||
if err := s.codes.Del(ctx, passwordRecoveryCodeKey(userID)); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
settings, err := s.GetPasswordWithoutRefresh(ctx, userID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
settings.EmailUnconfirmedPattern = ""
|
||||
settings.RecoveryCode = ""
|
||||
settings.RecoveryCodeExpiresAt = 0
|
||||
if s.passwords != nil {
|
||||
return s.passwords.Save(ctx, userID, settings)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func checkRecoveryCode(settings domain.PasswordSettings, code string) error {
|
||||
if settings.RecoveryCode == "" {
|
||||
if code == recoveryCode {
|
||||
func passwordRecoveryCodeKey(userID int64) string {
|
||||
return passwordRecoveryCodePrefix + fmt.Sprint(userID)
|
||||
}
|
||||
|
||||
func passwordRecoveryBinding(settings domain.PasswordSettings) string {
|
||||
sum := sha256.Sum256([]byte(fmt.Sprintf("%d\x00%s", settings.SRPID, normalizeLoginEmail(settings.RecoveryEmail))))
|
||||
return hex.EncodeToString(sum[:])
|
||||
}
|
||||
|
||||
func (s *Service) deletePasswordRecoveryCode(ctx context.Context, key, deliveryID string) error {
|
||||
for attempt := 0; attempt < passwordRecoveryCASRetries; attempt++ {
|
||||
snapshot, found, err := s.codes.GetSnapshot(ctx, key)
|
||||
if err != nil || !found {
|
||||
return err
|
||||
}
|
||||
if snapshot.Record.Channel != codeChannelPasswordRecovery || snapshot.Record.DeliveryID != deliveryID {
|
||||
return nil
|
||||
}
|
||||
return domain.ErrPasswordRecoveryNA
|
||||
deleted, err := s.codes.CompareAndDelete(ctx, key, snapshot.Revision)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if deleted {
|
||||
return nil
|
||||
}
|
||||
if err := ctx.Err(); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
if settings.RecoveryCodeExpiresAt > 0 && time.Now().Unix() > settings.RecoveryCodeExpiresAt {
|
||||
return domain.ErrEmailCodeInvalid
|
||||
return fmt.Errorf("delete password recovery code: concurrent state did not settle")
|
||||
}
|
||||
|
||||
// verifyPasswordRecoveryCode keeps check non-consuming while making the final
|
||||
// recovery a single-winner CAS. Wrong attempts are counted atomically and the
|
||||
// code is removed at the configured threshold.
|
||||
func (s *Service) verifyPasswordRecoveryCode(ctx context.Context, userID int64, binding, code string, consume bool) error {
|
||||
code = strings.TrimSpace(code)
|
||||
if code == "" {
|
||||
return domain.ErrRecoveryCodeEmpty
|
||||
}
|
||||
if subtle.ConstantTimeCompare([]byte(settings.RecoveryCode), []byte(code)) != 1 {
|
||||
return domain.ErrEmailCodeInvalid
|
||||
if s == nil || s.codes == nil || userID == 0 {
|
||||
return domain.ErrPasswordRecoveryExpired
|
||||
}
|
||||
return nil
|
||||
key := passwordRecoveryCodeKey(userID)
|
||||
for attempt := 0; attempt < passwordRecoveryCASRetries; attempt++ {
|
||||
snapshot, found, err := s.codes.GetSnapshot(ctx, key)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if !found {
|
||||
return domain.ErrPasswordRecoveryExpired
|
||||
}
|
||||
rec := snapshot.Record
|
||||
if rec.Channel != codeChannelPasswordRecovery || rec.UserID != userID || rec.RecoveryBinding == "" || rec.RecoveryBinding != binding {
|
||||
deleted, err := s.codes.CompareAndDelete(ctx, key, snapshot.Revision)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if deleted {
|
||||
return domain.ErrPasswordRecoveryExpired
|
||||
}
|
||||
continue
|
||||
}
|
||||
if subtle.ConstantTimeCompare([]byte(rec.Code), []byte(code)) == 1 {
|
||||
if !consume {
|
||||
return nil
|
||||
}
|
||||
deleted, err := s.codes.CompareAndDelete(ctx, key, snapshot.Revision)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if deleted {
|
||||
return nil
|
||||
}
|
||||
continue
|
||||
}
|
||||
maxAttempts := rec.MaxAttempts
|
||||
if maxAttempts <= 0 {
|
||||
maxAttempts = s.loginEmailCodeMaxAttempts
|
||||
}
|
||||
if maxAttempts <= 0 {
|
||||
maxAttempts = 1
|
||||
}
|
||||
rec.Attempts++
|
||||
var applied bool
|
||||
if rec.Attempts >= maxAttempts {
|
||||
applied, err = s.codes.CompareAndDelete(ctx, key, snapshot.Revision)
|
||||
} else {
|
||||
applied, err = s.codes.CompareAndUpdate(ctx, key, snapshot.Revision, rec)
|
||||
}
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if applied {
|
||||
return domain.ErrRecoveryCodeInvalid
|
||||
}
|
||||
if err := ctx.Err(); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return fmt.Errorf("verify password recovery code: concurrent state did not settle")
|
||||
}
|
||||
|
||||
func randomBytesOrDefault(n int, fallback []byte) []byte {
|
||||
|
|
@ -557,14 +702,24 @@ func randomDigits(n int) (string, error) {
|
|||
if n <= 0 {
|
||||
n = 6
|
||||
}
|
||||
b := make([]byte, n)
|
||||
if _, err := rand.Read(b); err != nil {
|
||||
return "", err
|
||||
}
|
||||
var out strings.Builder
|
||||
out.Grow(n)
|
||||
for _, v := range b {
|
||||
out.WriteByte(byte('0') + v%10)
|
||||
var buf [32]byte
|
||||
for out.Len() < n {
|
||||
if _, err := rand.Read(buf[:]); err != nil {
|
||||
return "", err
|
||||
}
|
||||
for _, v := range buf {
|
||||
// Reject the top six values so every digit has exactly 25 source
|
||||
// byte values instead of inheriting modulo bias from 256 %% 10.
|
||||
if v >= 250 {
|
||||
continue
|
||||
}
|
||||
out.WriteByte(byte('0') + v%10)
|
||||
if out.Len() == n {
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
return out.String(), nil
|
||||
}
|
||||
|
|
|
|||
|
|
@ -6,12 +6,14 @@ import (
|
|||
"crypto/sha512"
|
||||
"errors"
|
||||
"math/big"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"golang.org/x/crypto/pbkdf2"
|
||||
|
||||
"telesrv/internal/domain"
|
||||
"telesrv/internal/otpdelivery"
|
||||
"telesrv/internal/store/memory"
|
||||
)
|
||||
|
||||
|
|
@ -69,7 +71,9 @@ func TestPasswordSRPRoundTrip(t *testing.T) {
|
|||
func TestRecoverPasswordClearsTwoFactorPassword(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
const userID int64 = 1002
|
||||
svc := NewService(memory.NewPasswordStore())
|
||||
sender := &captureMailSender{}
|
||||
svc := NewService(memory.NewPasswordStore(),
|
||||
WithLoginEmailVerification(memory.NewCodeStore(), sender, time.Minute, 3, 6))
|
||||
|
||||
initial, err := svc.GetPassword(ctx, userID)
|
||||
if err != nil {
|
||||
|
|
@ -93,7 +97,13 @@ func TestRecoverPasswordClearsTwoFactorPassword(t *testing.T) {
|
|||
if pattern != "b***b@example.com" {
|
||||
t.Fatalf("recovery pattern = %q, want masked email", pattern)
|
||||
}
|
||||
if err := svc.RecoverPassword(ctx, userID, recoveryCode, nil); err != nil {
|
||||
if sender.to != "bob@example.com" || sender.code == "" || len(sender.requests) != 1 || sender.requests[0].Purpose != otpdelivery.PurposePasswordRecovery {
|
||||
t.Fatalf("recovery delivery = %+v, want one password-recovery email", sender)
|
||||
}
|
||||
if err := svc.CheckRecoveryPassword(ctx, userID, sender.code); err != nil {
|
||||
t.Fatalf("CheckRecoveryPassword: %v", err)
|
||||
}
|
||||
if err := svc.RecoverPassword(ctx, userID, sender.code, nil); err != nil {
|
||||
t.Fatalf("RecoverPassword clear: %v", err)
|
||||
}
|
||||
cleared, err := svc.GetPassword(ctx, userID)
|
||||
|
|
@ -105,6 +115,155 @@ func TestRecoverPasswordClearsTwoFactorPassword(t *testing.T) {
|
|||
}
|
||||
}
|
||||
|
||||
func TestPasswordRecoveryFailsClosedWithoutSenderOrIssuedCode(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
const userID int64 = 1012
|
||||
passwords := memory.NewPasswordStore()
|
||||
if err := passwords.Save(ctx, userID, domain.PasswordSettings{
|
||||
HasPassword: true, HasRecovery: true, RecoveryEmail: "owner@example.test",
|
||||
SRPID: 7, SRPVerifier: []byte{1, 2, 3},
|
||||
}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
svc := NewService(passwords)
|
||||
if _, err := svc.RequestPasswordRecovery(ctx, userID); !errors.Is(err, domain.ErrPasswordRecoveryNA) {
|
||||
t.Fatalf("RequestPasswordRecovery err=%v, want unavailable", err)
|
||||
}
|
||||
if err := svc.RecoverPassword(ctx, userID, "12345", nil); !errors.Is(err, domain.ErrPasswordRecoveryExpired) {
|
||||
t.Fatalf("standing fixed code err=%v, want expired", err)
|
||||
}
|
||||
if err := NewService(nil).RecoverPassword(ctx, userID, "12345", nil); !errors.Is(err, domain.ErrPasswordRecoveryExpired) {
|
||||
t.Fatalf("missing password store err=%v, want fail-closed expiry", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPasswordRecoveryAttemptLimitAndStateBinding(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
const userID int64 = 1013
|
||||
passwords := memory.NewPasswordStore()
|
||||
settings := domain.PasswordSettings{
|
||||
HasPassword: true, HasRecovery: true, RecoveryEmail: "owner@example.test",
|
||||
SRPID: 8, SRPVerifier: []byte{4, 5, 6},
|
||||
}
|
||||
if err := passwords.Save(ctx, userID, settings); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
sender := &captureMailSender{}
|
||||
svc := NewService(passwords,
|
||||
WithLoginEmailVerification(memory.NewCodeStore(), sender, time.Minute, 3, 6))
|
||||
if _, err := svc.RequestPasswordRecovery(ctx, userID); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for attempt := 1; attempt <= 3; attempt++ {
|
||||
if err := svc.CheckRecoveryPassword(ctx, userID, "000000"); !errors.Is(err, domain.ErrRecoveryCodeInvalid) {
|
||||
t.Fatalf("wrong attempt %d err=%v, want invalid", attempt, err)
|
||||
}
|
||||
}
|
||||
if err := svc.CheckRecoveryPassword(ctx, userID, sender.code); !errors.Is(err, domain.ErrPasswordRecoveryExpired) {
|
||||
t.Fatalf("code after attempt limit err=%v, want expired", err)
|
||||
}
|
||||
|
||||
if _, err := svc.RequestPasswordRecovery(ctx, userID); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
issued := sender.code
|
||||
settings.SRPID++
|
||||
if err := passwords.Save(ctx, userID, settings); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := svc.CheckRecoveryPassword(ctx, userID, issued); !errors.Is(err, domain.ErrPasswordRecoveryExpired) {
|
||||
t.Fatalf("code after 2FA state change err=%v, want expired", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestConcurrentPasswordRecoveryHasSingleConsumer(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
const userID int64 = 1014
|
||||
passwords := memory.NewPasswordStore()
|
||||
if err := passwords.Save(ctx, userID, domain.PasswordSettings{
|
||||
HasPassword: true, HasRecovery: true, RecoveryEmail: "owner@example.test",
|
||||
SRPID: 9, SRPVerifier: []byte{7, 8, 9},
|
||||
}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
sender := &captureMailSender{}
|
||||
svc := NewService(passwords,
|
||||
WithLoginEmailVerification(memory.NewCodeStore(), sender, time.Minute, 5, 6))
|
||||
if _, err := svc.RequestPasswordRecovery(ctx, userID); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
const workers = 24
|
||||
start := make(chan struct{})
|
||||
errs := make(chan error, workers)
|
||||
var wg sync.WaitGroup
|
||||
for i := 0; i < workers; i++ {
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
<-start
|
||||
errs <- svc.RecoverPassword(ctx, userID, sender.code, nil)
|
||||
}()
|
||||
}
|
||||
close(start)
|
||||
wg.Wait()
|
||||
close(errs)
|
||||
successes := 0
|
||||
for err := range errs {
|
||||
if err == nil {
|
||||
successes++
|
||||
continue
|
||||
}
|
||||
if !errors.Is(err, domain.ErrPasswordRecoveryExpired) {
|
||||
t.Fatalf("concurrent recovery err=%v", err)
|
||||
}
|
||||
}
|
||||
if successes != 1 {
|
||||
t.Fatalf("successful recoveries=%d, want 1", successes)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPasswordRecoveryDeliveryFailureRemovesIssuedCode(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
const userID int64 = 1015
|
||||
passwords := memory.NewPasswordStore()
|
||||
if err := passwords.Save(ctx, userID, domain.PasswordSettings{
|
||||
HasPassword: true, HasRecovery: true, RecoveryEmail: "owner@example.test",
|
||||
SRPID: 10, SRPVerifier: []byte{10},
|
||||
}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
sender := &captureMailSender{err: errors.New("provider rejected request")}
|
||||
svc := NewService(passwords,
|
||||
WithLoginEmailVerification(memory.NewCodeStore(), sender, time.Minute, 5, 6))
|
||||
if _, err := svc.RequestPasswordRecovery(ctx, userID); err == nil {
|
||||
t.Fatal("RequestPasswordRecovery succeeded after known delivery failure")
|
||||
}
|
||||
if err := svc.CheckRecoveryPassword(ctx, userID, sender.code); !errors.Is(err, domain.ErrPasswordRecoveryExpired) {
|
||||
t.Fatalf("undelivered code err=%v, want expired", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPasswordRecoveryUnknownDeliveryOutcomeKeepsIssuedCode(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
const userID int64 = 1016
|
||||
passwords := memory.NewPasswordStore()
|
||||
if err := passwords.Save(ctx, userID, domain.PasswordSettings{
|
||||
HasPassword: true, HasRecovery: true, RecoveryEmail: "owner@example.test",
|
||||
SRPID: 11, SRPVerifier: []byte{11},
|
||||
}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
sender := &captureMailSender{err: &otpdelivery.OutcomeUnknownError{Cause: errors.New("provider ACK lost")}}
|
||||
svc := NewService(passwords,
|
||||
WithLoginEmailVerification(memory.NewCodeStore(), sender, time.Minute, 5, 6))
|
||||
if _, err := svc.RequestPasswordRecovery(ctx, userID); err != nil {
|
||||
t.Fatalf("RequestPasswordRecovery outcome-unknown err=%v", err)
|
||||
}
|
||||
if err := svc.CheckRecoveryPassword(ctx, userID, sender.code); err != nil {
|
||||
t.Fatalf("outcome-unknown code was discarded: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResetPasswordWaitAndDecline(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
const userID int64 = 1003
|
||||
|
|
|
|||
|
|
@ -12,8 +12,6 @@ import (
|
|||
|
||||
const (
|
||||
passwordHashSize = 256
|
||||
recoveryCode = "12345"
|
||||
recoveryCodeTTL = 15 * 60
|
||||
)
|
||||
|
||||
var (
|
||||
|
|
|
|||
|
|
@ -775,10 +775,23 @@ func (s *Service) CancelCodeForAuthKey(ctx context.Context, authKeyID [8]byte, p
|
|||
return s.cancelCode(ctx, authKeyID, phone, phoneCodeHash)
|
||||
}
|
||||
|
||||
// LoginEmailResetAvailable reports whether this deployment can complete the
|
||||
// SMS fallback promised by auth.resetLoginEmail. Fixed development codes are
|
||||
// deliberately not a recovery channel.
|
||||
func (s *Service) LoginEmailResetAvailable() bool {
|
||||
return s != nil && s.phoneCodeSender != nil && s.codes != nil && s.users != nil
|
||||
}
|
||||
|
||||
// ConsumeLoginEmailReset authorizes auth.resetLoginEmail with the exact
|
||||
// email-login hash previously issued for this phone owner. Possession of only
|
||||
// a phone number is never sufficient to remove an authentication factor.
|
||||
func (s *Service) ConsumeLoginEmailReset(ctx context.Context, phone, phoneCodeHash string) (int64, error) {
|
||||
// Refuse before consuming the email proof or clearing any account state. If
|
||||
// there is no real SMS sender, the successor code would be the public
|
||||
// development code and could strip the login-email factor.
|
||||
if !s.LoginEmailResetAvailable() || s.codes == nil {
|
||||
return 0, ErrCodeInvalid
|
||||
}
|
||||
phone = normalizePhone(phone)
|
||||
rec, found, err := s.codes.Get(ctx, phoneCodeHash)
|
||||
if err != nil {
|
||||
|
|
|
|||
|
|
@ -399,7 +399,12 @@ func TestConsumeLoginEmailResetRequiresExactIssuedHash(t *testing.T) {
|
|||
users := &switchablePhoneOwnerStore{UserStore: baseUsers}
|
||||
codes := memory.NewCodeStore()
|
||||
delivery := &captureLoginCodeDelivery{}
|
||||
svc := NewService(users, memory.NewAuthorizationStore(), codes, nil, nil, "12345", WithLoginCodeDelivery(delivery))
|
||||
otp := &captureOTPSender{}
|
||||
svc := NewService(users, memory.NewAuthorizationStore(), codes, nil, nil, "12345",
|
||||
WithLoginCodeDelivery(delivery), WithPhoneCodeDelivery(otp, 5))
|
||||
if !svc.LoginEmailResetAvailable() {
|
||||
t.Fatal("LoginEmailResetAvailable=false with real SMS sender")
|
||||
}
|
||||
seed := func(hash, channel string) {
|
||||
t.Helper()
|
||||
if err := codes.Set(ctx, hash, store.PhoneCode{
|
||||
|
|
@ -451,11 +456,38 @@ func TestConsumeLoginEmailResetRequiresExactIssuedHash(t *testing.T) {
|
|||
if len(delivery.requests) != 1 || delivery.requests[0].UserID != owner.ID || delivery.requests[0].PhoneCodeHash != replacementHash {
|
||||
t.Fatalf("replacement delivery=%+v", delivery.requests)
|
||||
}
|
||||
if rec, found, err := codes.Get(ctx, replacementHash); err != nil || !found || rec.Version != store.PhoneCodeVersionCurrent || rec.IssuedUserID != owner.ID || rec.Channel != codeChannelPhone {
|
||||
if rec, found, err := codes.Get(ctx, replacementHash); err != nil || !found || rec.Version != store.PhoneCodeVersionCurrent || rec.IssuedUserID != owner.ID || rec.Channel != codeChannelSMS {
|
||||
t.Fatalf("replacement code=%+v found=%v err=%v", rec, found, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoginEmailResetUnavailableWithoutRealSMSSender(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
users := memory.NewUserStore()
|
||||
owner, err := users.Create(ctx, domain.User{Phone: "15550009339", FirstName: "Owner"})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
codes := memory.NewCodeStore()
|
||||
const hash = "unavailable-email-reset"
|
||||
if err := codes.Set(ctx, hash, store.PhoneCode{
|
||||
Version: store.PhoneCodeVersionCurrent, IssuedUserID: owner.ID,
|
||||
Phone: owner.Phone, Code: "654321", Channel: codeChannelEmailLogin,
|
||||
}, time.Minute); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
svc := NewService(users, memory.NewAuthorizationStore(), codes, nil, nil, "12345")
|
||||
if svc.LoginEmailResetAvailable() {
|
||||
t.Fatal("LoginEmailResetAvailable=true without real SMS sender")
|
||||
}
|
||||
if _, err := svc.ConsumeLoginEmailReset(ctx, owner.Phone, hash); !errors.Is(err, ErrCodeInvalid) {
|
||||
t.Fatalf("ConsumeLoginEmailReset err=%v, want invalid", err)
|
||||
}
|
||||
if _, found, err := codes.Get(ctx, hash); err != nil || !found {
|
||||
t.Fatalf("unavailable reset consumed proof found=%v err=%v", found, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestConcurrentLoginEmailResetHasSingleConsumer(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
users := memory.NewUserStore()
|
||||
|
|
@ -475,7 +507,8 @@ func TestConcurrentLoginEmailResetHasSingleConsumer(t *testing.T) {
|
|||
}, time.Minute); err != nil {
|
||||
t.Fatalf("seed code: %v", err)
|
||||
}
|
||||
svc := NewService(users, memory.NewAuthorizationStore(), codes, nil, nil, "12345")
|
||||
svc := NewService(users, memory.NewAuthorizationStore(), codes, nil, nil, "12345",
|
||||
WithPhoneCodeDelivery(&captureOTPSender{}, 5))
|
||||
const workers = 24
|
||||
start := make(chan struct{})
|
||||
errs := make(chan error, workers)
|
||||
|
|
@ -536,7 +569,8 @@ func TestLoginEmailResetLocksUserAcrossOwnerTransfer(t *testing.T) {
|
|||
t.Fatalf("seed reset code: %v", err)
|
||||
}
|
||||
delivery := &captureLoginCodeDelivery{}
|
||||
authSvc := NewService(users, memory.NewAuthorizationStore(), codes, nil, nil, "12345", WithLoginCodeDelivery(delivery))
|
||||
authSvc := NewService(users, memory.NewAuthorizationStore(), codes, nil, nil, "12345",
|
||||
WithLoginCodeDelivery(delivery), WithPhoneCodeDelivery(&captureOTPSender{}, 5))
|
||||
resetUserID, err := authSvc.ConsumeLoginEmailReset(ctx, ownerA.Phone, hash)
|
||||
if err != nil || resetUserID != ownerA.ID {
|
||||
t.Fatalf("ConsumeLoginEmailReset uid=%d err=%v", resetUserID, err)
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue