owpengram-server/internal/app/account/service_test.go
Astra 6435406690 account: stop rotating the SRP challenge on every getPassword read
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>
2026-09-14 12:06:46 +01:00

411 lines
15 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

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())
}