feat: sync configurable OTP delivery providers
This commit is contained in:
parent
c18f773701
commit
6af61f26ba
28 changed files with 2100 additions and 118 deletions
|
|
@ -9,6 +9,7 @@ import (
|
|||
"time"
|
||||
|
||||
"telesrv/internal/domain"
|
||||
"telesrv/internal/otpdelivery"
|
||||
"telesrv/internal/store"
|
||||
"telesrv/internal/store/memory"
|
||||
)
|
||||
|
|
@ -30,8 +31,10 @@ func createUser(t *testing.T, users *memory.UserStore, phone string) domain.User
|
|||
}
|
||||
|
||||
type captureMailSender struct {
|
||||
to string
|
||||
code string
|
||||
to string
|
||||
code string
|
||||
requests []otpdelivery.Request
|
||||
err error
|
||||
}
|
||||
|
||||
type blockingCodeCAS struct {
|
||||
|
|
@ -116,10 +119,102 @@ func (s *blockingCodeCAS) CompareAndDelete(ctx context.Context, key, revision st
|
|||
return s.CodeStore.CompareAndDelete(ctx, key, revision)
|
||||
}
|
||||
|
||||
func (s *captureMailSender) SendLoginCode(_ context.Context, to, code string, _ time.Duration) error {
|
||||
s.to = to
|
||||
s.code = code
|
||||
return nil
|
||||
func (s *captureMailSender) Deliver(_ context.Context, req otpdelivery.Request) (otpdelivery.Result, error) {
|
||||
s.to = req.Recipient
|
||||
s.code = req.Code
|
||||
s.requests = append(s.requests, req)
|
||||
return otpdelivery.Result{}, s.err
|
||||
}
|
||||
|
||||
func TestLoginEmailDeliveryCarriesPurposeAndStableID(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
users := memory.NewUserStore()
|
||||
codes := memory.NewCodeStore()
|
||||
sender := &captureMailSender{}
|
||||
svc := NewService(memory.NewPasswordStore(),
|
||||
WithUsers(users),
|
||||
WithLoginEmailVerification(codes, sender, time.Minute, 3, 6))
|
||||
u := createUser(t, users, "15550010150")
|
||||
|
||||
pattern, length, err := svc.SendLoginEmailCode(ctx, u.ID, "", "", "Alice@Example.Test", false)
|
||||
if err != nil {
|
||||
t.Fatalf("SendLoginEmailCode: %v", err)
|
||||
}
|
||||
if pattern == "" || length != 6 || len(sender.requests) != 1 {
|
||||
t.Fatalf("pattern=%q length=%d requests=%d", pattern, length, len(sender.requests))
|
||||
}
|
||||
req := sender.requests[0]
|
||||
if req.DeliveryID == "" || req.Purpose != otpdelivery.PurposeLoginEmailChange || req.Channel != otpdelivery.ChannelEmail ||
|
||||
req.Recipient != "alice@example.test" || len(req.Code) != 6 {
|
||||
t.Fatalf("request = %+v", req)
|
||||
}
|
||||
snapshot, found, err := codes.GetSnapshot(ctx, loginEmailVerifyChangePrefix+fmt.Sprint(u.ID))
|
||||
if err != nil || !found || snapshot.Record.DeliveryID != req.DeliveryID || snapshot.Record.Code != req.Code {
|
||||
t.Fatalf("snapshot=%+v found=%v err=%v", snapshot, found, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoginEmailSetupDeliveryUsesSetupPurpose(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
codes := memory.NewCodeStore()
|
||||
sender := &captureMailSender{}
|
||||
phone := "15550010151"
|
||||
phoneHash := "setup-purpose-hash"
|
||||
if err := codes.Set(ctx, phoneHash, store.PhoneCode{
|
||||
Version: store.PhoneCodeVersionCurrent,
|
||||
Phone: phone,
|
||||
Channel: codeChannelEmailSetupRequired,
|
||||
}, time.Minute); err != nil {
|
||||
t.Fatalf("seed setup code: %v", err)
|
||||
}
|
||||
svc := NewService(memory.NewPasswordStore(),
|
||||
WithUsers(memory.NewUserStore()),
|
||||
WithLoginEmailVerification(codes, sender, time.Minute, 3, 6))
|
||||
|
||||
if _, _, err := svc.SendLoginEmailCode(ctx, 0, phone, phoneHash, "new@example.test", true); err != nil {
|
||||
t.Fatalf("SendLoginEmailCode setup: %v", err)
|
||||
}
|
||||
if len(sender.requests) != 1 || sender.requests[0].Purpose != otpdelivery.PurposeLoginEmailSetup {
|
||||
t.Fatalf("requests = %+v", sender.requests)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoginEmailExplicitRejectionDeletesOnlyCurrentAttempt(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
users := memory.NewUserStore()
|
||||
codes := memory.NewCodeStore()
|
||||
sender := &captureMailSender{err: &otpdelivery.RejectedError{StatusCode: 400, Code: "RECIPIENT_INVALID"}}
|
||||
svc := NewService(memory.NewPasswordStore(),
|
||||
WithUsers(users),
|
||||
WithLoginEmailVerification(codes, sender, time.Minute, 3, 6))
|
||||
u := createUser(t, users, "15550010152")
|
||||
key := loginEmailVerifyChangePrefix + fmt.Sprint(u.ID)
|
||||
|
||||
if _, _, err := svc.SendLoginEmailCode(ctx, u.ID, "", "", "bad@example.test", false); err == nil {
|
||||
t.Fatal("explicit rejection succeeded")
|
||||
}
|
||||
if _, found, err := codes.Get(ctx, key); err != nil || found {
|
||||
t.Fatalf("rejected code found=%v err=%v", found, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoginEmailUnknownOutcomeReturnsSuccessAndKeepsCode(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
users := memory.NewUserStore()
|
||||
codes := memory.NewCodeStore()
|
||||
sender := &captureMailSender{err: &otpdelivery.OutcomeUnknownError{Cause: errors.New("ack lost")}}
|
||||
svc := NewService(memory.NewPasswordStore(),
|
||||
WithUsers(users),
|
||||
WithLoginEmailVerification(codes, sender, time.Minute, 3, 6))
|
||||
u := createUser(t, users, "15550010153")
|
||||
|
||||
if _, _, err := svc.SendLoginEmailCode(ctx, u.ID, "", "", "unknown@example.test", false); err != nil {
|
||||
t.Fatalf("unknown outcome: %v", err)
|
||||
}
|
||||
key := loginEmailVerifyChangePrefix + fmt.Sprint(u.ID)
|
||||
if rec, found, err := codes.Get(ctx, key); err != nil || !found || rec.Code != sender.code {
|
||||
t.Fatalf("unknown code=%+v found=%v err=%v", rec, found, err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestSetLoginEmailPersistsAndMasks 设置登录邮箱后,GetPassword 下发掩码 pattern,原始
|
||||
|
|
|
|||
|
|
@ -4,11 +4,13 @@ import (
|
|||
"context"
|
||||
"crypto/rand"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"telesrv/internal/domain"
|
||||
"telesrv/internal/otpdelivery"
|
||||
"telesrv/internal/store"
|
||||
)
|
||||
|
||||
|
|
@ -42,27 +44,60 @@ func (s *Service) SendChangePhoneCode(ctx context.Context, userID int64, authKey
|
|||
} else if found && existing.ID != 0 {
|
||||
return "", domain.AuthCodeDelivery{}, domain.ErrPhoneNumberOccupied
|
||||
}
|
||||
if s.codes == nil || strings.TrimSpace(s.phoneChangeCode) == "" {
|
||||
if s.codes == nil || (s.phoneCodeSender == nil && strings.TrimSpace(s.phoneChangeCode) == "") {
|
||||
return "", domain.AuthCodeDelivery{}, fmt.Errorf("phone change code service is not configured")
|
||||
}
|
||||
hash, err := phoneChangeHash()
|
||||
if err != nil {
|
||||
return "", domain.AuthCodeDelivery{}, err
|
||||
}
|
||||
code := s.phoneChangeCode
|
||||
channel := store.PhoneCodeChannelPhone
|
||||
deliveryID := ""
|
||||
if s.phoneCodeSender != nil {
|
||||
code, err = randomDigits(s.phoneCodeLength)
|
||||
if err != nil {
|
||||
return "", domain.AuthCodeDelivery{}, err
|
||||
}
|
||||
deliveryID, err = otpdelivery.NewDeliveryID()
|
||||
if err != nil {
|
||||
return "", domain.AuthCodeDelivery{}, err
|
||||
}
|
||||
channel = store.PhoneCodeChannelSMS
|
||||
}
|
||||
rec := store.PhoneCode{
|
||||
Version: store.PhoneCodeVersionCurrent,
|
||||
Phone: phone,
|
||||
Code: s.phoneChangeCode,
|
||||
Channel: store.PhoneCodeChannelPhone,
|
||||
Code: code,
|
||||
DeliveryID: deliveryID,
|
||||
Channel: channel,
|
||||
Purpose: store.PhoneCodePurposeChangePhone,
|
||||
UserID: userID,
|
||||
AuthKeyID: authKeyID,
|
||||
SessionID: sessionID,
|
||||
MaxAttempts: s.phoneChangeMaxAttempts,
|
||||
}
|
||||
expiresAt := time.Now().Add(s.phoneChangeCodeTTL)
|
||||
if err := s.codes.Set(ctx, hash, rec, s.phoneChangeCodeTTL); err != nil {
|
||||
return "", domain.AuthCodeDelivery{}, fmt.Errorf("store phone change code: %w", err)
|
||||
}
|
||||
if s.phoneCodeSender != nil {
|
||||
if err := deliverOTP(ctx, s.phoneCodeSender, otpdelivery.Request{
|
||||
DeliveryID: deliveryID,
|
||||
Purpose: otpdelivery.PurposeChangePhone,
|
||||
Channel: otpdelivery.ChannelSMS,
|
||||
Recipient: phone,
|
||||
Code: code,
|
||||
ExpiresAt: expiresAt,
|
||||
}); err != nil {
|
||||
cleanupCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), 2*time.Second)
|
||||
defer cancel()
|
||||
if _, _, cleanupErr := s.codes.ConsumeScoped(cleanupCtx, hash, rec.Scope()); cleanupErr != nil {
|
||||
return "", domain.AuthCodeDelivery{}, errors.Join(err, fmt.Errorf("rollback phone change code: %w", cleanupErr))
|
||||
}
|
||||
return "", domain.AuthCodeDelivery{}, err
|
||||
}
|
||||
}
|
||||
return hash, domain.AuthCodeDelivery{Kind: domain.AuthCodeDeliverySMS, Length: len(rec.Code)}, nil
|
||||
}
|
||||
|
||||
|
|
@ -107,7 +142,8 @@ func (s *Service) ChangePhone(ctx context.Context, userID int64, authKeyID, orig
|
|||
return domain.PhoneChangeResult{}, domain.ErrPhoneNumberOccupied
|
||||
}
|
||||
consumed := verified.Record
|
||||
if consumed.Version != store.PhoneCodeVersionCurrent || consumed.Scope() != scope || consumed.Channel != store.PhoneCodeChannelPhone {
|
||||
if consumed.Version != store.PhoneCodeVersionCurrent || consumed.Scope() != scope ||
|
||||
(consumed.Channel != store.PhoneCodeChannelPhone && consumed.Channel != store.PhoneCodeChannelSMS) {
|
||||
return domain.PhoneChangeResult{}, domain.ErrPhoneCodeInvalid
|
||||
}
|
||||
if date == 0 {
|
||||
|
|
|
|||
|
|
@ -8,6 +8,7 @@ import (
|
|||
"time"
|
||||
|
||||
"telesrv/internal/domain"
|
||||
"telesrv/internal/otpdelivery"
|
||||
"telesrv/internal/store"
|
||||
"telesrv/internal/store/memory"
|
||||
)
|
||||
|
|
@ -30,6 +31,16 @@ type recordingPhoneChangeStore struct {
|
|||
last domain.PhoneChangeRequest
|
||||
}
|
||||
|
||||
type trackingPhoneCodeStore struct {
|
||||
store.CodeStore
|
||||
lastHash string
|
||||
}
|
||||
|
||||
func (s *trackingPhoneCodeStore) Set(ctx context.Context, hash string, code store.PhoneCode, ttl time.Duration) error {
|
||||
s.lastHash = hash
|
||||
return s.CodeStore.Set(ctx, hash, code, ttl)
|
||||
}
|
||||
|
||||
func (s *recordingPhoneChangeStore) ChangePhone(ctx context.Context, req domain.PhoneChangeRequest) (domain.PhoneChangeResult, error) {
|
||||
s.mu.Lock()
|
||||
s.last = req
|
||||
|
|
@ -67,6 +78,53 @@ func newPhoneChangeFixture(t *testing.T) phoneChangeFixture {
|
|||
return phoneChangeFixture{ctx: ctx, service: service, users: users, auths: auths, codes: codes, events: events, user: u, authKeyID: authKeyID, changes: changes}
|
||||
}
|
||||
|
||||
func TestPhoneChangeWebhookDeliversRandomScopedCode(t *testing.T) {
|
||||
f := newPhoneChangeFixture(t)
|
||||
sender := &captureMailSender{}
|
||||
f.service.phoneCodeSender = sender
|
||||
f.service.phoneCodeLength = 6
|
||||
f.service.phoneChangeCode = ""
|
||||
|
||||
hash, delivery, err := f.service.SendChangePhoneCode(f.ctx, f.user.ID, f.authKeyID, 77, "15550012020")
|
||||
if err != nil {
|
||||
t.Fatalf("SendChangePhoneCode: %v", err)
|
||||
}
|
||||
if hash == "" || delivery.Kind != domain.AuthCodeDeliverySMS || delivery.Length != 6 || len(sender.requests) != 1 {
|
||||
t.Fatalf("hash=%q delivery=%+v requests=%d", hash, delivery, len(sender.requests))
|
||||
}
|
||||
req := sender.requests[0]
|
||||
if req.Purpose != otpdelivery.PurposeChangePhone || req.Channel != otpdelivery.ChannelSMS || req.Recipient != "15550012020" || req.DeliveryID == "" {
|
||||
t.Fatalf("request = %+v", req)
|
||||
}
|
||||
rec, found, err := f.codes.Get(f.ctx, hash)
|
||||
if err != nil || !found || rec.Channel != store.PhoneCodeChannelSMS || rec.DeliveryID != req.DeliveryID || rec.Code != req.Code {
|
||||
t.Fatalf("record=%+v found=%v err=%v", rec, found, err)
|
||||
}
|
||||
if _, err := f.service.ChangePhone(f.ctx, f.user.ID, f.authKeyID, f.authKeyID, 78, req.Recipient, hash, req.Code, 1700000000); err != nil {
|
||||
t.Fatalf("ChangePhone: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPhoneChangeWebhookRejectionRevokesScopedCode(t *testing.T) {
|
||||
f := newPhoneChangeFixture(t)
|
||||
sender := &captureMailSender{err: &otpdelivery.RejectedError{StatusCode: 503, Code: "UNAVAILABLE", Retryable: true}}
|
||||
tracked := &trackingPhoneCodeStore{CodeStore: f.codes}
|
||||
f.service.codes = tracked
|
||||
f.service.phoneCodeSender = sender
|
||||
f.service.phoneCodeLength = 5
|
||||
|
||||
hash, _, err := f.service.SendChangePhoneCode(f.ctx, f.user.ID, f.authKeyID, 77, "15550012021")
|
||||
if hash != "" || err == nil || len(sender.requests) != 1 {
|
||||
t.Fatalf("hash=%q err=%v requests=%d", hash, err, len(sender.requests))
|
||||
}
|
||||
if tracked.lastHash == "" {
|
||||
t.Fatal("code was not stored before delivery")
|
||||
}
|
||||
if rec, found, getErr := f.codes.Get(f.ctx, tracked.lastHash); getErr != nil || found || rec.Code != "" {
|
||||
t.Fatalf("post-rejection code rec=%+v found=%v err=%v", rec, found, getErr)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPhoneChangeScopesCodeAndPersistsDurableEvent(t *testing.T) {
|
||||
f := newPhoneChangeFixture(t)
|
||||
hash, delivery, err := f.service.SendChangePhoneCode(f.ctx, f.user.ID, f.authKeyID, 77, "+1 (555) 001-2002")
|
||||
|
|
|
|||
|
|
@ -4,13 +4,14 @@ import (
|
|||
"context"
|
||||
"crypto/rand"
|
||||
"crypto/subtle"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"telesrv/internal/domain"
|
||||
"telesrv/internal/links"
|
||||
"telesrv/internal/mail"
|
||||
"telesrv/internal/otpdelivery"
|
||||
"telesrv/internal/store"
|
||||
)
|
||||
|
||||
|
|
@ -47,7 +48,9 @@ type Service struct {
|
|||
phoneChangeCode string
|
||||
phoneChangeCodeTTL time.Duration
|
||||
phoneChangeMaxAttempts int
|
||||
loginEmailSender mail.Sender
|
||||
loginEmailSender otpdelivery.Sender
|
||||
phoneCodeSender otpdelivery.Sender
|
||||
phoneCodeLength int
|
||||
loginEmailCodeTTL time.Duration
|
||||
loginEmailCodeMaxAttempts int
|
||||
loginEmailCodeLength int
|
||||
|
|
@ -136,7 +139,7 @@ func WithPublicBaseURL(baseURL string) ServiceOption {
|
|||
}
|
||||
}
|
||||
|
||||
func WithLoginEmailVerification(codes store.CodeStore, sender mail.Sender, ttl time.Duration, maxAttempts, length int) ServiceOption {
|
||||
func WithLoginEmailVerification(codes store.CodeStore, sender otpdelivery.Sender, ttl time.Duration, maxAttempts, length int) ServiceOption {
|
||||
return func(s *Service) {
|
||||
s.codes = codes
|
||||
s.loginEmailSender = sender
|
||||
|
|
@ -152,6 +155,17 @@ func WithLoginEmailVerification(codes store.CodeStore, sender mail.Sender, ttl t
|
|||
}
|
||||
}
|
||||
|
||||
// WithPhoneCodeDelivery replaces the fixed development code used by the
|
||||
// change-phone flow with an externally delivered SMS code.
|
||||
func WithPhoneCodeDelivery(sender otpdelivery.Sender, length int) ServiceOption {
|
||||
return func(s *Service) {
|
||||
s.phoneCodeSender = sender
|
||||
if length > 0 {
|
||||
s.phoneCodeLength = length
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// NewService 创建 account 服务。
|
||||
func NewService(passwords store.PasswordStore, opts ...ServiceOption) *Service {
|
||||
s := &Service{
|
||||
|
|
@ -162,6 +176,7 @@ func NewService(passwords store.PasswordStore, opts ...ServiceOption) *Service {
|
|||
loginEmailCodeLength: 6,
|
||||
phoneChangeCodeTTL: 5 * time.Minute,
|
||||
phoneChangeMaxAttempts: 5,
|
||||
phoneCodeLength: 5,
|
||||
}
|
||||
for _, opt := range opts {
|
||||
opt(s)
|
||||
|
|
@ -616,18 +631,54 @@ func (s *Service) SendLoginEmailCode(ctx context.Context, userID int64, phone, p
|
|||
return "", 0, err
|
||||
}
|
||||
rec.Code = code
|
||||
deliveryID, err := otpdelivery.NewDeliveryID()
|
||||
if err != nil {
|
||||
return "", 0, err
|
||||
}
|
||||
rec.DeliveryID = deliveryID
|
||||
expiresAt := time.Now().Add(s.loginEmailCodeTTL)
|
||||
if err := s.codes.Set(ctx, key, rec, s.loginEmailCodeTTL); err != nil {
|
||||
return "", 0, err
|
||||
}
|
||||
if err := s.loginEmailSender.SendLoginCode(ctx, email, code, s.loginEmailCodeTTL); err != nil {
|
||||
// Set does not expose its generated revision. A blind Del here could
|
||||
// remove a newer concurrent resend; leave the unreachable random code
|
||||
// to expire or be replaced by the retry instead.
|
||||
snapshot, found, err := s.codes.GetSnapshot(ctx, key)
|
||||
if err != nil {
|
||||
return "", 0, err
|
||||
}
|
||||
if !found || snapshot.Record.DeliveryID != deliveryID {
|
||||
return "", 0, domain.ErrEmailCodeInvalid
|
||||
}
|
||||
purpose := otpdelivery.PurposeLoginEmailChange
|
||||
if setup {
|
||||
purpose = otpdelivery.PurposeLoginEmailSetup
|
||||
}
|
||||
if err := deliverOTP(ctx, s.loginEmailSender, otpdelivery.Request{
|
||||
DeliveryID: deliveryID,
|
||||
Purpose: purpose,
|
||||
Channel: otpdelivery.ChannelEmail,
|
||||
Recipient: email,
|
||||
Code: code,
|
||||
ExpiresAt: expiresAt,
|
||||
}); err != nil {
|
||||
cleanupCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), 2*time.Second)
|
||||
defer cancel()
|
||||
deleted, cleanupErr := s.codes.CompareAndDelete(cleanupCtx, key, snapshot.Revision)
|
||||
if cleanupErr != nil {
|
||||
return "", 0, fmt.Errorf("%w; rollback email code: %v", err, cleanupErr)
|
||||
}
|
||||
_ = deleted // false means a newer concurrent resend owns the key.
|
||||
return "", 0, err
|
||||
}
|
||||
return emailPattern(email), len(code), nil
|
||||
}
|
||||
|
||||
func deliverOTP(ctx context.Context, sender otpdelivery.Sender, req otpdelivery.Request) error {
|
||||
_, err := sender.Deliver(ctx, req)
|
||||
if errors.Is(err, otpdelivery.ErrOutcomeUnknown) {
|
||||
return nil
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *Service) VerifyLoginEmail(ctx context.Context, userID int64, phone, phoneCodeHash, code string, setup bool) (string, error) {
|
||||
if s == nil || s.codes == nil {
|
||||
return "", domain.ErrEmailNotAllowed
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue