auth: complete account authorization flows

(cherry picked from commit 04f4527df32ad5c35720cccc41d27fe51549612f)
This commit is contained in:
A 2026-06-08 21:42:32 +08:00
parent af41d18478
commit 6dc42942c8
21 changed files with 1780 additions and 50 deletions

View file

@ -2,6 +2,11 @@ package account
import (
"context"
"crypto/rand"
"crypto/subtle"
"fmt"
"strings"
"time"
"telesrv/internal/domain"
"telesrv/internal/store"
@ -9,6 +14,23 @@ import (
var defaultSecureRandom = []byte("telesrv-tdesktop-dev-secure-rand")
const (
passwordResetWait = 7 * 24 * time.Hour
passwordResetRetry = 24 * time.Hour
)
// EmailUnconfirmedError reports the dev recovery-code length expected by TDesktop.
type EmailUnconfirmedError struct {
Length int
}
func (e EmailUnconfirmedError) Error() string {
if e.Length <= 0 {
return "email unconfirmed"
}
return fmt.Sprintf("email unconfirmed: %d", e.Length)
}
// Service 提供账号安全配置查询。
type Service struct {
passwords store.PasswordStore
@ -46,14 +68,325 @@ func (s *Service) GetPassword(ctx context.Context, userID int64) (domain.Passwor
if !found {
return defaultPasswordSettings(), nil
}
if len(settings.SecureRandom) == 0 {
settings.SecureRandom = append([]byte(nil), defaultSecureRandom...)
settings = normalizePasswordSettings(settings)
if settings.HasPassword {
secret, b, err := makeSRPChallenge(settings.SRPVerifier)
if err != nil {
return domain.PasswordSettings{}, err
}
settings.SRPBSecret = secret
settings.SRPB = b
if settings.SRPID == 0 {
settings.SRPID, err = randomInt64()
if err != nil {
return domain.PasswordSettings{}, err
}
}
if err := s.passwords.Save(ctx, userID, settings); err != nil {
return domain.PasswordSettings{}, err
}
}
return settings, nil
}
func defaultPasswordSettings() domain.PasswordSettings {
return domain.PasswordSettings{SecureRandom: append([]byte(nil), defaultSecureRandom...)}
return normalizePasswordSettings(domain.PasswordSettings{SecureRandom: append([]byte(nil), defaultSecureRandom...)})
}
func normalizePasswordSettings(settings domain.PasswordSettings) domain.PasswordSettings {
if len(settings.SecureRandom) == 0 {
settings.SecureRandom = append([]byte(nil), defaultSecureRandom...)
}
if len(settings.NewAlgo.P) == 0 {
settings.NewAlgo = defaultPasswordAlgo()
}
if settings.NewSecureAlgo.Kind == "" {
settings.NewSecureAlgo = defaultSecureAlgo()
}
if settings.HasPassword && settings.CurrentAlgo == nil {
algo := settings.NewAlgo
settings.CurrentAlgo = &algo
}
if settings.RecoveryEmail != "" {
settings.HasRecovery = true
}
return settings
}
// CheckPassword validates the current account password check.
func (s *Service) CheckPassword(ctx context.Context, userID int64, check domain.PasswordCheck) error {
settings, err := s.GetPasswordWithoutRefresh(ctx, userID)
if err != nil {
return err
}
return checkSRP(settings, check)
}
// GetPasswordWithoutRefresh returns persisted settings without rotating the SRP challenge.
func (s *Service) GetPasswordWithoutRefresh(ctx context.Context, userID int64) (domain.PasswordSettings, error) {
if s == nil || s.passwords == nil || userID == 0 {
return defaultPasswordSettings(), nil
}
settings, found, err := s.passwords.GetByUser(ctx, userID)
if err != nil {
return domain.PasswordSettings{}, err
}
if !found {
return defaultPasswordSettings(), nil
}
return normalizePasswordSettings(settings), nil
}
// GetPasswordSettings validates the password and returns private 2FA settings.
func (s *Service) GetPasswordSettings(ctx context.Context, userID int64, check domain.PasswordCheck) (domain.PrivatePasswordSettings, error) {
settings, err := s.GetPasswordWithoutRefresh(ctx, userID)
if err != nil {
return domain.PrivatePasswordSettings{}, err
}
if err := checkSRP(settings, check); err != nil {
return domain.PrivatePasswordSettings{}, err
}
return domain.PrivatePasswordSettings{Email: settings.RecoveryEmail}, nil
}
// UpdatePasswordSettings sets, changes, clears, or updates the recovery email for 2FA.
func (s *Service) UpdatePasswordSettings(ctx context.Context, userID int64, check domain.PasswordCheck, input domain.PasswordInputSettings) error {
if s == nil || s.passwords == nil || userID == 0 {
return nil
}
settings, err := s.GetPasswordWithoutRefresh(ctx, userID)
if err != nil {
return err
}
if err := checkSRP(settings, check); err != nil {
return err
}
if len(input.NewPasswordHash) == 0 && !input.HasEmail {
settings = defaultPasswordSettings()
settings.SecureRandom = randomBytesOrDefault(passwordHashSize, settings.SecureRandom)
return s.passwords.Save(ctx, userID, settings)
}
if len(input.NewPasswordHash) > 0 {
if err := validateNewPasswordSettings(input); err != nil {
return err
}
srpID, err := randomInt64()
if err != nil {
return err
}
algo := *input.NewAlgo
settings.CurrentAlgo = &algo
settings.NewAlgo = defaultPasswordAlgo()
settings.SRPVerifier = padToHash(input.NewPasswordHash)
secret, b, err := makeSRPChallenge(settings.SRPVerifier)
if err != nil {
return err
}
settings.SRPBSecret = secret
settings.SRPB = b
settings.SRPID = srpID
settings.HasPassword = true
if input.HasHint {
settings.Hint = input.Hint
}
}
if input.HasEmail {
email := strings.TrimSpace(input.Email)
if email != "" && !strings.Contains(email, "@") {
return domain.ErrEmailInvalid
}
settings.RecoveryEmail = email
settings.HasRecovery = email != ""
settings.LoginEmailPattern = emailPattern(email)
settings.EmailUnconfirmedPattern = ""
}
settings.SecureRandom = randomBytesOrDefault(passwordHashSize, settings.SecureRandom)
return s.passwords.Save(ctx, userID, normalizePasswordSettings(settings))
}
func (s *Service) RequestPasswordRecovery(ctx context.Context, userID int64) (string, error) {
settings, err := s.GetPasswordWithoutRefresh(ctx, userID)
if err != nil {
return "", err
}
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
}
}
return emailPattern(settings.RecoveryEmail), nil
}
func (s *Service) CheckRecoveryPassword(ctx context.Context, userID int64, code string) error {
settings, err := s.GetPasswordWithoutRefresh(ctx, userID)
if err != nil {
return err
}
return checkRecoveryCode(settings, code)
}
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
}
settings, err := s.GetPasswordWithoutRefresh(ctx, userID)
if err != nil {
return err
}
if err := checkRecoveryCode(settings, code); err != nil {
return err
}
if input == nil || len(input.NewPasswordHash) == 0 {
settings = defaultPasswordSettings()
return s.passwords.Save(ctx, userID, settings)
}
if err := validateNewPasswordSettings(*input); err != nil {
return err
}
settings.CurrentAlgo = input.NewAlgo
settings.SRPVerifier = padToHash(input.NewPasswordHash)
settings.SRPID, err = randomInt64()
if err != nil {
return err
}
settings.SRPBSecret, settings.SRPB, err = makeSRPChallenge(settings.SRPVerifier)
if err != nil {
return err
}
settings.HasPassword = true
if input.HasHint {
settings.Hint = input.Hint
}
settings.RecoveryCode = ""
settings.RecoveryCodeExpiresAt = 0
return s.passwords.Save(ctx, userID, normalizePasswordSettings(settings))
}
func (s *Service) ResetPassword(ctx context.Context, userID int64) (domain.PasswordResetResult, error) {
if s == nil || s.passwords == nil || userID == 0 {
return domain.PasswordResetResult{Kind: domain.PasswordResetFailedWait, RetryDate: int(time.Now().Add(passwordResetRetry).Unix())}, nil
}
settings, err := s.GetPasswordWithoutRefresh(ctx, userID)
if err != nil {
return domain.PasswordResetResult{}, err
}
if !settings.HasPassword {
return domain.PasswordResetResult{Kind: domain.PasswordResetOK}, nil
}
if settings.HasRecovery {
return domain.PasswordResetResult{}, domain.ErrPasswordRecoveryNA
}
now := time.Now()
if settings.PendingResetDate > 0 {
if now.Unix() >= int64(settings.PendingResetDate) {
next := defaultPasswordSettings()
next.SecureRandom = randomBytesOrDefault(passwordHashSize, settings.SecureRandom)
if err := s.passwords.Save(ctx, userID, next); err != nil {
return domain.PasswordResetResult{}, err
}
return domain.PasswordResetResult{Kind: domain.PasswordResetOK}, nil
}
return domain.PasswordResetResult{Kind: domain.PasswordResetRequestedWait, UntilDate: settings.PendingResetDate}, nil
}
settings.PendingResetDate = int(now.Add(passwordResetWait).Unix())
if err := s.passwords.Save(ctx, userID, normalizePasswordSettings(settings)); err != nil {
return domain.PasswordResetResult{}, err
}
return domain.PasswordResetResult{Kind: domain.PasswordResetRequestedWait, UntilDate: settings.PendingResetDate}, nil
}
func (s *Service) DeclinePasswordReset(ctx context.Context, userID int64) error {
if s == nil || s.passwords == nil || userID == 0 {
return nil
}
settings, err := s.GetPasswordWithoutRefresh(ctx, userID)
if err != nil {
return err
}
settings.PendingResetDate = 0
return s.passwords.Save(ctx, userID, normalizePasswordSettings(settings))
}
func (s *Service) ConfirmPasswordEmail(ctx context.Context, userID int64, code string) error {
return s.CheckRecoveryPassword(ctx, userID, code)
}
func (s *Service) ResendPasswordEmail(ctx context.Context, userID int64) error {
_, err := s.RequestPasswordRecovery(ctx, userID)
return err
}
func (s *Service) CancelPasswordEmail(ctx context.Context, userID int64) error {
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 {
return nil
}
return domain.ErrPasswordRecoveryNA
}
if settings.RecoveryCodeExpiresAt > 0 && time.Now().Unix() > settings.RecoveryCodeExpiresAt {
return domain.ErrEmailCodeInvalid
}
if subtle.ConstantTimeCompare([]byte(settings.RecoveryCode), []byte(code)) != 1 {
return domain.ErrEmailCodeInvalid
}
return nil
}
func randomBytesOrDefault(n int, fallback []byte) []byte {
out := make([]byte, n)
if _, err := rand.Read(out); err != nil {
return append([]byte(nil), fallback...)
}
return out
}
func randomInt64() (int64, error) {
var b [8]byte
if _, err := rand.Read(b[:]); err != nil {
return 0, err
}
out := int64(0)
for _, v := range b {
out = (out << 8) | int64(v)
}
if out == 0 {
out = 1
}
if out < 0 {
out = -out
}
return out, nil
}
func emailPattern(email string) string {
if email == "" {
return ""
}
at := strings.Index(email, "@")
if at <= 1 {
return email
}
name := email[:at]
return name[:1] + "***" + name[len(name)-1:] + email[at:]
}
// GetReactionSettings returns account-level reaction preferences.

View file

@ -0,0 +1,185 @@
package account
import (
"bytes"
"context"
"errors"
"math/big"
"testing"
"time"
"telesrv/internal/domain"
"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)
}
}
func TestRecoverPasswordClearsTwoFactorPassword(t *testing.T) {
ctx := context.Background()
const userID int64 = 1002
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{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 err := svc.RecoverPassword(ctx, userID, recoveryCode, 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 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}
}

204
internal/app/account/srp.go Normal file
View file

@ -0,0 +1,204 @@
package account
import (
"bytes"
"crypto/rand"
"crypto/sha256"
"crypto/sha512"
"encoding/hex"
"math/big"
"golang.org/x/crypto/pbkdf2"
"telesrv/internal/domain"
)
const (
passwordHashSize = 256
recoveryCode = "12345"
recoveryCodeTTL = 15 * 60
)
var (
baseSalt1 = []byte{0xEC, 0xF8, 0x73, 0x76, 0x65, 0xBC, 0x77, 0x5A}
baseSalt2 = []byte{0xBE, 0xDE, 0x48, 0x88, 0x8C, 0x0F, 0x42, 0xAC, 0x34, 0xFF, 0xD1, 0xD4, 0x93, 0x5D, 0x8B, 0x21}
baseP = mustDecodeHex("c71caeb9c6b1c9048e6c522f70f13f73980d40238e3e21c14934d037563d930f48198a0aa7c14058229493d22530f4dbfa336f6e0ac925139543aed44cce7c3720fd51f69458705ac68cd4fe6b6b13abdc9746512969328454f18faf8c595f642477fe96bb2a941d5bcd1d4ac8cc49880708fa9b378e3c4f3a9060bee67cf9a4a4a695811051907e162753b56b0f6b410dba74d8a84b2a14b3144e0ef1284754fd17ed950d5965b4b9dd46582db1178d169c6bc465b0d6ff9ca3928fef5b9ae4e418fc15e83ebea0f87fa9ff5eed70050ded2849f47bf959d956850ce929851f0d8115f635b105ee2e4e15d04b2454bf6f4fadf034b10403119cd8e3b92fcc5b")
baseG = 3
)
func defaultPasswordAlgo() domain.PasswordKDFAlgo {
return domain.PasswordKDFAlgo{
Salt1: append([]byte(nil), baseSalt1...),
Salt2: append([]byte(nil), baseSalt2...),
G: baseG,
P: append([]byte(nil), baseP...),
}
}
func defaultSecureAlgo() domain.SecurePasswordKDFAlgo {
return domain.SecurePasswordKDFAlgo{
Kind: "pbkdf2_hmac_sha512_iter100000",
Salt: []byte{0x7D, 0x04, 0xB3, 0x4B, 0x94, 0x82, 0x8C, 0x3D},
}
}
func makeSRPChallenge(verifier []byte) (secret, b []byte, err error) {
secret = make([]byte, passwordHashSize)
if _, err := rand.Read(secret); err != nil {
return nil, nil, err
}
b, err = calcSRPB(secret, verifier)
if err != nil {
return nil, nil, err
}
return secret, b, nil
}
func calcSRPB(secret, verifier []byte) ([]byte, error) {
p := new(big.Int).SetBytes(baseP)
g := big.NewInt(int64(baseG))
v := new(big.Int).SetBytes(verifier)
b := new(big.Int).SetBytes(secret)
if v.Sign() <= 0 || v.Cmp(p) >= 0 {
return nil, domain.ErrPasswordHashInvalid
}
k := new(big.Int).SetBytes(hashBytes(padToHash(baseP), padToHash(g.Bytes())))
kv := new(big.Int).Mul(k, v)
kv.Mod(kv, p)
gb := new(big.Int).Exp(g, b, p)
out := new(big.Int).Add(kv, gb)
out.Mod(out, p)
return padToHash(out.Bytes()), nil
}
func checkSRP(settings domain.PasswordSettings, check domain.PasswordCheck) error {
if check.Empty {
if settings.HasPassword {
return domain.ErrPasswordHashInvalid
}
return nil
}
if !settings.HasPassword {
return domain.ErrPasswordHashInvalid
}
if settings.SRPID == 0 || settings.SRPID != check.SRPID {
return domain.ErrSRPIDInvalid
}
if len(settings.SRPVerifier) == 0 || len(settings.SRPBSecret) == 0 || len(settings.SRPB) == 0 {
return domain.ErrSRPPasswordChanged
}
got, err := calcSRPM1(settings, check.A)
if err != nil {
return err
}
if !bytes.Equal(got, check.M1) {
return domain.ErrPasswordHashInvalid
}
return nil
}
func calcSRPM1(settings domain.PasswordSettings, aBytes []byte) ([]byte, error) {
p := new(big.Int).SetBytes(baseP)
g := big.NewInt(int64(baseG))
a := new(big.Int).SetBytes(aBytes)
if !isGoodLarge(a, p) {
return nil, domain.ErrPasswordHashInvalid
}
v := new(big.Int).SetBytes(settings.SRPVerifier)
if !isGoodLarge(v, p) {
return nil, domain.ErrPasswordHashInvalid
}
b := new(big.Int).SetBytes(settings.SRPBSecret)
bForHash := padToHash(settings.SRPB)
aForHash := padToHash(aBytes)
u := new(big.Int).SetBytes(hashBytes(aForHash, bForHash))
if u.Sign() <= 0 {
return nil, domain.ErrPasswordHashInvalid
}
vu := new(big.Int).Exp(v, u, p)
sBase := new(big.Int).Mul(a, vu)
sBase.Mod(sBase, p)
s := new(big.Int).Exp(sBase, b, p)
k := hashBytes(padToHash(s.Bytes()))
salt1 := settings.NewAlgo.Salt1
if settings.CurrentAlgo != nil {
salt1 = settings.CurrentAlgo.Salt1
}
return hashBytes(
xorBytes(hashBytes(padToHash(baseP)), hashBytes(padToHash(g.Bytes()))),
hashBytes(salt1),
hashBytes(baseSalt2),
aForHash,
bForHash,
k,
), nil
}
func validateNewPasswordSettings(in domain.PasswordInputSettings) error {
if in.NewAlgo == nil || len(in.NewPasswordHash) == 0 {
return domain.ErrNewSettingsInvalid
}
algo := in.NewAlgo
if algo.G != baseG || !bytes.Equal(algo.P, baseP) || !bytes.Equal(algo.Salt2, baseSalt2) {
return domain.ErrNewSaltInvalid
}
if len(algo.Salt1) != len(baseSalt1)+32 || !bytes.Equal(algo.Salt1[:len(baseSalt1)], baseSalt1) {
return domain.ErrNewSaltInvalid
}
v := new(big.Int).SetBytes(in.NewPasswordHash)
if !isGoodLarge(v, new(big.Int).SetBytes(baseP)) {
return domain.ErrPasswordHashInvalid
}
return nil
}
func hashBytes(parts ...[]byte) []byte {
h := sha256.New()
for _, part := range parts {
_, _ = h.Write(part)
}
return h.Sum(nil)
}
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())
}
func padToHash(in []byte) []byte {
if len(in) >= passwordHashSize {
return append([]byte(nil), in[len(in)-passwordHashSize:]...)
}
out := make([]byte, passwordHashSize)
copy(out[passwordHashSize-len(in):], in)
return out
}
func xorBytes(a, b []byte) []byte {
out := make([]byte, len(a))
for i := range a {
out[i] = a[i] ^ b[i]
}
return out
}
func isGoodLarge(n, p *big.Int) bool {
return n.Sign() > 0 && new(big.Int).Sub(p, n).Sign() > 0
}
func mustDecodeHex(s string) []byte {
out, err := hex.DecodeString(s)
if err != nil {
panic(err)
}
return out
}

View file

@ -34,6 +34,7 @@ type Service struct {
codes store.CodeStore
authKeys store.AuthKeyStore
tempKeys store.TempAuthKeyBindingStore
passwords store.PasswordStore
messages store.MessageStore
dialogs store.DialogStore
fixedCode string
@ -51,6 +52,13 @@ func WithLoginMessages(messages store.MessageStore, dialogs store.DialogStore) O
}
}
// WithPasswords lets sign-in stop at SESSION_PASSWORD_NEEDED for 2FA accounts.
func WithPasswords(passwords store.PasswordStore) Option {
return func(s *Service) {
s.passwords = passwords
}
}
// NewService 创建登录服务。fixedCode 为开发固定验证码。
func NewService(users store.UserStore, auths store.AuthorizationStore, codes store.CodeStore, authKeys store.AuthKeyStore, tempKeys store.TempAuthKeyBindingStore, fixedCode string, opts ...Option) *Service {
s := &Service{users: users, auths: auths, codes: codes, authKeys: authKeys, tempKeys: tempKeys, fixedCode: fixedCode, codeTTL: 5 * time.Minute}
@ -114,6 +122,39 @@ func (s *Service) SendCode(ctx context.Context, phone string) (string, error) {
return hash, nil
}
// ResendCode invalidates an existing code hash and sends a fresh code to the same phone.
func (s *Service) ResendCode(ctx context.Context, phone, phoneCodeHash string) (string, error) {
phone = normalizePhone(phone)
rec, found, err := s.codes.Get(ctx, phoneCodeHash)
if err != nil {
return "", err
}
if !found {
return "", ErrCodeExpired
}
if rec.Phone != phone {
return "", ErrCodeInvalid
}
_ = s.codes.Del(ctx, phoneCodeHash)
return s.SendCode(ctx, phone)
}
// CancelCode invalidates a pending login code hash.
func (s *Service) CancelCode(ctx context.Context, phone, phoneCodeHash string) error {
phone = normalizePhone(phone)
rec, found, err := s.codes.Get(ctx, phoneCodeHash)
if err != nil {
return err
}
if !found {
return ErrCodeExpired
}
if rec.Phone != phone {
return ErrCodeInvalid
}
return s.codes.Del(ctx, phoneCodeHash)
}
// SignIn 校验验证码并尝试登录。
// needSignUp=true 表示验证码正确但用户不存在,调用方应引导注册(此时不删验证码,留给 SignUp
func (s *Service) SignIn(ctx context.Context, auth domain.Authorization, phone, phoneCodeHash, code string) (u domain.User, loginMessage domain.Message, needSignUp bool, err error) {
@ -139,6 +180,10 @@ func (s *Service) SignIn(ctx context.Context, auth domain.Authorization, phone,
if err := s.bind(ctx, auth, existing.ID); err != nil {
return domain.User{}, domain.Message{}, false, err
}
if s.passwordNeeded(ctx, existing.ID) {
_ = s.codes.Del(ctx, phoneCodeHash)
return existing, domain.Message{}, false, domain.ErrSessionPasswordNeeded
}
loginMessage, err = s.recordLoginMessage(ctx, existing.ID, rec.Code)
if err != nil {
return domain.User{}, domain.Message{}, false, err
@ -196,11 +241,40 @@ func (s *Service) LogOut(ctx context.Context, authKeyID [8]byte) error {
return s.auths.Delete(ctx, authKeyID)
}
func (s *Service) ListAuthorizations(ctx context.Context, userID int64) ([]domain.Authorization, error) {
if s == nil || s.auths == nil || userID == 0 {
return nil, nil
}
return s.auths.ListByUser(ctx, userID)
}
func (s *Service) ResetAuthorization(ctx context.Context, userID, hash int64) (domain.Authorization, bool, error) {
if s == nil || s.auths == nil || userID == 0 {
return domain.Authorization{}, false, nil
}
return s.auths.DeleteByHash(ctx, userID, hash)
}
func (s *Service) ResetAuthorizations(ctx context.Context, userID int64, keepAuthKeyID [8]byte) ([]domain.Authorization, error) {
if s == nil || s.auths == nil || userID == 0 {
return nil, nil
}
return s.auths.DeleteByUserExcept(ctx, userID, keepAuthKeyID)
}
func (s *Service) bind(ctx context.Context, auth domain.Authorization, userID int64) error {
auth.UserID = userID
return s.auths.Bind(ctx, auth)
}
func (s *Service) passwordNeeded(ctx context.Context, userID int64) bool {
if s.passwords == nil {
return false
}
settings, found, err := s.passwords.GetByUser(ctx, userID)
return err == nil && found && settings.HasPassword
}
const loginMessageTpl = `Login code: %s. Do not give this code to anyone, even if they say they are from Telegram!
This code can be used to log in to your Telegram account. We never ask it for anything else.

View file

@ -241,6 +241,45 @@ func TestSignUpWritesOfficialLoginMessage(t *testing.T) {
}
}
func TestSignInExistingTwoFactorAccountNeedsPassword(t *testing.T) {
ctx := context.Background()
passwords := memory.NewPasswordStore()
svc := NewService(memory.NewUserStore(), memory.NewAuthorizationStore(), memory.NewCodeStore(), nil, nil, "12345", WithPasswords(passwords))
var key [8]byte
key[0] = 7
hash, err := svc.SendCode(ctx, "+15550004312")
if err != nil {
t.Fatalf("SendCode signup: %v", err)
}
u, _, err := svc.SignUp(ctx, domain.Authorization{AuthKeyID: key}, "+15550004312", hash, "Two", "Factor")
if err != nil {
t.Fatalf("SignUp: %v", err)
}
if err := svc.LogOut(ctx, key); err != nil {
t.Fatalf("LogOut: %v", err)
}
if err := passwords.Save(ctx, u.ID, domain.PasswordSettings{HasPassword: true}); err != nil {
t.Fatalf("save password settings: %v", err)
}
hash, err = svc.SendCode(ctx, "+15550004312")
if err != nil {
t.Fatalf("SendCode signin: %v", err)
}
got, _, needSignUp, err := svc.SignIn(ctx, domain.Authorization{AuthKeyID: key}, "+15550004312", hash, "12345")
if !errors.Is(err, domain.ErrSessionPasswordNeeded) {
t.Fatalf("SignIn err = %v, want ErrSessionPasswordNeeded", err)
}
if needSignUp || got.ID != u.ID {
t.Fatalf("SignIn user=%+v needSignUp=%v, want existing 2FA user", got, needSignUp)
}
bound, found, err := svc.UserID(ctx, key)
if err != nil || !found || bound != u.ID {
t.Fatalf("UserID after password-needed = %d found=%v err=%v, want %d", bound, found, err, u.ID)
}
}
func testAuthKey(seed byte) mtcrypto.AuthKey {
var raw mtcrypto.Key
for i := range raw {