auth: complete account authorization flows
(cherry picked from commit 04f4527df32ad5c35720cccc41d27fe51549612f)
This commit is contained in:
parent
af41d18478
commit
6dc42942c8
21 changed files with 1780 additions and 50 deletions
|
|
@ -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.
|
||||
|
|
|
|||
185
internal/app/account/service_test.go
Normal file
185
internal/app/account/service_test.go
Normal 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
204
internal/app/account/srp.go
Normal 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
|
||||
}
|
||||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue