GetPassword minted a brand-new random SRP server secret and B on every call while only ever assigning SRPID once (when zero). Two account.getPassword calls in a row -- e.g. a settings screen refreshing state, then the transfer- ownership dialog's own cloudPassword().reload() moments later -- silently invalidated each other's B with no signal the client could detect (SRPID unchanged), so a password check built from the first response's B failed with PASSWORD_HASH_INVALID even though the typed password was correct. The challenge now stays stable across reads and only rotates when it's missing entirely; UpdatePasswordSettings/RecoverPassword already mint their own fresh challenge whenever the password actually changes. Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
411 lines
15 KiB
Go
411 lines
15 KiB
Go
package account
|
||
|
||
import (
|
||
"bytes"
|
||
"context"
|
||
"crypto/sha512"
|
||
"errors"
|
||
"math/big"
|
||
"sync"
|
||
"testing"
|
||
"time"
|
||
|
||
"golang.org/x/crypto/pbkdf2"
|
||
|
||
"telesrv/internal/domain"
|
||
"telesrv/internal/otpdelivery"
|
||
"telesrv/internal/store/memory"
|
||
)
|
||
|
||
func TestPasswordSRPRoundTrip(t *testing.T) {
|
||
ctx := context.Background()
|
||
const userID int64 = 1001
|
||
svc := NewService(memory.NewPasswordStore())
|
||
|
||
initial, err := svc.GetPassword(ctx, userID)
|
||
if err != nil {
|
||
t.Fatalf("GetPassword initial: %v", err)
|
||
}
|
||
algo := initial.NewAlgo
|
||
algo.Salt1 = append(append([]byte(nil), algo.Salt1...), bytes.Repeat([]byte{0xA5}, 32)...)
|
||
input := domain.PasswordInputSettings{
|
||
NewAlgo: &algo,
|
||
NewPasswordHash: verifierForPassword(algo, []byte("correct horse")),
|
||
Hint: "horse",
|
||
HasHint: true,
|
||
Email: "alice@example.com",
|
||
HasEmail: true,
|
||
}
|
||
if err := svc.UpdatePasswordSettings(ctx, userID, domain.PasswordCheck{Empty: true}, input); err != nil {
|
||
t.Fatalf("UpdatePasswordSettings set password: %v", err)
|
||
}
|
||
|
||
challenge, err := svc.GetPassword(ctx, userID)
|
||
if err != nil {
|
||
t.Fatalf("GetPassword challenge: %v", err)
|
||
}
|
||
if !challenge.HasPassword || challenge.SRPID == 0 || len(challenge.SRPB) == 0 {
|
||
t.Fatalf("challenge = %+v, want srp password challenge", challenge)
|
||
}
|
||
check := clientPasswordCheck(t, challenge, []byte("correct horse"))
|
||
if err := svc.CheckPassword(ctx, userID, check); err != nil {
|
||
t.Fatalf("CheckPassword valid SRP: %v", err)
|
||
}
|
||
|
||
private, err := svc.GetPasswordSettings(ctx, userID, check)
|
||
if err != nil {
|
||
t.Fatalf("GetPasswordSettings valid SRP: %v", err)
|
||
}
|
||
if private.Email != "alice@example.com" {
|
||
t.Fatalf("private email = %q, want alice@example.com", private.Email)
|
||
}
|
||
|
||
bad := check
|
||
bad.M1 = append([]byte(nil), check.M1...)
|
||
bad.M1[0] ^= 0xFF
|
||
if err := svc.CheckPassword(ctx, userID, bad); !errors.Is(err, domain.ErrPasswordHashInvalid) {
|
||
t.Fatalf("CheckPassword bad M1 err = %v, want ErrPasswordHashInvalid", err)
|
||
}
|
||
}
|
||
|
||
// TestGetPasswordChallengeStableAcrossReads guards against a real production
|
||
// bug: GetPassword used to mint a brand-new random SRP server secret/B on
|
||
// every single call while leaving SRPID untouched. Two account.getPassword
|
||
// calls in a row (e.g. a settings screen and, moments later, the transfer-
|
||
// ownership dialog's own cloudPassword().reload()) would silently invalidate
|
||
// each other's B with no signal the client could detect (SRPID unchanged),
|
||
// so a password check built from the first response's B failed with
|
||
// PASSWORD_HASH_INVALID even though the typed password was correct. The
|
||
// challenge must stay identical across reads until a real password change
|
||
// consumes it.
|
||
func TestGetPasswordChallengeStableAcrossReads(t *testing.T) {
|
||
ctx := context.Background()
|
||
const userID int64 = 1003
|
||
svc := NewService(memory.NewPasswordStore())
|
||
|
||
initial, err := svc.GetPassword(ctx, userID)
|
||
if err != nil {
|
||
t.Fatalf("GetPassword initial: %v", err)
|
||
}
|
||
algo := initial.NewAlgo
|
||
algo.Salt1 = append(append([]byte(nil), algo.Salt1...), bytes.Repeat([]byte{0x11}, 32)...)
|
||
if err := svc.UpdatePasswordSettings(ctx, userID, domain.PasswordCheck{Empty: true}, domain.PasswordInputSettings{
|
||
NewAlgo: &algo,
|
||
NewPasswordHash: verifierForPassword(algo, []byte("correct horse")),
|
||
}); err != nil {
|
||
t.Fatalf("UpdatePasswordSettings set password: %v", err)
|
||
}
|
||
|
||
first, err := svc.GetPassword(ctx, userID)
|
||
if err != nil {
|
||
t.Fatalf("GetPassword first: %v", err)
|
||
}
|
||
// Simulate a second, unrelated screen/dialog refreshing the same cloud
|
||
// password state before the user submits their check.
|
||
second, err := svc.GetPassword(ctx, userID)
|
||
if err != nil {
|
||
t.Fatalf("GetPassword second: %v", err)
|
||
}
|
||
if first.SRPID != second.SRPID || !bytes.Equal(first.SRPB, second.SRPB) {
|
||
t.Fatalf("challenge changed across reads: first=%+v second=%+v, want identical SRPID/SRPB", first, second)
|
||
}
|
||
|
||
check := clientPasswordCheck(t, first, []byte("correct horse"))
|
||
if err := svc.CheckPassword(ctx, userID, check); err != nil {
|
||
t.Fatalf("CheckPassword against first-read challenge after an intervening read: %v", err)
|
||
}
|
||
}
|
||
|
||
func TestRecoverPasswordClearsTwoFactorPassword(t *testing.T) {
|
||
ctx := context.Background()
|
||
const userID int64 = 1002
|
||
sender := &captureMailSender{}
|
||
svc := NewService(memory.NewPasswordStore(),
|
||
WithLoginEmailVerification(memory.NewCodeStore(), sender, time.Minute, 3, 6))
|
||
|
||
initial, err := svc.GetPassword(ctx, userID)
|
||
if err != nil {
|
||
t.Fatalf("GetPassword initial: %v", err)
|
||
}
|
||
algo := initial.NewAlgo
|
||
algo.Salt1 = append(append([]byte(nil), algo.Salt1...), bytes.Repeat([]byte{0x5C}, 32)...)
|
||
if err := svc.UpdatePasswordSettings(ctx, userID, domain.PasswordCheck{Empty: true}, domain.PasswordInputSettings{
|
||
NewAlgo: &algo,
|
||
NewPasswordHash: verifierForPassword(algo, []byte("old password")),
|
||
Email: "bob@example.com",
|
||
HasEmail: true,
|
||
}); err != nil {
|
||
t.Fatalf("UpdatePasswordSettings set password: %v", err)
|
||
}
|
||
|
||
pattern, err := svc.RequestPasswordRecovery(ctx, userID)
|
||
if err != nil {
|
||
t.Fatalf("RequestPasswordRecovery: %v", err)
|
||
}
|
||
if pattern != "b***b@example.com" {
|
||
t.Fatalf("recovery pattern = %q, want masked email", pattern)
|
||
}
|
||
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)
|
||
if err != nil {
|
||
t.Fatalf("GetPassword cleared: %v", err)
|
||
}
|
||
if cleared.HasPassword || cleared.HasRecovery {
|
||
t.Fatalf("cleared settings = %+v, want no password/recovery", cleared)
|
||
}
|
||
}
|
||
|
||
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
|
||
passwords := memory.NewPasswordStore()
|
||
svc := NewService(passwords)
|
||
|
||
if err := passwords.Save(ctx, userID, domain.PasswordSettings{HasPassword: true}); err != nil {
|
||
t.Fatalf("save password settings: %v", err)
|
||
}
|
||
result, err := svc.ResetPassword(ctx, userID)
|
||
if err != nil {
|
||
t.Fatalf("ResetPassword request: %v", err)
|
||
}
|
||
if result.Kind != domain.PasswordResetRequestedWait || result.UntilDate <= int(time.Now().Unix()) {
|
||
t.Fatalf("reset result = %+v, want requested future wait", result)
|
||
}
|
||
pending, _, err := passwords.GetByUser(ctx, userID)
|
||
if err != nil || pending.PendingResetDate != result.UntilDate {
|
||
t.Fatalf("pending reset = %+v found err=%v, want until date", pending, err)
|
||
}
|
||
|
||
if err := svc.DeclinePasswordReset(ctx, userID); err != nil {
|
||
t.Fatalf("DeclinePasswordReset: %v", err)
|
||
}
|
||
declined, _, err := passwords.GetByUser(ctx, userID)
|
||
if err != nil || declined.PendingResetDate != 0 {
|
||
t.Fatalf("declined settings = %+v err=%v, want no pending reset", declined, err)
|
||
}
|
||
|
||
declined.PendingResetDate = int(time.Now().Add(-time.Second).Unix())
|
||
if err := passwords.Save(ctx, userID, declined); err != nil {
|
||
t.Fatalf("save expired reset: %v", err)
|
||
}
|
||
result, err = svc.ResetPassword(ctx, userID)
|
||
if err != nil {
|
||
t.Fatalf("ResetPassword finalize: %v", err)
|
||
}
|
||
if result.Kind != domain.PasswordResetOK {
|
||
t.Fatalf("final reset result = %+v, want ok", result)
|
||
}
|
||
cleared, _, err := passwords.GetByUser(ctx, userID)
|
||
if err != nil || cleared.HasPassword || cleared.PendingResetDate != 0 {
|
||
t.Fatalf("cleared settings = %+v err=%v, want password cleared", cleared, err)
|
||
}
|
||
}
|
||
|
||
func clientPasswordCheck(t *testing.T, settings domain.PasswordSettings, password []byte) domain.PasswordCheck {
|
||
t.Helper()
|
||
algo := settings.NewAlgo
|
||
if settings.CurrentAlgo != nil {
|
||
algo = *settings.CurrentAlgo
|
||
}
|
||
p := new(big.Int).SetBytes(algo.P)
|
||
g := big.NewInt(int64(algo.G))
|
||
a := new(big.Int).SetBytes(bytes.Repeat([]byte{0x23}, passwordHashSize))
|
||
A := new(big.Int).Exp(g, a, p)
|
||
aForHash := padToHash(A.Bytes())
|
||
bForHash := padToHash(settings.SRPB)
|
||
x := new(big.Int).SetBytes(passwordDigest(algo, password))
|
||
u := new(big.Int).SetBytes(hashBytes(aForHash, bForHash))
|
||
k := new(big.Int).SetBytes(hashBytes(padToHash(algo.P), padToHash(g.Bytes())))
|
||
gx := new(big.Int).Exp(g, x, p)
|
||
kgx := new(big.Int).Mul(k, gx)
|
||
kgx.Mod(kgx, p)
|
||
b := new(big.Int).SetBytes(settings.SRPB)
|
||
base := new(big.Int).Sub(b, kgx)
|
||
base.Mod(base, p)
|
||
exp := new(big.Int).Mul(u, x)
|
||
exp.Add(exp, a)
|
||
s := new(big.Int).Exp(base, exp, p)
|
||
kBytes := hashBytes(padToHash(s.Bytes()))
|
||
m1 := hashBytes(
|
||
xorBytes(hashBytes(padToHash(algo.P)), hashBytes(padToHash(g.Bytes()))),
|
||
hashBytes(algo.Salt1),
|
||
hashBytes(algo.Salt2),
|
||
aForHash,
|
||
bForHash,
|
||
kBytes,
|
||
)
|
||
return domain.PasswordCheck{SRPID: settings.SRPID, A: aForHash, M1: m1}
|
||
}
|
||
|
||
// passwordDigest 与 verifierForPassword 是客户端侧(明文口令 → verifier)的模拟助手,
|
||
// 服务端从不执行明文口令路径,仅供这里的客户端 SRP helper 构造测试输入。
|
||
func passwordDigest(algo domain.PasswordKDFAlgo, password []byte) []byte {
|
||
hash1 := hashBytes(algo.Salt1, password, algo.Salt1)
|
||
hash2 := hashBytes(algo.Salt2, hash1, algo.Salt2)
|
||
hash3 := pbkdf2.Key(hash2, algo.Salt1, 100000, 64, sha512.New)
|
||
return hashBytes(algo.Salt2, hash3, algo.Salt2)
|
||
}
|
||
|
||
func verifierForPassword(algo domain.PasswordKDFAlgo, password []byte) []byte {
|
||
p := new(big.Int).SetBytes(algo.P)
|
||
g := big.NewInt(int64(algo.G))
|
||
x := new(big.Int).SetBytes(passwordDigest(algo, password))
|
||
return padToHash(new(big.Int).Exp(g, x, p).Bytes())
|
||
}
|