perf: sync protocol and core hardening updates
This commit is contained in:
parent
152fed3b87
commit
4390ebf5a9
283 changed files with 29231 additions and 2295 deletions
|
|
@ -3,6 +3,8 @@ package account
|
|||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
|
|
@ -32,6 +34,88 @@ type captureMailSender struct {
|
|||
code string
|
||||
}
|
||||
|
||||
type blockingCodeCAS struct {
|
||||
store.CodeStore
|
||||
mu sync.Mutex
|
||||
blockRevision string
|
||||
blockUpdate bool
|
||||
blockDelete bool
|
||||
entered chan struct{}
|
||||
release chan struct{}
|
||||
once sync.Once
|
||||
}
|
||||
|
||||
type switchableEmailOwnerStore struct {
|
||||
store.UserStore
|
||||
mu sync.RWMutex
|
||||
phone string
|
||||
override bool
|
||||
owner domain.User
|
||||
found bool
|
||||
}
|
||||
|
||||
func (s *switchableEmailOwnerStore) ByPhone(ctx context.Context, phone string) (domain.User, bool, error) {
|
||||
s.mu.RLock()
|
||||
if s.override && domain.NormalizePhone(phone) == s.phone {
|
||||
owner, found := s.owner, s.found
|
||||
s.mu.RUnlock()
|
||||
return owner, found, nil
|
||||
}
|
||||
s.mu.RUnlock()
|
||||
return s.UserStore.ByPhone(ctx, phone)
|
||||
}
|
||||
|
||||
func (s *switchableEmailOwnerStore) switchOwner(phone string, owner domain.User) {
|
||||
s.mu.Lock()
|
||||
s.phone = domain.NormalizePhone(phone)
|
||||
s.owner = owner
|
||||
s.found = true
|
||||
s.override = true
|
||||
s.mu.Unlock()
|
||||
}
|
||||
|
||||
type afterSavePasswordStore struct {
|
||||
store.PasswordStore
|
||||
once sync.Once
|
||||
afterSave func(userID int64, settings domain.PasswordSettings)
|
||||
}
|
||||
|
||||
func (s *afterSavePasswordStore) Save(ctx context.Context, userID int64, settings domain.PasswordSettings) error {
|
||||
if err := s.PasswordStore.Save(ctx, userID, settings); err != nil {
|
||||
return err
|
||||
}
|
||||
if s.afterSave != nil {
|
||||
s.once.Do(func() { s.afterSave(userID, settings) })
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *blockingCodeCAS) shouldBlock(revision string, update bool) bool {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
return revision == s.blockRevision && ((update && s.blockUpdate) || (!update && s.blockDelete))
|
||||
}
|
||||
|
||||
func (s *blockingCodeCAS) waitIfBlocked(revision string, update bool) {
|
||||
if !s.shouldBlock(revision, update) {
|
||||
return
|
||||
}
|
||||
s.once.Do(func() {
|
||||
close(s.entered)
|
||||
<-s.release
|
||||
})
|
||||
}
|
||||
|
||||
func (s *blockingCodeCAS) CompareAndUpdate(ctx context.Context, key, revision string, next store.PhoneCode) (bool, error) {
|
||||
s.waitIfBlocked(revision, true)
|
||||
return s.CodeStore.CompareAndUpdate(ctx, key, revision, next)
|
||||
}
|
||||
|
||||
func (s *blockingCodeCAS) CompareAndDelete(ctx context.Context, key, revision string) (bool, error) {
|
||||
s.waitIfBlocked(revision, false)
|
||||
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
|
||||
|
|
@ -66,22 +150,22 @@ func TestSetLoginEmailPersistsAndMasks(t *testing.T) {
|
|||
}
|
||||
}
|
||||
|
||||
// TestLoginEmailByPhoneAndClear 验证按手机号读取/清除登录邮箱(sendCode 检测 + reset 用)。
|
||||
// TestLoginEmailByPhoneAndClear 验证 sendCode 可按手机号读取,但 reset 只按已锁定 userID 清除。
|
||||
func TestLoginEmailByPhoneAndClear(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
svc, users := newLoginEmailService(t)
|
||||
createUser(t, users, "15550010002")
|
||||
u := createUser(t, users, "15550010002")
|
||||
|
||||
if err := svc.SetLoginEmailByPhone(ctx, "+1 555 001 0002", "bob@mail.com"); err != nil {
|
||||
t.Fatalf("SetLoginEmailByPhone: %v", err)
|
||||
if err := svc.SetLoginEmail(ctx, u.ID, "bob@mail.com"); err != nil {
|
||||
t.Fatalf("SetLoginEmail: %v", err)
|
||||
}
|
||||
email, found, err := svc.LoginEmailByPhone(ctx, "15550010002")
|
||||
if err != nil || !found || email != "bob@mail.com" {
|
||||
t.Fatalf("LoginEmailByPhone = %q found=%v err=%v", email, found, err)
|
||||
}
|
||||
|
||||
if err := svc.ClearLoginEmailByPhone(ctx, "15550010002"); err != nil {
|
||||
t.Fatalf("ClearLoginEmailByPhone: %v", err)
|
||||
if err := svc.ClearLoginEmail(ctx, u.ID); err != nil {
|
||||
t.Fatalf("ClearLoginEmail: %v", err)
|
||||
}
|
||||
if _, found, _ := svc.LoginEmailByPhone(ctx, "15550010002"); found {
|
||||
t.Fatal("login email still present after clear")
|
||||
|
|
@ -182,7 +266,7 @@ func TestLoginEmailSetupRejectsAlreadyOwnedEmailForNewPhone(t *testing.T) {
|
|||
if err := svc.SetLoginEmail(ctx, owner.ID, "owner@example.test"); err != nil {
|
||||
t.Fatalf("SetLoginEmail owner: %v", err)
|
||||
}
|
||||
if err := codes.Set(ctx, "new-phone-hash", store.PhoneCode{Phone: "15550010108", Channel: "email_setup_required", MaxAttempts: 2}, time.Minute); err != nil {
|
||||
if err := codes.Set(ctx, "new-phone-hash", store.PhoneCode{Version: store.PhoneCodeVersionCurrent, Phone: "15550010108", Channel: "email_setup_required", MaxAttempts: 2}, time.Minute); err != nil {
|
||||
t.Fatalf("seed phone code: %v", err)
|
||||
}
|
||||
|
||||
|
|
@ -234,7 +318,7 @@ func TestLoginEmailSetupStoresPendingEmailOnPhoneCodeHash(t *testing.T) {
|
|||
svc := NewService(memory.NewPasswordStore(),
|
||||
WithUsers(memory.NewUserStore()),
|
||||
WithLoginEmailVerification(codes, sender, time.Minute, 2, 6))
|
||||
if err := codes.Set(ctx, "phone-hash", store.PhoneCode{Phone: "15550010006", Channel: "email_setup_required", MaxAttempts: 2}, time.Minute); err != nil {
|
||||
if err := codes.Set(ctx, "phone-hash", store.PhoneCode{Version: store.PhoneCodeVersionCurrent, Phone: "15550010006", Channel: "email_setup_required", MaxAttempts: 2}, time.Minute); err != nil {
|
||||
t.Fatalf("seed phone code: %v", err)
|
||||
}
|
||||
|
||||
|
|
@ -252,7 +336,7 @@ func TestLoginEmailSetupStoresPendingEmailOnPhoneCodeHash(t *testing.T) {
|
|||
if err != nil || !found {
|
||||
t.Fatalf("phone code found=%v err=%v", found, err)
|
||||
}
|
||||
if rec.Channel != "email_login" || rec.Code != sender.code || rec.Email != "new@example.test" || !rec.VerifiedEmail || rec.PendingEmail != "new@example.test" {
|
||||
if rec.Channel != "email_login" || rec.Code != sender.code || rec.Email != "new@example.test" || !rec.VerifiedEmail || !rec.SignUpVerified || rec.PendingEmail != "new@example.test" {
|
||||
t.Fatalf("phone code after verify = %+v", rec)
|
||||
}
|
||||
}
|
||||
|
|
@ -291,3 +375,184 @@ func TestVerifyLoginEmailDeletesCodeAfterMaxAttempts(t *testing.T) {
|
|||
t.Fatal("login email was set after exhausted verification code")
|
||||
}
|
||||
}
|
||||
|
||||
func TestStaleLoginEmailVerificationCannotMutateResentCode(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
blockUpdate bool
|
||||
blockDelete bool
|
||||
verificationCode func(old string) string
|
||||
}{
|
||||
{
|
||||
name: "wrong-code-update",
|
||||
blockUpdate: true,
|
||||
verificationCode: func(old string) string {
|
||||
if old != "000000" {
|
||||
return "000000"
|
||||
}
|
||||
return "111111"
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "correct-code-delete",
|
||||
blockDelete: true,
|
||||
verificationCode: func(old string) string { return old },
|
||||
},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
users := memory.NewUserStore()
|
||||
baseCodes := memory.NewCodeStore()
|
||||
codes := &blockingCodeCAS{CodeStore: baseCodes}
|
||||
passwords := memory.NewPasswordStore()
|
||||
sender := &captureMailSender{}
|
||||
svc := NewService(passwords,
|
||||
WithUsers(users),
|
||||
WithLoginEmailVerification(codes, sender, time.Minute, 3, 6))
|
||||
u := createUser(t, users, "155500102"+fmt.Sprint(10+len(tc.name)))
|
||||
if _, _, err := svc.SendLoginEmailCode(ctx, u.ID, "", "", "cas@example.test", false); err != nil {
|
||||
t.Fatalf("first SendLoginEmailCode: %v", err)
|
||||
}
|
||||
oldCode := sender.code
|
||||
key := loginEmailVerifyChangePrefix + fmt.Sprint(u.ID)
|
||||
oldSnapshot, found, err := baseCodes.GetSnapshot(ctx, key)
|
||||
if err != nil || !found {
|
||||
t.Fatalf("old snapshot found=%v err=%v", found, err)
|
||||
}
|
||||
codes.blockRevision = oldSnapshot.Revision
|
||||
codes.blockUpdate = tc.blockUpdate
|
||||
codes.blockDelete = tc.blockDelete
|
||||
codes.entered = make(chan struct{})
|
||||
codes.release = make(chan struct{})
|
||||
|
||||
verifyErr := make(chan error, 1)
|
||||
go func() {
|
||||
_, err := svc.VerifyLoginEmail(ctx, u.ID, "", "", tc.verificationCode(oldCode), false)
|
||||
verifyErr <- err
|
||||
}()
|
||||
<-codes.entered
|
||||
for attempts := 0; attempts < 5; attempts++ {
|
||||
if _, _, err := svc.SendLoginEmailCode(ctx, u.ID, "", "", "cas@example.test", false); err != nil {
|
||||
t.Fatalf("resent SendLoginEmailCode: %v", err)
|
||||
}
|
||||
if sender.code != oldCode {
|
||||
break
|
||||
}
|
||||
}
|
||||
newCode := sender.code
|
||||
if newCode == oldCode {
|
||||
t.Fatal("random resend repeatedly produced the old code")
|
||||
}
|
||||
newSnapshot, found, err := baseCodes.GetSnapshot(ctx, key)
|
||||
if err != nil || !found || newSnapshot.Revision == oldSnapshot.Revision {
|
||||
t.Fatalf("new snapshot=%+v found=%v err=%v", newSnapshot, found, err)
|
||||
}
|
||||
close(codes.release)
|
||||
if err := <-verifyErr; !errors.Is(err, domain.ErrEmailCodeInvalid) {
|
||||
t.Fatalf("stale verification err=%v, want ErrEmailCodeInvalid", err)
|
||||
}
|
||||
current, found, err := baseCodes.GetSnapshot(ctx, key)
|
||||
if err != nil || !found || current.Revision != newSnapshot.Revision || current.Record.Code != newCode || current.Record.Attempts != 0 {
|
||||
t.Fatalf("current code after stale verifier=%+v found=%v err=%v", current, found, err)
|
||||
}
|
||||
if _, err := svc.VerifyLoginEmail(ctx, u.ID, "", "", newCode, false); err != nil {
|
||||
t.Fatalf("VerifyLoginEmail new code: %v", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestConcurrentWrongLoginEmailCodesNeverAuthorize(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
users := memory.NewUserStore()
|
||||
codes := memory.NewCodeStore()
|
||||
passwords := memory.NewPasswordStore()
|
||||
sender := &captureMailSender{}
|
||||
svc := NewService(passwords,
|
||||
WithUsers(users),
|
||||
WithLoginEmailVerification(codes, sender, time.Minute, 2, 6))
|
||||
u := createUser(t, users, "15550010231")
|
||||
if _, _, err := svc.SendLoginEmailCode(ctx, u.ID, "", "", "wrong@example.test", false); err != nil {
|
||||
t.Fatalf("SendLoginEmailCode: %v", err)
|
||||
}
|
||||
correct := sender.code
|
||||
wrong := "000000"
|
||||
if wrong == correct {
|
||||
wrong = "111111"
|
||||
}
|
||||
const workers = 32
|
||||
start := make(chan struct{})
|
||||
errs := make(chan error, workers)
|
||||
for i := 0; i < workers; i++ {
|
||||
go func() {
|
||||
<-start
|
||||
_, err := svc.VerifyLoginEmail(ctx, u.ID, "", "", wrong, false)
|
||||
errs <- err
|
||||
}()
|
||||
}
|
||||
close(start)
|
||||
for i := 0; i < workers; i++ {
|
||||
if err := <-errs; !errors.Is(err, domain.ErrEmailCodeInvalid) {
|
||||
t.Fatalf("wrong concurrent verification err=%v", err)
|
||||
}
|
||||
}
|
||||
if _, found, err := svc.LoginEmail(ctx, u.ID); err != nil || found {
|
||||
t.Fatalf("LoginEmail after wrong codes found=%v err=%v, want absent", found, err)
|
||||
}
|
||||
key := loginEmailVerifyChangePrefix + fmt.Sprint(u.ID)
|
||||
for attempts := 0; attempts < 3; attempts++ {
|
||||
if _, found, err := codes.GetSnapshot(ctx, key); err != nil {
|
||||
t.Fatalf("GetSnapshot after concurrent attempts: %v", err)
|
||||
} else if !found {
|
||||
break
|
||||
}
|
||||
if _, err := svc.VerifyLoginEmail(ctx, u.ID, "", "", wrong, false); !errors.Is(err, domain.ErrEmailCodeInvalid) {
|
||||
t.Fatalf("final wrong verification err=%v", err)
|
||||
}
|
||||
}
|
||||
if _, err := svc.VerifyLoginEmail(ctx, u.ID, "", "", correct, false); !errors.Is(err, domain.ErrEmailCodeInvalid) {
|
||||
t.Fatalf("correct code after exhausted attempts err=%v, want invalid", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEmailSetupOwnerTransferDuringSaveNeverWritesFactorToNewOwner(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
baseUsers := memory.NewUserStore()
|
||||
ownerA := createUser(t, baseUsers, "15550010241")
|
||||
ownerB := createUser(t, baseUsers, "15550010242")
|
||||
users := &switchableEmailOwnerStore{UserStore: baseUsers}
|
||||
basePasswords := memory.NewPasswordStore()
|
||||
passwords := &afterSavePasswordStore{PasswordStore: basePasswords}
|
||||
codes := memory.NewCodeStore()
|
||||
sender := &captureMailSender{}
|
||||
svc := NewService(passwords,
|
||||
WithUsers(users),
|
||||
WithLoginEmailVerification(codes, sender, time.Minute, 3, 6))
|
||||
hash := "owner-save-race"
|
||||
if err := codes.Set(ctx, hash, store.PhoneCode{
|
||||
Version: store.PhoneCodeVersionCurrent,
|
||||
IssuedUserID: ownerA.ID,
|
||||
Phone: ownerA.Phone,
|
||||
Channel: codeChannelEmailSetupRequired,
|
||||
MaxAttempts: 3,
|
||||
}, time.Minute); err != nil {
|
||||
t.Fatalf("seed phone code: %v", err)
|
||||
}
|
||||
if _, _, err := svc.SendLoginEmailCode(ctx, 0, ownerA.Phone, hash, "owner-a@example.test", true); err != nil {
|
||||
t.Fatalf("SendLoginEmailCode: %v", err)
|
||||
}
|
||||
passwords.afterSave = func(userID int64, settings domain.PasswordSettings) {
|
||||
if userID == ownerA.ID && settings.LoginEmail == "owner-a@example.test" {
|
||||
users.switchOwner(ownerA.Phone, ownerB)
|
||||
}
|
||||
}
|
||||
if _, err := svc.VerifyLoginEmail(ctx, 0, ownerA.Phone, hash, sender.code, true); !errors.Is(err, domain.ErrEmailCodeInvalid) {
|
||||
t.Fatalf("VerifyLoginEmail across save-time owner transfer err=%v, want invalid", err)
|
||||
}
|
||||
if settings, found, err := basePasswords.GetByUser(ctx, ownerB.ID); err != nil || (found && settings.LoginEmail != "") {
|
||||
t.Fatalf("new owner settings=%+v found=%v err=%v, SMTP factor leaked to B", settings, found, err)
|
||||
}
|
||||
if _, found, err := codes.GetSnapshot(ctx, hash); err != nil || found {
|
||||
t.Fatalf("owner-drift phone hash found=%v err=%v, want invalidated", found, err)
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -3,7 +3,6 @@ package account
|
|||
import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"crypto/subtle"
|
||||
"encoding/hex"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
|
@ -51,9 +50,10 @@ func (s *Service) SendChangePhoneCode(ctx context.Context, userID int64, authKey
|
|||
return "", domain.AuthCodeDelivery{}, err
|
||||
}
|
||||
rec := store.PhoneCode{
|
||||
Version: store.PhoneCodeVersionCurrent,
|
||||
Phone: phone,
|
||||
Code: s.phoneChangeCode,
|
||||
Channel: "phone",
|
||||
Channel: store.PhoneCodeChannelPhone,
|
||||
Purpose: store.PhoneCodePurposeChangePhone,
|
||||
UserID: userID,
|
||||
AuthKeyID: authKeyID,
|
||||
|
|
@ -68,7 +68,7 @@ func (s *Service) SendChangePhoneCode(ctx context.Context, userID int64, authKey
|
|||
|
||||
// ChangePhone 验证作用域和验证码后执行原子改号。返回事件用于当前 session 的
|
||||
// pts 簿记;其它 session 由 transactional outbox 投递 updateUserPhone。
|
||||
func (s *Service) ChangePhone(ctx context.Context, userID int64, authKeyID [8]byte, sessionID int64, phone, phoneCodeHash, code string, date int) (domain.PhoneChangeResult, error) {
|
||||
func (s *Service) ChangePhone(ctx context.Context, userID int64, authKeyID, originRawAuthKeyID [8]byte, sessionID int64, phone, phoneCodeHash, code string, date int) (domain.PhoneChangeResult, error) {
|
||||
if strings.TrimSpace(phoneCodeHash) == "" || strings.TrimSpace(code) == "" {
|
||||
return domain.PhoneChangeResult{}, domain.ErrPhoneCodeEmpty
|
||||
}
|
||||
|
|
@ -82,51 +82,45 @@ func (s *Service) ChangePhone(ctx context.Context, userID int64, authKeyID [8]by
|
|||
if s.codes == nil || s.phoneChanges == nil {
|
||||
return domain.PhoneChangeResult{}, fmt.Errorf("phone change service is not configured")
|
||||
}
|
||||
rec, found, err := s.codes.Get(ctx, phoneCodeHash)
|
||||
scope := store.PhoneCodeScope{
|
||||
Purpose: store.PhoneCodePurposeChangePhone,
|
||||
UserID: userID,
|
||||
AuthKeyID: authKeyID,
|
||||
Phone: phone,
|
||||
}
|
||||
verified, err := s.codes.VerifyScoped(ctx, phoneCodeHash, scope, strings.TrimSpace(code), s.phoneChangeMaxAttempts)
|
||||
if err != nil {
|
||||
return domain.PhoneChangeResult{}, err
|
||||
}
|
||||
if !found {
|
||||
switch verified.Status {
|
||||
case store.LoginCodeVerifyMissing:
|
||||
return domain.PhoneChangeResult{}, domain.ErrPhoneCodeExpired
|
||||
}
|
||||
if rec.Purpose != store.PhoneCodePurposeChangePhone || rec.Phone != phone || rec.UserID != userID || rec.AuthKeyID != authKeyID {
|
||||
case store.LoginCodeVerifyInvalid:
|
||||
return domain.PhoneChangeResult{}, domain.ErrPhoneCodeInvalid
|
||||
case store.LoginCodeVerifyAccepted:
|
||||
default:
|
||||
return domain.PhoneChangeResult{}, domain.ErrPhoneCodeInvalid
|
||||
}
|
||||
code = strings.TrimSpace(code)
|
||||
if subtle.ConstantTimeCompare([]byte(rec.Code), []byte(code)) != 1 {
|
||||
return domain.PhoneChangeResult{}, s.rejectPhoneChangeCode(ctx, phoneCodeHash, rec)
|
||||
}
|
||||
if existing, occupied, err := s.users.ByPhone(ctx, phone); err != nil {
|
||||
return domain.PhoneChangeResult{}, err
|
||||
} else if occupied && existing.ID != userID {
|
||||
return domain.PhoneChangeResult{}, domain.ErrPhoneNumberOccupied
|
||||
}
|
||||
// 正确 code 必须在进入持久化事务前原子消费。并发重放中只有一个请求能
|
||||
// 获得记录,其余请求不得再次推进 pts 或追加 user_phone event。
|
||||
consumed, found, err := s.codes.ConsumeScoped(ctx, phoneCodeHash, store.PhoneCodeScope{
|
||||
Purpose: store.PhoneCodePurposeChangePhone,
|
||||
UserID: userID,
|
||||
AuthKeyID: authKeyID,
|
||||
Phone: phone,
|
||||
})
|
||||
if err != nil {
|
||||
return domain.PhoneChangeResult{}, err
|
||||
}
|
||||
if !found {
|
||||
return domain.PhoneChangeResult{}, domain.ErrPhoneCodeExpired
|
||||
}
|
||||
if consumed.Purpose != rec.Purpose || consumed.UserID != rec.UserID || consumed.AuthKeyID != rec.AuthKeyID || consumed.Phone != rec.Phone ||
|
||||
subtle.ConstantTimeCompare([]byte(consumed.Code), []byte(code)) != 1 {
|
||||
consumed := verified.Record
|
||||
if consumed.Version != store.PhoneCodeVersionCurrent || consumed.Scope() != scope || consumed.Channel != store.PhoneCodeChannelPhone {
|
||||
return domain.PhoneChangeResult{}, domain.ErrPhoneCodeInvalid
|
||||
}
|
||||
if date == 0 {
|
||||
date = int(time.Now().Unix())
|
||||
}
|
||||
result, err := s.phoneChanges.ChangePhone(ctx, domain.PhoneChangeRequest{
|
||||
UserID: userID,
|
||||
Phone: phone,
|
||||
Date: date,
|
||||
ExcludeAuthKeyID: authKeyID,
|
||||
UserID: userID,
|
||||
Phone: phone,
|
||||
Date: date,
|
||||
// Authorization/code scope is the stable business (perm) key, while dispatch exclusion
|
||||
// must use the physical raw key. They differ on PFS/temp connections; conflating them
|
||||
// echoes updateUserPhone back to the initiating device and suppresses the wrong session.
|
||||
ExcludeAuthKeyID: originRawAuthKeyID,
|
||||
ExcludeSessionID: sessionID,
|
||||
})
|
||||
if err != nil {
|
||||
|
|
@ -162,20 +156,6 @@ func (s *Service) phoneChangeCaller(ctx context.Context, userID int64, authKeyID
|
|||
return u, nil
|
||||
}
|
||||
|
||||
func (s *Service) rejectPhoneChangeCode(ctx context.Context, hash string, rec store.PhoneCode) error {
|
||||
rec.Attempts++
|
||||
max := rec.MaxAttempts
|
||||
if max <= 0 {
|
||||
max = s.phoneChangeMaxAttempts
|
||||
}
|
||||
if max > 0 && rec.Attempts >= max {
|
||||
_ = s.codes.Del(ctx, hash)
|
||||
return domain.ErrPhoneCodeInvalid
|
||||
}
|
||||
_ = s.codes.Update(ctx, hash, rec)
|
||||
return domain.ErrPhoneCodeInvalid
|
||||
}
|
||||
|
||||
func phoneChangeHash() (string, error) {
|
||||
var raw [8]byte
|
||||
if _, err := rand.Read(raw[:]); err != nil {
|
||||
|
|
|
|||
|
|
@ -21,6 +21,26 @@ type phoneChangeFixture struct {
|
|||
events *memory.UpdateEventStore
|
||||
user domain.User
|
||||
authKeyID [8]byte
|
||||
changes *recordingPhoneChangeStore
|
||||
}
|
||||
|
||||
type recordingPhoneChangeStore struct {
|
||||
mu sync.Mutex
|
||||
inner store.PhoneChangeStore
|
||||
last domain.PhoneChangeRequest
|
||||
}
|
||||
|
||||
func (s *recordingPhoneChangeStore) ChangePhone(ctx context.Context, req domain.PhoneChangeRequest) (domain.PhoneChangeResult, error) {
|
||||
s.mu.Lock()
|
||||
s.last = req
|
||||
s.mu.Unlock()
|
||||
return s.inner.ChangePhone(ctx, req)
|
||||
}
|
||||
|
||||
func (s *recordingPhoneChangeStore) lastRequest() domain.PhoneChangeRequest {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
return s.last
|
||||
}
|
||||
|
||||
func newPhoneChangeFixture(t *testing.T) phoneChangeFixture {
|
||||
|
|
@ -38,12 +58,13 @@ func newPhoneChangeFixture(t *testing.T) phoneChangeFixture {
|
|||
if err := auths.Bind(ctx, domain.Authorization{AuthKeyID: authKeyID, UserID: u.ID, CreatedAt: time.Now().Add(-48 * time.Hour)}); err != nil {
|
||||
t.Fatalf("bind auth: %v", err)
|
||||
}
|
||||
changes := &recordingPhoneChangeStore{inner: memory.NewPhoneChangeStore(users, events)}
|
||||
service := NewService(
|
||||
memory.NewPasswordStore(),
|
||||
WithUsers(users),
|
||||
WithPhoneChange(memory.NewPhoneChangeStore(users, events), auths, codes, nil, "12345", time.Minute, 3),
|
||||
WithPhoneChange(changes, auths, codes, nil, "12345", time.Minute, 3),
|
||||
)
|
||||
return phoneChangeFixture{ctx: ctx, service: service, users: users, auths: auths, codes: codes, events: events, user: u, authKeyID: authKeyID}
|
||||
return phoneChangeFixture{ctx: ctx, service: service, users: users, auths: auths, codes: codes, events: events, user: u, authKeyID: authKeyID, changes: changes}
|
||||
}
|
||||
|
||||
func TestPhoneChangeScopesCodeAndPersistsDurableEvent(t *testing.T) {
|
||||
|
|
@ -59,17 +80,21 @@ func TestPhoneChangeScopesCodeAndPersistsDurableEvent(t *testing.T) {
|
|||
if err != nil || !found {
|
||||
t.Fatalf("load code found=%v err=%v", found, err)
|
||||
}
|
||||
if rec.Purpose != store.PhoneCodePurposeChangePhone || rec.Phone != "15550012002" || rec.UserID != f.user.ID || rec.AuthKeyID != f.authKeyID || rec.SessionID != 77 {
|
||||
if rec.Version != store.PhoneCodeVersionCurrent || rec.Purpose != store.PhoneCodePurposeChangePhone || rec.Phone != "15550012002" || rec.UserID != f.user.ID || rec.AuthKeyID != f.authKeyID || rec.SessionID != 77 {
|
||||
t.Fatalf("scoped code = %+v", rec)
|
||||
}
|
||||
|
||||
result, err := f.service.ChangePhone(f.ctx, f.user.ID, f.authKeyID, 88, "+1 555 001 2002", hash, "12345", 1700000000)
|
||||
rawAuthKeyID := [8]byte{8, 8, 8, 8}
|
||||
result, err := f.service.ChangePhone(f.ctx, f.user.ID, f.authKeyID, rawAuthKeyID, 88, "+1 555 001 2002", hash, "12345", 1700000000)
|
||||
if err != nil {
|
||||
t.Fatalf("change phone after session reconnect: %v", err)
|
||||
}
|
||||
if !result.Changed || result.User.Phone != "15550012002" || result.Event.Type != domain.UpdateEventUserPhone || result.Event.Phone != "15550012002" || result.Event.Pts != 1 {
|
||||
t.Fatalf("change result = %+v", result)
|
||||
}
|
||||
if got := f.changes.lastRequest().ExcludeAuthKeyID; got != rawAuthKeyID {
|
||||
t.Fatalf("outbox exclusion auth key = %x, want physical raw %x", got, rawAuthKeyID)
|
||||
}
|
||||
if _, found, _ := f.users.ByPhone(f.ctx, "15550012001"); found {
|
||||
t.Fatal("old phone still resolves")
|
||||
}
|
||||
|
|
@ -103,7 +128,7 @@ func TestPhoneChangeRejectsOccupiedAndCrossAuthCode(t *testing.T) {
|
|||
if err := f.auths.Bind(f.ctx, domain.Authorization{AuthKeyID: otherKey, UserID: occupied.ID}); err != nil {
|
||||
t.Fatalf("bind other auth: %v", err)
|
||||
}
|
||||
if _, err := f.service.ChangePhone(f.ctx, occupied.ID, otherKey, 99, "15550012004", hash, "12345", 0); !errors.Is(err, domain.ErrPhoneCodeInvalid) {
|
||||
if _, err := f.service.ChangePhone(f.ctx, occupied.ID, otherKey, otherKey, 99, "15550012004", hash, "12345", 0); !errors.Is(err, domain.ErrPhoneCodeExpired) {
|
||||
t.Fatalf("cross-auth change err = %v", err)
|
||||
}
|
||||
if got, found, _ := f.users.ByID(f.ctx, occupied.ID); !found || got.Phone != "15550012003" {
|
||||
|
|
@ -118,11 +143,11 @@ func TestPhoneChangeWrongCodeExhaustsAttempts(t *testing.T) {
|
|||
t.Fatalf("send code: %v", err)
|
||||
}
|
||||
for i := 0; i < 3; i++ {
|
||||
if _, err := f.service.ChangePhone(f.ctx, f.user.ID, f.authKeyID, 77, "15550012005", hash, "00000", 0); !errors.Is(err, domain.ErrPhoneCodeInvalid) {
|
||||
if _, err := f.service.ChangePhone(f.ctx, f.user.ID, f.authKeyID, f.authKeyID, 77, "15550012005", hash, "00000", 0); !errors.Is(err, domain.ErrPhoneCodeInvalid) {
|
||||
t.Fatalf("wrong attempt %d err = %v", i+1, err)
|
||||
}
|
||||
}
|
||||
if _, err := f.service.ChangePhone(f.ctx, f.user.ID, f.authKeyID, 77, "15550012005", hash, "12345", 0); !errors.Is(err, domain.ErrPhoneCodeExpired) {
|
||||
if _, err := f.service.ChangePhone(f.ctx, f.user.ID, f.authKeyID, f.authKeyID, 77, "15550012005", hash, "12345", 0); !errors.Is(err, domain.ErrPhoneCodeExpired) {
|
||||
t.Fatalf("exhausted code err = %v", err)
|
||||
}
|
||||
if got, _, _ := f.users.ByID(f.ctx, f.user.ID); got.Phone != "15550012001" {
|
||||
|
|
@ -143,10 +168,10 @@ func TestPhoneChangeNewSendInvalidatesPreviousHash(t *testing.T) {
|
|||
if oldHash == newHash {
|
||||
t.Fatalf("hash was not rotated: %q", oldHash)
|
||||
}
|
||||
if _, err := f.service.ChangePhone(f.ctx, f.user.ID, f.authKeyID, 99, "15550012006", oldHash, "12345", 1700000001); !errors.Is(err, domain.ErrPhoneCodeExpired) {
|
||||
if _, err := f.service.ChangePhone(f.ctx, f.user.ID, f.authKeyID, f.authKeyID, 99, "15550012006", oldHash, "12345", 1700000001); !errors.Is(err, domain.ErrPhoneCodeExpired) {
|
||||
t.Fatalf("old hash replay err = %v", err)
|
||||
}
|
||||
if _, err := f.service.ChangePhone(f.ctx, f.user.ID, f.authKeyID, 99, "15550012006", newHash, "12345", 1700000002); err != nil {
|
||||
if _, err := f.service.ChangePhone(f.ctx, f.user.ID, f.authKeyID, f.authKeyID, 99, "15550012006", newHash, "12345", 1700000002); err != nil {
|
||||
t.Fatalf("new hash change: %v", err)
|
||||
}
|
||||
events, err := f.events.ListAfter(f.ctx, f.user.ID, 0, 10)
|
||||
|
|
@ -168,7 +193,7 @@ func TestPhoneChangeConcurrentReplayAppendsOneEvent(t *testing.T) {
|
|||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
_, err := f.service.ChangePhone(f.ctx, f.user.ID, f.authKeyID, 88, "15550012007", hash, "12345", 1700000003)
|
||||
_, err := f.service.ChangePhone(f.ctx, f.user.ID, f.authKeyID, f.authKeyID, 88, "15550012007", hash, "12345", 1700000003)
|
||||
errs <- err
|
||||
}()
|
||||
}
|
||||
|
|
|
|||
|
|
@ -17,13 +17,14 @@ import (
|
|||
var defaultSecureRandom = []byte("telesrv-tdesktop-dev-secure-rand")
|
||||
|
||||
const (
|
||||
passwordResetWait = 7 * 24 * time.Hour
|
||||
passwordResetRetry = 24 * time.Hour
|
||||
loginEmailVerifyChangePrefix = "login-email-change:"
|
||||
loginEmailVerifySetupPrefix = "login-email-setup:"
|
||||
codeChannelEmailSetup = "email_setup"
|
||||
codeChannelEmailChange = "email_change"
|
||||
codeChannelEmailLogin = "email_login"
|
||||
passwordResetWait = 7 * 24 * time.Hour
|
||||
passwordResetRetry = 24 * time.Hour
|
||||
loginEmailVerifyChangePrefix = "login-email-change:"
|
||||
loginEmailVerifySetupPrefix = "login-email-setup:"
|
||||
codeChannelEmailSetup = "email_setup"
|
||||
codeChannelEmailChange = "email_change"
|
||||
codeChannelEmailLogin = "email_login"
|
||||
codeChannelEmailSetupRequired = "email_setup_required"
|
||||
)
|
||||
|
||||
// Service 提供账号安全配置查询。
|
||||
|
|
@ -568,12 +569,16 @@ func (s *Service) SendLoginEmailCode(ctx context.Context, userID int64, phone, p
|
|||
}
|
||||
key := loginEmailVerifyChangePrefix + fmt.Sprint(userID)
|
||||
rec := store.PhoneCode{
|
||||
Version: store.PhoneCodeVersionCurrent,
|
||||
Code: "",
|
||||
Channel: codeChannelEmailChange,
|
||||
PendingEmail: email,
|
||||
MaxAttempts: s.loginEmailCodeMaxAttempts,
|
||||
}
|
||||
if setup {
|
||||
if s.users == nil {
|
||||
return "", 0, domain.ErrEmailNotAllowed
|
||||
}
|
||||
phone = domain.NormalizePhone(phone)
|
||||
phoneRec, found, err := s.codes.Get(ctx, phoneCodeHash)
|
||||
if err != nil {
|
||||
|
|
@ -582,7 +587,8 @@ func (s *Service) SendLoginEmailCode(ctx context.Context, userID int64, phone, p
|
|||
if !found {
|
||||
return "", 0, domain.ErrEmailCodeInvalid
|
||||
}
|
||||
if phoneRec.Phone != phone {
|
||||
if phoneRec.Version != store.PhoneCodeVersionCurrent || phoneRec.Purpose != "" || phoneRec.Phone != phone ||
|
||||
phoneRec.Channel != codeChannelEmailSetupRequired || phoneRec.SignUpVerified {
|
||||
return "", 0, domain.ErrEmailInvalid
|
||||
}
|
||||
targetUserID := int64(0)
|
||||
|
|
@ -591,6 +597,9 @@ func (s *Service) SendLoginEmailCode(ctx context.Context, userID int64, phone, p
|
|||
} else if found {
|
||||
targetUserID = existingUserID
|
||||
}
|
||||
if phoneRec.IssuedUserID != targetUserID {
|
||||
return "", 0, domain.ErrEmailInvalid
|
||||
}
|
||||
if err := s.ensureLoginEmailAvailable(ctx, targetUserID, email); err != nil {
|
||||
return "", 0, err
|
||||
}
|
||||
|
|
@ -611,7 +620,9 @@ func (s *Service) SendLoginEmailCode(ctx context.Context, userID int64, phone, p
|
|||
return "", 0, err
|
||||
}
|
||||
if err := s.loginEmailSender.SendLoginCode(ctx, email, code, s.loginEmailCodeTTL); err != nil {
|
||||
_ = s.codes.Del(ctx, key)
|
||||
// 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.
|
||||
return "", 0, err
|
||||
}
|
||||
return emailPattern(email), len(code), nil
|
||||
|
|
@ -625,28 +636,43 @@ func (s *Service) VerifyLoginEmail(ctx context.Context, userID int64, phone, pho
|
|||
if setup {
|
||||
key = loginEmailVerifySetupPrefix + phoneCodeHash
|
||||
}
|
||||
rec, found, err := s.codes.Get(ctx, key)
|
||||
snapshot, found, err := s.codes.GetSnapshot(ctx, key)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if !found {
|
||||
return "", domain.ErrEmailCodeInvalid
|
||||
}
|
||||
rec := snapshot.Record
|
||||
if strings.TrimSpace(code) == "" || subtle.ConstantTimeCompare([]byte(rec.Code), []byte(strings.TrimSpace(code))) != 1 {
|
||||
return "", s.rejectEmailCode(ctx, key, rec)
|
||||
return "", s.rejectEmailCode(ctx, key, snapshot)
|
||||
}
|
||||
email := normalizeLoginEmail(rec.PendingEmail)
|
||||
if !validLoginEmail(email) {
|
||||
_ = s.codes.Del(ctx, key)
|
||||
applied, deleteErr := s.codes.CompareAndDelete(ctx, key, snapshot.Revision)
|
||||
if deleteErr != nil {
|
||||
return "", deleteErr
|
||||
}
|
||||
if !applied {
|
||||
return "", domain.ErrEmailCodeInvalid
|
||||
}
|
||||
return "", domain.ErrEmailInvalid
|
||||
}
|
||||
if setup {
|
||||
if s.users == nil || rec.Channel != codeChannelEmailSetup {
|
||||
return "", domain.ErrEmailCodeInvalid
|
||||
}
|
||||
phone = domain.NormalizePhone(phone)
|
||||
phoneRec, found, err := s.codes.Get(ctx, phoneCodeHash)
|
||||
if rec.Phone != phone {
|
||||
return "", domain.ErrEmailCodeInvalid
|
||||
}
|
||||
phoneSnapshot, found, err := s.codes.GetSnapshot(ctx, phoneCodeHash)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if !found || phoneRec.Phone != phone {
|
||||
phoneRec := phoneSnapshot.Record
|
||||
if !found || phoneRec.Version != store.PhoneCodeVersionCurrent || phoneRec.Purpose != "" ||
|
||||
phoneRec.Phone != phone || phoneRec.Channel != codeChannelEmailSetupRequired || phoneRec.SignUpVerified {
|
||||
return "", domain.ErrEmailCodeInvalid
|
||||
}
|
||||
targetUserID := int64(0)
|
||||
|
|
@ -655,11 +681,23 @@ func (s *Service) VerifyLoginEmail(ctx context.Context, userID int64, phone, pho
|
|||
} else if found {
|
||||
targetUserID = existingUserID
|
||||
}
|
||||
if phoneRec.IssuedUserID != targetUserID {
|
||||
s.invalidateLoginCode(ctx, phoneCodeHash, phone)
|
||||
return "", domain.ErrEmailCodeInvalid
|
||||
}
|
||||
if err := s.ensureLoginEmailAvailable(ctx, targetUserID, email); err != nil {
|
||||
_ = s.codes.Del(ctx, key)
|
||||
return "", err
|
||||
}
|
||||
_ = s.codes.Del(ctx, key)
|
||||
// Claim this exact email-code revision before mutating the phone login
|
||||
// state. A concurrent resend rotates the revision, so an old verifier
|
||||
// can neither consume the new code nor authorize the phone hash.
|
||||
claimed, err := s.codes.CompareAndDelete(ctx, key, snapshot.Revision)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if !claimed {
|
||||
return "", domain.ErrEmailCodeInvalid
|
||||
}
|
||||
phoneRec.Channel = codeChannelEmailLogin
|
||||
phoneRec.Code = strings.TrimSpace(code)
|
||||
phoneRec.Email = email
|
||||
|
|
@ -667,43 +705,94 @@ func (s *Service) VerifyLoginEmail(ctx context.Context, userID int64, phone, pho
|
|||
phoneRec.VerifiedEmail = true
|
||||
phoneRec.Attempts = 0
|
||||
phoneRec.MaxAttempts = s.loginEmailCodeMaxAttempts
|
||||
if err := s.codes.Update(ctx, phoneCodeHash, phoneRec); err != nil {
|
||||
updated, err := s.codes.CompareAndUpdate(ctx, phoneCodeHash, phoneSnapshot.Revision, phoneRec)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if _, found, err := s.userIDByPhone(ctx, phone); err != nil {
|
||||
if !updated {
|
||||
return "", domain.ErrEmailCodeInvalid
|
||||
}
|
||||
if targetUserID == 0 {
|
||||
verified, err := s.codes.VerifyLogin(ctx, phoneCodeHash, phone, phoneRec.Code, true, s.loginEmailCodeMaxAttempts)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if verified.Status != store.LoginCodeVerifyAccepted || verified.Record.IssuedUserID != 0 || !verified.Record.SignUpVerified {
|
||||
return "", domain.ErrEmailCodeInvalid
|
||||
}
|
||||
phoneRec = verified.Record
|
||||
}
|
||||
afterUserID := int64(0)
|
||||
if existingUserID, found, err := s.userIDByPhone(ctx, phone); err != nil {
|
||||
return "", err
|
||||
} else if found {
|
||||
if err := s.SetLoginEmailByPhone(ctx, phone, email); err != nil {
|
||||
afterUserID = existingUserID
|
||||
}
|
||||
if afterUserID != targetUserID || phoneRec.IssuedUserID != afterUserID {
|
||||
s.invalidateLoginCode(ctx, phoneCodeHash, phone)
|
||||
return "", domain.ErrEmailCodeInvalid
|
||||
}
|
||||
if targetUserID != 0 {
|
||||
// Keep the identity selected before SMTP verification. Re-resolving
|
||||
// phone at this write boundary would let an A→B transfer attach A's
|
||||
// verified factor to B.
|
||||
if err := s.SetLoginEmail(ctx, targetUserID, email); err != nil {
|
||||
return "", err
|
||||
}
|
||||
finalUserID := int64(0)
|
||||
if existingUserID, found, err := s.userIDByPhone(ctx, phone); err != nil {
|
||||
return "", err
|
||||
} else if found {
|
||||
finalUserID = existingUserID
|
||||
}
|
||||
if finalUserID != targetUserID {
|
||||
s.invalidateLoginCode(ctx, phoneCodeHash, phone)
|
||||
return "", domain.ErrEmailCodeInvalid
|
||||
}
|
||||
}
|
||||
return email, nil
|
||||
}
|
||||
if err := s.ensureLoginEmailAvailable(ctx, userID, email); err != nil {
|
||||
_ = s.codes.Del(ctx, key)
|
||||
return "", err
|
||||
}
|
||||
_ = s.codes.Del(ctx, key)
|
||||
claimed, err := s.codes.CompareAndDelete(ctx, key, snapshot.Revision)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if !claimed {
|
||||
return "", domain.ErrEmailCodeInvalid
|
||||
}
|
||||
if err := s.SetLoginEmail(ctx, userID, email); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return email, nil
|
||||
}
|
||||
|
||||
func (s *Service) rejectEmailCode(ctx context.Context, key string, rec store.PhoneCode) error {
|
||||
func (s *Service) rejectEmailCode(ctx context.Context, key string, snapshot store.PhoneCodeSnapshot) error {
|
||||
rec := snapshot.Record
|
||||
rec.Attempts++
|
||||
max := rec.MaxAttempts
|
||||
if max <= 0 {
|
||||
max = s.loginEmailCodeMaxAttempts
|
||||
}
|
||||
if max > 0 && rec.Attempts >= max {
|
||||
_ = s.codes.Del(ctx, key)
|
||||
if _, err := s.codes.CompareAndDelete(ctx, key, snapshot.Revision); err != nil {
|
||||
return err
|
||||
}
|
||||
return domain.ErrEmailCodeInvalid
|
||||
}
|
||||
_ = s.codes.Update(ctx, key, rec)
|
||||
if _, err := s.codes.CompareAndUpdate(ctx, key, snapshot.Revision, rec); err != nil {
|
||||
return err
|
||||
}
|
||||
return domain.ErrEmailCodeInvalid
|
||||
}
|
||||
|
||||
func (s *Service) invalidateLoginCode(ctx context.Context, hash, phone string) {
|
||||
cleanupCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), 2*time.Second)
|
||||
defer cancel()
|
||||
_, _ = s.codes.InvalidateLoginCode(cleanupCtx, hash, phone)
|
||||
}
|
||||
|
||||
// SetLoginEmail 为已登录用户写入登录邮箱(authed 的 emailVerifyPurposeLoginChange)。
|
||||
// 账号无 2FA 也可设置:account_passwords 行可在 has_password=false 下仅承载登录邮箱。
|
||||
func (s *Service) SetLoginEmail(ctx context.Context, userID int64, email string) error {
|
||||
|
|
@ -726,19 +815,6 @@ func (s *Service) SetLoginEmail(ctx context.Context, userID int64, email string)
|
|||
return s.passwords.Save(ctx, userID, settings)
|
||||
}
|
||||
|
||||
// SetLoginEmailByPhone 为某手机号对应的账号写入登录邮箱(登录流程中的
|
||||
// emailVerifyPurposeLoginSetup,此时尚未鉴权,只能凭 phone 定位用户)。
|
||||
func (s *Service) SetLoginEmailByPhone(ctx context.Context, phone, email string) error {
|
||||
userID, found, err := s.userIDByPhone(ctx, phone)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if !found {
|
||||
return domain.ErrEmailInvalid
|
||||
}
|
||||
return s.SetLoginEmail(ctx, userID, email)
|
||||
}
|
||||
|
||||
// LoginEmail 返回已登录用户的登录邮箱原始地址(用于 verifyEmail 回显 emailVerified.email)。
|
||||
func (s *Service) LoginEmail(ctx context.Context, userID int64) (string, bool, error) {
|
||||
if s == nil || s.passwords == nil || userID == 0 {
|
||||
|
|
@ -754,8 +830,7 @@ func (s *Service) LoginEmail(ctx context.Context, userID int64) (string, bool, e
|
|||
return normalizeLoginEmail(settings.LoginEmail), true, nil
|
||||
}
|
||||
|
||||
// LoginEmailByPhone 按手机号返回登录邮箱原始地址(供 auth.sendCode 检测是否改投邮箱、
|
||||
// login-setup 回显、reset 回显使用)。
|
||||
// LoginEmailByPhone 按手机号返回登录邮箱原始地址,供 auth.sendCode 检测是否改投邮箱。
|
||||
func (s *Service) LoginEmailByPhone(ctx context.Context, phone string) (string, bool, error) {
|
||||
userID, found, err := s.userIDByPhone(ctx, phone)
|
||||
if err != nil || !found {
|
||||
|
|
@ -764,14 +839,12 @@ func (s *Service) LoginEmailByPhone(ctx context.Context, phone string) (string,
|
|||
return s.LoginEmail(ctx, userID)
|
||||
}
|
||||
|
||||
// ClearLoginEmailByPhone 清除某手机号账号的登录邮箱(auth.resetLoginEmail)。
|
||||
func (s *Service) ClearLoginEmailByPhone(ctx context.Context, phone string) error {
|
||||
userID, found, err := s.userIDByPhone(ctx, phone)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if !found {
|
||||
return nil
|
||||
// ClearLoginEmail clears the factor on the exact account selected by the
|
||||
// preceding reset-code consume. Authentication factors must never be mutated
|
||||
// through a second phone→user lookup.
|
||||
func (s *Service) ClearLoginEmail(ctx context.Context, userID int64) error {
|
||||
if s == nil || s.passwords == nil || userID == 0 {
|
||||
return domain.ErrEmailInvalid
|
||||
}
|
||||
settings, found, err := s.passwords.GetByUser(ctx, userID)
|
||||
if err != nil || !found {
|
||||
|
|
|
|||
|
|
@ -15,6 +15,7 @@ func TestResendCodePreservesChangePhoneScopeAndSMSDelivery(t *testing.T) {
|
|||
codes := memory.NewCodeStore()
|
||||
authKeyID := [8]byte{8, 7, 6}
|
||||
rec := store.PhoneCode{
|
||||
Version: store.PhoneCodeVersionCurrent,
|
||||
Phone: "15550014001",
|
||||
Code: "old",
|
||||
Channel: codeChannelPhone,
|
||||
|
|
@ -69,3 +70,34 @@ func TestResendCodePreservesChangePhoneScopeAndSMSDelivery(t *testing.T) {
|
|||
t.Fatal("scoped cancel left hash valid")
|
||||
}
|
||||
}
|
||||
|
||||
func TestResendAndCancelRejectLegacyChangePhoneCode(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
codes := memory.NewCodeStore()
|
||||
authKeyID := [8]byte{8, 8, 8}
|
||||
legacy := store.PhoneCode{
|
||||
Version: 0, Phone: "15550014002", Code: "12345", Channel: codeChannelPhone,
|
||||
Purpose: store.PhoneCodePurposeChangePhone, UserID: 43, AuthKeyID: authKeyID,
|
||||
}
|
||||
svc := NewService(memory.NewUserStore(), memory.NewAuthorizationStore(), codes, nil, nil, "12345", WithCodeTTL(time.Minute))
|
||||
|
||||
if err := codes.Set(ctx, "legacy-resend", legacy, time.Minute); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := svc.ResendCodeForAuthKey(ctx, authKeyID, legacy.Phone, "legacy-resend"); err != ErrCodeExpired {
|
||||
t.Fatalf("legacy resend err=%v, want ErrCodeExpired", err)
|
||||
}
|
||||
if _, found, _ := codes.Get(ctx, "legacy-resend"); found {
|
||||
t.Fatal("legacy resend left code active")
|
||||
}
|
||||
|
||||
if err := codes.Set(ctx, "legacy-cancel", legacy, time.Minute); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := svc.CancelCodeForAuthKey(ctx, authKeyID, legacy.Phone, "legacy-cancel"); err != ErrCodeExpired {
|
||||
t.Fatalf("legacy cancel err=%v, want ErrCodeExpired", err)
|
||||
}
|
||||
if _, found, _ := codes.Get(ctx, "legacy-cancel"); found {
|
||||
t.Fatal("legacy cancel left code active")
|
||||
}
|
||||
}
|
||||
|
|
|
|||
366
internal/app/auth/login_code_delivery_test.go
Normal file
366
internal/app/auth/login_code_delivery_test.go
Normal file
|
|
@ -0,0 +1,366 @@
|
|||
package auth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"telesrv/internal/domain"
|
||||
"telesrv/internal/store"
|
||||
"telesrv/internal/store/memory"
|
||||
)
|
||||
|
||||
type captureLoginCodeDelivery struct {
|
||||
requests []domain.LoginCodeDeliveryRequest
|
||||
result domain.LoginCodeDeliveryResult
|
||||
err error
|
||||
failAt int
|
||||
}
|
||||
|
||||
func (d *captureLoginCodeDelivery) DeliverLoginCodeMessage(_ context.Context, req domain.LoginCodeDeliveryRequest) (domain.LoginCodeDeliveryResult, error) {
|
||||
d.requests = append(d.requests, req)
|
||||
if d.err != nil && (d.failAt == 0 || len(d.requests) == d.failAt) {
|
||||
return domain.LoginCodeDeliveryResult{}, d.err
|
||||
}
|
||||
return d.result, nil
|
||||
}
|
||||
|
||||
type trackingCodeStore struct {
|
||||
store.CodeStore
|
||||
lastSetHash string
|
||||
deleted []string
|
||||
deleteCtx []error
|
||||
deleteErr error
|
||||
}
|
||||
|
||||
func (s *trackingCodeStore) Set(ctx context.Context, hash string, code store.PhoneCode, ttl time.Duration) error {
|
||||
s.lastSetHash = hash
|
||||
return s.CodeStore.Set(ctx, hash, code, ttl)
|
||||
}
|
||||
|
||||
func (s *trackingCodeStore) Del(ctx context.Context, hash string) error {
|
||||
s.deleted = append(s.deleted, hash)
|
||||
s.deleteCtx = append(s.deleteCtx, ctx.Err())
|
||||
if s.deleteErr != nil {
|
||||
return s.deleteErr
|
||||
}
|
||||
return s.CodeStore.Del(ctx, hash)
|
||||
}
|
||||
|
||||
func TestExistingAccountSendCodeDeliversBeforeSignInAndDoesNotRedeliver(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
users := memory.NewUserStore()
|
||||
authz := memory.NewAuthorizationStore()
|
||||
codes := memory.NewCodeStore()
|
||||
u, err := users.Create(ctx, domain.User{Phone: "15550009201", FirstName: "Existing"})
|
||||
if err != nil {
|
||||
t.Fatalf("create user: %v", err)
|
||||
}
|
||||
delivery := &captureLoginCodeDelivery{result: domain.LoginCodeDeliveryResult{Created: true}}
|
||||
svc := NewService(users, authz, codes, nil, nil, "12345", WithLoginCodeDelivery(delivery))
|
||||
|
||||
before := int(time.Now().Unix())
|
||||
hash, err := svc.SendCode(ctx, "+1 555 000 9201")
|
||||
if err != nil {
|
||||
t.Fatalf("SendCode: %v", err)
|
||||
}
|
||||
if hash == "" || len(delivery.requests) != 1 {
|
||||
t.Fatalf("SendCode hash=%q delivery calls=%d, want non-empty/1", hash, len(delivery.requests))
|
||||
}
|
||||
req := delivery.requests[0]
|
||||
if req.UserID != u.ID || req.PhoneCodeHash != hash || req.Code != "12345" || req.Date < before || req.ExpiresAt < int64(before)+int64((5*time.Minute)/time.Second)-1 {
|
||||
t.Fatalf("delivery request = %+v, want user=%d hash=%q code=12345 date>=%d", req, u.ID, hash, before)
|
||||
}
|
||||
if rec, found, err := codes.Get(ctx, hash); err != nil || !found || rec.Code != "12345" {
|
||||
t.Fatalf("code after synchronous delivery = %+v found=%v err=%v", rec, found, err)
|
||||
}
|
||||
|
||||
var key [8]byte
|
||||
key[0] = 0x92
|
||||
got, lateMessage, needSignUp, err := svc.SignIn(ctx, domain.Authorization{AuthKeyID: key}, "+15550009201", hash, "12345")
|
||||
if err != nil || needSignUp || got.ID != u.ID {
|
||||
t.Fatalf("SignIn user=%d needSignUp=%v err=%v, want %d/false", got.ID, needSignUp, err, u.ID)
|
||||
}
|
||||
if lateMessage.ID != 0 || len(delivery.requests) != 1 {
|
||||
t.Fatalf("SignIn lateMessage=%+v delivery calls=%d, want zero/unchanged", lateMessage, len(delivery.requests))
|
||||
}
|
||||
}
|
||||
|
||||
func TestDeliveredLoginCodeSurvivesWrongSignInAndCancelWithoutDuplicate(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
users := memory.NewUserStore()
|
||||
u, err := users.Create(ctx, domain.User{Phone: "15550009208", FirstName: "Cancel"})
|
||||
if err != nil {
|
||||
t.Fatalf("create user: %v", err)
|
||||
}
|
||||
codes := memory.NewCodeStore()
|
||||
dialogs := memory.NewDialogStore()
|
||||
messages := memory.NewMessageStore(dialogs)
|
||||
events := memory.NewUpdateEventStore()
|
||||
delivery := memory.NewLoginCodeDeliveryStore(messages, events)
|
||||
svc := NewService(users, memory.NewAuthorizationStore(), codes, nil, nil, "12345", WithLoginCodeDelivery(delivery))
|
||||
|
||||
hash, err := svc.SendCode(ctx, "15550009208")
|
||||
if err != nil {
|
||||
t.Fatalf("SendCode: %v", err)
|
||||
}
|
||||
assertFacts := func(stage string) {
|
||||
t.Helper()
|
||||
history, historyErr := messages.ListByUser(ctx, u.ID, domain.MessageFilter{
|
||||
HasPeer: true,
|
||||
Peer: domain.Peer{Type: domain.PeerTypeUser, ID: domain.OfficialSystemUserID},
|
||||
Limit: 10,
|
||||
})
|
||||
durable, eventErr := events.ListAfter(ctx, u.ID, 0, 10)
|
||||
if historyErr != nil || eventErr != nil || len(history.Messages) != 1 || len(durable) != 1 {
|
||||
t.Fatalf("%s messages=%d events=%d historyErr=%v eventErr=%v, want 1/1", stage, len(history.Messages), len(durable), historyErr, eventErr)
|
||||
}
|
||||
}
|
||||
assertFacts("after SendCode")
|
||||
|
||||
if _, late, _, err := svc.SignIn(ctx, domain.Authorization{}, "15550009208", hash, "00000"); !errors.Is(err, ErrCodeInvalid) || late.ID != 0 {
|
||||
t.Fatalf("wrong SignIn late=%+v err=%v, want ErrCodeInvalid/no message", late, err)
|
||||
}
|
||||
assertFacts("after wrong SignIn")
|
||||
|
||||
if err := svc.CancelCode(ctx, "15550009208", hash); err != nil {
|
||||
t.Fatalf("CancelCode: %v", err)
|
||||
}
|
||||
assertFacts("after CancelCode")
|
||||
if _, late, _, err := svc.SignIn(ctx, domain.Authorization{}, "15550009208", hash, "12345"); !errors.Is(err, ErrCodeExpired) || late.ID != 0 {
|
||||
t.Fatalf("SignIn after cancel late=%+v err=%v, want ErrCodeExpired/no message", late, err)
|
||||
}
|
||||
assertFacts("after canceled SignIn")
|
||||
}
|
||||
|
||||
func TestCodeIssuedBeforeConcurrentOwnerCreationIsRejected(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
users := memory.NewUserStore()
|
||||
authz := memory.NewAuthorizationStore()
|
||||
codes := memory.NewCodeStore()
|
||||
delivery := &captureLoginCodeDelivery{}
|
||||
svc := NewService(users, authz, codes, nil, nil, "12345", WithLoginCodeDelivery(delivery))
|
||||
|
||||
hash, err := svc.SendCode(ctx, "15550009209")
|
||||
if err != nil {
|
||||
t.Fatalf("SendCode before signup: %v", err)
|
||||
}
|
||||
rec, found, err := codes.Get(ctx, hash)
|
||||
if err != nil || !found || rec.Version != store.PhoneCodeVersionCurrent || rec.IssuedUserID != 0 || rec.SignUpVerified || len(delivery.requests) != 0 {
|
||||
t.Fatalf("pre-signup code=%+v found=%v err=%v deliveries=%d", rec, found, err, len(delivery.requests))
|
||||
}
|
||||
u, err := users.Create(ctx, domain.User{Phone: "15550009209", FirstName: "Concurrent"})
|
||||
if err != nil {
|
||||
t.Fatalf("concurrent create user: %v", err)
|
||||
}
|
||||
var key [8]byte
|
||||
key[0] = 0x93
|
||||
got, lateMessage, needSignUp, err := svc.SignIn(ctx, domain.Authorization{AuthKeyID: key}, "15550009209", hash, "12345")
|
||||
if !errors.Is(err, ErrCodeInvalid) || needSignUp || got.ID != 0 || lateMessage.ID != 0 {
|
||||
t.Fatalf("SignIn after owner creation got=%+v late=%+v needSignUp=%v err=%v, want invalid", got, lateMessage, needSignUp, err)
|
||||
}
|
||||
if len(delivery.requests) != 0 {
|
||||
t.Fatalf("owner-transfer code was delivered to new owner: %+v", delivery.requests)
|
||||
}
|
||||
if bound, ok, err := svc.UserID(ctx, key); err != nil || ok || bound != 0 {
|
||||
t.Fatalf("bound user=%d ok=%v err=%v, want no authorization (created uid=%d)", bound, ok, err, u.ID)
|
||||
}
|
||||
}
|
||||
|
||||
func TestExistingAccountRepeatedSendCodeDeliversEachIssuedHashWithoutSignIn(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
users := memory.NewUserStore()
|
||||
if _, err := users.Create(ctx, domain.User{Phone: "15550009202"}); err != nil {
|
||||
t.Fatalf("create user: %v", err)
|
||||
}
|
||||
delivery := &captureLoginCodeDelivery{}
|
||||
svc := NewService(users, memory.NewAuthorizationStore(), memory.NewCodeStore(), nil, nil, "12345", WithLoginCodeDelivery(delivery))
|
||||
|
||||
first, err := svc.SendCode(ctx, "15550009202")
|
||||
if err != nil {
|
||||
t.Fatalf("first SendCode: %v", err)
|
||||
}
|
||||
second, err := svc.SendCode(ctx, "15550009202")
|
||||
if err != nil {
|
||||
t.Fatalf("second SendCode: %v", err)
|
||||
}
|
||||
if first == second || len(delivery.requests) != 2 {
|
||||
t.Fatalf("hashes=%q/%q delivery calls=%d, want distinct/2", first, second, len(delivery.requests))
|
||||
}
|
||||
if delivery.requests[0].PhoneCodeHash != first || delivery.requests[1].PhoneCodeHash != second {
|
||||
t.Fatalf("delivery hashes = %q/%q, want %q/%q", delivery.requests[0].PhoneCodeHash, delivery.requests[1].PhoneCodeHash, first, second)
|
||||
}
|
||||
}
|
||||
|
||||
func TestExistingAccountSendCodeDeliveryFailureRevokesCode(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
users := memory.NewUserStore()
|
||||
if _, err := users.Create(ctx, domain.User{Phone: "15550009203"}); err != nil {
|
||||
t.Fatalf("create user: %v", err)
|
||||
}
|
||||
baseCodes := memory.NewCodeStore()
|
||||
codes := &trackingCodeStore{CodeStore: baseCodes}
|
||||
deliveryCause := errors.New("durable write failed")
|
||||
delivery := &captureLoginCodeDelivery{err: deliveryCause}
|
||||
svc := NewService(users, memory.NewAuthorizationStore(), codes, nil, nil, "12345", WithLoginCodeDelivery(delivery))
|
||||
|
||||
hash, err := svc.SendCode(ctx, "15550009203")
|
||||
if hash != "" || !errors.Is(err, ErrLoginCodeDeliveryFailed) || !errors.Is(err, deliveryCause) {
|
||||
t.Fatalf("SendCode hash=%q err=%v, want empty ErrLoginCodeDeliveryFailed+cause", hash, err)
|
||||
}
|
||||
if codes.lastSetHash == "" || len(codes.deleted) != 1 || codes.deleted[0] != codes.lastSetHash {
|
||||
t.Fatalf("set hash=%q deleted=%v, want exact rollback", codes.lastSetHash, codes.deleted)
|
||||
}
|
||||
if _, found, getErr := baseCodes.Get(ctx, codes.lastSetHash); getErr != nil || found {
|
||||
t.Fatalf("rolled-back hash found=%v err=%v", found, getErr)
|
||||
}
|
||||
}
|
||||
|
||||
func TestExistingAccountAmbiguousDeliveryPreservesCodeForIdempotentRetry(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
users := memory.NewUserStore()
|
||||
if _, err := users.Create(ctx, domain.User{Phone: "15550009213"}); err != nil {
|
||||
t.Fatalf("create user: %v", err)
|
||||
}
|
||||
baseCodes := memory.NewCodeStore()
|
||||
codes := &trackingCodeStore{CodeStore: baseCodes}
|
||||
delivery := &captureLoginCodeDelivery{err: domain.ErrLoginCodeDeliveryCommitAmbiguous}
|
||||
svc := NewService(users, memory.NewAuthorizationStore(), codes, nil, nil, "12345", WithLoginCodeDelivery(delivery))
|
||||
|
||||
hash, err := svc.SendCode(ctx, "15550009213")
|
||||
if hash != "" || !errors.Is(err, ErrLoginCodeDeliveryFailed) || !errors.Is(err, domain.ErrLoginCodeDeliveryCommitAmbiguous) {
|
||||
t.Fatalf("SendCode hash=%q err=%v, want ambiguous delivery failure", hash, err)
|
||||
}
|
||||
if codes.lastSetHash == "" || len(codes.deleted) != 0 {
|
||||
t.Fatalf("ambiguous delivery set=%q deleted=%v, want code preserved", codes.lastSetHash, codes.deleted)
|
||||
}
|
||||
if rec, found, getErr := baseCodes.Get(ctx, codes.lastSetHash); getErr != nil || !found || rec.Code != "12345" {
|
||||
t.Fatalf("ambiguous delivery code=%+v found=%v err=%v", rec, found, getErr)
|
||||
}
|
||||
}
|
||||
|
||||
func TestExplicitDeliveryFailureRollsBackWithDetachedContext(t *testing.T) {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
users := memory.NewUserStore()
|
||||
if _, err := users.Create(context.Background(), domain.User{Phone: "15550009214"}); err != nil {
|
||||
t.Fatalf("create user: %v", err)
|
||||
}
|
||||
baseCodes := memory.NewCodeStore()
|
||||
codes := &trackingCodeStore{CodeStore: baseCodes}
|
||||
delivery := &captureLoginCodeDelivery{err: errors.New("definite rollback")}
|
||||
svc := NewService(users, memory.NewAuthorizationStore(), codes, nil, nil, "12345", WithLoginCodeDelivery(delivery))
|
||||
|
||||
if hash, err := svc.SendCode(ctx, "15550009214"); hash != "" || !errors.Is(err, ErrLoginCodeDeliveryFailed) {
|
||||
t.Fatalf("SendCode hash=%q err=%v, want definite failure", hash, err)
|
||||
}
|
||||
if len(codes.deleted) != 1 || len(codes.deleteCtx) != 1 || codes.deleteCtx[0] != nil {
|
||||
t.Fatalf("rollback deleted=%v ctxErr=%v, want one detached delete", codes.deleted, codes.deleteCtx)
|
||||
}
|
||||
if _, found, err := baseCodes.Get(context.Background(), codes.lastSetHash); err != nil || found {
|
||||
t.Fatalf("detached rollback found=%v err=%v", found, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestExistingAccountMissingDeliveryFailsClosedAndRevokesCode(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
users := memory.NewUserStore()
|
||||
if _, err := users.Create(ctx, domain.User{Phone: "15550009204"}); err != nil {
|
||||
t.Fatalf("create user: %v", err)
|
||||
}
|
||||
baseCodes := memory.NewCodeStore()
|
||||
codes := &trackingCodeStore{CodeStore: baseCodes}
|
||||
svc := NewService(users, memory.NewAuthorizationStore(), codes, nil, nil, "12345")
|
||||
|
||||
hash, err := svc.SendCode(ctx, "15550009204")
|
||||
if hash != "" || !errors.Is(err, ErrLoginCodeDeliveryUnavailable) {
|
||||
t.Fatalf("SendCode hash=%q err=%v, want unavailable", hash, err)
|
||||
}
|
||||
if codes.lastSetHash == "" {
|
||||
t.Fatal("missing delivery was checked before code creation; want rollback path covered")
|
||||
}
|
||||
if _, found, getErr := baseCodes.Get(ctx, codes.lastSetHash); getErr != nil || found {
|
||||
t.Fatalf("unavailable delivery hash found=%v err=%v", found, getErr)
|
||||
}
|
||||
}
|
||||
|
||||
func TestExistingAccountResendDeliversNewHashAndInvalidatesOld(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
users := memory.NewUserStore()
|
||||
if _, err := users.Create(ctx, domain.User{Phone: "15550009205"}); err != nil {
|
||||
t.Fatalf("create user: %v", err)
|
||||
}
|
||||
codes := memory.NewCodeStore()
|
||||
delivery := &captureLoginCodeDelivery{}
|
||||
svc := NewService(users, memory.NewAuthorizationStore(), codes, nil, nil, "12345", WithLoginCodeDelivery(delivery))
|
||||
|
||||
oldHash, err := svc.SendCode(ctx, "15550009205")
|
||||
if err != nil {
|
||||
t.Fatalf("SendCode: %v", err)
|
||||
}
|
||||
newHash, err := svc.ResendCode(ctx, "15550009205", oldHash)
|
||||
if err != nil {
|
||||
t.Fatalf("ResendCode: %v", err)
|
||||
}
|
||||
if oldHash == newHash || len(delivery.requests) != 2 || delivery.requests[1].PhoneCodeHash != newHash {
|
||||
t.Fatalf("old/new=%q/%q deliveries=%+v", oldHash, newHash, delivery.requests)
|
||||
}
|
||||
if _, found, err := codes.Get(ctx, oldHash); err != nil || found {
|
||||
t.Fatalf("old code found=%v err=%v", found, err)
|
||||
}
|
||||
if _, found, err := codes.Get(ctx, newHash); err != nil || !found {
|
||||
t.Fatalf("new code found=%v err=%v", found, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestExistingAccountResendDeliveryFailureLeavesNoUsableCode(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
users := memory.NewUserStore()
|
||||
if _, err := users.Create(ctx, domain.User{Phone: "15550009206"}); err != nil {
|
||||
t.Fatalf("create user: %v", err)
|
||||
}
|
||||
codes := memory.NewCodeStore()
|
||||
delivery := &captureLoginCodeDelivery{err: errors.New("second delivery failed"), failAt: 2}
|
||||
svc := NewService(users, memory.NewAuthorizationStore(), codes, nil, nil, "12345", WithLoginCodeDelivery(delivery))
|
||||
|
||||
oldHash, err := svc.SendCode(ctx, "15550009206")
|
||||
if err != nil {
|
||||
t.Fatalf("SendCode: %v", err)
|
||||
}
|
||||
newHash, err := svc.ResendCode(ctx, "15550009206", oldHash)
|
||||
if newHash != "" || !errors.Is(err, ErrLoginCodeDeliveryFailed) || len(delivery.requests) != 2 {
|
||||
t.Fatalf("ResendCode hash=%q err=%v deliveries=%d", newHash, err, len(delivery.requests))
|
||||
}
|
||||
failedHash := delivery.requests[1].PhoneCodeHash
|
||||
for _, hash := range []string{oldHash, failedHash} {
|
||||
if _, found, getErr := codes.Get(ctx, hash); getErr != nil || found {
|
||||
t.Fatalf("failed resend hash %q found=%v err=%v", hash, found, getErr)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestConfiguredEmailLoginDoesNotLeakCodeThroughAppDelivery(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
users := memory.NewUserStore()
|
||||
if _, err := users.Create(ctx, domain.User{Phone: "15550009207"}); err != nil {
|
||||
t.Fatalf("create user: %v", err)
|
||||
}
|
||||
emails := &testLoginEmailStore{emails: map[string]string{"15550009207": "secure@example.test"}}
|
||||
mailSender := &testMailSender{}
|
||||
delivery := &captureLoginCodeDelivery{}
|
||||
svc := NewService(users, memory.NewAuthorizationStore(), memory.NewCodeStore(), nil, nil, "12345",
|
||||
WithLoginEmail(LoginEmailOptions{Enabled: true, CodeLength: 6, Store: emails, Sender: mailSender}),
|
||||
WithLoginCodeDelivery(delivery),
|
||||
)
|
||||
|
||||
if _, err := svc.SendCode(ctx, "15550009207"); err != nil {
|
||||
t.Fatalf("SendCode: %v", err)
|
||||
}
|
||||
if mailSender.to != "secure@example.test" || mailSender.code == "" {
|
||||
t.Fatalf("email delivery = %q/%q", mailSender.to, mailSender.code)
|
||||
}
|
||||
if len(delivery.requests) != 0 {
|
||||
t.Fatalf("email code leaked into app delivery: %+v", delivery.requests)
|
||||
}
|
||||
}
|
||||
|
|
@ -19,8 +19,7 @@ func (s *testLoginEmailStore) LoginEmailByPhone(_ context.Context, phone string)
|
|||
return email, ok, nil
|
||||
}
|
||||
|
||||
func (s *testLoginEmailStore) SetLoginEmailByPhone(_ context.Context, phone, email string) error {
|
||||
s.emails[domain.NormalizePhone(phone)] = email
|
||||
func (s *testLoginEmailStore) SetLoginEmail(_ context.Context, _ int64, _ string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -9,13 +9,17 @@ import (
|
|||
"telesrv/internal/store/memory"
|
||||
)
|
||||
|
||||
// TestSignInWithEmailCompletesLogin 验证带 email_verification 的登录:注册账号→登出→
|
||||
// 重新 sendCode→用任意邮箱验证码经 SignInWithEmail 完成登录。
|
||||
// TestSignInWithEmailCompletesLogin 验证旧客户端把 phone channel 放进
|
||||
// email_verification 时仍可登录,但验证码必须精确匹配,不能用任意非空值绕过。
|
||||
func TestSignInWithEmailCompletesLogin(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
users := memory.NewUserStore()
|
||||
authz := memory.NewAuthorizationStore()
|
||||
svc := NewService(users, authz, memory.NewCodeStore(), nil, nil, "12345")
|
||||
dialogs := memory.NewDialogStore()
|
||||
messages := memory.NewMessageStore(dialogs)
|
||||
svc := NewService(users, authz, memory.NewCodeStore(), nil, nil, "12345",
|
||||
WithLoginCodeDelivery(memory.NewLoginCodeDeliveryStore(messages, memory.NewUpdateEventStore())),
|
||||
)
|
||||
var key [8]byte
|
||||
key[0] = 0x42
|
||||
|
||||
|
|
@ -23,6 +27,7 @@ func TestSignInWithEmailCompletesLogin(t *testing.T) {
|
|||
if err != nil {
|
||||
t.Fatalf("SendCode signup: %v", err)
|
||||
}
|
||||
verifyCodeForSignUp(t, svc, "+15550009001", hash, "12345")
|
||||
u, _, err := svc.SignUp(ctx, domain.Authorization{AuthKeyID: key}, "+15550009001", hash, "Email", "Login")
|
||||
if err != nil {
|
||||
t.Fatalf("SignUp: %v", err)
|
||||
|
|
@ -35,7 +40,10 @@ func TestSignInWithEmailCompletesLogin(t *testing.T) {
|
|||
if err != nil {
|
||||
t.Fatalf("SendCode signin: %v", err)
|
||||
}
|
||||
got, _, needSignUp, err := svc.SignInWithEmail(ctx, domain.Authorization{AuthKeyID: key}, "+15550009001", hash, "anything-goes")
|
||||
if _, _, _, err := svc.SignInWithEmail(ctx, domain.Authorization{AuthKeyID: key}, "+15550009001", hash, "anything-goes"); !errors.Is(err, ErrCodeInvalid) {
|
||||
t.Fatalf("SignInWithEmail arbitrary nonempty code err=%v, want ErrCodeInvalid", err)
|
||||
}
|
||||
got, _, needSignUp, err := svc.SignInWithEmail(ctx, domain.Authorization{AuthKeyID: key}, "+15550009001", hash, "12345")
|
||||
if err != nil {
|
||||
t.Fatalf("SignInWithEmail: %v", err)
|
||||
}
|
||||
|
|
@ -48,7 +56,7 @@ func TestSignInWithEmailCompletesLogin(t *testing.T) {
|
|||
}
|
||||
}
|
||||
|
||||
// TestSignInWithEmailRejectsEmptyCode 空邮箱验证码必须被拒(即使开发环境码任意,也不能空)。
|
||||
// TestSignInWithEmailRejectsEmptyCode 空邮箱验证码必须被拒。
|
||||
func TestSignInWithEmailRejectsEmptyCode(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
svc := NewService(memory.NewUserStore(), memory.NewAuthorizationStore(), memory.NewCodeStore(), nil, nil, "12345")
|
||||
|
|
@ -66,7 +74,12 @@ func TestSignInWithEmailRejectsEmptyCode(t *testing.T) {
|
|||
func TestSignInWithEmailStillHonorsTwoFactor(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
passwords := memory.NewPasswordStore()
|
||||
svc := NewService(memory.NewUserStore(), memory.NewAuthorizationStore(), memory.NewCodeStore(), nil, nil, "12345", WithPasswords(passwords))
|
||||
dialogs := memory.NewDialogStore()
|
||||
messages := memory.NewMessageStore(dialogs)
|
||||
svc := NewService(memory.NewUserStore(), memory.NewAuthorizationStore(), memory.NewCodeStore(), nil, nil, "12345",
|
||||
WithPasswords(passwords),
|
||||
WithLoginCodeDelivery(memory.NewLoginCodeDeliveryStore(messages, memory.NewUpdateEventStore())),
|
||||
)
|
||||
var key [8]byte
|
||||
key[0] = 0x43
|
||||
|
||||
|
|
@ -74,6 +87,7 @@ func TestSignInWithEmailStillHonorsTwoFactor(t *testing.T) {
|
|||
if err != nil {
|
||||
t.Fatalf("SendCode signup: %v", err)
|
||||
}
|
||||
verifyCodeForSignUp(t, svc, "+15550009003", hash, "12345")
|
||||
u, _, err := svc.SignUp(ctx, domain.Authorization{AuthKeyID: key}, "+15550009003", hash, "Two", "Factor")
|
||||
if err != nil {
|
||||
t.Fatalf("SignUp: %v", err)
|
||||
|
|
@ -89,7 +103,7 @@ func TestSignInWithEmailStillHonorsTwoFactor(t *testing.T) {
|
|||
if err != nil {
|
||||
t.Fatalf("SendCode signin: %v", err)
|
||||
}
|
||||
got, _, _, err := svc.SignInWithEmail(ctx, domain.Authorization{AuthKeyID: key}, "+15550009003", hash, "any-email-code")
|
||||
got, _, _, err := svc.SignInWithEmail(ctx, domain.Authorization{AuthKeyID: key}, "+15550009003", hash, "12345")
|
||||
if !errors.Is(err, domain.ErrSessionPasswordNeeded) {
|
||||
t.Fatalf("SignInWithEmail err = %v, want ErrSessionPasswordNeeded", err)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -19,6 +19,7 @@ func TestSignUpPremiumGrant(t *testing.T) {
|
|||
if err != nil {
|
||||
t.Fatalf("SendCode: %v", err)
|
||||
}
|
||||
verifyCodeForSignUp(t, svc, "+15550004401", hash, "12345")
|
||||
u, _, err := svc.SignUp(ctx, domain.Authorization{}, "+15550004401", hash, "Prem", "User")
|
||||
if err != nil {
|
||||
t.Fatalf("SignUp: %v", err)
|
||||
|
|
@ -41,6 +42,7 @@ func TestSignUpPremiumGrantDisabled(t *testing.T) {
|
|||
if err != nil {
|
||||
t.Fatalf("SendCode: %v", err)
|
||||
}
|
||||
verifyCodeForSignUp(t, svc, "+15550004402", hash, "12345")
|
||||
u, _, err := svc.SignUp(ctx, domain.Authorization{}, "+15550004402", hash, "Free", "User")
|
||||
if err != nil {
|
||||
t.Fatalf("SignUp: %v", err)
|
||||
|
|
|
|||
|
|
@ -27,6 +27,13 @@ var (
|
|||
ErrCodeExpired = errors.New("phone code expired or not found")
|
||||
ErrCodeInvalid = errors.New("phone code invalid")
|
||||
ErrEncryptedMessageInvalid = errors.New("encrypted message invalid")
|
||||
// ErrLoginCodeDeliveryUnavailable 表示已有账号的 app-code 没有可用的
|
||||
// durable message/event/outbox 投递边界。这是服务端配置错误,不能降级成
|
||||
// “继续返回 sentCode,等 signIn 后补发”。
|
||||
ErrLoginCodeDeliveryUnavailable = errors.New("login code durable delivery unavailable")
|
||||
// ErrLoginCodeDeliveryFailed 表示 durable 投递未成功。SendCode/ResendCode
|
||||
// 必须同时撤销刚写入的 CodeStore hash,防止客户拿到无法送达的码。
|
||||
ErrLoginCodeDeliveryFailed = errors.New("login code durable delivery failed")
|
||||
// ErrPhoneNumberInvalid 表示手机号为空或非纯数字/长度越界。
|
||||
// 0090 把 users.phone 唯一约束改为忽略空串的部分索引(bot 行 phone=''),
|
||||
// 因此 phone 校验必须前移到 auth 入口,否则 sendCode/signUp 可无限铸造
|
||||
|
|
@ -40,6 +47,7 @@ const (
|
|||
codeChannelPhone = "phone"
|
||||
codeChannelEmailLogin = "email_login"
|
||||
codeChannelEmailSetupRequired = "email_setup_required"
|
||||
loginCodeRollbackTimeout = 2 * time.Second
|
||||
)
|
||||
|
||||
// validPhone 校验规范化后的手机号:5-32 位纯数字(上限对齐 users.phone 列宽)。
|
||||
|
|
@ -68,6 +76,7 @@ type Service struct {
|
|||
passwords store.PasswordStore
|
||||
messages store.MessageStore
|
||||
dialogs store.DialogStore
|
||||
loginCodeDelivery store.LoginCodeDeliveryStore
|
||||
bots store.BotStore
|
||||
fixedCode string
|
||||
codeTTL time.Duration
|
||||
|
|
@ -83,7 +92,7 @@ type Service struct {
|
|||
|
||||
type loginEmailStore interface {
|
||||
LoginEmailByPhone(ctx context.Context, phone string) (string, bool, error)
|
||||
SetLoginEmailByPhone(ctx context.Context, phone, email string) error
|
||||
SetLoginEmail(ctx context.Context, userID int64, email string) error
|
||||
}
|
||||
|
||||
type LoginEmailOptions struct {
|
||||
|
|
@ -102,7 +111,9 @@ type authorizationRevoker interface {
|
|||
// Option 调整登录服务的可选依赖。
|
||||
type Option func(*Service)
|
||||
|
||||
// WithLoginMessages 在登录成功后写入官方系统账号的登录消息与会话摘要。
|
||||
// WithLoginMessages 在新用户注册成功后写入官方系统账号的首条登录消息与会话摘要。
|
||||
// 已有账号的 app 验证码必须在 auth.sendCode/resendCode 阶段通过
|
||||
// WithLoginCodeDelivery 持久化,禁止在 signIn 成功后补发。
|
||||
func WithLoginMessages(messages store.MessageStore, dialogs store.DialogStore) Option {
|
||||
return func(s *Service) {
|
||||
s.messages = messages
|
||||
|
|
@ -110,6 +121,15 @@ func WithLoginMessages(messages store.MessageStore, dialogs store.DialogStore) O
|
|||
}
|
||||
}
|
||||
|
||||
// WithLoginCodeDelivery 注入已有账号 app-code 的 durable 投递边界。
|
||||
// 实现必须以 user_id + phone_code_hash 幂等,并原子写入 777000
|
||||
// message/dialog/user update event/dispatch outbox。
|
||||
func WithLoginCodeDelivery(delivery store.LoginCodeDeliveryStore) Option {
|
||||
return func(s *Service) {
|
||||
s.loginCodeDelivery = delivery
|
||||
}
|
||||
}
|
||||
|
||||
// WithPasswords lets sign-in stop at SESSION_PASSWORD_NEEDED for 2FA accounts.
|
||||
func WithPasswords(passwords store.PasswordStore) Option {
|
||||
return func(s *Service) {
|
||||
|
|
@ -271,53 +291,157 @@ func (s *Service) SendCode(ctx context.Context, phone string) (string, error) {
|
|||
if systemLoginPhoneForbidden(phone) {
|
||||
return "", ErrSystemUserLoginForbidden
|
||||
}
|
||||
existing, found, err := s.currentPhoneOwner(ctx, phone)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("lookup login-code recipient: %w", err)
|
||||
}
|
||||
if found && systemUserLoginForbidden(existing) {
|
||||
return "", ErrSystemUserLoginForbidden
|
||||
}
|
||||
issuedUserID := int64(0)
|
||||
if found {
|
||||
issuedUserID = existing.ID
|
||||
}
|
||||
if s.loginEmailEnabled && s.loginEmails != nil {
|
||||
email, found, err := s.loginEmails.LoginEmailByPhone(ctx, phone)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if found && strings.TrimSpace(email) != "" {
|
||||
return s.createEmailLoginCode(ctx, phone, email)
|
||||
return s.createEmailLoginCode(ctx, phone, email, issuedUserID)
|
||||
}
|
||||
if s.loginEmailRequireSetup {
|
||||
return s.createSetupRequiredCode(ctx, phone)
|
||||
return s.createSetupRequiredCode(ctx, phone, issuedUserID)
|
||||
}
|
||||
}
|
||||
return s.createPhoneCode(ctx, phone)
|
||||
return s.createPhoneCode(ctx, phone, issuedUserID)
|
||||
}
|
||||
|
||||
func (s *Service) createPhoneCode(ctx context.Context, phone string) (string, error) {
|
||||
func (s *Service) currentPhoneOwner(ctx context.Context, phone string) (domain.User, bool, error) {
|
||||
if s == nil || s.users == nil {
|
||||
return domain.User{}, false, fmt.Errorf("user store is not configured")
|
||||
}
|
||||
return s.users.ByPhone(ctx, phone)
|
||||
}
|
||||
|
||||
func (s *Service) issuedOwnerMatches(ctx context.Context, phone string, issuedUserID int64) (bool, error) {
|
||||
current, found, err := s.currentPhoneOwner(ctx, phone)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
currentUserID := int64(0)
|
||||
if found {
|
||||
currentUserID = current.ID
|
||||
}
|
||||
return currentUserID == issuedUserID, nil
|
||||
}
|
||||
|
||||
func (s *Service) ensureIssuedOwnerAfterSet(ctx context.Context, hash string, rec store.PhoneCode) error {
|
||||
matches, err := s.issuedOwnerMatches(ctx, rec.Phone, rec.IssuedUserID)
|
||||
if err == nil && matches {
|
||||
return nil
|
||||
}
|
||||
cause := err
|
||||
if cause == nil {
|
||||
cause = ErrCodeInvalid
|
||||
}
|
||||
return s.rollbackUndeliveredCode(ctx, hash, cause)
|
||||
}
|
||||
|
||||
func (s *Service) invalidateLoginCodeDetached(ctx context.Context, hash, phone string) {
|
||||
cleanupCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), loginCodeRollbackTimeout)
|
||||
defer cancel()
|
||||
_, _ = s.codes.InvalidateLoginCode(cleanupCtx, hash, phone)
|
||||
}
|
||||
|
||||
func (s *Service) createPhoneCode(ctx context.Context, phone string, existingUserID int64) (string, error) {
|
||||
hash, err := randomHex(8)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if err := s.codes.Set(ctx, hash, store.PhoneCode{
|
||||
Phone: phone,
|
||||
Code: s.fixedCode,
|
||||
Channel: codeChannelPhone,
|
||||
MaxAttempts: s.codeMaxAttempts,
|
||||
Version: store.PhoneCodeVersionCurrent,
|
||||
IssuedUserID: existingUserID,
|
||||
Phone: phone,
|
||||
Code: s.fixedCode,
|
||||
Channel: codeChannelPhone,
|
||||
MaxAttempts: s.codeMaxAttempts,
|
||||
}, s.codeTTL); err != nil {
|
||||
return "", fmt.Errorf("store code: %w", err)
|
||||
}
|
||||
rec := store.PhoneCode{Phone: phone, IssuedUserID: existingUserID}
|
||||
if err := s.ensureIssuedOwnerAfterSet(ctx, hash, rec); err != nil {
|
||||
return "", err
|
||||
}
|
||||
// 新手机号还没有 owner/dialog,只能在 SignUp 创建用户后写第一条
|
||||
// 777000 消息。已有账号则必须在 sendCode RPC 返回前把 app-code
|
||||
// 作为普通 incoming message + durable update/outbox 提交;登录成功不再补发。
|
||||
if existingUserID == 0 {
|
||||
return hash, nil
|
||||
}
|
||||
if err := s.deliverLoginCode(ctx, existingUserID, hash, s.fixedCode); err != nil {
|
||||
return "", s.rollbackUndeliveredCode(ctx, hash, err)
|
||||
}
|
||||
if err := s.ensureIssuedOwnerAfterSet(ctx, hash, rec); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return hash, nil
|
||||
}
|
||||
|
||||
func (s *Service) createSetupRequiredCode(ctx context.Context, phone string) (string, error) {
|
||||
func (s *Service) deliverLoginCode(ctx context.Context, userID int64, phoneCodeHash, code string) error {
|
||||
if s.loginCodeDelivery == nil {
|
||||
return ErrLoginCodeDeliveryUnavailable
|
||||
}
|
||||
now := time.Now()
|
||||
if _, err := s.loginCodeDelivery.DeliverLoginCodeMessage(ctx, domain.LoginCodeDeliveryRequest{
|
||||
UserID: userID,
|
||||
PhoneCodeHash: phoneCodeHash,
|
||||
Code: code,
|
||||
Date: int(now.Unix()),
|
||||
ExpiresAt: now.Add(s.codeTTL).Unix(),
|
||||
}); err != nil {
|
||||
return errors.Join(ErrLoginCodeDeliveryFailed, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *Service) rollbackUndeliveredCode(ctx context.Context, phoneCodeHash string, cause error) error {
|
||||
// lib/pq can report an I/O failure after COMMIT reached PostgreSQL. In that
|
||||
// state deleting the code could turn an already delivered 777000 message
|
||||
// into an unusable login attempt. Preserve it and let the delivery receipt
|
||||
// make the retry idempotent.
|
||||
if errors.Is(cause, domain.ErrLoginCodeDeliveryCommitAmbiguous) {
|
||||
return cause
|
||||
}
|
||||
cleanupCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), loginCodeRollbackTimeout)
|
||||
defer cancel()
|
||||
if err := s.codes.Del(cleanupCtx, phoneCodeHash); err != nil {
|
||||
return errors.Join(cause, fmt.Errorf("rollback undelivered login code: %w", err))
|
||||
}
|
||||
return cause
|
||||
}
|
||||
|
||||
func (s *Service) createSetupRequiredCode(ctx context.Context, phone string, issuedUserID int64) (string, error) {
|
||||
hash, err := randomHex(8)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if err := s.codes.Set(ctx, hash, store.PhoneCode{
|
||||
Phone: phone,
|
||||
Channel: codeChannelEmailSetupRequired,
|
||||
MaxAttempts: s.codeMaxAttempts,
|
||||
Version: store.PhoneCodeVersionCurrent,
|
||||
IssuedUserID: issuedUserID,
|
||||
Phone: phone,
|
||||
Channel: codeChannelEmailSetupRequired,
|
||||
MaxAttempts: s.codeMaxAttempts,
|
||||
}, s.codeTTL); err != nil {
|
||||
return "", fmt.Errorf("store code: %w", err)
|
||||
}
|
||||
if err := s.ensureIssuedOwnerAfterSet(ctx, hash, store.PhoneCode{Phone: phone, IssuedUserID: issuedUserID}); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return hash, nil
|
||||
}
|
||||
|
||||
func (s *Service) createEmailLoginCode(ctx context.Context, phone, email string) (string, error) {
|
||||
func (s *Service) createEmailLoginCode(ctx context.Context, phone, email string, issuedUserID int64) (string, error) {
|
||||
hash, err := randomHex(8)
|
||||
if err != nil {
|
||||
return "", err
|
||||
|
|
@ -327,22 +451,28 @@ func (s *Service) createEmailLoginCode(ctx context.Context, phone, email string)
|
|||
return "", err
|
||||
}
|
||||
rec := store.PhoneCode{
|
||||
Phone: phone,
|
||||
Code: code,
|
||||
Channel: codeChannelEmailLogin,
|
||||
Email: strings.TrimSpace(email),
|
||||
MaxAttempts: s.codeMaxAttempts,
|
||||
Version: store.PhoneCodeVersionCurrent,
|
||||
IssuedUserID: issuedUserID,
|
||||
Phone: phone,
|
||||
Code: code,
|
||||
Channel: codeChannelEmailLogin,
|
||||
Email: strings.TrimSpace(email),
|
||||
MaxAttempts: s.codeMaxAttempts,
|
||||
}
|
||||
if err := s.codes.Set(ctx, hash, rec, s.codeTTL); err != nil {
|
||||
return "", fmt.Errorf("store email code: %w", err)
|
||||
}
|
||||
if err := s.ensureIssuedOwnerAfterSet(ctx, hash, rec); err != nil {
|
||||
return "", err
|
||||
}
|
||||
if s.loginEmailSender == nil {
|
||||
_ = s.codes.Del(ctx, hash)
|
||||
return "", fmt.Errorf("login email sender is not configured")
|
||||
return "", s.rollbackUndeliveredCode(ctx, hash, fmt.Errorf("login email sender is not configured"))
|
||||
}
|
||||
if err := s.loginEmailSender.SendLoginCode(ctx, rec.Email, code, s.codeTTL); err != nil {
|
||||
_ = s.codes.Del(ctx, hash)
|
||||
return "", fmt.Errorf("send login email code: %w", err)
|
||||
return "", s.rollbackUndeliveredCode(ctx, hash, fmt.Errorf("send login email code: %w", err))
|
||||
}
|
||||
if err := s.ensureIssuedOwnerAfterSet(ctx, hash, rec); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return hash, nil
|
||||
}
|
||||
|
|
@ -396,20 +526,53 @@ func (s *Service) resendCode(ctx context.Context, authKeyID [8]byte, phone, phon
|
|||
if rec.Phone != phone {
|
||||
return "", ErrCodeInvalid
|
||||
}
|
||||
if rec.Purpose == store.PhoneCodePurposeChangePhone && (authKeyID == ([8]byte{}) || rec.AuthKeyID != authKeyID) {
|
||||
if rec.Purpose == store.PhoneCodePurposeChangePhone {
|
||||
if authKeyID == ([8]byte{}) || rec.AuthKeyID != authKeyID {
|
||||
return "", ErrCodeInvalid
|
||||
}
|
||||
consumed, ok, err := s.codes.ConsumeScoped(ctx, phoneCodeHash, rec.Scope())
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if !ok {
|
||||
return "", ErrCodeExpired
|
||||
}
|
||||
return s.recreateChangePhoneCode(ctx, consumed)
|
||||
}
|
||||
if rec.Version != store.PhoneCodeVersionCurrent {
|
||||
_, _, _ = s.codes.TakeLoginCode(ctx, phoneCodeHash, phone)
|
||||
return "", ErrCodeExpired
|
||||
}
|
||||
if matches, err := s.issuedOwnerMatches(ctx, phone, rec.IssuedUserID); err != nil {
|
||||
return "", err
|
||||
} else if !matches {
|
||||
s.invalidateLoginCodeDetached(ctx, phoneCodeHash, phone)
|
||||
return "", ErrCodeInvalid
|
||||
}
|
||||
_ = s.codes.Del(ctx, phoneCodeHash)
|
||||
if rec.Purpose == store.PhoneCodePurposeChangePhone {
|
||||
return s.recreateChangePhoneCode(ctx, rec)
|
||||
consumed, ok, err := s.codes.TakeLoginCode(ctx, phoneCodeHash, phone)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if !ok {
|
||||
return "", ErrCodeExpired
|
||||
}
|
||||
rec = consumed
|
||||
if matches, err := s.issuedOwnerMatches(ctx, phone, rec.IssuedUserID); err != nil {
|
||||
return "", err
|
||||
} else if !matches {
|
||||
s.invalidateLoginCodeDetached(ctx, phoneCodeHash, phone)
|
||||
return "", ErrCodeInvalid
|
||||
}
|
||||
if rec.Channel == codeChannelEmailLogin && strings.TrimSpace(rec.Email) != "" {
|
||||
return s.createEmailLoginCode(ctx, phone, rec.Email)
|
||||
return s.createEmailLoginCode(ctx, phone, rec.Email, rec.IssuedUserID)
|
||||
}
|
||||
if rec.Channel == codeChannelEmailSetupRequired {
|
||||
return s.createSetupRequiredCode(ctx, phone)
|
||||
return s.createSetupRequiredCode(ctx, phone, rec.IssuedUserID)
|
||||
}
|
||||
return s.SendCode(ctx, phone)
|
||||
if rec.Channel != codeChannelPhone {
|
||||
return "", ErrCodeInvalid
|
||||
}
|
||||
return s.createPhoneCode(ctx, phone, rec.IssuedUserID)
|
||||
}
|
||||
|
||||
func (s *Service) recreateChangePhoneCode(ctx context.Context, rec store.PhoneCode) (string, error) {
|
||||
|
|
@ -439,6 +602,75 @@ func (s *Service) CancelCodeForAuthKey(ctx context.Context, authKeyID [8]byte, p
|
|||
return s.cancelCode(ctx, authKeyID, phone, phoneCodeHash)
|
||||
}
|
||||
|
||||
// ConsumeLoginEmailReset authorizes auth.resetLoginEmail with the exact
|
||||
// email-login hash previously issued for this phone owner. Possession of only
|
||||
// a phone number is never sufficient to remove an authentication factor.
|
||||
func (s *Service) ConsumeLoginEmailReset(ctx context.Context, phone, phoneCodeHash string) (int64, error) {
|
||||
phone = normalizePhone(phone)
|
||||
rec, found, err := s.codes.Get(ctx, phoneCodeHash)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
if !found {
|
||||
return 0, ErrCodeExpired
|
||||
}
|
||||
if rec.Version != store.PhoneCodeVersionCurrent {
|
||||
_, _, _ = s.codes.TakeLoginCode(ctx, phoneCodeHash, phone)
|
||||
return 0, ErrCodeExpired
|
||||
}
|
||||
if rec.Purpose != "" || rec.Phone != phone || rec.Channel != codeChannelEmailLogin || rec.SignUpVerified {
|
||||
return 0, ErrCodeInvalid
|
||||
}
|
||||
before, beforeFound, err := s.currentPhoneOwner(ctx, phone)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
if !beforeFound || systemUserLoginForbidden(before) || rec.IssuedUserID == 0 || rec.IssuedUserID != before.ID {
|
||||
s.invalidateLoginCodeDetached(ctx, phoneCodeHash, phone)
|
||||
return 0, ErrCodeInvalid
|
||||
}
|
||||
consumed, consumedOK, err := s.codes.TakeLoginCode(ctx, phoneCodeHash, phone)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
if !consumedOK {
|
||||
return 0, ErrCodeExpired
|
||||
}
|
||||
if consumed.Channel != codeChannelEmailLogin || consumed.IssuedUserID != before.ID {
|
||||
return 0, ErrCodeInvalid
|
||||
}
|
||||
after, afterFound, err := s.currentPhoneOwner(ctx, phone)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
if !afterFound || after.ID != before.ID {
|
||||
return 0, ErrCodeInvalid
|
||||
}
|
||||
return before.ID, nil
|
||||
}
|
||||
|
||||
// SendPhoneCodeAfterLoginEmailReset issues the replacement app code only for
|
||||
// the exact user selected by ConsumeLoginEmailReset. It deliberately bypasses
|
||||
// SendCode's phone→owner reclassification so an A→B transfer cannot send B a
|
||||
// code and return that hash to A's reset flow.
|
||||
func (s *Service) SendPhoneCodeAfterLoginEmailReset(ctx context.Context, phone string, expectedUserID int64) (string, error) {
|
||||
phone = normalizePhone(phone)
|
||||
if !validPhone(phone) {
|
||||
return "", ErrPhoneNumberInvalid
|
||||
}
|
||||
if expectedUserID == 0 || systemLoginPhoneForbidden(phone) {
|
||||
return "", ErrCodeInvalid
|
||||
}
|
||||
owner, found, err := s.currentPhoneOwner(ctx, phone)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if !found || owner.ID != expectedUserID || systemUserLoginForbidden(owner) {
|
||||
return "", ErrCodeInvalid
|
||||
}
|
||||
return s.createPhoneCode(ctx, phone, expectedUserID)
|
||||
}
|
||||
|
||||
func (s *Service) cancelCode(ctx context.Context, authKeyID [8]byte, phone, phoneCodeHash string) error {
|
||||
phone = normalizePhone(phone)
|
||||
rec, found, err := s.codes.Get(ctx, phoneCodeHash)
|
||||
|
|
@ -451,10 +683,37 @@ func (s *Service) cancelCode(ctx context.Context, authKeyID [8]byte, phone, phon
|
|||
if rec.Phone != phone {
|
||||
return ErrCodeInvalid
|
||||
}
|
||||
if rec.Purpose == store.PhoneCodePurposeChangePhone && (authKeyID == ([8]byte{}) || rec.AuthKeyID != authKeyID) {
|
||||
if rec.Purpose == store.PhoneCodePurposeChangePhone {
|
||||
if authKeyID == ([8]byte{}) || rec.AuthKeyID != authKeyID {
|
||||
return ErrCodeInvalid
|
||||
}
|
||||
_, consumed, err := s.codes.ConsumeScoped(ctx, phoneCodeHash, rec.Scope())
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if !consumed {
|
||||
return ErrCodeExpired
|
||||
}
|
||||
return nil
|
||||
}
|
||||
if rec.Version != store.PhoneCodeVersionCurrent {
|
||||
_, _, _ = s.codes.TakeLoginCode(ctx, phoneCodeHash, phone)
|
||||
return ErrCodeExpired
|
||||
}
|
||||
if matches, err := s.issuedOwnerMatches(ctx, phone, rec.IssuedUserID); err != nil {
|
||||
return err
|
||||
} else if !matches {
|
||||
s.invalidateLoginCodeDetached(ctx, phoneCodeHash, phone)
|
||||
return ErrCodeInvalid
|
||||
}
|
||||
return s.codes.Del(ctx, phoneCodeHash)
|
||||
_, consumed, err := s.codes.TakeLoginCode(ctx, phoneCodeHash, phone)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if !consumed {
|
||||
return ErrCodeExpired
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// SignIn 校验验证码并尝试登录。
|
||||
|
|
@ -464,102 +723,164 @@ func (s *Service) SignIn(ctx context.Context, auth domain.Authorization, phone,
|
|||
if systemLoginPhoneForbidden(phone) {
|
||||
return domain.User{}, domain.Message{}, false, ErrSystemUserLoginForbidden
|
||||
}
|
||||
rec, found, err := s.codes.Get(ctx, phoneCodeHash)
|
||||
if err != nil {
|
||||
return domain.User{}, domain.Message{}, false, err
|
||||
}
|
||||
if !found {
|
||||
return domain.User{}, domain.Message{}, false, ErrCodeExpired
|
||||
}
|
||||
if rec.Phone != phone || rec.Channel == codeChannelEmailSetupRequired {
|
||||
return domain.User{}, domain.Message{}, false, ErrCodeInvalid
|
||||
}
|
||||
if rec.Channel == codeChannelEmailLogin {
|
||||
return domain.User{}, domain.Message{}, false, ErrCodeInvalid
|
||||
}
|
||||
if rec.Code != code {
|
||||
return domain.User{}, domain.Message{}, false, s.rejectCode(ctx, phoneCodeHash, rec, ErrCodeInvalid)
|
||||
}
|
||||
|
||||
existing, found, err := s.users.ByPhone(ctx, phone)
|
||||
if err != nil {
|
||||
return domain.User{}, domain.Message{}, false, err
|
||||
}
|
||||
if !found {
|
||||
return domain.User{}, domain.Message{}, true, nil // 验证码对、但需注册
|
||||
}
|
||||
return s.finishSignIn(ctx, auth, existing, phoneCodeHash, rec.Code)
|
||||
}
|
||||
|
||||
// SignInWithEmail 处理带 email_verification 的 auth.signIn:账号设置了登录邮箱后,新设备
|
||||
// 的验证码改投递到邮箱,客户端凭邮箱码(而非短信码)登录。开启真实登录邮箱后必须匹配
|
||||
// 随机邮箱码;未开启该特性时仅保留旧开发路径的任意非空兼容。仍校验 phone_code_hash
|
||||
// 有效、手机号匹配,并与短信登录共用 2FA 门控——即便走邮箱验证,开启了两步验证的账号
|
||||
// 同样会停在 SESSION_PASSWORD_NEEDED。
|
||||
func (s *Service) SignInWithEmail(ctx context.Context, auth domain.Authorization, phone, phoneCodeHash, code string) (domain.User, domain.Message, bool, error) {
|
||||
phone = normalizePhone(phone)
|
||||
if systemLoginPhoneForbidden(phone) {
|
||||
return domain.User{}, domain.Message{}, false, ErrSystemUserLoginForbidden
|
||||
}
|
||||
rec, found, err := s.codes.Get(ctx, phoneCodeHash)
|
||||
if err != nil {
|
||||
return domain.User{}, domain.Message{}, false, err
|
||||
}
|
||||
if !found {
|
||||
return domain.User{}, domain.Message{}, false, ErrCodeExpired
|
||||
}
|
||||
if rec.Phone != phone {
|
||||
return domain.User{}, domain.Message{}, false, ErrCodeInvalid
|
||||
}
|
||||
if rec.Channel != codeChannelEmailLogin {
|
||||
if s.loginEmailEnabled {
|
||||
return domain.User{}, domain.Message{}, false, ErrCodeInvalid
|
||||
}
|
||||
if strings.TrimSpace(code) == "" {
|
||||
return domain.User{}, domain.Message{}, false, ErrCodeInvalid
|
||||
}
|
||||
} else if rec.Code != strings.TrimSpace(code) {
|
||||
return domain.User{}, domain.Message{}, false, s.rejectCode(ctx, phoneCodeHash, rec, ErrCodeInvalid)
|
||||
}
|
||||
existing, found, err := s.users.ByPhone(ctx, phone)
|
||||
_, existing, found, err := s.verifyLoginCode(ctx, phone, phoneCodeHash, code, false)
|
||||
if err != nil {
|
||||
return domain.User{}, domain.Message{}, false, err
|
||||
}
|
||||
if !found {
|
||||
return domain.User{}, domain.Message{}, true, nil
|
||||
}
|
||||
return s.finishSignIn(ctx, auth, existing, phoneCodeHash, rec.Code)
|
||||
return s.finishSignIn(ctx, auth, existing)
|
||||
}
|
||||
|
||||
// SignInWithEmail 处理带 email_verification 的 auth.signIn:账号设置了登录邮箱后,新设备
|
||||
// 的验证码改投递到邮箱,客户端凭邮箱码(而非短信码)登录。开启真实登录邮箱后必须匹配
|
||||
// 随机邮箱码;未开启该特性时仍允许旧客户端把 phone channel 放进
|
||||
// email_verification,但必须精确匹配该 phone code,不能再接受任意非空值。
|
||||
// 两条路径共用 owner 绑定、原子尝试计数与 2FA 门控。
|
||||
func (s *Service) SignInWithEmail(ctx context.Context, auth domain.Authorization, phone, phoneCodeHash, code string) (domain.User, domain.Message, bool, error) {
|
||||
phone = normalizePhone(phone)
|
||||
if systemLoginPhoneForbidden(phone) {
|
||||
return domain.User{}, domain.Message{}, false, ErrSystemUserLoginForbidden
|
||||
}
|
||||
_, existing, found, err := s.verifyLoginCode(ctx, phone, phoneCodeHash, strings.TrimSpace(code), true)
|
||||
if err != nil {
|
||||
return domain.User{}, domain.Message{}, false, err
|
||||
}
|
||||
if !found {
|
||||
return domain.User{}, domain.Message{}, true, nil
|
||||
}
|
||||
return s.finishSignIn(ctx, auth, existing)
|
||||
}
|
||||
|
||||
// verifyLoginCode closes the login-code state transition around one atomic
|
||||
// CodeStore verification. The phone owner is read both before and after that
|
||||
// linearization point. A hash issued for an unregistered number therefore can
|
||||
// never authorize whichever account happens to acquire that number later.
|
||||
func (s *Service) verifyLoginCode(ctx context.Context, phone, phoneCodeHash, code string, emailPath bool) (store.PhoneCode, domain.User, bool, error) {
|
||||
rec, found, err := s.codes.Get(ctx, phoneCodeHash)
|
||||
if err != nil {
|
||||
return store.PhoneCode{}, domain.User{}, false, err
|
||||
}
|
||||
if !found {
|
||||
return store.PhoneCode{}, domain.User{}, false, ErrCodeExpired
|
||||
}
|
||||
if rec.Version != store.PhoneCodeVersionCurrent {
|
||||
_, _, _ = s.codes.TakeLoginCode(ctx, phoneCodeHash, phone)
|
||||
return store.PhoneCode{}, domain.User{}, false, ErrCodeExpired
|
||||
}
|
||||
if rec.Phone != phone || rec.Purpose != "" {
|
||||
return store.PhoneCode{}, domain.User{}, false, ErrCodeInvalid
|
||||
}
|
||||
channelAllowed := rec.Channel == codeChannelPhone && !emailPath
|
||||
if emailPath {
|
||||
channelAllowed = rec.Channel == codeChannelEmailLogin || (!s.loginEmailEnabled && rec.Channel == codeChannelPhone)
|
||||
}
|
||||
if !channelAllowed {
|
||||
return store.PhoneCode{}, domain.User{}, false, ErrCodeInvalid
|
||||
}
|
||||
|
||||
before, beforeFound, err := s.currentPhoneOwner(ctx, phone)
|
||||
if err != nil {
|
||||
return store.PhoneCode{}, domain.User{}, false, err
|
||||
}
|
||||
beforeUserID := int64(0)
|
||||
if beforeFound {
|
||||
if systemUserLoginForbidden(before) {
|
||||
return store.PhoneCode{}, domain.User{}, false, ErrSystemUserLoginForbidden
|
||||
}
|
||||
beforeUserID = before.ID
|
||||
}
|
||||
if rec.IssuedUserID != beforeUserID {
|
||||
s.invalidateLoginCodeDetached(ctx, phoneCodeHash, phone)
|
||||
return store.PhoneCode{}, domain.User{}, false, ErrCodeInvalid
|
||||
}
|
||||
// A verified sign-up marker may precede auth.signIn on the email-setup
|
||||
// path, and a normal signIn response can be lost and retried. The marker is
|
||||
// already the durable authorization fact; return signUpRequired
|
||||
// idempotently without asking CodeStore to verify it a second time.
|
||||
if rec.SignUpVerified {
|
||||
if beforeFound || rec.IssuedUserID != 0 || subtle.ConstantTimeCompare([]byte(rec.Code), []byte(code)) != 1 {
|
||||
return store.PhoneCode{}, domain.User{}, false, ErrCodeInvalid
|
||||
}
|
||||
after, afterFound, err := s.currentPhoneOwner(ctx, phone)
|
||||
if err != nil {
|
||||
return store.PhoneCode{}, domain.User{}, false, err
|
||||
}
|
||||
if afterFound || after.ID != 0 {
|
||||
s.invalidateLoginCodeDetached(ctx, phoneCodeHash, phone)
|
||||
return store.PhoneCode{}, domain.User{}, false, ErrCodeInvalid
|
||||
}
|
||||
return rec, domain.User{}, false, nil
|
||||
}
|
||||
|
||||
result, err := s.codes.VerifyLogin(ctx, phoneCodeHash, phone, code, !beforeFound, s.codeMaxAttempts)
|
||||
if err != nil {
|
||||
return store.PhoneCode{}, domain.User{}, false, err
|
||||
}
|
||||
after, afterFound, ownerErr := s.currentPhoneOwner(ctx, phone)
|
||||
if ownerErr != nil {
|
||||
return store.PhoneCode{}, domain.User{}, false, ownerErr
|
||||
}
|
||||
afterUserID := int64(0)
|
||||
if afterFound {
|
||||
afterUserID = after.ID
|
||||
}
|
||||
recordOwnerMismatch := result.Status != store.LoginCodeVerifyMissing && result.Record.IssuedUserID != rec.IssuedUserID
|
||||
if beforeUserID != afterUserID || recordOwnerMismatch {
|
||||
// keepForSignUp may have left a verified marker behind. Remove it on
|
||||
// owner drift so a later transfer-back cannot resurrect authorization.
|
||||
s.invalidateLoginCodeDetached(ctx, phoneCodeHash, phone)
|
||||
return store.PhoneCode{}, domain.User{}, false, ErrCodeInvalid
|
||||
}
|
||||
switch result.Status {
|
||||
case store.LoginCodeVerifyMissing:
|
||||
return store.PhoneCode{}, domain.User{}, false, ErrCodeExpired
|
||||
case store.LoginCodeVerifyInvalid:
|
||||
return store.PhoneCode{}, domain.User{}, false, ErrCodeInvalid
|
||||
case store.LoginCodeVerifyAccepted:
|
||||
if result.Record.Version != store.PhoneCodeVersionCurrent || result.Record.Phone != phone || result.Record.IssuedUserID != afterUserID {
|
||||
return store.PhoneCode{}, domain.User{}, false, ErrCodeInvalid
|
||||
}
|
||||
if afterFound && systemUserLoginForbidden(after) {
|
||||
return store.PhoneCode{}, domain.User{}, false, ErrSystemUserLoginForbidden
|
||||
}
|
||||
return result.Record, after, afterFound, nil
|
||||
default:
|
||||
return store.PhoneCode{}, domain.User{}, false, ErrCodeInvalid
|
||||
}
|
||||
}
|
||||
|
||||
// finishSignIn 是短信/邮箱两条登录路径在「验证码已通过、用户已存在」之后的共用收尾:
|
||||
// 处理 2FA password_pending 绑定、写登录消息、消费验证码。
|
||||
func (s *Service) finishSignIn(ctx context.Context, auth domain.Authorization, existing domain.User, phoneCodeHash, loginCode string) (domain.User, domain.Message, bool, error) {
|
||||
// 验证码已由 VerifyLogin 原子消费;这里只处理 2FA password_pending 绑定。已有账号的 app-code 消息已在
|
||||
// SendCode/ResendCode 返回前持久化与入 outbox,这里绝不能再创建或补发;
|
||||
// 否则未完成登录/2FA 的真实验证码反而不会及时到达旧设备。
|
||||
func (s *Service) finishSignIn(ctx context.Context, auth domain.Authorization, existing domain.User) (domain.User, domain.Message, bool, error) {
|
||||
if systemUserLoginForbidden(existing) {
|
||||
_ = s.codes.Del(ctx, phoneCodeHash)
|
||||
return domain.User{}, domain.Message{}, false, ErrSystemUserLoginForbidden
|
||||
}
|
||||
// 开启两步验证的账号:把授权标记为 password_pending 再写入,业务鉴权据此拒绝该 auth_key,
|
||||
// 直到 auth.checkPassword 通过。绝不能先以完全授权写入再返回 SESSION_PASSWORD_NEEDED,
|
||||
// 否则客户端忽略该错误即可直接调用业务 RPC 绕过两步验证。
|
||||
passwordNeeded := s.passwordNeeded(ctx, existing.ID)
|
||||
passwordNeeded, err := s.passwordNeeded(ctx, existing.ID)
|
||||
if err != nil {
|
||||
// Password state is part of the authentication decision. Treat store
|
||||
// failures as fail-closed and leave the auth key entirely unbound.
|
||||
return domain.User{}, domain.Message{}, false, err
|
||||
}
|
||||
auth.PasswordPending = passwordNeeded
|
||||
if err := s.bind(ctx, auth, existing.ID); err != nil {
|
||||
return domain.User{}, domain.Message{}, false, err
|
||||
}
|
||||
if passwordNeeded {
|
||||
_ = s.codes.Del(ctx, phoneCodeHash)
|
||||
return existing, domain.Message{}, false, domain.ErrSessionPasswordNeeded
|
||||
}
|
||||
loginMessage, err := s.recordLoginMessage(ctx, existing.ID, loginCode)
|
||||
if err != nil {
|
||||
return domain.User{}, domain.Message{}, false, err
|
||||
}
|
||||
_ = s.codes.Del(ctx, phoneCodeHash)
|
||||
return existing, loginMessage, false, nil
|
||||
return existing, domain.Message{}, false, nil
|
||||
}
|
||||
|
||||
// SignUp 在 SignIn 判定需注册后创建用户并绑定授权。
|
||||
// signUp 的 TL 请求不带验证码,这里校验 phone_code_hash 仍有效且手机号匹配。
|
||||
// signUp 的 TL 请求不带验证码,因此只消费由正确 SignIn/email setup 原子
|
||||
// 标记过的 hash。直接 SendCode→SignUp 永远不能创建账号。
|
||||
func (s *Service) SignUp(ctx context.Context, auth domain.Authorization, phone, phoneCodeHash, firstName, lastName string) (domain.User, domain.Message, error) {
|
||||
phone = normalizePhone(phone)
|
||||
if !validPhone(phone) {
|
||||
|
|
@ -580,15 +901,48 @@ func (s *Service) SignUp(ctx context.Context, auth domain.Authorization, phone,
|
|||
if !found {
|
||||
return domain.User{}, domain.Message{}, ErrCodeExpired
|
||||
}
|
||||
if rec.Phone != phone {
|
||||
if rec.Version != store.PhoneCodeVersionCurrent {
|
||||
_, _, _ = s.codes.ConsumeSignUpVerified(ctx, phoneCodeHash, phone)
|
||||
return domain.User{}, domain.Message{}, ErrCodeExpired
|
||||
}
|
||||
if rec.Phone != phone || rec.Purpose != "" {
|
||||
return domain.User{}, domain.Message{}, ErrCodeInvalid
|
||||
}
|
||||
if rec.Channel == codeChannelEmailSetupRequired {
|
||||
if !rec.SignUpVerified {
|
||||
return domain.User{}, domain.Message{}, ErrCodeInvalid
|
||||
}
|
||||
if rec.IssuedUserID != 0 {
|
||||
s.invalidateLoginCodeDetached(ctx, phoneCodeHash, phone)
|
||||
return domain.User{}, domain.Message{}, ErrCodeInvalid
|
||||
}
|
||||
if rec.Channel != codeChannelPhone && rec.Channel != codeChannelEmailLogin {
|
||||
return domain.User{}, domain.Message{}, ErrCodeInvalid
|
||||
}
|
||||
if s.loginEmailRequireSetup && !rec.VerifiedEmail && strings.TrimSpace(rec.PendingEmail) == "" {
|
||||
return domain.User{}, domain.Message{}, ErrCodeInvalid
|
||||
}
|
||||
if current, currentFound, err := s.currentPhoneOwner(ctx, phone); err != nil {
|
||||
return domain.User{}, domain.Message{}, err
|
||||
} else if currentFound || current.ID != 0 {
|
||||
s.invalidateLoginCodeDetached(ctx, phoneCodeHash, phone)
|
||||
return domain.User{}, domain.Message{}, ErrCodeInvalid
|
||||
}
|
||||
consumed, consumedOK, err := s.codes.ConsumeSignUpVerified(ctx, phoneCodeHash, phone)
|
||||
if err != nil {
|
||||
return domain.User{}, domain.Message{}, err
|
||||
}
|
||||
if !consumedOK {
|
||||
return domain.User{}, domain.Message{}, ErrCodeExpired
|
||||
}
|
||||
rec = consumed
|
||||
if rec.IssuedUserID != 0 || !rec.SignUpVerified || (rec.Channel != codeChannelPhone && rec.Channel != codeChannelEmailLogin) {
|
||||
return domain.User{}, domain.Message{}, ErrCodeInvalid
|
||||
}
|
||||
if current, currentFound, err := s.currentPhoneOwner(ctx, phone); err != nil {
|
||||
return domain.User{}, domain.Message{}, err
|
||||
} else if currentFound || current.ID != 0 {
|
||||
return domain.User{}, domain.Message{}, ErrCodeInvalid
|
||||
}
|
||||
|
||||
accessHash, err := randomInt64()
|
||||
if err != nil {
|
||||
|
|
@ -610,18 +964,22 @@ func (s *Service) SignUp(ctx context.Context, auth domain.Authorization, phone,
|
|||
return domain.User{}, domain.Message{}, err
|
||||
}
|
||||
if rec.VerifiedEmail && strings.TrimSpace(rec.PendingEmail) != "" && s.loginEmails != nil {
|
||||
if err := s.loginEmails.SetLoginEmailByPhone(ctx, phone, rec.PendingEmail); err != nil {
|
||||
if err := s.loginEmails.SetLoginEmail(ctx, u.ID, rec.PendingEmail); err != nil {
|
||||
return domain.User{}, domain.Message{}, err
|
||||
}
|
||||
}
|
||||
if err := s.bind(ctx, auth, u.ID); err != nil {
|
||||
return domain.User{}, domain.Message{}, err
|
||||
}
|
||||
loginMessage, err := s.recordLoginMessage(ctx, u.ID, rec.Code)
|
||||
if err != nil {
|
||||
return domain.User{}, domain.Message{}, err
|
||||
loginMessage := domain.Message{}
|
||||
// SMTP setup/login codes are secret factors, not 777000 app messages. Only
|
||||
// the normal phone/app-code registration path creates the bootstrap dialog.
|
||||
if rec.Channel == codeChannelPhone {
|
||||
loginMessage, err = s.recordLoginMessage(ctx, u.ID, rec.Code)
|
||||
if err != nil {
|
||||
return domain.User{}, domain.Message{}, err
|
||||
}
|
||||
}
|
||||
_ = s.codes.Del(ctx, phoneCodeHash)
|
||||
return u, loginMessage, nil
|
||||
}
|
||||
|
||||
|
|
@ -865,15 +1223,21 @@ func (s *Service) authorizationsByUserExcept(ctx context.Context, userID int64,
|
|||
|
||||
func (s *Service) bind(ctx context.Context, auth domain.Authorization, userID int64) error {
|
||||
auth.UserID = userID
|
||||
// Bind 是授权切换的持久化状态边界:生产 store 会先清同 auth key 的旧用户
|
||||
// update state,再原子建立新用户 baseline。RPC 层不得在 Bind 成功后清整个 key,
|
||||
// 否则会把刚建立的 retained-floor checkpoint 一并删除。
|
||||
return s.auths.Bind(ctx, auth)
|
||||
}
|
||||
|
||||
func (s *Service) passwordNeeded(ctx context.Context, userID int64) bool {
|
||||
func (s *Service) passwordNeeded(ctx context.Context, userID int64) (bool, error) {
|
||||
if s.passwords == nil {
|
||||
return false
|
||||
return false, nil
|
||||
}
|
||||
settings, found, err := s.passwords.GetByUser(ctx, userID)
|
||||
return err == nil && found && settings.HasPassword
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
return found && settings.HasPassword, nil
|
||||
}
|
||||
|
||||
const loginMessageTpl = `Login code: %s. Do not give this code to anyone, even if they say they are from Telegram!
|
||||
|
|
@ -1005,20 +1369,6 @@ func authKeyIDInt64(id [8]byte) int64 {
|
|||
return int64(binary.LittleEndian.Uint64(id[:]))
|
||||
}
|
||||
|
||||
func (s *Service) rejectCode(ctx context.Context, hash string, rec store.PhoneCode, ret error) error {
|
||||
rec.Attempts++
|
||||
max := rec.MaxAttempts
|
||||
if max <= 0 {
|
||||
max = s.codeMaxAttempts
|
||||
}
|
||||
if max > 0 && rec.Attempts >= max {
|
||||
_ = s.codes.Del(ctx, hash)
|
||||
return ret
|
||||
}
|
||||
_ = s.codes.Update(ctx, hash, rec)
|
||||
return ret
|
||||
}
|
||||
|
||||
func normalizePhone(phone string) string {
|
||||
return domain.NormalizePhone(phone)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -178,6 +178,14 @@ func TestPhoneCodeAcceptsTDesktopDigitsOnlySignIn(t *testing.T) {
|
|||
}
|
||||
}
|
||||
|
||||
func verifyCodeForSignUp(t *testing.T, svc *Service, phone, hash, code string) {
|
||||
t.Helper()
|
||||
got, msg, needSignUp, err := svc.SignIn(context.Background(), domain.Authorization{}, phone, hash, code)
|
||||
if err != nil || !needSignUp || got.ID != 0 || msg.ID != 0 {
|
||||
t.Fatalf("SignIn before SignUp user=%+v message=%+v needSignUp=%v err=%v, want empty/empty/true/nil", got, msg, needSignUp, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSystemUserPhoneCannotLoginOrSignUp(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
codes := memory.NewCodeStore()
|
||||
|
|
@ -257,6 +265,7 @@ func TestMultipleAuthKeysKeepSeparateUsers(t *testing.T) {
|
|||
if err != nil {
|
||||
t.Fatalf("SendCode user1: %v", err)
|
||||
}
|
||||
verifyCodeForSignUp(t, svc, "+15550005001", hash1, "12345")
|
||||
user1, _, err := svc.SignUp(ctx, domain.Authorization{AuthKeyID: key1}, "+15550005001", hash1, "One", "")
|
||||
if err != nil {
|
||||
t.Fatalf("SignUp user1: %v", err)
|
||||
|
|
@ -265,6 +274,7 @@ func TestMultipleAuthKeysKeepSeparateUsers(t *testing.T) {
|
|||
if err != nil {
|
||||
t.Fatalf("SendCode user2: %v", err)
|
||||
}
|
||||
verifyCodeForSignUp(t, svc, "+15550005002", hash2, "12345")
|
||||
user2, _, err := svc.SignUp(ctx, domain.Authorization{AuthKeyID: key2}, "+15550005002", hash2, "Two", "")
|
||||
if err != nil {
|
||||
t.Fatalf("SignUp user2: %v", err)
|
||||
|
|
@ -294,6 +304,7 @@ func TestLogOutThenSignInSameAuthKeySwitchesUser(t *testing.T) {
|
|||
if err != nil {
|
||||
t.Fatalf("SendCode user1: %v", err)
|
||||
}
|
||||
verifyCodeForSignUp(t, svc, "+15550006001", hash1, "12345")
|
||||
user1, _, err := svc.SignUp(ctx, domain.Authorization{AuthKeyID: key}, "+15550006001", hash1, "One", "")
|
||||
if err != nil {
|
||||
t.Fatalf("SignUp user1: %v", err)
|
||||
|
|
@ -312,6 +323,7 @@ func TestLogOutThenSignInSameAuthKeySwitchesUser(t *testing.T) {
|
|||
if err != nil {
|
||||
t.Fatalf("SendCode user2: %v", err)
|
||||
}
|
||||
verifyCodeForSignUp(t, svc, "+15550006002", hash2, "12345")
|
||||
user2, _, err := svc.SignUp(ctx, domain.Authorization{AuthKeyID: key}, "+15550006002", hash2, "Two", "")
|
||||
if err != nil {
|
||||
t.Fatalf("SignUp user2: %v", err)
|
||||
|
|
@ -337,6 +349,7 @@ func TestResetAuthorizationDeletesProtocolAuthKey(t *testing.T) {
|
|||
if err != nil {
|
||||
t.Fatalf("SendCode: %v", err)
|
||||
}
|
||||
verifyCodeForSignUp(t, svc, "+15550007001", hash, "12345")
|
||||
u, _, err := svc.SignUp(ctx, domain.Authorization{AuthKeyID: key}, "+15550007001", hash, "One", "")
|
||||
if err != nil {
|
||||
t.Fatalf("SignUp: %v", err)
|
||||
|
|
@ -375,6 +388,7 @@ func TestResetAuthorizationsDeletesOnlyRevokedProtocolAuthKeys(t *testing.T) {
|
|||
if err != nil {
|
||||
t.Fatalf("SendCode: %v", err)
|
||||
}
|
||||
verifyCodeForSignUp(t, svc, "+15550007002", hash, "12345")
|
||||
u, _, err := svc.SignUp(ctx, domain.Authorization{AuthKeyID: keep}, "+15550007002", hash, "Two", "")
|
||||
if err != nil {
|
||||
t.Fatalf("SignUp: %v", err)
|
||||
|
|
@ -399,16 +413,27 @@ func TestSignUpWritesOfficialLoginMessage(t *testing.T) {
|
|||
ctx := context.Background()
|
||||
dialogs := memory.NewDialogStore()
|
||||
messages := memory.NewMessageStore(dialogs)
|
||||
svc := NewService(memory.NewUserStore(), memory.NewAuthorizationStore(), memory.NewCodeStore(), nil, nil, "12345", WithLoginMessages(messages, dialogs))
|
||||
delivery := &captureLoginCodeDelivery{}
|
||||
svc := NewService(memory.NewUserStore(), memory.NewAuthorizationStore(), memory.NewCodeStore(), nil, nil, "12345",
|
||||
WithLoginMessages(messages, dialogs),
|
||||
WithLoginCodeDelivery(delivery),
|
||||
)
|
||||
|
||||
hash, err := svc.SendCode(ctx, "+15550004311")
|
||||
if err != nil {
|
||||
t.Fatalf("SendCode: %v", err)
|
||||
}
|
||||
if len(delivery.requests) != 0 {
|
||||
t.Fatalf("unregistered SendCode delivered before user exists: %+v", delivery.requests)
|
||||
}
|
||||
verifyCodeForSignUp(t, svc, "+15550004311", hash, "12345")
|
||||
u, msg, err := svc.SignUp(ctx, domain.Authorization{}, "+15550004311", hash, "Test", "User")
|
||||
if err != nil {
|
||||
t.Fatalf("SignUp: %v", err)
|
||||
}
|
||||
if len(delivery.requests) != 0 {
|
||||
t.Fatalf("SignUp unexpectedly used existing-account delivery: %+v", delivery.requests)
|
||||
}
|
||||
|
||||
list, err := dialogs.ListByUser(ctx, u.ID, domain.DialogFilter{Limit: 10})
|
||||
if err != nil {
|
||||
|
|
@ -431,17 +456,23 @@ func TestSignUpWritesOfficialLoginMessage(t *testing.T) {
|
|||
}
|
||||
}
|
||||
|
||||
func TestSignInLoginMessagePreservesOfficialDialogReadWatermark(t *testing.T) {
|
||||
func TestSendCodeLoginMessagePreservesOfficialDialogReadWatermark(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
dialogs := memory.NewDialogStore()
|
||||
messages := memory.NewMessageStore(dialogs)
|
||||
svc := NewService(memory.NewUserStore(), memory.NewAuthorizationStore(), memory.NewCodeStore(), nil, nil, "12345", WithLoginMessages(messages, dialogs))
|
||||
events := memory.NewUpdateEventStore()
|
||||
delivery := memory.NewLoginCodeDeliveryStore(messages, events)
|
||||
svc := NewService(memory.NewUserStore(), memory.NewAuthorizationStore(), memory.NewCodeStore(), nil, nil, "12345",
|
||||
WithLoginMessages(messages, dialogs),
|
||||
WithLoginCodeDelivery(delivery),
|
||||
)
|
||||
phone := "+15550004312"
|
||||
|
||||
hash, err := svc.SendCode(ctx, phone)
|
||||
if err != nil {
|
||||
t.Fatalf("SendCode signup: %v", err)
|
||||
}
|
||||
verifyCodeForSignUp(t, svc, phone, hash, "12345")
|
||||
u, first, err := svc.SignUp(ctx, domain.Authorization{}, phone, hash, "Test", "User")
|
||||
if err != nil {
|
||||
t.Fatalf("SignUp: %v", err)
|
||||
|
|
@ -452,15 +483,6 @@ func TestSignInLoginMessagePreservesOfficialDialogReadWatermark(t *testing.T) {
|
|||
} else if read.MaxID != first.ID || read.StillUnreadCount != 0 {
|
||||
t.Fatalf("read first login message = %+v, want max_id %d unread 0", read, first.ID)
|
||||
}
|
||||
|
||||
hash, err = svc.SendCode(ctx, phone)
|
||||
if err != nil {
|
||||
t.Fatalf("SendCode signin second: %v", err)
|
||||
}
|
||||
_, second, needSignUp, err := svc.SignIn(ctx, domain.Authorization{}, phone, hash, "12345")
|
||||
if err != nil || needSignUp {
|
||||
t.Fatalf("SignIn second needSignUp=%v err=%v", needSignUp, err)
|
||||
}
|
||||
assertOfficialDialog := func(wantTop, wantRead, wantUnread int) {
|
||||
t.Helper()
|
||||
list, err := dialogs.ListByUser(ctx, u.ID, domain.DialogFilter{Limit: 10})
|
||||
|
|
@ -475,23 +497,67 @@ func TestSignInLoginMessagePreservesOfficialDialogReadWatermark(t *testing.T) {
|
|||
t.Fatalf("dialog = %+v, want top=%d read=%d unread=%d", got, wantTop, wantRead, wantUnread)
|
||||
}
|
||||
}
|
||||
latestLoginMessage := func(wantCount int) domain.Message {
|
||||
t.Helper()
|
||||
history, err := messages.ListByUser(ctx, u.ID, domain.MessageFilter{
|
||||
HasPeer: true,
|
||||
Peer: peer,
|
||||
Limit: 10,
|
||||
})
|
||||
if err != nil || len(history.Messages) != wantCount {
|
||||
t.Fatalf("official history count=%d err=%v, want %d", len(history.Messages), err, wantCount)
|
||||
}
|
||||
latest := history.Messages[0]
|
||||
for _, msg := range history.Messages[1:] {
|
||||
if msg.ID > latest.ID {
|
||||
latest = msg
|
||||
}
|
||||
}
|
||||
return latest
|
||||
}
|
||||
|
||||
hash, err = svc.SendCode(ctx, phone)
|
||||
if err != nil {
|
||||
t.Fatalf("SendCode signin second: %v", err)
|
||||
}
|
||||
second := latestLoginMessage(2)
|
||||
// 核心时序:SendCode 返回时 message/dialog/unread 已提交,尚未 SignIn。
|
||||
assertOfficialDialog(second.ID, first.ID, 1)
|
||||
_, signInMessage, needSignUp, err := svc.SignIn(ctx, domain.Authorization{}, phone, hash, "12345")
|
||||
if err != nil || needSignUp {
|
||||
t.Fatalf("SignIn second needSignUp=%v err=%v", needSignUp, err)
|
||||
}
|
||||
if signInMessage.ID != 0 {
|
||||
t.Fatalf("SignIn second returned a late login message %+v", signInMessage)
|
||||
}
|
||||
assertOfficialDialog(second.ID, first.ID, 1)
|
||||
|
||||
hash, err = svc.SendCode(ctx, phone)
|
||||
if err != nil {
|
||||
t.Fatalf("SendCode signin third: %v", err)
|
||||
}
|
||||
_, third, needSignUp, err := svc.SignIn(ctx, domain.Authorization{}, phone, hash, "12345")
|
||||
third := latestLoginMessage(3)
|
||||
assertOfficialDialog(third.ID, first.ID, 2)
|
||||
_, signInMessage, needSignUp, err = svc.SignIn(ctx, domain.Authorization{}, phone, hash, "12345")
|
||||
if err != nil || needSignUp {
|
||||
t.Fatalf("SignIn third needSignUp=%v err=%v", needSignUp, err)
|
||||
}
|
||||
if signInMessage.ID != 0 {
|
||||
t.Fatalf("SignIn third returned a late login message %+v", signInMessage)
|
||||
}
|
||||
assertOfficialDialog(third.ID, first.ID, 2)
|
||||
}
|
||||
|
||||
func TestSignInExistingTwoFactorAccountNeedsPassword(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
passwords := memory.NewPasswordStore()
|
||||
svc := NewService(memory.NewUserStore(), memory.NewAuthorizationStore(), memory.NewCodeStore(), nil, nil, "12345", WithPasswords(passwords))
|
||||
dialogs := memory.NewDialogStore()
|
||||
messages := memory.NewMessageStore(dialogs)
|
||||
delivery := memory.NewLoginCodeDeliveryStore(messages, memory.NewUpdateEventStore())
|
||||
svc := NewService(memory.NewUserStore(), memory.NewAuthorizationStore(), memory.NewCodeStore(), nil, nil, "12345",
|
||||
WithPasswords(passwords),
|
||||
WithLoginCodeDelivery(delivery),
|
||||
)
|
||||
var key [8]byte
|
||||
key[0] = 7
|
||||
|
||||
|
|
@ -499,6 +565,7 @@ func TestSignInExistingTwoFactorAccountNeedsPassword(t *testing.T) {
|
|||
if err != nil {
|
||||
t.Fatalf("SendCode signup: %v", err)
|
||||
}
|
||||
verifyCodeForSignUp(t, svc, "+15550004312", hash, "12345")
|
||||
u, _, err := svc.SignUp(ctx, domain.Authorization{AuthKeyID: key}, "+15550004312", hash, "Two", "Factor")
|
||||
if err != nil {
|
||||
t.Fatalf("SignUp: %v", err)
|
||||
|
|
@ -514,13 +581,16 @@ func TestSignInExistingTwoFactorAccountNeedsPassword(t *testing.T) {
|
|||
if err != nil {
|
||||
t.Fatalf("SendCode signin: %v", err)
|
||||
}
|
||||
got, _, needSignUp, err := svc.SignIn(ctx, domain.Authorization{AuthKeyID: key}, "+15550004312", hash, "12345")
|
||||
got, signInMessage, 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)
|
||||
}
|
||||
if signInMessage.ID != 0 {
|
||||
t.Fatalf("2FA SignIn returned a late login message %+v", signInMessage)
|
||||
}
|
||||
// 两步验证未完成:业务鉴权(UserID)必须视为未登录,避免绕过 2FA。
|
||||
bound, found, err := svc.UserID(ctx, key)
|
||||
if err != nil || found || bound != 0 {
|
||||
|
|
@ -539,6 +609,10 @@ func TestSignInExistingTwoFactorAccountNeedsPassword(t *testing.T) {
|
|||
if err != nil || !found || bound != u.ID {
|
||||
t.Fatalf("UserID after 2FA passed = %d found=%v err=%v, want %d", bound, found, err, u.ID)
|
||||
}
|
||||
list, err := dialogs.ListByUser(ctx, u.ID, domain.DialogFilter{Limit: 10})
|
||||
if err != nil || len(list.Messages) != 1 {
|
||||
t.Fatalf("2FA login-code messages after password = %+v err=%v, want exactly the SendCode message", list.Messages, err)
|
||||
}
|
||||
}
|
||||
|
||||
func testAuthKey(seed byte) mtcrypto.AuthKey {
|
||||
|
|
|
|||
554
internal/app/auth/signup_state_test.go
Normal file
554
internal/app/auth/signup_state_test.go
Normal file
|
|
@ -0,0 +1,554 @@
|
|||
package auth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
accountapp "telesrv/internal/app/account"
|
||||
"telesrv/internal/domain"
|
||||
"telesrv/internal/store"
|
||||
"telesrv/internal/store/memory"
|
||||
)
|
||||
|
||||
func TestSignUpRequiresCorrectSignInAndConsumesMarkerOnce(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
users := memory.NewUserStore()
|
||||
codes := memory.NewCodeStore()
|
||||
svc := NewService(users, memory.NewAuthorizationStore(), codes, nil, nil, "12345")
|
||||
phone := "15550009301"
|
||||
|
||||
hash, err := svc.SendCode(ctx, phone)
|
||||
if err != nil {
|
||||
t.Fatalf("SendCode: %v", err)
|
||||
}
|
||||
if _, _, err := svc.SignUp(ctx, domain.Authorization{}, phone, hash, "Direct", "Bypass"); !errors.Is(err, ErrCodeInvalid) {
|
||||
t.Fatalf("direct SignUp err=%v, want ErrCodeInvalid", err)
|
||||
}
|
||||
if _, _, _, err := svc.SignIn(ctx, domain.Authorization{}, phone, hash, "00000"); !errors.Is(err, ErrCodeInvalid) {
|
||||
t.Fatalf("wrong SignIn err=%v, want ErrCodeInvalid", err)
|
||||
}
|
||||
if rec, found, err := codes.Get(ctx, hash); err != nil || !found || rec.SignUpVerified {
|
||||
t.Fatalf("wrong code marker=%v found=%v err=%v, want live/unverified", rec.SignUpVerified, found, err)
|
||||
}
|
||||
if _, _, err := svc.SignUp(ctx, domain.Authorization{}, phone, hash, "Wrong", "Code"); !errors.Is(err, ErrCodeInvalid) {
|
||||
t.Fatalf("SignUp after wrong code err=%v, want ErrCodeInvalid", err)
|
||||
}
|
||||
|
||||
verifyCodeForSignUp(t, svc, phone, hash, "12345")
|
||||
if rec, found, err := codes.Get(ctx, hash); err != nil || !found || !rec.SignUpVerified || rec.IssuedUserID != 0 {
|
||||
t.Fatalf("verified record=%+v found=%v err=%v", rec, found, err)
|
||||
}
|
||||
if _, msg, needSignUp, err := svc.SignIn(ctx, domain.Authorization{}, phone, hash, "12345"); err != nil || !needSignUp || msg.ID != 0 {
|
||||
t.Fatalf("idempotent SignIn needSignUp=%v message=%+v err=%v", needSignUp, msg, err)
|
||||
}
|
||||
u, _, err := svc.SignUp(ctx, domain.Authorization{}, phone, hash, "Verified", "User")
|
||||
if err != nil || u.Phone != phone {
|
||||
t.Fatalf("verified SignUp user=%+v err=%v", u, err)
|
||||
}
|
||||
if _, _, err := svc.SignUp(ctx, domain.Authorization{}, phone, hash, "Replay", "User"); !errors.Is(err, ErrCodeExpired) {
|
||||
t.Fatalf("replayed SignUp err=%v, want ErrCodeExpired", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestConcurrentSignUpConsumesVerifiedHashExactlyOnce(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
users := memory.NewUserStore()
|
||||
svc := NewService(users, memory.NewAuthorizationStore(), memory.NewCodeStore(), nil, nil, "12345")
|
||||
phone := "15550009302"
|
||||
hash, err := svc.SendCode(ctx, phone)
|
||||
if err != nil {
|
||||
t.Fatalf("SendCode: %v", err)
|
||||
}
|
||||
verifyCodeForSignUp(t, svc, phone, hash, "12345")
|
||||
|
||||
const workers = 16
|
||||
start := make(chan struct{})
|
||||
errs := make(chan error, workers)
|
||||
for i := 0; i < workers; i++ {
|
||||
go func(i int) {
|
||||
<-start
|
||||
var key [8]byte
|
||||
key[0] = byte(i + 1)
|
||||
_, _, err := svc.SignUp(ctx, domain.Authorization{AuthKeyID: key}, phone, hash, "Concurrent", "User")
|
||||
errs <- err
|
||||
}(i)
|
||||
}
|
||||
close(start)
|
||||
successes := 0
|
||||
for i := 0; i < workers; i++ {
|
||||
err := <-errs
|
||||
switch {
|
||||
case err == nil:
|
||||
successes++
|
||||
case errors.Is(err, ErrCodeExpired), errors.Is(err, ErrCodeInvalid):
|
||||
default:
|
||||
t.Fatalf("concurrent SignUp err=%v", err)
|
||||
}
|
||||
}
|
||||
if successes != 1 {
|
||||
t.Fatalf("successful SignUp calls=%d, want 1", successes)
|
||||
}
|
||||
}
|
||||
|
||||
type afterVerifyCodeStore struct {
|
||||
store.CodeStore
|
||||
once sync.Once
|
||||
afterVerify func()
|
||||
}
|
||||
|
||||
type failingPasswordStore struct {
|
||||
store.PasswordStore
|
||||
err error
|
||||
}
|
||||
|
||||
func (s *failingPasswordStore) GetByUser(context.Context, int64) (domain.PasswordSettings, bool, error) {
|
||||
return domain.PasswordSettings{}, false, s.err
|
||||
}
|
||||
|
||||
type switchablePhoneOwnerStore struct {
|
||||
store.UserStore
|
||||
mu sync.RWMutex
|
||||
phone string
|
||||
override bool
|
||||
owner domain.User
|
||||
found bool
|
||||
}
|
||||
|
||||
func (s *switchablePhoneOwnerStore) ByPhone(ctx context.Context, phone string) (domain.User, bool, error) {
|
||||
s.mu.RLock()
|
||||
if s.override && domain.NormalizePhone(phone) == s.phone {
|
||||
owner, found := s.owner, s.found
|
||||
s.mu.RUnlock()
|
||||
return owner, found, nil
|
||||
}
|
||||
s.mu.RUnlock()
|
||||
return s.UserStore.ByPhone(ctx, phone)
|
||||
}
|
||||
|
||||
func (s *switchablePhoneOwnerStore) setOwnerView(phone string, owner domain.User, found bool) {
|
||||
s.mu.Lock()
|
||||
s.phone = domain.NormalizePhone(phone)
|
||||
s.owner = owner
|
||||
s.found = found
|
||||
s.override = true
|
||||
s.mu.Unlock()
|
||||
}
|
||||
|
||||
func (s *switchablePhoneOwnerStore) resetOwnerView() {
|
||||
s.mu.Lock()
|
||||
s.override = false
|
||||
s.mu.Unlock()
|
||||
}
|
||||
|
||||
func (s *afterVerifyCodeStore) VerifyLogin(ctx context.Context, hash, phone, code string, keep bool, maxAttempts int) (store.LoginCodeVerifyResult, error) {
|
||||
result, err := s.CodeStore.VerifyLogin(ctx, hash, phone, code, keep, maxAttempts)
|
||||
if err == nil && result.Status == store.LoginCodeVerifyAccepted && s.afterVerify != nil {
|
||||
s.once.Do(s.afterVerify)
|
||||
}
|
||||
return result, err
|
||||
}
|
||||
|
||||
func TestOwnerTransferAcrossVerifyInvalidatesHashPermanently(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
users := memory.NewUserStore()
|
||||
baseCodes := memory.NewCodeStore()
|
||||
var createErr error
|
||||
codes := &afterVerifyCodeStore{CodeStore: baseCodes}
|
||||
svc := NewService(users, memory.NewAuthorizationStore(), codes, nil, nil, "12345")
|
||||
phone := "15550009303"
|
||||
hash, err := svc.SendCode(ctx, phone)
|
||||
if err != nil {
|
||||
t.Fatalf("SendCode: %v", err)
|
||||
}
|
||||
codes.afterVerify = func() {
|
||||
_, createErr = users.Create(ctx, domain.User{Phone: phone, FirstName: "NewOwner"})
|
||||
}
|
||||
if _, _, needSignUp, err := svc.SignIn(ctx, domain.Authorization{}, phone, hash, "12345"); !errors.Is(err, ErrCodeInvalid) || needSignUp {
|
||||
t.Fatalf("SignIn across owner transfer needSignUp=%v err=%v, want invalid", needSignUp, err)
|
||||
}
|
||||
if createErr != nil {
|
||||
t.Fatalf("create concurrent owner: %v", createErr)
|
||||
}
|
||||
if _, found, err := baseCodes.Get(ctx, hash); err != nil || found {
|
||||
t.Fatalf("owner-drift hash found=%v err=%v, want invalidated", found, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPasswordLookupFailureNeverCreatesOrChangesAuthorization(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
users := memory.NewUserStore()
|
||||
target, err := users.Create(ctx, domain.User{Phone: "15550009320", FirstName: "Target"})
|
||||
if err != nil {
|
||||
t.Fatalf("create target: %v", err)
|
||||
}
|
||||
previous, err := users.Create(ctx, domain.User{Phone: "15550009321", FirstName: "Previous"})
|
||||
if err != nil {
|
||||
t.Fatalf("create previous: %v", err)
|
||||
}
|
||||
authz := memory.NewAuthorizationStore()
|
||||
lookupErr := errors.New("password store unavailable")
|
||||
passwords := &failingPasswordStore{PasswordStore: memory.NewPasswordStore(), err: lookupErr}
|
||||
svc := NewService(users, authz, memory.NewCodeStore(), nil, nil, "12345",
|
||||
WithPasswords(passwords),
|
||||
WithLoginCodeDelivery(&captureLoginCodeDelivery{}),
|
||||
)
|
||||
|
||||
t.Run("unbound-key-remains-unbound", func(t *testing.T) {
|
||||
key := [8]byte{0xC1}
|
||||
hash, err := svc.SendCode(ctx, target.Phone)
|
||||
if err != nil {
|
||||
t.Fatalf("SendCode: %v", err)
|
||||
}
|
||||
if _, _, _, err := svc.SignIn(ctx, domain.Authorization{AuthKeyID: key}, target.Phone, hash, "12345"); !errors.Is(err, lookupErr) {
|
||||
t.Fatalf("SignIn err=%v, want password lookup failure", err)
|
||||
}
|
||||
if got, found, err := authz.ByAuthKey(ctx, key); err != nil || found {
|
||||
t.Fatalf("authorization=%+v found=%v err=%v, want absent", got, found, err)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("previous-binding-remains-unchanged", func(t *testing.T) {
|
||||
key := [8]byte{0xC2}
|
||||
original := domain.Authorization{AuthKeyID: key, UserID: previous.ID, Hash: 987654321}
|
||||
if err := authz.Bind(ctx, original); err != nil {
|
||||
t.Fatalf("bind previous authorization: %v", err)
|
||||
}
|
||||
hash, err := svc.SendCode(ctx, target.Phone)
|
||||
if err != nil {
|
||||
t.Fatalf("SendCode: %v", err)
|
||||
}
|
||||
if _, _, _, err := svc.SignIn(ctx, domain.Authorization{AuthKeyID: key}, target.Phone, hash, "12345"); !errors.Is(err, lookupErr) {
|
||||
t.Fatalf("SignIn err=%v, want password lookup failure", err)
|
||||
}
|
||||
got, found, err := authz.ByAuthKey(ctx, key)
|
||||
if err != nil || !found || got.UserID != previous.ID || got.Hash != original.Hash || got.PasswordPending != original.PasswordPending {
|
||||
t.Fatalf("authorization after failure=%+v found=%v err=%v, want unchanged %+v", got, found, err, original)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestOwnerTransferAwayAndBackCannotReviveLoginHash(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
t.Run("unregistered-signin", func(t *testing.T) {
|
||||
baseUsers := memory.NewUserStore()
|
||||
other, err := baseUsers.Create(ctx, domain.User{Phone: "15550009311", FirstName: "Other"})
|
||||
if err != nil {
|
||||
t.Fatalf("create other owner: %v", err)
|
||||
}
|
||||
users := &switchablePhoneOwnerStore{UserStore: baseUsers}
|
||||
codes := memory.NewCodeStore()
|
||||
svc := NewService(users, memory.NewAuthorizationStore(), codes, nil, nil, "12345")
|
||||
phone := "15550009310"
|
||||
hash, err := svc.SendCode(ctx, phone)
|
||||
if err != nil {
|
||||
t.Fatalf("SendCode: %v", err)
|
||||
}
|
||||
users.setOwnerView(phone, other, true)
|
||||
if _, _, _, err := svc.SignIn(ctx, domain.Authorization{}, phone, hash, "12345"); !errors.Is(err, ErrCodeInvalid) {
|
||||
t.Fatalf("SignIn after 0->B owner transfer err=%v, want invalid", err)
|
||||
}
|
||||
users.resetOwnerView()
|
||||
if _, _, _, err := svc.SignIn(ctx, domain.Authorization{}, phone, hash, "12345"); !errors.Is(err, ErrCodeExpired) {
|
||||
t.Fatalf("SignIn after 0->B->0 err=%v, want expired", err)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("existing-resend", func(t *testing.T) {
|
||||
baseUsers := memory.NewUserStore()
|
||||
ownerA, err := baseUsers.Create(ctx, domain.User{Phone: "15550009312", FirstName: "A"})
|
||||
if err != nil {
|
||||
t.Fatalf("create owner A: %v", err)
|
||||
}
|
||||
ownerB, err := baseUsers.Create(ctx, domain.User{Phone: "15550009313", FirstName: "B"})
|
||||
if err != nil {
|
||||
t.Fatalf("create owner B: %v", err)
|
||||
}
|
||||
users := &switchablePhoneOwnerStore{UserStore: baseUsers}
|
||||
codes := memory.NewCodeStore()
|
||||
svc := NewService(users, memory.NewAuthorizationStore(), codes, nil, nil, "12345", WithLoginCodeDelivery(&captureLoginCodeDelivery{}))
|
||||
hash, err := svc.SendCode(ctx, ownerA.Phone)
|
||||
if err != nil {
|
||||
t.Fatalf("SendCode: %v", err)
|
||||
}
|
||||
users.setOwnerView(ownerA.Phone, ownerB, true)
|
||||
if _, err := svc.ResendCode(ctx, ownerA.Phone, hash); !errors.Is(err, ErrCodeInvalid) {
|
||||
t.Fatalf("ResendCode after A->B err=%v, want invalid", err)
|
||||
}
|
||||
users.resetOwnerView()
|
||||
if _, _, _, err := svc.SignIn(ctx, domain.Authorization{}, ownerA.Phone, hash, "12345"); !errors.Is(err, ErrCodeExpired) {
|
||||
t.Fatalf("SignIn after A->B->A err=%v, want expired", err)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("existing-cancel", func(t *testing.T) {
|
||||
baseUsers := memory.NewUserStore()
|
||||
ownerA, err := baseUsers.Create(ctx, domain.User{Phone: "15550009314", FirstName: "A"})
|
||||
if err != nil {
|
||||
t.Fatalf("create owner A: %v", err)
|
||||
}
|
||||
ownerB, err := baseUsers.Create(ctx, domain.User{Phone: "15550009315", FirstName: "B"})
|
||||
if err != nil {
|
||||
t.Fatalf("create owner B: %v", err)
|
||||
}
|
||||
users := &switchablePhoneOwnerStore{UserStore: baseUsers}
|
||||
codes := memory.NewCodeStore()
|
||||
svc := NewService(users, memory.NewAuthorizationStore(), codes, nil, nil, "12345", WithLoginCodeDelivery(&captureLoginCodeDelivery{}))
|
||||
hash, err := svc.SendCode(ctx, ownerA.Phone)
|
||||
if err != nil {
|
||||
t.Fatalf("SendCode: %v", err)
|
||||
}
|
||||
users.setOwnerView(ownerA.Phone, ownerB, true)
|
||||
if err := svc.CancelCode(ctx, ownerA.Phone, hash); !errors.Is(err, ErrCodeInvalid) {
|
||||
t.Fatalf("CancelCode after A->B err=%v, want invalid", err)
|
||||
}
|
||||
users.resetOwnerView()
|
||||
if _, _, _, err := svc.SignIn(ctx, domain.Authorization{}, ownerA.Phone, hash, "12345"); !errors.Is(err, ErrCodeExpired) {
|
||||
t.Fatalf("SignIn after canceled A->B->A err=%v, want expired", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestEmailSetupVerificationAuthorizesSignUpWithout777000Message(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
users := memory.NewUserStore()
|
||||
codes := memory.NewCodeStore()
|
||||
passwords := memory.NewPasswordStore()
|
||||
sender := &testMailSender{}
|
||||
accountSvc := accountapp.NewService(passwords,
|
||||
accountapp.WithUsers(users),
|
||||
accountapp.WithLoginEmailVerification(codes, sender, time.Minute, 3, 6),
|
||||
)
|
||||
dialogs := memory.NewDialogStore()
|
||||
messages := memory.NewMessageStore(dialogs)
|
||||
authSvc := NewService(users, memory.NewAuthorizationStore(), codes, nil, nil, "12345",
|
||||
WithLoginMessages(messages, dialogs),
|
||||
WithLoginEmail(LoginEmailOptions{
|
||||
Enabled: true,
|
||||
RequireSetup: true,
|
||||
CodeLength: 6,
|
||||
Store: accountSvc,
|
||||
Sender: sender,
|
||||
}),
|
||||
)
|
||||
phone := "15550009304"
|
||||
hash, err := authSvc.SendCode(ctx, phone)
|
||||
if err != nil {
|
||||
t.Fatalf("SendCode: %v", err)
|
||||
}
|
||||
if _, _, err := authSvc.SignUp(ctx, domain.Authorization{}, phone, hash, "Direct", "Email"); !errors.Is(err, ErrCodeInvalid) {
|
||||
t.Fatalf("SignUp before email setup err=%v, want ErrCodeInvalid", err)
|
||||
}
|
||||
if _, _, err := accountSvc.SendLoginEmailCode(ctx, 0, phone, hash, "new@example.test", true); err != nil {
|
||||
t.Fatalf("SendLoginEmailCode: %v", err)
|
||||
}
|
||||
bad := wrongCode(sender.code, '0')
|
||||
if _, err := accountSvc.VerifyLoginEmail(ctx, 0, phone, hash, bad, true); !errors.Is(err, domain.ErrEmailCodeInvalid) {
|
||||
t.Fatalf("wrong VerifyLoginEmail err=%v, want ErrEmailCodeInvalid", err)
|
||||
}
|
||||
if rec, found, err := codes.Get(ctx, hash); err != nil || !found || rec.SignUpVerified {
|
||||
t.Fatalf("wrong SMTP code marker=%v found=%v err=%v", rec.SignUpVerified, found, err)
|
||||
}
|
||||
if _, err := accountSvc.VerifyLoginEmail(ctx, 0, phone, hash, sender.code, true); err != nil {
|
||||
t.Fatalf("VerifyLoginEmail: %v", err)
|
||||
}
|
||||
if rec, found, err := codes.Get(ctx, hash); err != nil || !found || !rec.SignUpVerified || rec.Channel != codeChannelEmailLogin {
|
||||
t.Fatalf("email-verified phone code=%+v found=%v err=%v", rec, found, err)
|
||||
}
|
||||
if _, msg, needSignUp, err := authSvc.SignInWithEmail(ctx, domain.Authorization{}, phone, hash, sender.code); err != nil || !needSignUp || msg.ID != 0 {
|
||||
t.Fatalf("SignInWithEmail after setup needSignUp=%v message=%+v err=%v", needSignUp, msg, err)
|
||||
}
|
||||
u, msg, err := authSvc.SignUp(ctx, domain.Authorization{}, phone, hash, "Email", "User")
|
||||
if err != nil {
|
||||
t.Fatalf("SignUp after email setup: %v", err)
|
||||
}
|
||||
if msg.ID != 0 || msg.Body != "" {
|
||||
t.Fatalf("email SignUp returned SMTP code message: %+v", msg)
|
||||
}
|
||||
list, err := dialogs.ListByUser(ctx, u.ID, domain.DialogFilter{Limit: 10})
|
||||
if err != nil {
|
||||
t.Fatalf("ListByUser: %v", err)
|
||||
}
|
||||
if len(list.Dialogs) != 0 || len(list.Messages) != 0 {
|
||||
t.Fatalf("email SignUp created 777000 bootstrap state: dialogs=%+v messages=%+v", list.Dialogs, list.Messages)
|
||||
}
|
||||
if email, found, err := accountSvc.LoginEmailByPhone(ctx, phone); err != nil || !found || email != "new@example.test" {
|
||||
t.Fatalf("LoginEmailByPhone email=%q found=%v err=%v", email, found, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestConsumeLoginEmailResetRequiresExactIssuedHash(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
baseUsers := memory.NewUserStore()
|
||||
owner, err := baseUsers.Create(ctx, domain.User{Phone: "15550009330", FirstName: "Owner"})
|
||||
if err != nil {
|
||||
t.Fatalf("create owner: %v", err)
|
||||
}
|
||||
other, err := baseUsers.Create(ctx, domain.User{Phone: "15550009331", FirstName: "Other"})
|
||||
if err != nil {
|
||||
t.Fatalf("create other: %v", err)
|
||||
}
|
||||
users := &switchablePhoneOwnerStore{UserStore: baseUsers}
|
||||
codes := memory.NewCodeStore()
|
||||
delivery := &captureLoginCodeDelivery{}
|
||||
svc := NewService(users, memory.NewAuthorizationStore(), codes, nil, nil, "12345", WithLoginCodeDelivery(delivery))
|
||||
seed := func(hash, channel string) {
|
||||
t.Helper()
|
||||
if err := codes.Set(ctx, hash, store.PhoneCode{
|
||||
Version: store.PhoneCodeVersionCurrent,
|
||||
IssuedUserID: owner.ID,
|
||||
Phone: owner.Phone,
|
||||
Code: "654321",
|
||||
Channel: channel,
|
||||
MaxAttempts: 5,
|
||||
}, time.Minute); err != nil {
|
||||
t.Fatalf("seed %s: %v", hash, err)
|
||||
}
|
||||
}
|
||||
|
||||
if _, err := svc.ConsumeLoginEmailReset(ctx, owner.Phone, "arbitrary-missing"); !errors.Is(err, ErrCodeExpired) {
|
||||
t.Fatalf("arbitrary hash err=%v, want expired", err)
|
||||
}
|
||||
seed("wrong-phone", codeChannelEmailLogin)
|
||||
if _, err := svc.ConsumeLoginEmailReset(ctx, other.Phone, "wrong-phone"); !errors.Is(err, ErrCodeInvalid) {
|
||||
t.Fatalf("wrong phone err=%v, want invalid", err)
|
||||
}
|
||||
if _, found, err := codes.Get(ctx, "wrong-phone"); err != nil || !found {
|
||||
t.Fatalf("wrong-phone probe destroyed valid hash found=%v err=%v", found, err)
|
||||
}
|
||||
seed("wrong-channel", codeChannelPhone)
|
||||
if _, err := svc.ConsumeLoginEmailReset(ctx, owner.Phone, "wrong-channel"); !errors.Is(err, ErrCodeInvalid) {
|
||||
t.Fatalf("wrong channel err=%v, want invalid", err)
|
||||
}
|
||||
|
||||
seed("owner-drift", codeChannelEmailLogin)
|
||||
users.setOwnerView(owner.Phone, other, true)
|
||||
if _, err := svc.ConsumeLoginEmailReset(ctx, owner.Phone, "owner-drift"); !errors.Is(err, ErrCodeInvalid) {
|
||||
t.Fatalf("A->B reset err=%v, want invalid", err)
|
||||
}
|
||||
users.resetOwnerView()
|
||||
if _, err := svc.ConsumeLoginEmailReset(ctx, owner.Phone, "owner-drift"); !errors.Is(err, ErrCodeExpired) {
|
||||
t.Fatalf("A->B->A reset err=%v, want expired", err)
|
||||
}
|
||||
|
||||
seed("successful-reset", codeChannelEmailLogin)
|
||||
resetUserID, err := svc.ConsumeLoginEmailReset(ctx, owner.Phone, "successful-reset")
|
||||
if err != nil || resetUserID != owner.ID {
|
||||
t.Fatalf("successful reset consume uid=%d err=%v", resetUserID, err)
|
||||
}
|
||||
replacementHash, err := svc.SendPhoneCodeAfterLoginEmailReset(ctx, owner.Phone, resetUserID)
|
||||
if err != nil || replacementHash == "" {
|
||||
t.Fatalf("replacement hash=%q err=%v", replacementHash, err)
|
||||
}
|
||||
if len(delivery.requests) != 1 || delivery.requests[0].UserID != owner.ID || delivery.requests[0].PhoneCodeHash != replacementHash {
|
||||
t.Fatalf("replacement delivery=%+v", delivery.requests)
|
||||
}
|
||||
if rec, found, err := codes.Get(ctx, replacementHash); err != nil || !found || rec.Version != store.PhoneCodeVersionCurrent || rec.IssuedUserID != owner.ID || rec.Channel != codeChannelPhone {
|
||||
t.Fatalf("replacement code=%+v found=%v err=%v", rec, found, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestConcurrentLoginEmailResetHasSingleConsumer(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
users := memory.NewUserStore()
|
||||
owner, err := users.Create(ctx, domain.User{Phone: "15550009332", FirstName: "Owner"})
|
||||
if err != nil {
|
||||
t.Fatalf("create owner: %v", err)
|
||||
}
|
||||
codes := memory.NewCodeStore()
|
||||
hash := "concurrent-email-reset"
|
||||
if err := codes.Set(ctx, hash, store.PhoneCode{
|
||||
Version: store.PhoneCodeVersionCurrent,
|
||||
IssuedUserID: owner.ID,
|
||||
Phone: owner.Phone,
|
||||
Code: "654321",
|
||||
Channel: codeChannelEmailLogin,
|
||||
MaxAttempts: 5,
|
||||
}, time.Minute); err != nil {
|
||||
t.Fatalf("seed code: %v", err)
|
||||
}
|
||||
svc := NewService(users, memory.NewAuthorizationStore(), codes, nil, nil, "12345")
|
||||
const workers = 24
|
||||
start := make(chan struct{})
|
||||
errs := make(chan error, workers)
|
||||
for i := 0; i < workers; i++ {
|
||||
go func() {
|
||||
<-start
|
||||
_, err := svc.ConsumeLoginEmailReset(ctx, owner.Phone, hash)
|
||||
errs <- err
|
||||
}()
|
||||
}
|
||||
close(start)
|
||||
successes := 0
|
||||
for i := 0; i < workers; i++ {
|
||||
err := <-errs
|
||||
if err == nil {
|
||||
successes++
|
||||
continue
|
||||
}
|
||||
if !errors.Is(err, ErrCodeExpired) {
|
||||
t.Fatalf("concurrent reset err=%v", err)
|
||||
}
|
||||
}
|
||||
if successes != 1 {
|
||||
t.Fatalf("successful reset consumers=%d, want 1", successes)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoginEmailResetLocksUserAcrossOwnerTransfer(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
baseUsers := memory.NewUserStore()
|
||||
ownerA, err := baseUsers.Create(ctx, domain.User{Phone: "15550009340", FirstName: "A"})
|
||||
if err != nil {
|
||||
t.Fatalf("create A: %v", err)
|
||||
}
|
||||
ownerB, err := baseUsers.Create(ctx, domain.User{Phone: "15550009341", FirstName: "B"})
|
||||
if err != nil {
|
||||
t.Fatalf("create B: %v", err)
|
||||
}
|
||||
users := &switchablePhoneOwnerStore{UserStore: baseUsers}
|
||||
passwords := memory.NewPasswordStore()
|
||||
accountSvc := accountapp.NewService(passwords, accountapp.WithUsers(users))
|
||||
if err := accountSvc.SetLoginEmail(ctx, ownerA.ID, "a@example.test"); err != nil {
|
||||
t.Fatalf("SetLoginEmail A: %v", err)
|
||||
}
|
||||
if err := accountSvc.SetLoginEmail(ctx, ownerB.ID, "b@example.test"); err != nil {
|
||||
t.Fatalf("SetLoginEmail B: %v", err)
|
||||
}
|
||||
codes := memory.NewCodeStore()
|
||||
hash := "locked-reset-user"
|
||||
if err := codes.Set(ctx, hash, store.PhoneCode{
|
||||
Version: store.PhoneCodeVersionCurrent,
|
||||
IssuedUserID: ownerA.ID,
|
||||
Phone: ownerA.Phone,
|
||||
Code: "654321",
|
||||
Channel: codeChannelEmailLogin,
|
||||
MaxAttempts: 5,
|
||||
}, time.Minute); err != nil {
|
||||
t.Fatalf("seed reset code: %v", err)
|
||||
}
|
||||
delivery := &captureLoginCodeDelivery{}
|
||||
authSvc := NewService(users, memory.NewAuthorizationStore(), codes, nil, nil, "12345", WithLoginCodeDelivery(delivery))
|
||||
resetUserID, err := authSvc.ConsumeLoginEmailReset(ctx, ownerA.Phone, hash)
|
||||
if err != nil || resetUserID != ownerA.ID {
|
||||
t.Fatalf("ConsumeLoginEmailReset uid=%d err=%v", resetUserID, err)
|
||||
}
|
||||
users.setOwnerView(ownerA.Phone, ownerB, true)
|
||||
if err := accountSvc.ClearLoginEmail(ctx, resetUserID); err != nil {
|
||||
t.Fatalf("ClearLoginEmail exact A: %v", err)
|
||||
}
|
||||
if _, err := authSvc.SendPhoneCodeAfterLoginEmailReset(ctx, ownerA.Phone, resetUserID); !errors.Is(err, ErrCodeInvalid) {
|
||||
t.Fatalf("SendPhoneCodeAfterLoginEmailReset across A->B err=%v, want invalid", err)
|
||||
}
|
||||
if _, found, err := accountSvc.LoginEmail(ctx, ownerA.ID); err != nil || found {
|
||||
t.Fatalf("A login email found=%v err=%v, want cleared", found, err)
|
||||
}
|
||||
if email, found, err := accountSvc.LoginEmail(ctx, ownerB.ID); err != nil || !found || email != "b@example.test" {
|
||||
t.Fatalf("B login email=%q found=%v err=%v, want unchanged", email, found, err)
|
||||
}
|
||||
if len(delivery.requests) != 0 {
|
||||
t.Fatalf("owner B received reset replacement code: %+v", delivery.requests)
|
||||
}
|
||||
}
|
||||
|
|
@ -1242,6 +1242,25 @@ func (s *Service) SendMessage(ctx context.Context, userID int64, req domain.Send
|
|||
if req.UserID != userID {
|
||||
return domain.SendChannelMessageResult{}, domain.ErrChannelInvalid
|
||||
}
|
||||
if req.RandomID != 0 && !req.IdempotencyPreflighted {
|
||||
fingerprint, err := store.ChannelSendFingerprint(req)
|
||||
if err != nil {
|
||||
return domain.SendChannelMessageResult{}, err
|
||||
}
|
||||
req.IdempotencyFingerprint = fingerprint
|
||||
if replayStore, ok := s.channels.(store.ChannelSendReplayStore); ok {
|
||||
replay, found, err := replayStore.LookupChannelSendReplay(ctx, domain.ChannelSendReplayRequest{
|
||||
ChannelID: req.ChannelID,
|
||||
SenderUserID: req.UserID,
|
||||
RandomID: req.RandomID,
|
||||
IdempotencyFingerprint: fingerprint,
|
||||
})
|
||||
if err != nil || found {
|
||||
return replay, err
|
||||
}
|
||||
req.IdempotencyPreflighted = true
|
||||
}
|
||||
}
|
||||
if err := s.ensureCanSend(ctx, req.UserID); err != nil {
|
||||
return domain.SendChannelMessageResult{}, err
|
||||
}
|
||||
|
|
@ -1253,6 +1272,25 @@ func (s *Service) SendMessage(ctx context.Context, userID int64, req domain.Send
|
|||
return s.channels.SendChannelMessage(ctx, req)
|
||||
}
|
||||
|
||||
// LookupChannelSendReplay reads a regular-channel or monoforum receipt without current
|
||||
// membership/send-gate checks. The authenticated caller remains bound to SenderUserID.
|
||||
func (s *Service) LookupChannelSendReplay(ctx context.Context, userID int64, req domain.ChannelSendReplayRequest) (domain.SendChannelMessageResult, bool, error) {
|
||||
if s == nil || s.channels == nil || userID == 0 {
|
||||
return domain.SendChannelMessageResult{}, false, nil
|
||||
}
|
||||
if req.SenderUserID == 0 {
|
||||
req.SenderUserID = userID
|
||||
}
|
||||
if req.SenderUserID != userID || req.ChannelID == 0 || req.RandomID == 0 {
|
||||
return domain.SendChannelMessageResult{}, false, domain.ErrChannelInvalid
|
||||
}
|
||||
replayStore, ok := s.channels.(store.ChannelSendReplayStore)
|
||||
if !ok {
|
||||
return domain.SendChannelMessageResult{}, false, nil
|
||||
}
|
||||
return replayStore.LookupChannelSendReplay(ctx, req)
|
||||
}
|
||||
|
||||
func (s *Service) ensureCanSend(ctx context.Context, userID int64) error {
|
||||
if s == nil || s.sendGate == nil || userID == 0 {
|
||||
return nil
|
||||
|
|
@ -1754,6 +1792,26 @@ func (s *Service) SendMonoforumMessage(ctx context.Context, req domain.SendMonof
|
|||
if s == nil || s.channels == nil || req.MonoforumID == 0 || req.SenderUserID == 0 || req.SavedPeer.ID == 0 {
|
||||
return domain.SendChannelMessageResult{}, domain.ErrChannelInvalid
|
||||
}
|
||||
if req.RandomID != 0 && !req.IdempotencyPreflighted {
|
||||
fingerprint, err := store.MonoforumSendFingerprint(req)
|
||||
if err != nil {
|
||||
return domain.SendChannelMessageResult{}, err
|
||||
}
|
||||
req.IdempotencyFingerprint = fingerprint
|
||||
if replayStore, ok := s.channels.(store.ChannelSendReplayStore); ok {
|
||||
replay, found, err := replayStore.LookupChannelSendReplay(ctx, domain.ChannelSendReplayRequest{
|
||||
ChannelID: req.MonoforumID,
|
||||
SenderUserID: req.SenderUserID,
|
||||
SavedPeer: req.SavedPeer,
|
||||
RandomID: req.RandomID,
|
||||
IdempotencyFingerprint: fingerprint,
|
||||
})
|
||||
if err != nil || found {
|
||||
return replay, err
|
||||
}
|
||||
req.IdempotencyPreflighted = true
|
||||
}
|
||||
}
|
||||
if err := s.ensureCanSend(ctx, req.SenderUserID); err != nil {
|
||||
return domain.SendChannelMessageResult{}, err
|
||||
}
|
||||
|
|
@ -2021,6 +2079,26 @@ func (s *Service) DirtyActiveChannelsForUser(ctx context.Context, userID int64,
|
|||
return s.channels.ListDirtyActiveChannelsForUser(ctx, userID, sinceDate, afterChannelID, limit)
|
||||
}
|
||||
|
||||
// MaxChannelPts returns the durable channel watermark used by the fan-out saturation recovery
|
||||
// sweep. It intentionally performs no viewer access check: target visibility is derived from the
|
||||
// process-local joined-membership index, while getChannelDifference performs authoritative access
|
||||
// validation when a client consumes the nudge.
|
||||
func (s *Service) MaxChannelPts(ctx context.Context, channelID int64) (int, error) {
|
||||
if s == nil || s.channels == nil || channelID == 0 {
|
||||
return 0, domain.ErrChannelInvalid
|
||||
}
|
||||
return s.channels.MaxChannelPts(ctx, channelID)
|
||||
}
|
||||
|
||||
// MaxChannelPtsBatch reloads a bounded recovery page in one store call. Missing ids are omitted:
|
||||
// they represent channels deleted after the process-local online-membership snapshot was taken.
|
||||
func (s *Service) MaxChannelPtsBatch(ctx context.Context, channelIDs []int64) (map[int64]int, error) {
|
||||
if s == nil || s.channels == nil {
|
||||
return nil, domain.ErrChannelInvalid
|
||||
}
|
||||
return s.channels.MaxChannelPtsBatch(ctx, channelIDs)
|
||||
}
|
||||
|
||||
// ActiveMemberIDs returns a bounded list for transient online fanout such as typing.
|
||||
func (s *Service) ActiveMemberIDs(ctx context.Context, userID, channelID int64, limit int) ([]int64, error) {
|
||||
if s == nil || s.channels == nil || userID == 0 || channelID == 0 {
|
||||
|
|
|
|||
|
|
@ -30,6 +30,46 @@ func TestServiceSendMessageHonorsSendPermissionGate(t *testing.T) {
|
|||
}
|
||||
}
|
||||
|
||||
func TestServiceChannelReplayPrecedesCurrentSendPermissionGate(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
channels := memory.NewChannelStore()
|
||||
created, err := channels.CreateChannel(ctx, domain.CreateChannelRequest{
|
||||
CreatorUserID: 1001,
|
||||
Title: "replay gate",
|
||||
Megagroup: true,
|
||||
Date: 1_700_000_000,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateChannel: %v", err)
|
||||
}
|
||||
req := domain.SendChannelMessageRequest{
|
||||
ChannelID: created.Channel.ID,
|
||||
RandomID: 92,
|
||||
Message: "committed before restriction",
|
||||
Date: 1_700_000_001,
|
||||
}
|
||||
allowed := NewService(channels)
|
||||
first, err := allowed.SendMessage(ctx, 1001, req)
|
||||
if err != nil {
|
||||
t.Fatalf("first SendMessage: %v", err)
|
||||
}
|
||||
|
||||
denied := NewService(channels, WithSendPermissionChecker(channelDenySendChecker{}))
|
||||
req.Date++
|
||||
replay, err := denied.SendMessage(ctx, 1001, req)
|
||||
if err != nil {
|
||||
t.Fatalf("replay through denied gate: %v", err)
|
||||
}
|
||||
if !replay.Duplicate || replay.Message.ID != first.Message.ID {
|
||||
t.Fatalf("replay = %+v, want committed duplicate %d", replay, first.Message.ID)
|
||||
}
|
||||
|
||||
req.Message = "different intent"
|
||||
if _, err := denied.SendMessage(ctx, 1001, req); !errors.Is(err, domain.ErrMessageRandomIDDuplicate) {
|
||||
t.Fatalf("conflicting replay err=%v, want ErrMessageRandomIDDuplicate before send gate", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestServiceSendMonoforumMessageHonorsSendPermissionGate(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
svc := NewService(memory.NewChannelStore(), WithSendPermissionChecker(channelDenySendChecker{}))
|
||||
|
|
@ -44,6 +84,52 @@ func TestServiceSendMonoforumMessageHonorsSendPermissionGate(t *testing.T) {
|
|||
}
|
||||
}
|
||||
|
||||
func TestServiceMonoforumReplayPrecedesCurrentSendPermissionGate(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
channels := memory.NewChannelStore()
|
||||
parent, err := channels.CreateChannel(ctx, domain.CreateChannelRequest{
|
||||
CreatorUserID: 1001,
|
||||
Title: "direct messages",
|
||||
Broadcast: true,
|
||||
Date: 1_700_000_010,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateChannel: %v", err)
|
||||
}
|
||||
enabled, err := channels.SetPaidMessagesPrice(ctx, 1001, parent.Channel.ID, 0, true)
|
||||
if err != nil {
|
||||
t.Fatalf("SetPaidMessagesPrice: %v", err)
|
||||
}
|
||||
req := domain.SendMonoforumMessageRequest{
|
||||
MonoforumID: enabled.Channel.LinkedMonoforumID,
|
||||
SenderUserID: 1002,
|
||||
SavedPeer: domain.Peer{Type: domain.PeerTypeUser, ID: 1002},
|
||||
RandomID: 93,
|
||||
Message: "committed direct message",
|
||||
Date: 1_700_000_011,
|
||||
}
|
||||
allowed := NewService(channels)
|
||||
first, err := allowed.SendMonoforumMessage(ctx, req)
|
||||
if err != nil {
|
||||
t.Fatalf("first SendMonoforumMessage: %v", err)
|
||||
}
|
||||
|
||||
denied := NewService(channels, WithSendPermissionChecker(channelDenySendChecker{}))
|
||||
req.Date++
|
||||
replay, err := denied.SendMonoforumMessage(ctx, req)
|
||||
if err != nil {
|
||||
t.Fatalf("monoforum replay through denied gate: %v", err)
|
||||
}
|
||||
if !replay.Duplicate || replay.Message.ID != first.Message.ID {
|
||||
t.Fatalf("monoforum replay = %+v, want committed duplicate %d", replay, first.Message.ID)
|
||||
}
|
||||
|
||||
req.Message = "different intent"
|
||||
if _, err := denied.SendMonoforumMessage(ctx, req); !errors.Is(err, domain.ErrMessageRandomIDDuplicate) {
|
||||
t.Fatalf("conflicting monoforum replay err=%v, want ErrMessageRandomIDDuplicate before send gate", err)
|
||||
}
|
||||
}
|
||||
|
||||
type channelDenySendChecker struct{}
|
||||
|
||||
func (channelDenySendChecker) CanSendMessages(context.Context, int64) error {
|
||||
|
|
@ -813,7 +899,8 @@ func TestCreateChatCreatesMegagroupWithChannelPts(t *testing.T) {
|
|||
duplicate, err := service.SendMessage(ctx, 1001, domain.SendChannelMessageRequest{
|
||||
ChannelID: created.Channel.ID,
|
||||
RandomID: 99,
|
||||
Message: "hello again",
|
||||
Message: "hello",
|
||||
ViaBotID: 1003,
|
||||
Date: 12,
|
||||
})
|
||||
if err != nil {
|
||||
|
|
@ -2032,12 +2119,12 @@ func TestChannelEditDeleteAndLocalClearUseChannelPts(t *testing.T) {
|
|||
if edited.Event.Type != domain.ChannelUpdateEditMessage || edited.Event.Pts != 4 || edited.Event.PtsCount != 1 {
|
||||
t.Fatalf("edit event = %+v, want channel edit pts=4 count=1", edited.Event)
|
||||
}
|
||||
duplicate, err := service.SendMessage(ctx, 1002, domain.SendChannelMessageRequest{ChannelID: created.Channel.ID, RandomID: 2, Message: "two retry", Date: 13})
|
||||
duplicate, err := service.SendMessage(ctx, 1002, domain.SendChannelMessageRequest{ChannelID: created.Channel.ID, RandomID: 2, Message: "two", Date: 13})
|
||||
if err != nil {
|
||||
t.Fatalf("duplicate SendMessage after edit: %v", err)
|
||||
}
|
||||
if !duplicate.Duplicate || duplicate.Event.Type != domain.ChannelUpdateNewMessage || duplicate.Message.Body != "two" || duplicate.Event.Message.Body != "two" {
|
||||
t.Fatalf("duplicate after edit = %+v, want original new-message snapshot", duplicate)
|
||||
if !duplicate.Duplicate || duplicate.Event.Type != domain.ChannelUpdateNewMessage || duplicate.Message.Body != "two edited" || duplicate.Event.Message.Body != "two edited" {
|
||||
t.Fatalf("duplicate after edit = %+v, want current message in new-message replay", duplicate)
|
||||
}
|
||||
|
||||
deleted, err := service.DeleteMessages(ctx, 1001, domain.DeleteChannelMessagesRequest{
|
||||
|
|
|
|||
|
|
@ -8,10 +8,10 @@ import (
|
|||
"fmt"
|
||||
"image"
|
||||
"image/color"
|
||||
"io"
|
||||
stddraw "image/draw"
|
||||
_ "image/jpeg" // 注册 jpeg DecodeConfig,用于读取上传头像/图片尺寸
|
||||
"image/png"
|
||||
"io"
|
||||
"math"
|
||||
"strings"
|
||||
"time"
|
||||
|
|
@ -48,14 +48,40 @@ func (s *Service) UploadProfilePhotoKind(ctx context.Context, ownerType domain.P
|
|||
|
||||
// CreatePhotoFromUpload 把已上传文件组装成 Photo(不绑定 profile_photos),用于频道头像 / 图片消息。
|
||||
func (s *Service) CreatePhotoFromUpload(ctx context.Context, file domain.UploadedFileRef) (domain.Photo, error) {
|
||||
data, err := s.assembleUpload(ctx, file.OwnerUserID, file.FileID, file.Parts)
|
||||
intentHash, err := uploadedMediaIntentHash(domain.UploadedMediaPhoto, file, nil)
|
||||
if err != nil {
|
||||
return domain.Photo{}, err
|
||||
}
|
||||
if photo, found, err := s.replayUploadedPhoto(ctx, file, intentHash); err != nil || found {
|
||||
return photo, err
|
||||
}
|
||||
data, err := s.readUploadBytes(ctx, file.OwnerUserID, file.FileID, file.Parts)
|
||||
if err != nil {
|
||||
return domain.Photo{}, err
|
||||
}
|
||||
if len(data) == 0 {
|
||||
return domain.Photo{}, domain.ErrPhotoInvalid
|
||||
}
|
||||
return s.createPhoto(ctx, data, photoSizeSpecsForMessage(data))
|
||||
photo, err := s.createPhoto(ctx, data, photoSizeSpecsForMessage(data))
|
||||
if err != nil {
|
||||
return domain.Photo{}, err
|
||||
}
|
||||
receipt, err := s.commitUploadedMediaReceipt(ctx, file, domain.UploadedMediaPhoto, intentHash, photo.ID)
|
||||
if err != nil {
|
||||
return domain.Photo{}, err
|
||||
}
|
||||
if receipt.MediaID != photo.ID {
|
||||
winner, found, err := s.media.GetPhoto(ctx, receipt.MediaID)
|
||||
if err != nil {
|
||||
return domain.Photo{}, err
|
||||
}
|
||||
if !found {
|
||||
return domain.Photo{}, fmt.Errorf("concurrent upload receipt references missing photo %d", receipt.MediaID)
|
||||
}
|
||||
photo = winner
|
||||
}
|
||||
s.cleanupMaterializedUpload(ctx, file, "photo materialized")
|
||||
return photo, nil
|
||||
}
|
||||
|
||||
// CreatePhotoFromBytes stores already-fetched image bytes as a message Photo.
|
||||
|
|
@ -202,6 +228,13 @@ func validateAvatarMarkupSize(size domain.PhotoSize) error {
|
|||
|
||||
// CreateDocumentFromUpload 把已上传文件组装成 Document(文件/视频/音频/gif/贴纸消息),落 blob + documents。
|
||||
func (s *Service) CreateDocumentFromUpload(ctx context.Context, file domain.UploadedFileRef, spec domain.DocumentSpec) (domain.Document, error) {
|
||||
intentHash, err := uploadedMediaIntentHash(domain.UploadedMediaDocument, file, &spec)
|
||||
if err != nil {
|
||||
return domain.Document{}, err
|
||||
}
|
||||
if doc, found, err := s.replayUploadedDocument(ctx, file, intentHash); err != nil || found {
|
||||
return doc, err
|
||||
}
|
||||
body, err := s.assembleUploadBlob(ctx, file.OwnerUserID, file.FileID, file.Parts)
|
||||
if err != nil {
|
||||
return domain.Document{}, err
|
||||
|
|
@ -234,11 +267,13 @@ func (s *Service) CreateDocumentFromUpload(ctx context.Context, file domain.Uplo
|
|||
DCID: s.dc,
|
||||
Attributes: spec.Attributes,
|
||||
}
|
||||
thumbMaterialized := false
|
||||
if spec.Thumb != nil {
|
||||
thumbData, err := s.assembleUpload(ctx, spec.Thumb.OwnerUserID, spec.Thumb.FileID, spec.Thumb.Parts)
|
||||
thumbData, err := s.readUploadBytes(ctx, spec.Thumb.OwnerUserID, spec.Thumb.FileID, spec.Thumb.Parts)
|
||||
if err == nil && len(thumbData) > 0 {
|
||||
if thumb, err := s.putDocumentThumb(ctx, docID, thumbData); err == nil {
|
||||
doc.Thumbs = []domain.PhotoSize{thumb}
|
||||
thumbMaterialized = true
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -250,12 +285,23 @@ func (s *Service) CreateDocumentFromUpload(ctx context.Context, file domain.Uplo
|
|||
if err := s.media.PutDocument(ctx, doc); err != nil {
|
||||
return domain.Document{}, err
|
||||
}
|
||||
if err := s.cleanupUploadParts(ctx, file.OwnerUserID, file.FileID); err != nil {
|
||||
s.log.Warn("cleanup assembled document upload parts failed",
|
||||
zap.Int64("owner_user_id", file.OwnerUserID),
|
||||
zap.Int64("file_id", file.FileID),
|
||||
zap.Int64("document_id", docID),
|
||||
zap.Error(err))
|
||||
receipt, err := s.commitUploadedMediaReceipt(ctx, file, domain.UploadedMediaDocument, intentHash, doc.ID)
|
||||
if err != nil {
|
||||
return domain.Document{}, err
|
||||
}
|
||||
if receipt.MediaID != doc.ID {
|
||||
winner, found, err := s.media.GetDocument(ctx, receipt.MediaID)
|
||||
if err != nil {
|
||||
return domain.Document{}, err
|
||||
}
|
||||
if !found {
|
||||
return domain.Document{}, fmt.Errorf("concurrent upload receipt references missing document %d", receipt.MediaID)
|
||||
}
|
||||
doc = winner
|
||||
}
|
||||
s.cleanupMaterializedUpload(ctx, file, "document materialized")
|
||||
if spec.Thumb != nil && thumbMaterialized {
|
||||
s.cleanupMaterializedUpload(ctx, *spec.Thumb, "document thumbnail materialized")
|
||||
}
|
||||
return doc, nil
|
||||
}
|
||||
|
|
@ -277,6 +323,7 @@ var faststartVideoMimes = map[string]bool{
|
|||
// 此时只发生几次 16 字节读,不读整段媒体。
|
||||
// 2. 仅 moov 在末尾时才重写;且优先走流式(仅 ftyp+moov 进内存,mdat 大块分块流式拼接),
|
||||
// 不把整段视频 2× 驻留内存。moov 非末尾的罕见排布回退到全量重排。
|
||||
//
|
||||
// 任何不适用/失败都返回原 body,绝不让上传失败或损坏数据。
|
||||
func (s *Service) maybeFaststartVideoBlob(ctx context.Context, mimeType string, body assembledUploadBlob) assembledUploadBlob {
|
||||
if !faststartVideoMimes[strings.ToLower(strings.TrimSpace(mimeType))] {
|
||||
|
|
|
|||
|
|
@ -26,6 +26,7 @@ type fakeMediaStore struct {
|
|||
parts map[string][]domain.UploadPart
|
||||
webPages map[int64]domain.MessageWebPage
|
||||
seedState map[string]string
|
||||
receipts map[string]domain.UploadedMediaReceipt
|
||||
}
|
||||
|
||||
func newFakeMediaStore() *fakeMediaStore {
|
||||
|
|
@ -36,9 +37,35 @@ func newFakeMediaStore() *fakeMediaStore {
|
|||
sets: map[int64]domain.StickerSet{},
|
||||
parts: map[string][]domain.UploadPart{},
|
||||
seedState: map[string]string{},
|
||||
receipts: map[string]domain.UploadedMediaReceipt{},
|
||||
}
|
||||
}
|
||||
|
||||
func fakeUploadReceiptKey(ownerUserID, fileID int64) string {
|
||||
return fmt.Sprintf("%d/%d", ownerUserID, fileID)
|
||||
}
|
||||
|
||||
func (f *fakeMediaStore) GetUploadedMediaReceipt(_ context.Context, ownerUserID, fileID int64) (domain.UploadedMediaReceipt, bool, error) {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
receipt, ok := f.receipts[fakeUploadReceiptKey(ownerUserID, fileID)]
|
||||
receipt.IntentHash = append([]byte(nil), receipt.IntentHash...)
|
||||
return receipt, ok, nil
|
||||
}
|
||||
|
||||
func (f *fakeMediaStore) PutUploadedMediaReceipt(_ context.Context, receipt domain.UploadedMediaReceipt) (domain.UploadedMediaReceipt, bool, error) {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
key := fakeUploadReceiptKey(receipt.OwnerUserID, receipt.FileID)
|
||||
if stored, ok := f.receipts[key]; ok {
|
||||
stored.IntentHash = append([]byte(nil), stored.IntentHash...)
|
||||
return stored, false, nil
|
||||
}
|
||||
receipt.IntentHash = append([]byte(nil), receipt.IntentHash...)
|
||||
f.receipts[key] = receipt
|
||||
return receipt, true, nil
|
||||
}
|
||||
|
||||
func (f *fakeMediaStore) SaveFilePart(_ context.Context, part domain.UploadPart) error {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
|
|
|
|||
|
|
@ -514,6 +514,20 @@ func orderDocuments(docs []domain.Document, ids []int64) []domain.Document {
|
|||
// assembleUpload 把已上传分片按 part 顺序拼成完整字节,并清理分片。
|
||||
// expectedParts>0 时校验分片连续且齐全。
|
||||
func (s *Service) assembleUpload(ctx context.Context, ownerUserID, fileID int64, expectedParts int) ([]byte, error) {
|
||||
buf, err := s.readUploadBytes(ctx, ownerUserID, fileID, expectedParts)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := s.cleanupUploadParts(ctx, ownerUserID, fileID); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return buf, nil
|
||||
}
|
||||
|
||||
// readUploadBytes validates and reads all parts without consuming them. Message-media
|
||||
// materialization persists an upload receipt before cleanup; callers that do not need replayability
|
||||
// continue to use assembleUpload.
|
||||
func (s *Service) readUploadBytes(ctx context.Context, ownerUserID, fileID int64, expectedParts int) ([]byte, error) {
|
||||
parts, _, err := s.loadAndValidateUploadParts(ctx, ownerUserID, fileID, expectedParts)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
|
|
@ -532,9 +546,6 @@ func (s *Service) assembleUpload(ctx context.Context, ownerUserID, fileID int64,
|
|||
}
|
||||
buf = append(buf, data...)
|
||||
}
|
||||
if err := s.cleanupUploadParts(ctx, ownerUserID, fileID); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return buf, nil
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -124,6 +124,51 @@ func TestCreateDocumentFromUploadStreamsBodyAndCleansParts(t *testing.T) {
|
|||
if string(body) != strings.Join(parts, "") {
|
||||
t.Fatalf("body blob mismatch")
|
||||
}
|
||||
replayed, err := svc.CreateDocumentFromUpload(ctx,
|
||||
domain.UploadedFileRef{OwnerUserID: 10, FileID: 200, Parts: len(parts), Name: "large.bin", Big: true},
|
||||
domain.DocumentSpec{MimeType: "application/octet-stream"},
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("replay CreateDocumentFromUpload after part cleanup: %v", err)
|
||||
}
|
||||
if replayed.ID != doc.ID || replayed.AccessHash != doc.AccessHash {
|
||||
t.Fatalf("replayed document = %d/%d, want original %d/%d", replayed.ID, replayed.AccessHash, doc.ID, doc.AccessHash)
|
||||
}
|
||||
if _, err := svc.CreateDocumentFromUpload(ctx,
|
||||
domain.UploadedFileRef{OwnerUserID: 10, FileID: 200, Parts: len(parts), Name: "large.bin", Big: true},
|
||||
domain.DocumentSpec{MimeType: "text/plain"},
|
||||
); !errors.Is(err, domain.ErrFilePartsInvalid) {
|
||||
t.Fatalf("changed materialization intent err = %v, want ErrFilePartsInvalid", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreatePhotoFromUploadReceiptReplaysAfterPartCleanup(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
media := newFakeMediaStore()
|
||||
svc, _ := newUploadPartTestService(t, media, domain.UploadPartQuota{})
|
||||
file := domain.UploadedFileRef{OwnerUserID: 10, FileID: 201, Parts: 1, Name: "photo.jpg"}
|
||||
if _, err := svc.SaveFilePart(ctx, file.OwnerUserID, file.FileID, 0, []byte("image-bytes")); err != nil {
|
||||
t.Fatalf("SaveFilePart: %v", err)
|
||||
}
|
||||
first, err := svc.CreatePhotoFromUpload(ctx, file)
|
||||
if err != nil {
|
||||
t.Fatalf("CreatePhotoFromUpload: %v", err)
|
||||
}
|
||||
if remaining, err := media.LoadFileParts(ctx, file.OwnerUserID, file.FileID); err != nil || len(remaining) != 0 {
|
||||
t.Fatalf("upload parts after photo materialization = %+v err=%v", remaining, err)
|
||||
}
|
||||
replayed, err := svc.CreatePhotoFromUpload(ctx, file)
|
||||
if err != nil {
|
||||
t.Fatalf("replay CreatePhotoFromUpload: %v", err)
|
||||
}
|
||||
if replayed.ID != first.ID || replayed.AccessHash != first.AccessHash {
|
||||
t.Fatalf("replayed photo = %d/%d, want original %d/%d", replayed.ID, replayed.AccessHash, first.ID, first.AccessHash)
|
||||
}
|
||||
changed := file
|
||||
changed.Name = "different.jpg"
|
||||
if _, err := svc.CreatePhotoFromUpload(ctx, changed); !errors.Is(err, domain.ErrFilePartsInvalid) {
|
||||
t.Fatalf("changed photo intent err = %v, want ErrFilePartsInvalid", err)
|
||||
}
|
||||
}
|
||||
|
||||
type countingUploadPartBackend struct {
|
||||
|
|
|
|||
106
internal/app/files/upload_receipt.go
Normal file
106
internal/app/files/upload_receipt.go
Normal file
|
|
@ -0,0 +1,106 @@
|
|||
package files
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
|
||||
"go.uber.org/zap"
|
||||
|
||||
"telesrv/internal/domain"
|
||||
)
|
||||
|
||||
const uploadedMediaIntentVersion = 1
|
||||
|
||||
type uploadedMediaIntent struct {
|
||||
Version int `json:"version"`
|
||||
Kind domain.UploadedMediaKind `json:"kind"`
|
||||
File domain.UploadedFileRef `json:"file"`
|
||||
Spec *domain.DocumentSpec `json:"spec,omitempty"`
|
||||
}
|
||||
|
||||
func uploadedMediaIntentHash(kind domain.UploadedMediaKind, file domain.UploadedFileRef, spec *domain.DocumentSpec) ([]byte, error) {
|
||||
payload, err := json.Marshal(uploadedMediaIntent{
|
||||
Version: uploadedMediaIntentVersion,
|
||||
Kind: kind,
|
||||
File: file,
|
||||
Spec: spec,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("marshal uploaded media intent: %w", err)
|
||||
}
|
||||
sum := sha256.Sum256(payload)
|
||||
return sum[:], nil
|
||||
}
|
||||
|
||||
func sameUploadedMediaReceipt(receipt domain.UploadedMediaReceipt, kind domain.UploadedMediaKind, intentHash []byte) bool {
|
||||
return receipt.Kind == kind && len(intentHash) == sha256.Size && bytes.Equal(receipt.IntentHash, intentHash)
|
||||
}
|
||||
|
||||
func (s *Service) replayUploadedPhoto(ctx context.Context, file domain.UploadedFileRef, intentHash []byte) (domain.Photo, bool, error) {
|
||||
receipt, found, err := s.media.GetUploadedMediaReceipt(ctx, file.OwnerUserID, file.FileID)
|
||||
if err != nil || !found {
|
||||
return domain.Photo{}, false, err
|
||||
}
|
||||
if !sameUploadedMediaReceipt(receipt, domain.UploadedMediaPhoto, intentHash) {
|
||||
return domain.Photo{}, false, domain.ErrFilePartsInvalid
|
||||
}
|
||||
photo, found, err := s.media.GetPhoto(ctx, receipt.MediaID)
|
||||
if err != nil {
|
||||
return domain.Photo{}, false, err
|
||||
}
|
||||
if !found {
|
||||
return domain.Photo{}, false, fmt.Errorf("uploaded photo receipt %d/%d references missing photo %d", file.OwnerUserID, file.FileID, receipt.MediaID)
|
||||
}
|
||||
s.cleanupMaterializedUpload(ctx, file, "photo replay")
|
||||
return photo, true, nil
|
||||
}
|
||||
|
||||
func (s *Service) replayUploadedDocument(ctx context.Context, file domain.UploadedFileRef, intentHash []byte) (domain.Document, bool, error) {
|
||||
receipt, found, err := s.media.GetUploadedMediaReceipt(ctx, file.OwnerUserID, file.FileID)
|
||||
if err != nil || !found {
|
||||
return domain.Document{}, false, err
|
||||
}
|
||||
if !sameUploadedMediaReceipt(receipt, domain.UploadedMediaDocument, intentHash) {
|
||||
return domain.Document{}, false, domain.ErrFilePartsInvalid
|
||||
}
|
||||
doc, found, err := s.media.GetDocument(ctx, receipt.MediaID)
|
||||
if err != nil {
|
||||
return domain.Document{}, false, err
|
||||
}
|
||||
if !found {
|
||||
return domain.Document{}, false, fmt.Errorf("uploaded document receipt %d/%d references missing document %d", file.OwnerUserID, file.FileID, receipt.MediaID)
|
||||
}
|
||||
s.cleanupMaterializedUpload(ctx, file, "document replay")
|
||||
return doc, true, nil
|
||||
}
|
||||
|
||||
func (s *Service) commitUploadedMediaReceipt(ctx context.Context, file domain.UploadedFileRef, kind domain.UploadedMediaKind, intentHash []byte, mediaID int64) (domain.UploadedMediaReceipt, error) {
|
||||
receipt, _, err := s.media.PutUploadedMediaReceipt(ctx, domain.UploadedMediaReceipt{
|
||||
OwnerUserID: file.OwnerUserID,
|
||||
FileID: file.FileID,
|
||||
IntentHash: intentHash,
|
||||
Kind: kind,
|
||||
MediaID: mediaID,
|
||||
})
|
||||
if err != nil {
|
||||
return domain.UploadedMediaReceipt{}, err
|
||||
}
|
||||
if !sameUploadedMediaReceipt(receipt, kind, intentHash) {
|
||||
return domain.UploadedMediaReceipt{}, domain.ErrFilePartsInvalid
|
||||
}
|
||||
return receipt, nil
|
||||
}
|
||||
|
||||
func (s *Service) cleanupMaterializedUpload(ctx context.Context, file domain.UploadedFileRef, reason string) {
|
||||
if err := s.cleanupUploadParts(ctx, file.OwnerUserID, file.FileID); err != nil {
|
||||
s.log.Warn("cleanup materialized upload parts failed",
|
||||
zap.String("reason", reason),
|
||||
zap.Int64("owner_user_id", file.OwnerUserID),
|
||||
zap.Int64("file_id", file.FileID),
|
||||
zap.Error(err),
|
||||
)
|
||||
}
|
||||
}
|
||||
|
|
@ -17,12 +17,49 @@ type TempAuthKeyRetentionStore interface {
|
|||
DeleteExpired(ctx context.Context, expiredBefore int64, limit int) (int, error)
|
||||
}
|
||||
|
||||
// OrphanAuthKeyRetentionStore 回收从未形成授权/temp binding 的旧握手 key。
|
||||
// protected 是当前连接注册表实际使用的 raw auth_key_id 快照。
|
||||
type OrphanAuthKeyRetentionStore interface {
|
||||
DeleteOrphaned(ctx context.Context, olderThan time.Duration, limit int, protected [][8]byte) (int, error)
|
||||
}
|
||||
|
||||
type ActiveRawAuthKeyProvider interface {
|
||||
ActiveRawAuthKeyIDs() [][8]byte
|
||||
}
|
||||
|
||||
// ActiveAuthKeyHeartbeatStore 把本实例仍在使用的 raw auth key 活性持久化。多实例下
|
||||
// orphan GC 不能只看当前进程的 active 快照;其它实例的 heartbeat 会推进数据库
|
||||
// last_used_at,使它们不会被误判为孤儿。
|
||||
type ActiveAuthKeyHeartbeatStore interface {
|
||||
TouchActiveRawAuthKeys(ctx context.Context, ids [][8]byte) error
|
||||
}
|
||||
|
||||
// BotAPIUpdateRetentionStore 回收 Bot API getUpdates 投递队列的死行(性能审计 H1):
|
||||
// 已确认且超过宽限期的行 + 按消息 date 超过保留期的行(官方 Bot API updates 最多保留 24h)。
|
||||
type BotAPIUpdateRetentionStore interface {
|
||||
DeleteDeliveredOrExpired(ctx context.Context, confirmedGrace, maxAge time.Duration, limit int) (int, error)
|
||||
}
|
||||
|
||||
// UserUpdateEventRetentionStore 只回收所有当前授权设备都明确确认过的账号事件前缀。
|
||||
// 它不是普通 TTL:任一授权缺 state 时确认水位为 0,不得删除该设备可能仍需的事件。
|
||||
type UserUpdateEventRetentionStore interface {
|
||||
DeleteConfirmedPrefix(ctx context.Context, olderThan time.Duration, limit int) (int, error)
|
||||
}
|
||||
|
||||
// ChannelUpdateEventRetentionStore 回收超过保留期的 channel durable update 连续前缀。
|
||||
// 具体 store 必须在同一事务内删除事件并推进 retained floor,低于 floor 的客户端由
|
||||
// updates.getChannelDifference 走 channelDifferenceTooLong 快照恢复。
|
||||
type ChannelUpdateEventRetentionStore interface {
|
||||
DeleteExpiredChannelUpdateEvents(ctx context.Context, olderThan time.Duration, limit int) (int, error)
|
||||
}
|
||||
|
||||
// LoginCodeDeliveryRetentionStore reclaims only compact idempotency receipts
|
||||
// after their associated opaque code lifetime. It must not delete the message,
|
||||
// durable update event, or outbox facts created by the delivery transaction.
|
||||
type LoginCodeDeliveryRetentionStore interface {
|
||||
DeleteExpiredLoginCodeDeliveries(ctx context.Context, expiredBefore time.Time, limit int) (int, error)
|
||||
}
|
||||
|
||||
// botAPIConfirmedGrace 是已确认 Bot API update 行的删除宽限:确认水位之下的行不会再被
|
||||
// getUpdates 读取(fromID 恒 > confirmed),宽限仅防御 offset 回拨调试;回收目标是清堆积。
|
||||
const botAPIConfirmedGrace = 15 * time.Minute
|
||||
|
|
@ -32,23 +69,39 @@ const botAPIConfirmedGrace = 15 * time.Minute
|
|||
// 连接;回收目标是清堆积,晚一天无妨。
|
||||
const tempAuthKeyExpiryGrace = 24 * time.Hour
|
||||
|
||||
const (
|
||||
// terminal failed outbox 只承担短期诊断隔离;它不是 durable update log。
|
||||
// 删除该任务会由 head trigger 立即放行同账号下一 pts,而 user_update_events
|
||||
// 继续保留,在线漏推由正常 difference 路径补偿。
|
||||
defaultOutboxPoisonRetention = time.Minute
|
||||
defaultOutboxPoisonInterval = 15 * time.Second
|
||||
)
|
||||
|
||||
// RetentionWorker 周期性回收存储中的死数据。
|
||||
//
|
||||
// 注意:本 worker 刻意不清理 user_update_events —— pts log 永久保留。原因:TDesktop 不支持
|
||||
// 账号级 updates.differenceTooLong(api_updates.cpp 收到该响应只打一行日志,且漏掉
|
||||
// setRequesting(false),会永久锁死整个 update 引擎),服务端因此无法让"落后超过保留期"的
|
||||
// 客户端整库重置;一旦裁剪 events,落后客户端的 getDifference 会拿到不完整的事件链而静默
|
||||
// 丢消息。详见 docs/performance-audit.md 与 docs/compatibility-matrix.md。user_update_events
|
||||
// 长期膨胀作为已知 todo。
|
||||
// 注意:TDesktop 不支持账号级 updates.differenceTooLong(api_updates.cpp 收到该响应只
|
||||
// 记录日志且不清 requesting,会永久锁死 update 引擎),因此绝不能按普通 TTL 硬裁剪
|
||||
// user_update_events。本 worker 只允许 store 删除“所有当前授权设备都明确确认”的连续安全
|
||||
// 前缀;落后或缺 state 的任一设备都会把 floor 压回 0。客户端偶然带回已确认前的旧 pts 时,
|
||||
// updates 服务通过普通 differenceSlice checkpoint 推进,不发送 differenceTooLong。
|
||||
type RetentionWorker struct {
|
||||
outbox DispatchOutboxRetentionStore
|
||||
tempKeys TempAuthKeyRetentionStore // 可为 nil(不回收 temp key 绑定)
|
||||
botAPIUpdates BotAPIUpdateRetentionStore // 可为 nil(不回收 Bot API 队列)
|
||||
logger *zap.Logger
|
||||
retention time.Duration
|
||||
botAPIRetention time.Duration
|
||||
interval time.Duration
|
||||
batch int
|
||||
outbox DispatchOutboxRetentionStore
|
||||
tempKeys TempAuthKeyRetentionStore // 可为 nil(不回收 temp key 绑定)
|
||||
botAPIUpdates BotAPIUpdateRetentionStore // 可为 nil(不回收 Bot API 队列)
|
||||
userUpdates UserUpdateEventRetentionStore
|
||||
channelUpdates ChannelUpdateEventRetentionStore
|
||||
loginCodeDeliveries LoginCodeDeliveryRetentionStore
|
||||
orphanAuthKeys OrphanAuthKeyRetentionStore
|
||||
activeAuthKeys ActiveRawAuthKeyProvider
|
||||
activeAuthKeyHeartbeat ActiveAuthKeyHeartbeatStore
|
||||
logger *zap.Logger
|
||||
retention time.Duration
|
||||
botAPIRetention time.Duration
|
||||
orphanRetention time.Duration
|
||||
outboxPoisonRetention time.Duration
|
||||
outboxPoisonInterval time.Duration
|
||||
interval time.Duration
|
||||
batch int
|
||||
}
|
||||
|
||||
func NewRetentionWorker(outbox DispatchOutboxRetentionStore, tempKeys TempAuthKeyRetentionStore, logger *zap.Logger, retention, interval time.Duration, batch int) *RetentionWorker {
|
||||
|
|
@ -65,15 +118,32 @@ func NewRetentionWorker(outbox DispatchOutboxRetentionStore, tempKeys TempAuthKe
|
|||
batch = 10000
|
||||
}
|
||||
return &RetentionWorker{
|
||||
outbox: outbox,
|
||||
tempKeys: tempKeys,
|
||||
logger: logger,
|
||||
retention: retention,
|
||||
interval: interval,
|
||||
batch: batch,
|
||||
outbox: outbox,
|
||||
tempKeys: tempKeys,
|
||||
logger: logger,
|
||||
retention: retention,
|
||||
outboxPoisonRetention: defaultOutboxPoisonRetention,
|
||||
outboxPoisonInterval: defaultOutboxPoisonInterval,
|
||||
interval: interval,
|
||||
batch: batch,
|
||||
}
|
||||
}
|
||||
|
||||
// WithDispatchOutboxPoisonPolicy 配置 terminal failed head 的独立短隔离与清理周期。
|
||||
// 该周期不能复用 durable update 的周级保留期,否则一条确定性构造错误会冻结该
|
||||
// 用户整条在线 pts lane。<=0 分别回退到 1m/15s 的安全默认值。
|
||||
func (w *RetentionWorker) WithDispatchOutboxPoisonPolicy(retention, interval time.Duration) *RetentionWorker {
|
||||
if retention <= 0 {
|
||||
retention = defaultOutboxPoisonRetention
|
||||
}
|
||||
if interval <= 0 {
|
||||
interval = defaultOutboxPoisonInterval
|
||||
}
|
||||
w.outboxPoisonRetention = retention
|
||||
w.outboxPoisonInterval = interval
|
||||
return w
|
||||
}
|
||||
|
||||
// WithBotAPIUpdateRetention 启用 bot_api_updates 队列回收;retention <=0 时用官方语义默认 24h。
|
||||
func (w *RetentionWorker) WithBotAPIUpdateRetention(store BotAPIUpdateRetentionStore, retention time.Duration) *RetentionWorker {
|
||||
if retention <= 0 {
|
||||
|
|
@ -84,26 +154,102 @@ func (w *RetentionWorker) WithBotAPIUpdateRetention(store BotAPIUpdateRetentionS
|
|||
return w
|
||||
}
|
||||
|
||||
// WithUserUpdateRetention 启用账号 update 的共同确认安全前缀回收。TDesktop 不支持
|
||||
// account differenceTooLong,具体 store 必须保证未确认前缀永不删除。
|
||||
func (w *RetentionWorker) WithUserUpdateRetention(store UserUpdateEventRetentionStore) *RetentionWorker {
|
||||
w.userUpdates = store
|
||||
return w
|
||||
}
|
||||
|
||||
// WithChannelUpdateRetention 启用 channel durable update 的有界 TTL 回收;复用 worker 的
|
||||
// retention/interval/batch,并由 store 的 retained floor 保证旧 pts 不会读到静默空洞。
|
||||
func (w *RetentionWorker) WithChannelUpdateRetention(store ChannelUpdateEventRetentionStore) *RetentionWorker {
|
||||
w.channelUpdates = store
|
||||
return w
|
||||
}
|
||||
|
||||
// WithLoginCodeDeliveryRetention enables bounded seek cleanup for compact
|
||||
// phone_code_hash receipts. Each row carries its own expiry derived from the
|
||||
// code TTL, so this cleanup intentionally does not reuse update-log retention.
|
||||
func (w *RetentionWorker) WithLoginCodeDeliveryRetention(store LoginCodeDeliveryRetentionStore) *RetentionWorker {
|
||||
w.loginCodeDeliveries = store
|
||||
return w
|
||||
}
|
||||
|
||||
// WithOrphanAuthKeyRetention 启用未授权握手 key 的有界回收。active 必须提供 raw key,
|
||||
// 不能提供 temp→perm business key;否则未登录或 PFS 连接会被误判为 orphan。
|
||||
func (w *RetentionWorker) WithOrphanAuthKeyRetention(store OrphanAuthKeyRetentionStore, active ActiveRawAuthKeyProvider, retention time.Duration) *RetentionWorker {
|
||||
w.orphanAuthKeys = store
|
||||
w.activeAuthKeys = active
|
||||
w.activeAuthKeyHeartbeat, _ = store.(ActiveAuthKeyHeartbeatStore)
|
||||
w.orphanRetention = retention
|
||||
return w
|
||||
}
|
||||
|
||||
func (w *RetentionWorker) Run(ctx context.Context) {
|
||||
w.runOnce(ctx)
|
||||
ticker := time.NewTicker(w.interval)
|
||||
defer ticker.Stop()
|
||||
retentionTicker := time.NewTicker(w.interval)
|
||||
defer retentionTicker.Stop()
|
||||
poisonTicker := time.NewTicker(w.outboxPoisonInterval)
|
||||
defer poisonTicker.Stop()
|
||||
var (
|
||||
heartbeatTicker *time.Ticker
|
||||
heartbeatC <-chan time.Time
|
||||
)
|
||||
if interval := w.orphanHeartbeatInterval(); interval > 0 {
|
||||
heartbeatTicker = time.NewTicker(interval)
|
||||
heartbeatC = heartbeatTicker.C
|
||||
defer heartbeatTicker.Stop()
|
||||
}
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-ticker.C:
|
||||
w.runOnce(ctx)
|
||||
case <-retentionTicker.C:
|
||||
w.runRetentionOnce(ctx)
|
||||
case <-poisonTicker.C:
|
||||
w.runOutboxPoisonOnce(ctx)
|
||||
case <-heartbeatC:
|
||||
w.heartbeatActiveAuthKeys(ctx)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (w *RetentionWorker) runOnce(ctx context.Context) {
|
||||
outboxDeleted, err := w.outbox.DeleteFailed(ctx, w.retention, w.batch)
|
||||
w.runOutboxPoisonOnce(ctx)
|
||||
w.runRetentionOnce(ctx)
|
||||
}
|
||||
|
||||
func (w *RetentionWorker) runOutboxPoisonOnce(ctx context.Context) {
|
||||
if w.outbox == nil {
|
||||
return
|
||||
}
|
||||
outboxDeleted, err := w.outbox.DeleteFailed(ctx, w.outboxPoisonRetention, w.batch)
|
||||
if err != nil {
|
||||
w.logger.Warn("清理 failed dispatch_outbox 失败", zap.Error(err))
|
||||
w.logger.Error("清理 terminal failed dispatch_outbox 失败",
|
||||
zap.String("signal", "dispatch_outbox_poison_cleanup_failed"),
|
||||
zap.Duration("quarantine", w.outboxPoisonRetention),
|
||||
zap.Error(err),
|
||||
)
|
||||
} else if outboxDeleted > 0 {
|
||||
w.logger.Info("清理 failed dispatch_outbox 完成", zap.Int("deleted", outboxDeleted))
|
||||
// Error 级结构化信号刻意保留:发生 terminal failed 代表确定性编码、事件缺失
|
||||
// 或其它不可自动重试故障。任务删除只解冻在线 lane,不会删除 durable event。
|
||||
w.logger.Error("terminal failed dispatch_outbox 已结束隔离并释放用户 lane",
|
||||
zap.String("signal", "dispatch_outbox_poison_released"),
|
||||
zap.Int("deleted", outboxDeleted),
|
||||
zap.Duration("quarantine", w.outboxPoisonRetention),
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
func (w *RetentionWorker) runRetentionOnce(ctx context.Context) {
|
||||
if w.loginCodeDeliveries != nil {
|
||||
deleted, err := w.loginCodeDeliveries.DeleteExpiredLoginCodeDeliveries(ctx, time.Now(), w.batch)
|
||||
if err != nil {
|
||||
w.logger.Warn("回收过期 login-code delivery 回执失败", zap.Error(err))
|
||||
} else if deleted > 0 {
|
||||
w.logger.Info("回收过期 login-code delivery 回执完成", zap.Int("deleted", deleted))
|
||||
}
|
||||
}
|
||||
if w.tempKeys != nil {
|
||||
expiredBefore := time.Now().Add(-tempAuthKeyExpiryGrace).Unix()
|
||||
|
|
@ -114,6 +260,24 @@ func (w *RetentionWorker) runOnce(ctx context.Context) {
|
|||
w.logger.Info("回收过期 temp auth key 绑定完成", zap.Int("deleted", tempDeleted))
|
||||
}
|
||||
}
|
||||
if w.orphanAuthKeys != nil && w.orphanRetention > 0 {
|
||||
var protected [][8]byte
|
||||
if w.activeAuthKeys != nil {
|
||||
protected = w.activeAuthKeys.ActiveRawAuthKeyIDs()
|
||||
}
|
||||
if !w.touchActiveAuthKeys(ctx, protected) {
|
||||
// Fail safe: if this instance cannot publish its own active set, deleting against a
|
||||
// stale database heartbeat could evict keys used by another instance too. Keep all
|
||||
// candidates for this pass and retry after the next heartbeat.
|
||||
} else {
|
||||
orphanDeleted, err := w.orphanAuthKeys.DeleteOrphaned(ctx, w.orphanRetention, w.batch, protected)
|
||||
if err != nil {
|
||||
w.logger.Warn("回收未授权 orphan auth key 失败", zap.Error(err))
|
||||
} else if orphanDeleted > 0 {
|
||||
w.logger.Info("回收未授权 orphan auth key 完成", zap.Int("deleted", orphanDeleted))
|
||||
}
|
||||
}
|
||||
}
|
||||
if w.botAPIUpdates != nil {
|
||||
botAPIDeleted, err := w.botAPIUpdates.DeleteDeliveredOrExpired(ctx, botAPIConfirmedGrace, w.botAPIRetention, w.batch)
|
||||
if err != nil {
|
||||
|
|
@ -122,4 +286,63 @@ func (w *RetentionWorker) runOnce(ctx context.Context) {
|
|||
w.logger.Info("回收 bot_api_updates 队列完成", zap.Int("deleted", botAPIDeleted))
|
||||
}
|
||||
}
|
||||
if w.userUpdates != nil {
|
||||
userDeleted, err := w.userUpdates.DeleteConfirmedPrefix(ctx, w.retention, w.batch)
|
||||
if err != nil {
|
||||
w.logger.Warn("回收已共同确认的 user_update_events 前缀失败", zap.Error(err))
|
||||
} else if userDeleted > 0 {
|
||||
w.logger.Info("回收已共同确认的 user_update_events 前缀完成", zap.Int("deleted", userDeleted))
|
||||
}
|
||||
}
|
||||
if w.channelUpdates != nil {
|
||||
channelDeleted, err := w.channelUpdates.DeleteExpiredChannelUpdateEvents(ctx, w.retention, w.batch)
|
||||
if err != nil {
|
||||
// store 会逐频道隔离坏 gap 后继续本轮;deleted 可能非零,必须同时记录,
|
||||
// 既不能把全局 pass 伪装成完全失败,也不能吞掉不变量错误。
|
||||
w.logger.Warn("回收过期 channel_update_events 存在隔离频道",
|
||||
zap.Int("deleted", channelDeleted),
|
||||
zap.Error(err),
|
||||
)
|
||||
} else if channelDeleted > 0 {
|
||||
w.logger.Info("回收过期 channel_update_events 连续前缀完成", zap.Int("deleted", channelDeleted))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (w *RetentionWorker) orphanHeartbeatInterval() time.Duration {
|
||||
if w.activeAuthKeyHeartbeat == nil || w.activeAuthKeys == nil || w.orphanRetention <= 0 {
|
||||
return 0
|
||||
}
|
||||
interval := w.orphanRetention / 3
|
||||
if interval <= 0 {
|
||||
interval = time.Nanosecond
|
||||
}
|
||||
if w.interval > 0 && w.interval < interval {
|
||||
interval = w.interval
|
||||
}
|
||||
return interval
|
||||
}
|
||||
|
||||
func (w *RetentionWorker) heartbeatActiveAuthKeys(ctx context.Context) {
|
||||
if w.activeAuthKeys == nil {
|
||||
return
|
||||
}
|
||||
w.touchActiveAuthKeys(ctx, w.activeAuthKeys.ActiveRawAuthKeyIDs())
|
||||
}
|
||||
|
||||
// touchActiveAuthKeys returns false only when a configured durable heartbeat failed. A store that
|
||||
// predates the optional heartbeat interface keeps single-instance behavior.
|
||||
func (w *RetentionWorker) touchActiveAuthKeys(ctx context.Context, protected [][8]byte) bool {
|
||||
if w.activeAuthKeyHeartbeat == nil {
|
||||
return true
|
||||
}
|
||||
if err := w.activeAuthKeyHeartbeat.TouchActiveRawAuthKeys(ctx, protected); err != nil {
|
||||
w.logger.Error("刷新 active raw auth key heartbeat 失败,本轮跳过 orphan GC",
|
||||
zap.String("signal", "auth_key_heartbeat_failed"),
|
||||
zap.Int("active_keys", len(protected)),
|
||||
zap.Error(err),
|
||||
)
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
|
|
|||
|
|
@ -2,19 +2,57 @@ package maintenance
|
|||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go.uber.org/zap"
|
||||
"go.uber.org/zap/zapcore"
|
||||
"go.uber.org/zap/zaptest/observer"
|
||||
)
|
||||
|
||||
type fakeOutboxRetention struct {
|
||||
calls int
|
||||
calls int
|
||||
olderThan time.Duration
|
||||
limit int
|
||||
deleted int
|
||||
}
|
||||
|
||||
func (f *fakeOutboxRetention) DeleteFailed(context.Context, time.Duration, int) (int, error) {
|
||||
func (f *fakeOutboxRetention) DeleteFailed(_ context.Context, olderThan time.Duration, limit int) (int, error) {
|
||||
f.calls++
|
||||
return 0, nil
|
||||
f.olderThan = olderThan
|
||||
f.limit = limit
|
||||
return f.deleted, nil
|
||||
}
|
||||
|
||||
func TestRetentionWorkerUsesIndependentOutboxPoisonPolicyAndSignalsRelease(t *testing.T) {
|
||||
core, logs := observer.New(zapcore.ErrorLevel)
|
||||
outbox := &fakeOutboxRetention{deleted: 2}
|
||||
w := NewRetentionWorker(outbox, nil, zap.New(core), 7*24*time.Hour, time.Hour, 73).
|
||||
WithDispatchOutboxPoisonPolicy(2*time.Minute, 7*time.Second)
|
||||
|
||||
w.runOnce(context.Background())
|
||||
|
||||
if outbox.calls != 1 || outbox.olderThan != 2*time.Minute || outbox.limit != 73 {
|
||||
t.Fatalf("outbox poison calls/args = %d/%v/%d, want 1/2m/73", outbox.calls, outbox.olderThan, outbox.limit)
|
||||
}
|
||||
entries := logs.FilterMessage("terminal failed dispatch_outbox 已结束隔离并释放用户 lane").All()
|
||||
if len(entries) != 1 {
|
||||
t.Fatalf("poison release error signals = %d, want 1", len(entries))
|
||||
}
|
||||
if got := entries[0].ContextMap()["signal"]; got != "dispatch_outbox_poison_released" {
|
||||
t.Fatalf("poison signal = %v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRetentionWorkerOutboxPoisonPolicyDefaultsAreShort(t *testing.T) {
|
||||
outbox := &fakeOutboxRetention{}
|
||||
w := NewRetentionWorker(outbox, nil, zap.NewNop(), 168*time.Hour, time.Hour, 100).
|
||||
WithDispatchOutboxPoisonPolicy(0, 0)
|
||||
w.runOutboxPoisonOnce(context.Background())
|
||||
if outbox.olderThan != defaultOutboxPoisonRetention || w.outboxPoisonInterval != defaultOutboxPoisonInterval {
|
||||
t.Fatalf("default poison policy = %v/%v, want %v/%v", outbox.olderThan, w.outboxPoisonInterval, defaultOutboxPoisonRetention, defaultOutboxPoisonInterval)
|
||||
}
|
||||
}
|
||||
|
||||
type fakeTempKeyRetention struct {
|
||||
|
|
@ -65,6 +103,34 @@ type fakeBotAPIRetention struct {
|
|||
limit int
|
||||
}
|
||||
|
||||
type fakeLoginCodeDeliveryRetention struct {
|
||||
calls int
|
||||
expiredBefore time.Time
|
||||
limit int
|
||||
}
|
||||
|
||||
func (f *fakeLoginCodeDeliveryRetention) DeleteExpiredLoginCodeDeliveries(_ context.Context, expiredBefore time.Time, limit int) (int, error) {
|
||||
f.calls++
|
||||
f.expiredBefore = expiredBefore
|
||||
f.limit = limit
|
||||
return 4, nil
|
||||
}
|
||||
|
||||
func TestRetentionWorkerReclaimsExpiredLoginCodeDeliveryReceipts(t *testing.T) {
|
||||
loginCodes := &fakeLoginCodeDeliveryRetention{}
|
||||
w := NewRetentionWorker(&fakeOutboxRetention{}, nil, zap.NewNop(), 168*time.Hour, time.Hour, 83).
|
||||
WithLoginCodeDeliveryRetention(loginCodes)
|
||||
before := time.Now()
|
||||
w.runRetentionOnce(context.Background())
|
||||
after := time.Now()
|
||||
if loginCodes.calls != 1 || loginCodes.limit != 83 {
|
||||
t.Fatalf("login-code retention calls/limit = %d/%d, want 1/83", loginCodes.calls, loginCodes.limit)
|
||||
}
|
||||
if loginCodes.expiredBefore.Before(before) || loginCodes.expiredBefore.After(after) {
|
||||
t.Fatalf("login-code expiry boundary = %v, want within [%v,%v]", loginCodes.expiredBefore, before, after)
|
||||
}
|
||||
}
|
||||
|
||||
func (f *fakeBotAPIRetention) DeleteDeliveredOrExpired(_ context.Context, confirmedGrace, maxAge time.Duration, limit int) (int, error) {
|
||||
f.calls++
|
||||
f.confirmedGrace = confirmedGrace
|
||||
|
|
@ -99,3 +165,151 @@ func TestRetentionWorkerBotAPIRetentionDefaultsTo24h(t *testing.T) {
|
|||
t.Fatalf("default bot api retention = %v, want 24h", botAPI.maxAge)
|
||||
}
|
||||
}
|
||||
|
||||
type fakeUserUpdateRetention struct {
|
||||
calls int
|
||||
olderThan time.Duration
|
||||
limit int
|
||||
}
|
||||
|
||||
func (f *fakeUserUpdateRetention) DeleteConfirmedPrefix(_ context.Context, olderThan time.Duration, limit int) (int, error) {
|
||||
f.calls++
|
||||
f.olderThan = olderThan
|
||||
f.limit = limit
|
||||
return 9, nil
|
||||
}
|
||||
|
||||
func TestRetentionWorkerReclaimsOnlyConfirmedUserUpdatePrefix(t *testing.T) {
|
||||
const retention = 7 * 24 * time.Hour
|
||||
store := &fakeUserUpdateRetention{}
|
||||
w := NewRetentionWorker(&fakeOutboxRetention{}, nil, zap.NewNop(), retention, time.Hour, 91).
|
||||
WithUserUpdateRetention(store)
|
||||
w.runOnce(context.Background())
|
||||
if store.calls != 1 || store.olderThan != retention || store.limit != 91 {
|
||||
t.Fatalf("user update retention calls/args = %d/%v/%d, want 1/%v/91", store.calls, store.olderThan, store.limit, retention)
|
||||
}
|
||||
}
|
||||
|
||||
type fakeChannelUpdateRetention struct {
|
||||
calls int
|
||||
olderThan time.Duration
|
||||
limit int
|
||||
}
|
||||
|
||||
func (f *fakeChannelUpdateRetention) DeleteExpiredChannelUpdateEvents(_ context.Context, olderThan time.Duration, limit int) (int, error) {
|
||||
f.calls++
|
||||
f.olderThan = olderThan
|
||||
f.limit = limit
|
||||
return 7, nil
|
||||
}
|
||||
|
||||
func TestRetentionWorkerReclaimsChannelUpdates(t *testing.T) {
|
||||
const (
|
||||
retention = 14 * 24 * time.Hour
|
||||
batch = 321
|
||||
)
|
||||
channelUpdates := &fakeChannelUpdateRetention{}
|
||||
w := NewRetentionWorker(&fakeOutboxRetention{}, nil, zap.NewNop(), retention, time.Hour, batch).
|
||||
WithChannelUpdateRetention(channelUpdates)
|
||||
|
||||
w.runOnce(context.Background())
|
||||
|
||||
if channelUpdates.calls != 1 {
|
||||
t.Fatalf("channel update retention calls = %d, want 1", channelUpdates.calls)
|
||||
}
|
||||
if channelUpdates.olderThan != retention || channelUpdates.limit != batch {
|
||||
t.Fatalf("channel update retention args = (%v, %d), want (%v, %d)",
|
||||
channelUpdates.olderThan, channelUpdates.limit, retention, batch)
|
||||
}
|
||||
}
|
||||
|
||||
type fakeOrphanAuthKeyRetention struct {
|
||||
calls int
|
||||
olderThan time.Duration
|
||||
limit int
|
||||
protected [][8]byte
|
||||
}
|
||||
|
||||
func (f *fakeOrphanAuthKeyRetention) DeleteOrphaned(_ context.Context, olderThan time.Duration, limit int, protected [][8]byte) (int, error) {
|
||||
f.calls++
|
||||
f.olderThan = olderThan
|
||||
f.limit = limit
|
||||
f.protected = append([][8]byte(nil), protected...)
|
||||
return 2, nil
|
||||
}
|
||||
|
||||
type fakeActiveRawAuthKeys struct{ ids [][8]byte }
|
||||
|
||||
func (f fakeActiveRawAuthKeys) ActiveRawAuthKeyIDs() [][8]byte {
|
||||
return append([][8]byte(nil), f.ids...)
|
||||
}
|
||||
|
||||
func TestRetentionWorkerProtectsActiveRawAuthKeysFromOrphanGC(t *testing.T) {
|
||||
store := &fakeOrphanAuthKeyRetention{}
|
||||
active := fakeActiveRawAuthKeys{ids: [][8]byte{{1}, {2}}}
|
||||
w := NewRetentionWorker(&fakeOutboxRetention{}, nil, zap.NewNop(), time.Hour, time.Hour, 73).
|
||||
WithOrphanAuthKeyRetention(store, active, 24*time.Hour)
|
||||
|
||||
w.runOnce(context.Background())
|
||||
|
||||
if store.calls != 1 || store.olderThan != 24*time.Hour || store.limit != 73 {
|
||||
t.Fatalf("orphan retention calls/args = %d/%v/%d, want 1/24h/73", store.calls, store.olderThan, store.limit)
|
||||
}
|
||||
if len(store.protected) != 2 || store.protected[0] != ([8]byte{1}) || store.protected[1] != ([8]byte{2}) {
|
||||
t.Fatalf("protected raw auth keys = %v, want {1},{2}", store.protected)
|
||||
}
|
||||
}
|
||||
|
||||
type fakeHeartbeatOrphanRetention struct {
|
||||
fakeOrphanAuthKeyRetention
|
||||
heartbeatCalls int
|
||||
heartbeatIDs [][8]byte
|
||||
heartbeatErr error
|
||||
}
|
||||
|
||||
func (f *fakeHeartbeatOrphanRetention) TouchActiveRawAuthKeys(_ context.Context, ids [][8]byte) error {
|
||||
f.heartbeatCalls++
|
||||
f.heartbeatIDs = append([][8]byte(nil), ids...)
|
||||
return f.heartbeatErr
|
||||
}
|
||||
|
||||
func TestRetentionWorkerHeartbeatsActiveKeysBeforeOrphanDelete(t *testing.T) {
|
||||
store := &fakeHeartbeatOrphanRetention{}
|
||||
active := fakeActiveRawAuthKeys{ids: [][8]byte{{3}, {4}}}
|
||||
w := NewRetentionWorker(&fakeOutboxRetention{}, nil, zap.NewNop(), time.Hour, 2*time.Hour, 19).
|
||||
WithOrphanAuthKeyRetention(store, active, 3*time.Hour)
|
||||
|
||||
w.runRetentionOnce(context.Background())
|
||||
|
||||
if store.heartbeatCalls != 1 || store.calls != 1 {
|
||||
t.Fatalf("heartbeat/delete calls = %d/%d, want 1/1", store.heartbeatCalls, store.calls)
|
||||
}
|
||||
if len(store.heartbeatIDs) != 2 || store.heartbeatIDs[0] != ([8]byte{3}) || store.heartbeatIDs[1] != ([8]byte{4}) {
|
||||
t.Fatalf("heartbeat ids = %v, want {3},{4}", store.heartbeatIDs)
|
||||
}
|
||||
// min(retention worker interval=2h, orphan retention/3=1h)
|
||||
if got := w.orphanHeartbeatInterval(); got != time.Hour {
|
||||
t.Fatalf("heartbeat interval = %v, want 1h", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRetentionWorkerSkipsOrphanDeleteWhenHeartbeatFails(t *testing.T) {
|
||||
core, logs := observer.New(zapcore.ErrorLevel)
|
||||
store := &fakeHeartbeatOrphanRetention{heartbeatErr: errors.New("db unavailable")}
|
||||
active := fakeActiveRawAuthKeys{ids: [][8]byte{{5}}}
|
||||
w := NewRetentionWorker(&fakeOutboxRetention{}, nil, zap.New(core), time.Hour, 30*time.Minute, 11).
|
||||
WithOrphanAuthKeyRetention(store, active, 24*time.Hour)
|
||||
if got := w.orphanHeartbeatInterval(); got != 30*time.Minute {
|
||||
t.Fatalf("heartbeat interval = %v, want worker interval 30m", got)
|
||||
}
|
||||
|
||||
w.runRetentionOnce(context.Background())
|
||||
|
||||
if store.heartbeatCalls != 1 || store.calls != 0 {
|
||||
t.Fatalf("heartbeat/delete calls = %d/%d, want 1/0", store.heartbeatCalls, store.calls)
|
||||
}
|
||||
entries := logs.FilterMessage("刷新 active raw auth key heartbeat 失败,本轮跳过 orphan GC").All()
|
||||
if len(entries) != 1 || entries[0].ContextMap()["signal"] != "auth_key_heartbeat_failed" {
|
||||
t.Fatalf("heartbeat failure signals = %+v", entries)
|
||||
}
|
||||
}
|
||||
|
|
|
|||
28
internal/app/messages/album_group.go
Normal file
28
internal/app/messages/album_group.go
Normal file
|
|
@ -0,0 +1,28 @@
|
|||
package messages
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
|
||||
"telesrv/internal/domain"
|
||||
"telesrv/internal/store"
|
||||
)
|
||||
|
||||
// ReserveAlbumGroup 把 RPC 已验证的一批 album item 交给持久层原子预留。
|
||||
// 该能力只传 domain DTO;上传媒体解析与 tg 类型仍停留在 RPC edge。
|
||||
func (s *Service) ReserveAlbumGroup(ctx context.Context, userID int64, req domain.AlbumGroupReservationRequest) (int64, error) {
|
||||
if s == nil || s.messages == nil || userID <= 0 {
|
||||
return 0, domain.ErrAlbumGroupReservationInvalid
|
||||
}
|
||||
if req.SenderUserID == 0 {
|
||||
req.SenderUserID = userID
|
||||
}
|
||||
if req.SenderUserID != userID {
|
||||
return 0, domain.ErrAlbumGroupReservationInvalid
|
||||
}
|
||||
reservations, ok := s.messages.(store.AlbumGroupStore)
|
||||
if !ok {
|
||||
return 0, errors.New("message store does not support album group reservations")
|
||||
}
|
||||
return reservations.ReserveAlbumGroup(ctx, req)
|
||||
}
|
||||
|
|
@ -2,6 +2,7 @@ package messages
|
|||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
|
||||
"telesrv/internal/app/userprojection"
|
||||
"telesrv/internal/domain"
|
||||
|
|
@ -96,6 +97,28 @@ func (s *Service) SendPrivateText(ctx context.Context, userID int64, req domain.
|
|||
if req.SenderUserID == 0 {
|
||||
req.SenderUserID = userID
|
||||
}
|
||||
if req.SenderUserID != userID {
|
||||
return domain.SendPrivateTextResult{}, domain.ErrUserSendRestricted
|
||||
}
|
||||
if req.RandomID != 0 && !req.IdempotencyPreflighted {
|
||||
fingerprint, err := store.PrivateSendFingerprint(req)
|
||||
if err != nil {
|
||||
return domain.SendPrivateTextResult{}, err
|
||||
}
|
||||
req.IdempotencyFingerprint = fingerprint
|
||||
if replayStore, ok := s.messages.(store.PrivateSendReplayStore); ok {
|
||||
replay, found, err := replayStore.LookupPrivateSendReplay(ctx, domain.PrivateSendReplayRequest{
|
||||
SenderUserID: req.SenderUserID,
|
||||
RecipientUserID: req.RecipientUserID,
|
||||
RandomID: req.RandomID,
|
||||
IdempotencyFingerprint: fingerprint,
|
||||
})
|
||||
if err != nil || found {
|
||||
return replay, err
|
||||
}
|
||||
req.IdempotencyPreflighted = true
|
||||
}
|
||||
}
|
||||
if err := s.ensureCanSend(ctx, req.SenderUserID); err != nil {
|
||||
return domain.SendPrivateTextResult{}, err
|
||||
}
|
||||
|
|
@ -113,6 +136,26 @@ func (s *Service) SendPrivateText(ctx context.Context, userID int64, req domain.
|
|||
return res, err
|
||||
}
|
||||
|
||||
// LookupPrivateSendReplay exposes the immutable receipt to the RPC boundary without executing
|
||||
// send permission checks, business automation or bot responders. Sender identity is still bound
|
||||
// to the authenticated app-service caller.
|
||||
func (s *Service) LookupPrivateSendReplay(ctx context.Context, userID int64, req domain.PrivateSendReplayRequest) (domain.SendPrivateTextResult, bool, error) {
|
||||
if s == nil || s.messages == nil || userID == 0 {
|
||||
return domain.SendPrivateTextResult{}, false, nil
|
||||
}
|
||||
if req.SenderUserID == 0 {
|
||||
req.SenderUserID = userID
|
||||
}
|
||||
if req.SenderUserID != userID || req.RecipientUserID == 0 || req.RandomID == 0 {
|
||||
return domain.SendPrivateTextResult{}, false, fmt.Errorf("private send replay: invalid authenticated scope")
|
||||
}
|
||||
replayStore, ok := s.messages.(store.PrivateSendReplayStore)
|
||||
if !ok {
|
||||
return domain.SendPrivateTextResult{}, false, nil
|
||||
}
|
||||
return replayStore.LookupPrivateSendReplay(ctx, req)
|
||||
}
|
||||
|
||||
func (s *Service) ensureCanSend(ctx context.Context, userID int64) error {
|
||||
if s == nil || s.sendGate == nil || userID == 0 {
|
||||
return nil
|
||||
|
|
|
|||
|
|
@ -29,6 +29,38 @@ func TestServiceSendPrivateTextHonorsSendPermissionGate(t *testing.T) {
|
|||
}
|
||||
}
|
||||
|
||||
func TestServicePrivateReplayPrecedesCurrentSendPermissionGate(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
messages := memory.NewMessageStore()
|
||||
allowed := NewService(messages, nil)
|
||||
req := domain.SendPrivateTextRequest{
|
||||
SenderUserID: 1001,
|
||||
RecipientUserID: 1002,
|
||||
RandomID: 91,
|
||||
Message: "committed before restriction",
|
||||
Date: 1_700_000_000,
|
||||
}
|
||||
first, err := allowed.SendPrivateText(ctx, 1001, req)
|
||||
if err != nil {
|
||||
t.Fatalf("first SendPrivateText: %v", err)
|
||||
}
|
||||
|
||||
denied := NewService(messages, nil, WithSendPermissionChecker(denySendChecker{}))
|
||||
req.Date++ // execution time is not part of the immutable send intent.
|
||||
replay, err := denied.SendPrivateText(ctx, 1001, req)
|
||||
if err != nil {
|
||||
t.Fatalf("replay through denied gate: %v", err)
|
||||
}
|
||||
if !replay.Duplicate || replay.SenderMessage.ID != first.SenderMessage.ID {
|
||||
t.Fatalf("replay = %+v, want committed duplicate %d", replay, first.SenderMessage.ID)
|
||||
}
|
||||
|
||||
req.Message = "different intent"
|
||||
if _, err := denied.SendPrivateText(ctx, 1001, req); !errors.Is(err, domain.ErrMessageRandomIDDuplicate) {
|
||||
t.Fatalf("conflicting replay err=%v, want ErrMessageRandomIDDuplicate before send gate", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestServiceForwardPrivateMessagesHonorsSendPermissionGate(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
store := &gateMessageStore{}
|
||||
|
|
|
|||
|
|
@ -27,6 +27,10 @@ type newMessageEventFinder interface {
|
|||
FindNewMessageEvent(ctx context.Context, userID int64, messageBoxID int) (domain.UpdateEvent, bool, error)
|
||||
}
|
||||
|
||||
type userUpdateRetentionCheckpointStore interface {
|
||||
UserUpdateRetentionCheckpoint(ctx context.Context, authKeyID [8]byte, userID int64) (pts, date int, ok bool, err error)
|
||||
}
|
||||
|
||||
// ServiceOption 调整 updates 服务的运行时依赖。
|
||||
type ServiceOption func(*Service)
|
||||
|
||||
|
|
@ -146,6 +150,11 @@ func (s *Service) AcknowledgeCurrentState(ctx context.Context, authKeyID [8]byte
|
|||
if err := s.saveConfirmedState(ctx, authKeyID, userID, st); err != nil {
|
||||
return domain.UpdateState{}, err
|
||||
}
|
||||
// getState 明确建立“从当前快照开始同步”的 baseline;即使响应丢失,客户端也会
|
||||
// 重试 getState/重新拉 snapshot,而不会依赖 baseline 之前的 durable event。
|
||||
if err := s.observeClientState(ctx, authKeyID, userID, st); err != nil {
|
||||
return domain.UpdateState{}, err
|
||||
}
|
||||
return st, nil
|
||||
}
|
||||
|
||||
|
|
@ -162,6 +171,27 @@ func (s *Service) GetDifference(ctx context.Context, authKeyID [8]byte, userID i
|
|||
if err != nil {
|
||||
return domain.UpdateDifference{}, err
|
||||
}
|
||||
// 只把客户端在本次请求中实际带回的 cursor 记为 observed。绝不能把本次将要
|
||||
// 返回的 State 当确认:响应可能在 socket/进程故障中丢失。恶意/损坏客户端带来的
|
||||
// 超前 pts 钳到账号当前连续水位,避免把 retention 安全边界推过 durable truth。
|
||||
observed := from
|
||||
if observed.Pts < 0 {
|
||||
observed.Pts = 0
|
||||
}
|
||||
if observed.Pts > st.Pts {
|
||||
observed.Pts = st.Pts
|
||||
}
|
||||
if err := s.observeClientState(ctx, authKeyID, userID, observed); err != nil {
|
||||
return domain.UpdateDifference{}, err
|
||||
}
|
||||
// TDesktop 不支持账号级 updates.differenceTooLong。retention 只能删除所有授权
|
||||
// 设备都已确认的共同前缀;当前设备若仍带更旧 pts,用一个空的普通
|
||||
// differenceSlice 把 IntermediateState 推进到已确认 checkpoint,再从 live tail 续拉。
|
||||
if checkpoint, found, err := s.retainedPrefixCheckpoint(ctx, authKeyID, userID, from, st); err != nil {
|
||||
return domain.UpdateDifference{}, err
|
||||
} else if found {
|
||||
return checkpoint, nil
|
||||
}
|
||||
if s.events == nil || from.Pts >= st.Pts {
|
||||
if from.Date != 0 {
|
||||
st.Date = from.Date
|
||||
|
|
@ -176,6 +206,17 @@ func (s *Service) GetDifference(ctx context.Context, authKeyID [8]byte, userID i
|
|||
return domain.UpdateDifference{}, err
|
||||
}
|
||||
contiguous, gapEvent, expectedPts := contiguousPrefixAndGap(events, from.Pts)
|
||||
// Retention may advance after the pre-read checkpoint probe and before ListAfter obtains its
|
||||
// statement snapshot. If it removed the whole requested prefix, the read is empty or starts at a
|
||||
// gap. Re-read the checkpoint before returning a non-advancing empty difference; otherwise a
|
||||
// client can believe synchronization completed while retaining a cursor below deleted history.
|
||||
if len(contiguous) == 0 && from.Pts < st.Pts {
|
||||
if checkpoint, found, err := s.retainedPrefixCheckpoint(ctx, authKeyID, userID, from, st); err != nil {
|
||||
return domain.UpdateDifference{}, err
|
||||
} else if found {
|
||||
return checkpoint, nil
|
||||
}
|
||||
}
|
||||
last := from.Pts
|
||||
if len(contiguous) > 0 {
|
||||
last = contiguous[len(contiguous)-1].Pts
|
||||
|
|
@ -215,6 +256,32 @@ func (s *Service) GetDifference(ctx context.Context, authKeyID [8]byte, userID i
|
|||
}, nil
|
||||
}
|
||||
|
||||
func (s *Service) retainedPrefixCheckpoint(ctx context.Context, authKeyID [8]byte, userID int64, from, current domain.UpdateState) (domain.UpdateDifference, bool, error) {
|
||||
checkpoints, ok := s.events.(userUpdateRetentionCheckpointStore)
|
||||
if !ok {
|
||||
return domain.UpdateDifference{}, false, nil
|
||||
}
|
||||
pts, date, found, err := checkpoints.UserUpdateRetentionCheckpoint(ctx, authKeyID, userID)
|
||||
if err != nil {
|
||||
return domain.UpdateDifference{}, false, err
|
||||
}
|
||||
if !found || from.Pts >= pts {
|
||||
return domain.UpdateDifference{}, false, nil
|
||||
}
|
||||
checkpoint := from
|
||||
checkpoint.Pts = pts
|
||||
checkpoint.Seq = 0
|
||||
if date > 0 {
|
||||
checkpoint.Date = date
|
||||
} else if checkpoint.Date == 0 {
|
||||
checkpoint.Date = current.Date
|
||||
}
|
||||
if err := s.saveConfirmedState(ctx, authKeyID, userID, checkpoint); err != nil {
|
||||
return domain.UpdateDifference{}, false, err
|
||||
}
|
||||
return domain.UpdateDifference{State: checkpoint, Partial: true}, true, nil
|
||||
}
|
||||
|
||||
func (s *Service) currentState(ctx context.Context, userID int64) (domain.UpdateState, error) {
|
||||
current, err := s.currentPts(ctx, userID)
|
||||
if err != nil {
|
||||
|
|
@ -235,6 +302,14 @@ func (s *Service) saveConfirmedState(ctx context.Context, authKeyID [8]byte, use
|
|||
return s.states.Save(ctx, authKeyID, userID, st)
|
||||
}
|
||||
|
||||
func (s *Service) observeClientState(ctx context.Context, authKeyID [8]byte, userID int64, st domain.UpdateState) error {
|
||||
if s.states == nil {
|
||||
return nil
|
||||
}
|
||||
st.Seq = 0
|
||||
return s.states.ObserveClientState(ctx, authKeyID, userID, st)
|
||||
}
|
||||
|
||||
// contiguousPrefix 返回从 from 起 pts 严格连续(from+1, from+2, ...)的事件前缀。
|
||||
// 先按 pts 升序排序以兼容存储返回顺序,遇到空洞即停。
|
||||
func contiguousPrefix(events []domain.UpdateEvent, from int) []domain.UpdateEvent {
|
||||
|
|
@ -287,7 +362,7 @@ func (s *Service) RecordNewMessage(ctx context.Context, authKeyID [8]byte, userI
|
|||
if date == 0 {
|
||||
date = int(time.Now().Unix())
|
||||
}
|
||||
return s.recordEvent(ctx, authKeyID, userID, domain.UpdateEvent{
|
||||
return s.recordEvent(ctx, authKeyID, [8]byte{}, userID, domain.UpdateEvent{
|
||||
Type: domain.UpdateEventNewMessage,
|
||||
Date: date,
|
||||
Message: msg,
|
||||
|
|
@ -318,7 +393,7 @@ func (s *Service) PublishNewMessage(ctx context.Context, userID int64, msg domai
|
|||
if date == 0 {
|
||||
date = int(time.Now().Unix())
|
||||
}
|
||||
return s.recordEventCore(ctx, [8]byte{}, userID, domain.UpdateEvent{
|
||||
return s.recordEventCore(ctx, [8]byte{}, [8]byte{}, userID, domain.UpdateEvent{
|
||||
Type: domain.UpdateEventNewMessage,
|
||||
Date: date,
|
||||
Message: msg,
|
||||
|
|
@ -369,11 +444,11 @@ func (s *Service) RecordMessagePoll(ctx context.Context, authKeyID [8]byte, user
|
|||
}
|
||||
|
||||
// RecordStory records a story snapshot change for offline difference replay.
|
||||
func (s *Service) RecordStory(ctx context.Context, authKeyID [8]byte, userID int64, story domain.Story, excludeSessionID int64) (domain.UpdateEvent, domain.UpdateState, error) {
|
||||
func (s *Service) RecordStory(ctx context.Context, stateAuthKeyID [8]byte, userID int64, story domain.Story, excludeAuthKeyID [8]byte, excludeSessionID int64) (domain.UpdateEvent, domain.UpdateState, error) {
|
||||
if userID == 0 && story.Owner.Type == domain.PeerTypeUser {
|
||||
userID = story.Owner.ID
|
||||
}
|
||||
return s.recordEvent(ctx, authKeyID, userID, domain.UpdateEvent{
|
||||
return s.recordEvent(ctx, stateAuthKeyID, excludeAuthKeyID, userID, domain.UpdateEvent{
|
||||
Type: domain.UpdateEventStory,
|
||||
Date: story.Date,
|
||||
Peer: story.Owner,
|
||||
|
|
@ -389,7 +464,7 @@ func (s *Service) RecordStoryFanout(ctx context.Context, userID int64, story dom
|
|||
if userID == 0 {
|
||||
return domain.UpdateEvent{}, domain.UpdateState{}, domain.ErrStoryPeerInvalid
|
||||
}
|
||||
return s.recordEventCore(ctx, [8]byte{}, userID, domain.UpdateEvent{
|
||||
return s.recordEventCore(ctx, [8]byte{}, [8]byte{}, userID, domain.UpdateEvent{
|
||||
Type: domain.UpdateEventStory,
|
||||
Date: story.Date,
|
||||
Peer: story.Owner,
|
||||
|
|
@ -399,11 +474,11 @@ func (s *Service) RecordStoryFanout(ctx context.Context, userID int64, story dom
|
|||
}
|
||||
|
||||
// RecordReadStories records a read boundary update for multi-device sync.
|
||||
func (s *Service) RecordReadStories(ctx context.Context, authKeyID [8]byte, userID int64, read domain.StoryReadResult, excludeSessionID int64) (domain.UpdateEvent, domain.UpdateState, error) {
|
||||
func (s *Service) RecordReadStories(ctx context.Context, stateAuthKeyID [8]byte, userID int64, read domain.StoryReadResult, excludeAuthKeyID [8]byte, excludeSessionID int64) (domain.UpdateEvent, domain.UpdateState, error) {
|
||||
if userID == 0 {
|
||||
userID = read.ViewerID
|
||||
}
|
||||
return s.recordEvent(ctx, authKeyID, userID, domain.UpdateEvent{
|
||||
return s.recordEvent(ctx, stateAuthKeyID, excludeAuthKeyID, userID, domain.UpdateEvent{
|
||||
Type: domain.UpdateEventReadStories,
|
||||
Date: read.Date,
|
||||
Peer: read.Peer,
|
||||
|
|
@ -413,11 +488,11 @@ func (s *Service) RecordReadStories(ctx context.Context, authKeyID [8]byte, user
|
|||
}
|
||||
|
||||
// RecordSentStoryReaction records the current user's story reaction for multi-device sync.
|
||||
func (s *Service) RecordSentStoryReaction(ctx context.Context, authKeyID [8]byte, userID int64, reaction domain.StoryReactionResult, excludeSessionID int64) (domain.UpdateEvent, domain.UpdateState, error) {
|
||||
func (s *Service) RecordSentStoryReaction(ctx context.Context, stateAuthKeyID [8]byte, userID int64, reaction domain.StoryReactionResult, excludeAuthKeyID [8]byte, excludeSessionID int64) (domain.UpdateEvent, domain.UpdateState, error) {
|
||||
if userID == 0 {
|
||||
userID = reaction.ViewerID
|
||||
}
|
||||
return s.recordEvent(ctx, authKeyID, userID, domain.UpdateEvent{
|
||||
return s.recordEvent(ctx, stateAuthKeyID, excludeAuthKeyID, userID, domain.UpdateEvent{
|
||||
Type: domain.UpdateEventSentStoryReaction,
|
||||
Date: reaction.Date,
|
||||
Peer: reaction.Peer,
|
||||
|
|
@ -432,7 +507,7 @@ func (s *Service) RecordSentStoryReaction(ctx context.Context, authKeyID [8]byte
|
|||
// sent by another user. It does not advance any owner device confirmation state:
|
||||
// the owner did not initiate the RPC, but online outbox and offline difference
|
||||
// must still see the durable event.
|
||||
func (s *Service) RecordNewStoryReaction(ctx context.Context, authKeyID [8]byte, ownerUserID int64, reaction domain.StoryReactionResult, excludeSessionID int64) (domain.UpdateEvent, domain.UpdateState, error) {
|
||||
func (s *Service) RecordNewStoryReaction(ctx context.Context, stateAuthKeyID [8]byte, ownerUserID int64, reaction domain.StoryReactionResult, excludeAuthKeyID [8]byte, excludeSessionID int64) (domain.UpdateEvent, domain.UpdateState, error) {
|
||||
if ownerUserID == 0 && reaction.Story.Owner.Type == domain.PeerTypeUser {
|
||||
ownerUserID = reaction.Story.Owner.ID
|
||||
}
|
||||
|
|
@ -442,7 +517,7 @@ func (s *Service) RecordNewStoryReaction(ctx context.Context, authKeyID [8]byte,
|
|||
if ownerUserID == 0 || reaction.ViewerID == 0 || reaction.Reaction == nil {
|
||||
return domain.UpdateEvent{}, domain.UpdateState{}, domain.ErrStoryPeerInvalid
|
||||
}
|
||||
return s.recordEventCore(ctx, authKeyID, ownerUserID, domain.UpdateEvent{
|
||||
return s.recordEventCore(ctx, stateAuthKeyID, excludeAuthKeyID, ownerUserID, domain.UpdateEvent{
|
||||
Type: domain.UpdateEventNewStoryReaction,
|
||||
Date: reaction.Date,
|
||||
Peer: domain.Peer{Type: domain.PeerTypeUser, ID: reaction.ViewerID},
|
||||
|
|
@ -456,7 +531,7 @@ func (s *Service) RecordNewStoryReaction(ctx context.Context, authKeyID [8]byte,
|
|||
// RecordQuickReplyMutation records account-local quick reply state changes for
|
||||
// multi-device sync. Quick-reply TL updates do not carry pts, so outbox appends
|
||||
// auxiliary pts bookkeeping just like other account settings events.
|
||||
func (s *Service) RecordQuickReplyMutation(ctx context.Context, authKeyID [8]byte, userID int64, mutation domain.QuickReplyMutation, excludeSessionID int64) (domain.UpdateEvent, domain.UpdateState, error) {
|
||||
func (s *Service) RecordQuickReplyMutation(ctx context.Context, stateAuthKeyID [8]byte, userID int64, mutation domain.QuickReplyMutation, excludeAuthKeyID [8]byte, excludeSessionID int64) (domain.UpdateEvent, domain.UpdateState, error) {
|
||||
if userID == 0 {
|
||||
userID = mutation.List.OwnerUserID
|
||||
}
|
||||
|
|
@ -481,16 +556,16 @@ func (s *Service) RecordQuickReplyMutation(ctx context.Context, authKeyID [8]byt
|
|||
default:
|
||||
event.Type = domain.UpdateEventQuickReplies
|
||||
}
|
||||
return s.recordEvent(ctx, authKeyID, userID, event, true, excludeSessionID)
|
||||
return s.recordEvent(ctx, stateAuthKeyID, excludeAuthKeyID, userID, event, true, excludeSessionID)
|
||||
}
|
||||
|
||||
// RecordReadHistory 推进 update 状态并追加一条 read_history_inbox 事件。
|
||||
func (s *Service) RecordReadHistory(ctx context.Context, authKeyID [8]byte, userID int64, read domain.ReadHistoryResult, excludeSessionID int64) (domain.UpdateEvent, domain.UpdateState, error) {
|
||||
func (s *Service) RecordReadHistory(ctx context.Context, stateAuthKeyID [8]byte, userID int64, read domain.ReadHistoryResult, excludeAuthKeyID [8]byte, excludeSessionID int64) (domain.UpdateEvent, domain.UpdateState, error) {
|
||||
if userID == 0 {
|
||||
userID = read.OwnerUserID
|
||||
}
|
||||
date := int(time.Now().Unix())
|
||||
return s.recordEvent(ctx, authKeyID, userID, domain.UpdateEvent{
|
||||
return s.recordEvent(ctx, stateAuthKeyID, excludeAuthKeyID, userID, domain.UpdateEvent{
|
||||
Type: domain.UpdateEventReadHistoryInbox,
|
||||
Date: date,
|
||||
Peer: read.Peer,
|
||||
|
|
@ -503,8 +578,8 @@ func (s *Service) RecordReadHistory(ctx context.Context, authKeyID [8]byte, user
|
|||
|
||||
// RecordChannelState 记录当前账号与某频道成员关系变化(leave/kick),
|
||||
// 离线设备经 difference 收到 updateChannel 后重拉 channel 状态。
|
||||
func (s *Service) RecordChannelState(ctx context.Context, authKeyID [8]byte, userID, channelID int64, excludeSessionID int64) (domain.UpdateEvent, domain.UpdateState, error) {
|
||||
return s.recordEvent(ctx, authKeyID, userID, domain.UpdateEvent{
|
||||
func (s *Service) RecordChannelState(ctx context.Context, stateAuthKeyID [8]byte, userID, channelID int64, excludeAuthKeyID [8]byte, excludeSessionID int64) (domain.UpdateEvent, domain.UpdateState, error) {
|
||||
return s.recordEvent(ctx, stateAuthKeyID, excludeAuthKeyID, userID, domain.UpdateEvent{
|
||||
Type: domain.UpdateEventChannelState,
|
||||
Peer: domain.Peer{Type: domain.PeerTypeChannel, ID: channelID},
|
||||
PtsCount: 1,
|
||||
|
|
@ -512,8 +587,8 @@ func (s *Service) RecordChannelState(ctx context.Context, authKeyID [8]byte, use
|
|||
}
|
||||
|
||||
// RecordContactsReset 记录通讯录视角变化,供离线设备通过 updates.getDifference 触发重拉。
|
||||
func (s *Service) RecordContactsReset(ctx context.Context, authKeyID [8]byte, userID int64, excludeSessionID int64) (domain.UpdateEvent, domain.UpdateState, error) {
|
||||
return s.recordEvent(ctx, authKeyID, userID, domain.UpdateEvent{
|
||||
func (s *Service) RecordContactsReset(ctx context.Context, stateAuthKeyID [8]byte, userID int64, excludeAuthKeyID [8]byte, excludeSessionID int64) (domain.UpdateEvent, domain.UpdateState, error) {
|
||||
return s.recordEvent(ctx, stateAuthKeyID, excludeAuthKeyID, userID, domain.UpdateEvent{
|
||||
Type: domain.UpdateEventContactsReset,
|
||||
PtsCount: 1,
|
||||
}, true, excludeSessionID)
|
||||
|
|
@ -522,8 +597,8 @@ func (s *Service) RecordContactsReset(ctx context.Context, authKeyID [8]byte, us
|
|||
// RecordDraftMessage 记录某会话云草稿变化(保存/清空都是同一事件——草稿是绝对
|
||||
// 状态,重放时按 peer 重载当前值)。updateDraftMessage 无 pts 字段,走 LacksWirePts
|
||||
// aux 簿记;topMsgID 是 forum 话题草稿键(复用 MaxID 列持久化)。
|
||||
func (s *Service) RecordDraftMessage(ctx context.Context, authKeyID [8]byte, userID int64, peer domain.Peer, topMsgID int, excludeSessionID int64) (domain.UpdateEvent, domain.UpdateState, error) {
|
||||
return s.recordEvent(ctx, authKeyID, userID, domain.UpdateEvent{
|
||||
func (s *Service) RecordDraftMessage(ctx context.Context, stateAuthKeyID [8]byte, userID int64, peer domain.Peer, topMsgID int, excludeAuthKeyID [8]byte, excludeSessionID int64) (domain.UpdateEvent, domain.UpdateState, error) {
|
||||
return s.recordEvent(ctx, stateAuthKeyID, excludeAuthKeyID, userID, domain.UpdateEvent{
|
||||
Type: domain.UpdateEventDraftMessage,
|
||||
Peer: peer,
|
||||
MaxID: topMsgID,
|
||||
|
|
@ -533,8 +608,8 @@ func (s *Service) RecordDraftMessage(ctx context.Context, authKeyID [8]byte, use
|
|||
|
||||
// RecordDialogPinned 记录单个会话置顶状态变化;folderID 是会话所在 folder
|
||||
// (0 主列表/1 归档),缺失会让离线设备把归档内置顶重放到主列表。
|
||||
func (s *Service) RecordDialogPinned(ctx context.Context, authKeyID [8]byte, userID int64, peer domain.Peer, pinned bool, folderID int, excludeSessionID int64) (domain.UpdateEvent, domain.UpdateState, error) {
|
||||
return s.recordEvent(ctx, authKeyID, userID, domain.UpdateEvent{
|
||||
func (s *Service) RecordDialogPinned(ctx context.Context, stateAuthKeyID [8]byte, userID int64, peer domain.Peer, pinned bool, folderID int, excludeAuthKeyID [8]byte, excludeSessionID int64) (domain.UpdateEvent, domain.UpdateState, error) {
|
||||
return s.recordEvent(ctx, stateAuthKeyID, excludeAuthKeyID, userID, domain.UpdateEvent{
|
||||
Type: domain.UpdateEventDialogPinned,
|
||||
Peer: peer,
|
||||
Bool: pinned,
|
||||
|
|
@ -544,8 +619,8 @@ func (s *Service) RecordDialogPinned(ctx context.Context, authKeyID [8]byte, use
|
|||
}
|
||||
|
||||
// RecordPinnedDialogs 记录指定 folder 内置顶顺序变化,并把新顺序持久化给 getDifference/outbox。
|
||||
func (s *Service) RecordPinnedDialogs(ctx context.Context, authKeyID [8]byte, userID int64, folderID int, order []domain.Peer, excludeSessionID int64) (domain.UpdateEvent, domain.UpdateState, error) {
|
||||
return s.recordEvent(ctx, authKeyID, userID, domain.UpdateEvent{
|
||||
func (s *Service) RecordPinnedDialogs(ctx context.Context, stateAuthKeyID [8]byte, userID int64, folderID int, order []domain.Peer, excludeAuthKeyID [8]byte, excludeSessionID int64) (domain.UpdateEvent, domain.UpdateState, error) {
|
||||
return s.recordEvent(ctx, stateAuthKeyID, excludeAuthKeyID, userID, domain.UpdateEvent{
|
||||
Type: domain.UpdateEventPinnedDialogs,
|
||||
Peers: append([]domain.Peer(nil), order...),
|
||||
FolderID: folderID,
|
||||
|
|
@ -554,8 +629,8 @@ func (s *Service) RecordPinnedDialogs(ctx context.Context, authKeyID [8]byte, us
|
|||
}
|
||||
|
||||
// RecordSavedDialogPinned 记录收藏夹单个子会话置顶状态变化。
|
||||
func (s *Service) RecordSavedDialogPinned(ctx context.Context, authKeyID [8]byte, userID int64, peer domain.Peer, pinned bool, excludeSessionID int64) (domain.UpdateEvent, domain.UpdateState, error) {
|
||||
return s.recordEvent(ctx, authKeyID, userID, domain.UpdateEvent{
|
||||
func (s *Service) RecordSavedDialogPinned(ctx context.Context, stateAuthKeyID [8]byte, userID int64, peer domain.Peer, pinned bool, excludeAuthKeyID [8]byte, excludeSessionID int64) (domain.UpdateEvent, domain.UpdateState, error) {
|
||||
return s.recordEvent(ctx, stateAuthKeyID, excludeAuthKeyID, userID, domain.UpdateEvent{
|
||||
Type: domain.UpdateEventSavedDialogPinned,
|
||||
Peer: peer,
|
||||
Bool: pinned,
|
||||
|
|
@ -564,8 +639,8 @@ func (s *Service) RecordSavedDialogPinned(ctx context.Context, authKeyID [8]byte
|
|||
}
|
||||
|
||||
// RecordPinnedSavedDialogs 记录收藏夹置顶顺序变化,新顺序持久化给 getDifference/outbox。
|
||||
func (s *Service) RecordPinnedSavedDialogs(ctx context.Context, authKeyID [8]byte, userID int64, order []domain.Peer, excludeSessionID int64) (domain.UpdateEvent, domain.UpdateState, error) {
|
||||
return s.recordEvent(ctx, authKeyID, userID, domain.UpdateEvent{
|
||||
func (s *Service) RecordPinnedSavedDialogs(ctx context.Context, stateAuthKeyID [8]byte, userID int64, order []domain.Peer, excludeAuthKeyID [8]byte, excludeSessionID int64) (domain.UpdateEvent, domain.UpdateState, error) {
|
||||
return s.recordEvent(ctx, stateAuthKeyID, excludeAuthKeyID, userID, domain.UpdateEvent{
|
||||
Type: domain.UpdateEventPinnedSavedDialogs,
|
||||
Peers: append([]domain.Peer(nil), order...),
|
||||
PtsCount: 1,
|
||||
|
|
@ -573,8 +648,8 @@ func (s *Service) RecordPinnedSavedDialogs(ctx context.Context, authKeyID [8]byt
|
|||
}
|
||||
|
||||
// RecordDialogUnreadMark 记录手动未读标记变化。
|
||||
func (s *Service) RecordDialogUnreadMark(ctx context.Context, authKeyID [8]byte, userID int64, peer domain.Peer, unread bool, excludeSessionID int64) (domain.UpdateEvent, domain.UpdateState, error) {
|
||||
return s.recordEvent(ctx, authKeyID, userID, domain.UpdateEvent{
|
||||
func (s *Service) RecordDialogUnreadMark(ctx context.Context, stateAuthKeyID [8]byte, userID int64, peer domain.Peer, unread bool, excludeAuthKeyID [8]byte, excludeSessionID int64) (domain.UpdateEvent, domain.UpdateState, error) {
|
||||
return s.recordEvent(ctx, stateAuthKeyID, excludeAuthKeyID, userID, domain.UpdateEvent{
|
||||
Type: domain.UpdateEventDialogUnreadMark,
|
||||
Peer: peer,
|
||||
Bool: unread,
|
||||
|
|
@ -583,8 +658,8 @@ func (s *Service) RecordDialogUnreadMark(ctx context.Context, authKeyID [8]byte,
|
|||
}
|
||||
|
||||
// RecordChannelViewForumAsMessages records a per-account forum presentation state change.
|
||||
func (s *Service) RecordChannelViewForumAsMessages(ctx context.Context, authKeyID [8]byte, userID, channelID int64, enabled bool, excludeSessionID int64) (domain.UpdateEvent, domain.UpdateState, error) {
|
||||
return s.recordEvent(ctx, authKeyID, userID, domain.UpdateEvent{
|
||||
func (s *Service) RecordChannelViewForumAsMessages(ctx context.Context, stateAuthKeyID [8]byte, userID, channelID int64, enabled bool, excludeAuthKeyID [8]byte, excludeSessionID int64) (domain.UpdateEvent, domain.UpdateState, error) {
|
||||
return s.recordEvent(ctx, stateAuthKeyID, excludeAuthKeyID, userID, domain.UpdateEvent{
|
||||
Type: domain.UpdateEventChannelViewForum,
|
||||
Peer: domain.Peer{Type: domain.PeerTypeChannel, ID: channelID},
|
||||
Bool: enabled,
|
||||
|
|
@ -594,8 +669,8 @@ func (s *Service) RecordChannelViewForumAsMessages(ctx context.Context, authKeyI
|
|||
|
||||
// RecordChannelDiscussionInbox 记录 forum 话题级已读(updateReadChannelDiscussionInbox),
|
||||
// 占一个账号 pts(LacksWirePts),供自己其它设备在线同步与离线差分恢复。
|
||||
func (s *Service) RecordChannelDiscussionInbox(ctx context.Context, authKeyID [8]byte, userID, channelID int64, topicID, maxID int, excludeSessionID int64) (domain.UpdateEvent, domain.UpdateState, error) {
|
||||
return s.recordEvent(ctx, authKeyID, userID, domain.UpdateEvent{
|
||||
func (s *Service) RecordChannelDiscussionInbox(ctx context.Context, stateAuthKeyID [8]byte, userID, channelID int64, topicID, maxID int, excludeAuthKeyID [8]byte, excludeSessionID int64) (domain.UpdateEvent, domain.UpdateState, error) {
|
||||
return s.recordEvent(ctx, stateAuthKeyID, excludeAuthKeyID, userID, domain.UpdateEvent{
|
||||
Type: domain.UpdateEventReadChannelDiscussionInbox,
|
||||
Peer: domain.Peer{Type: domain.PeerTypeChannel, ID: channelID},
|
||||
TopMsgID: topicID,
|
||||
|
|
@ -605,8 +680,8 @@ func (s *Service) RecordChannelDiscussionInbox(ctx context.Context, authKeyID [8
|
|||
}
|
||||
|
||||
// RecordPeerSettings 记录 peer settings 变化。
|
||||
func (s *Service) RecordPeerSettings(ctx context.Context, authKeyID [8]byte, userID int64, peer domain.Peer, settings domain.PeerSettings, excludeSessionID int64) (domain.UpdateEvent, domain.UpdateState, error) {
|
||||
return s.recordEvent(ctx, authKeyID, userID, domain.UpdateEvent{
|
||||
func (s *Service) RecordPeerSettings(ctx context.Context, stateAuthKeyID [8]byte, userID int64, peer domain.Peer, settings domain.PeerSettings, excludeAuthKeyID [8]byte, excludeSessionID int64) (domain.UpdateEvent, domain.UpdateState, error) {
|
||||
return s.recordEvent(ctx, stateAuthKeyID, excludeAuthKeyID, userID, domain.UpdateEvent{
|
||||
Type: domain.UpdateEventPeerSettings,
|
||||
Peer: peer,
|
||||
Settings: settings,
|
||||
|
|
@ -615,8 +690,8 @@ func (s *Service) RecordPeerSettings(ctx context.Context, authKeyID [8]byte, use
|
|||
}
|
||||
|
||||
// RecordPeerStoryBlocked 记录当前账号 story blocklist 对某个 peer 的可见状态变化。
|
||||
func (s *Service) RecordPeerStoryBlocked(ctx context.Context, authKeyID [8]byte, userID int64, peer domain.Peer, blocked bool, excludeSessionID int64) (domain.UpdateEvent, domain.UpdateState, error) {
|
||||
return s.recordEvent(ctx, authKeyID, userID, domain.UpdateEvent{
|
||||
func (s *Service) RecordPeerStoryBlocked(ctx context.Context, stateAuthKeyID [8]byte, userID int64, peer domain.Peer, blocked bool, excludeAuthKeyID [8]byte, excludeSessionID int64) (domain.UpdateEvent, domain.UpdateState, error) {
|
||||
return s.recordEvent(ctx, stateAuthKeyID, excludeAuthKeyID, userID, domain.UpdateEvent{
|
||||
Type: domain.UpdateEventPeerStoryBlocked,
|
||||
Peer: peer,
|
||||
Bool: blocked,
|
||||
|
|
@ -625,13 +700,13 @@ func (s *Service) RecordPeerStoryBlocked(ctx context.Context, authKeyID [8]byte,
|
|||
}
|
||||
|
||||
// RecordDialogFilter 记录单个 filter 的创建、更新或删除;folder 为 nil 表示删除。
|
||||
func (s *Service) RecordDialogFilter(ctx context.Context, authKeyID [8]byte, userID int64, folderID int, folder *domain.DialogFolder, excludeSessionID int64) (domain.UpdateEvent, domain.UpdateState, error) {
|
||||
func (s *Service) RecordDialogFilter(ctx context.Context, stateAuthKeyID [8]byte, userID int64, folderID int, folder *domain.DialogFolder, excludeAuthKeyID [8]byte, excludeSessionID int64) (domain.UpdateEvent, domain.UpdateState, error) {
|
||||
var copyFolder *domain.DialogFolder
|
||||
if folder != nil {
|
||||
f := *folder
|
||||
copyFolder = &f
|
||||
}
|
||||
return s.recordEvent(ctx, authKeyID, userID, domain.UpdateEvent{
|
||||
return s.recordEvent(ctx, stateAuthKeyID, excludeAuthKeyID, userID, domain.UpdateEvent{
|
||||
Type: domain.UpdateEventDialogFilter,
|
||||
FilterID: folderID,
|
||||
DialogFilter: copyFolder,
|
||||
|
|
@ -640,8 +715,8 @@ func (s *Service) RecordDialogFilter(ctx context.Context, authKeyID [8]byte, use
|
|||
}
|
||||
|
||||
// RecordDialogFilterOrder 记录 filter 顺序变化。
|
||||
func (s *Service) RecordDialogFilterOrder(ctx context.Context, authKeyID [8]byte, userID int64, order []int, excludeSessionID int64) (domain.UpdateEvent, domain.UpdateState, error) {
|
||||
return s.recordEvent(ctx, authKeyID, userID, domain.UpdateEvent{
|
||||
func (s *Service) RecordDialogFilterOrder(ctx context.Context, stateAuthKeyID [8]byte, userID int64, order []int, excludeAuthKeyID [8]byte, excludeSessionID int64) (domain.UpdateEvent, domain.UpdateState, error) {
|
||||
return s.recordEvent(ctx, stateAuthKeyID, excludeAuthKeyID, userID, domain.UpdateEvent{
|
||||
Type: domain.UpdateEventDialogFilterOrder,
|
||||
FilterOrder: append([]int(nil), order...),
|
||||
PtsCount: 1,
|
||||
|
|
@ -649,16 +724,16 @@ func (s *Service) RecordDialogFilterOrder(ctx context.Context, authKeyID [8]byte
|
|||
}
|
||||
|
||||
// RecordDialogFiltersReload 通知其他设备重新拉取 filter 列表。
|
||||
func (s *Service) RecordDialogFiltersReload(ctx context.Context, authKeyID [8]byte, userID int64, excludeSessionID int64) (domain.UpdateEvent, domain.UpdateState, error) {
|
||||
return s.recordEvent(ctx, authKeyID, userID, domain.UpdateEvent{
|
||||
func (s *Service) RecordDialogFiltersReload(ctx context.Context, stateAuthKeyID [8]byte, userID int64, excludeAuthKeyID [8]byte, excludeSessionID int64) (domain.UpdateEvent, domain.UpdateState, error) {
|
||||
return s.recordEvent(ctx, stateAuthKeyID, excludeAuthKeyID, userID, domain.UpdateEvent{
|
||||
Type: domain.UpdateEventDialogFilters,
|
||||
PtsCount: 1,
|
||||
}, true, excludeSessionID)
|
||||
}
|
||||
|
||||
// RecordFolderPeers 记录归档/还原会话的 folder_id 变化。
|
||||
func (s *Service) RecordFolderPeers(ctx context.Context, authKeyID [8]byte, userID int64, peers []domain.FolderPeerUpdate, excludeSessionID int64) (domain.UpdateEvent, domain.UpdateState, error) {
|
||||
return s.recordEvent(ctx, authKeyID, userID, domain.UpdateEvent{
|
||||
func (s *Service) RecordFolderPeers(ctx context.Context, stateAuthKeyID [8]byte, userID int64, peers []domain.FolderPeerUpdate, excludeAuthKeyID [8]byte, excludeSessionID int64) (domain.UpdateEvent, domain.UpdateState, error) {
|
||||
return s.recordEvent(ctx, stateAuthKeyID, excludeAuthKeyID, userID, domain.UpdateEvent{
|
||||
Type: domain.UpdateEventFolderPeers,
|
||||
FolderPeers: append([]domain.FolderPeerUpdate(nil), peers...),
|
||||
PtsCount: 1,
|
||||
|
|
@ -666,8 +741,8 @@ func (s *Service) RecordFolderPeers(ctx context.Context, authKeyID [8]byte, user
|
|||
}
|
||||
|
||||
// RecordChannelAvailableMessages records a local channel history clear for multi-device sync.
|
||||
func (s *Service) RecordChannelAvailableMessages(ctx context.Context, authKeyID [8]byte, userID, channelID int64, availableMinID int, excludeSessionID int64) (domain.UpdateEvent, domain.UpdateState, error) {
|
||||
return s.recordEvent(ctx, authKeyID, userID, domain.UpdateEvent{
|
||||
func (s *Service) RecordChannelAvailableMessages(ctx context.Context, stateAuthKeyID [8]byte, userID, channelID int64, availableMinID int, excludeAuthKeyID [8]byte, excludeSessionID int64) (domain.UpdateEvent, domain.UpdateState, error) {
|
||||
return s.recordEvent(ctx, stateAuthKeyID, excludeAuthKeyID, userID, domain.UpdateEvent{
|
||||
Type: domain.UpdateEventChannelAvailable,
|
||||
Peer: domain.Peer{Type: domain.PeerTypeChannel, ID: channelID},
|
||||
MaxID: availableMinID,
|
||||
|
|
@ -675,15 +750,15 @@ func (s *Service) RecordChannelAvailableMessages(ctx context.Context, authKeyID
|
|||
}, true, excludeSessionID)
|
||||
}
|
||||
|
||||
func (s *Service) recordEvent(ctx context.Context, authKeyID [8]byte, userID int64, event domain.UpdateEvent, dispatch bool, excludeSessionID int64) (domain.UpdateEvent, domain.UpdateState, error) {
|
||||
return s.recordEventCore(ctx, authKeyID, userID, event, dispatch, excludeSessionID, true)
|
||||
func (s *Service) recordEvent(ctx context.Context, stateAuthKeyID, excludeAuthKeyID [8]byte, userID int64, event domain.UpdateEvent, dispatch bool, excludeSessionID int64) (domain.UpdateEvent, domain.UpdateState, error) {
|
||||
return s.recordEventCore(ctx, stateAuthKeyID, excludeAuthKeyID, userID, event, dispatch, excludeSessionID, true)
|
||||
}
|
||||
|
||||
func (s *Service) recordEventWithoutState(ctx context.Context, userID int64, event domain.UpdateEvent) (domain.UpdateEvent, domain.UpdateState, error) {
|
||||
return s.recordEventCore(ctx, [8]byte{}, userID, event, false, 0, false)
|
||||
return s.recordEventCore(ctx, [8]byte{}, [8]byte{}, userID, event, false, 0, false)
|
||||
}
|
||||
|
||||
func (s *Service) recordEventCore(ctx context.Context, authKeyID [8]byte, userID int64, event domain.UpdateEvent, dispatch bool, excludeSessionID int64, saveState bool) (domain.UpdateEvent, domain.UpdateState, error) {
|
||||
func (s *Service) recordEventCore(ctx context.Context, stateAuthKeyID, excludeAuthKeyID [8]byte, userID int64, event domain.UpdateEvent, dispatch bool, excludeSessionID int64, saveState bool) (domain.UpdateEvent, domain.UpdateState, error) {
|
||||
date := event.Date
|
||||
if date == 0 {
|
||||
date = int(time.Now().Unix())
|
||||
|
|
@ -698,7 +773,7 @@ func (s *Service) recordEventCore(ctx context.Context, authKeyID [8]byte, userID
|
|||
var err error
|
||||
if dispatch {
|
||||
if appender, ok := s.events.(dispatchingEventAppender); ok {
|
||||
event, err = appender.AppendAllocatedWithDispatch(ctx, userID, event, authKeyID, excludeSessionID)
|
||||
event, err = appender.AppendAllocatedWithDispatch(ctx, userID, event, excludeAuthKeyID, excludeSessionID)
|
||||
} else {
|
||||
event, err = s.events.AppendAllocated(ctx, userID, event)
|
||||
}
|
||||
|
|
@ -735,7 +810,7 @@ func (s *Service) recordEventCore(ctx context.Context, authKeyID [8]byte, userID
|
|||
st.Pts = event.Pts
|
||||
}
|
||||
if saveState && s.states != nil {
|
||||
if err := s.states.Save(ctx, authKeyID, userID, st); err != nil {
|
||||
if err := s.states.Save(ctx, stateAuthKeyID, userID, st); err != nil {
|
||||
return domain.UpdateEvent{}, domain.UpdateState{}, err
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -100,7 +100,7 @@ func TestRecordReadHistoryFeedsGetDifference(t *testing.T) {
|
|||
Peer: peer,
|
||||
MaxID: 10,
|
||||
Changed: true,
|
||||
}, 0)
|
||||
}, [8]byte{}, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("RecordReadHistory: %v", err)
|
||||
}
|
||||
|
|
@ -132,7 +132,7 @@ func TestRecordChannelReadHistoryKeepsChannelPtsPayload(t *testing.T) {
|
|||
StillUnreadCount: 3,
|
||||
ChannelPts: 77,
|
||||
Changed: true,
|
||||
}, 0)
|
||||
}, [8]byte{}, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("RecordReadHistory: %v", err)
|
||||
}
|
||||
|
|
@ -157,24 +157,24 @@ func TestRecordSettingsEventsFeedGetDifference(t *testing.T) {
|
|||
ownerUserID := int64(1000000001)
|
||||
peer := domain.Peer{Type: domain.PeerTypeUser, ID: 1000000002}
|
||||
|
||||
if _, _, err := svc.RecordContactsReset(ctx, authKeyID, ownerUserID, 0); err != nil {
|
||||
if _, _, err := svc.RecordContactsReset(ctx, authKeyID, ownerUserID, [8]byte{}, 0); err != nil {
|
||||
t.Fatalf("RecordContactsReset: %v", err)
|
||||
}
|
||||
if _, _, err := svc.RecordDialogPinned(ctx, authKeyID, ownerUserID, peer, true, 0, 0); err != nil {
|
||||
if _, _, err := svc.RecordDialogPinned(ctx, authKeyID, ownerUserID, peer, true, 0, [8]byte{}, 0); err != nil {
|
||||
t.Fatalf("RecordDialogPinned: %v", err)
|
||||
}
|
||||
order := []domain.Peer{peer}
|
||||
if _, _, err := svc.RecordPinnedDialogs(ctx, authKeyID, ownerUserID, 0, order, 0); err != nil {
|
||||
if _, _, err := svc.RecordPinnedDialogs(ctx, authKeyID, ownerUserID, 0, order, [8]byte{}, 0); err != nil {
|
||||
t.Fatalf("RecordPinnedDialogs: %v", err)
|
||||
}
|
||||
if _, _, err := svc.RecordDialogUnreadMark(ctx, authKeyID, ownerUserID, peer, false, 0); err != nil {
|
||||
if _, _, err := svc.RecordDialogUnreadMark(ctx, authKeyID, ownerUserID, peer, false, [8]byte{}, 0); err != nil {
|
||||
t.Fatalf("RecordDialogUnreadMark: %v", err)
|
||||
}
|
||||
settings := domain.PeerSettings{ShareContact: true}
|
||||
if _, _, err := svc.RecordPeerSettings(ctx, authKeyID, ownerUserID, peer, settings, 0); err != nil {
|
||||
if _, _, err := svc.RecordPeerSettings(ctx, authKeyID, ownerUserID, peer, settings, [8]byte{}, 0); err != nil {
|
||||
t.Fatalf("RecordPeerSettings: %v", err)
|
||||
}
|
||||
stateEvent, state, err := svc.RecordPeerStoryBlocked(ctx, authKeyID, ownerUserID, peer, true, 0)
|
||||
stateEvent, state, err := svc.RecordPeerStoryBlocked(ctx, authKeyID, ownerUserID, peer, true, [8]byte{}, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("RecordPeerStoryBlocked: %v", err)
|
||||
}
|
||||
|
|
@ -221,22 +221,29 @@ func TestRecordSettingsEventsFeedGetDifference(t *testing.T) {
|
|||
|
||||
func TestRecordSettingsEventUsesDispatchAppender(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
var authKeyID [8]byte
|
||||
authKeyID[0] = 4
|
||||
authKeyID := [8]byte{4}
|
||||
rawAuthKeyID := [8]byte{4, 9}
|
||||
events := &captureDispatchAppender{UpdateEventStore: memory.NewUpdateEventStore()}
|
||||
svc := NewService(memory.NewUpdateStateStore(), events)
|
||||
states := &captureStateStore{}
|
||||
svc := NewService(states, events)
|
||||
peer := domain.Peer{Type: domain.PeerTypeUser, ID: 1000000002}
|
||||
|
||||
event, state, err := svc.RecordDialogPinned(ctx, authKeyID, 1000000001, peer, true, 0, 42)
|
||||
event, state, err := svc.RecordDialogPinned(ctx, authKeyID, 1000000001, peer, true, 0, rawAuthKeyID, 42)
|
||||
if err != nil {
|
||||
t.Fatalf("RecordDialogPinned: %v", err)
|
||||
}
|
||||
if event.Pts != 1 || state.Pts != 1 {
|
||||
t.Fatalf("event/state = %+v / %+v, want first pts", event, state)
|
||||
}
|
||||
if !events.dispatched || events.excludeAuthKeyID != authKeyID || events.excludeSessionID != 42 || events.event.Type != domain.UpdateEventDialogPinned || events.event.Peer != peer {
|
||||
if !events.dispatched || events.excludeAuthKeyID != rawAuthKeyID || events.excludeSessionID != 42 || events.event.Type != domain.UpdateEventDialogPinned || events.event.Peer != peer {
|
||||
t.Fatalf("dispatch capture = %+v exclude_auth=%v exclude_session=%d dispatched=%v, want dialog_pinned outbox", events.event, events.excludeAuthKeyID, events.excludeSessionID, events.dispatched)
|
||||
}
|
||||
if states.lastSaveAuthKeyID != authKeyID {
|
||||
t.Fatalf("device state auth key = %x, want business/perm %x", states.lastSaveAuthKeyID, authKeyID)
|
||||
}
|
||||
if _, found, err := states.Get(ctx, rawAuthKeyID, 1000000001); err != nil || found {
|
||||
t.Fatalf("raw temp key unexpectedly owns device state: found=%v err=%v", found, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRecordSettingsEventDispatchFailureDoesNotRecordEvent(t *testing.T) {
|
||||
|
|
@ -246,7 +253,7 @@ func TestRecordSettingsEventDispatchFailureDoesNotRecordEvent(t *testing.T) {
|
|||
events := &failingDispatchAppender{UpdateEventStore: memory.NewUpdateEventStore()}
|
||||
svc := NewService(memory.NewUpdateStateStore(), events)
|
||||
|
||||
_, _, err := svc.RecordDialogPinned(ctx, authKeyID, 1000000001, domain.Peer{Type: domain.PeerTypeUser, ID: 1000000002}, true, 0, 42)
|
||||
_, _, err := svc.RecordDialogPinned(ctx, authKeyID, 1000000001, domain.Peer{Type: domain.PeerTypeUser, ID: 1000000002}, true, 0, authKeyID, 42)
|
||||
if !errors.Is(err, errDispatchFailed) {
|
||||
t.Fatalf("RecordDialogPinned err = %v, want dispatch failure", err)
|
||||
}
|
||||
|
|
@ -267,7 +274,7 @@ func TestRecordPeerStoryBlockedUsesDispatchAppender(t *testing.T) {
|
|||
svc := NewService(memory.NewUpdateStateStore(), events)
|
||||
peer := domain.Peer{Type: domain.PeerTypeUser, ID: 1000000002}
|
||||
|
||||
event, state, err := svc.RecordPeerStoryBlocked(ctx, authKeyID, 1000000001, peer, true, 91)
|
||||
event, state, err := svc.RecordPeerStoryBlocked(ctx, authKeyID, 1000000001, peer, true, authKeyID, 91)
|
||||
if err != nil {
|
||||
t.Fatalf("RecordPeerStoryBlocked: %v", err)
|
||||
}
|
||||
|
|
@ -294,7 +301,7 @@ func TestRecordStoryUsesDispatchAppenderExcludeCurrentSession(t *testing.T) {
|
|||
Caption: "owner story",
|
||||
}
|
||||
|
||||
event, state, err := svc.RecordStory(ctx, authKeyID, owner.ID, story, 1234)
|
||||
event, state, err := svc.RecordStory(ctx, authKeyID, owner.ID, story, authKeyID, 1234)
|
||||
if err != nil {
|
||||
t.Fatalf("RecordStory: %v", err)
|
||||
}
|
||||
|
|
@ -330,7 +337,7 @@ func TestRecordStoryReadAndSentReactionExcludeCurrentSession(t *testing.T) {
|
|||
MaxReadID: story.ID,
|
||||
Advanced: true,
|
||||
Date: 1700000201,
|
||||
}, 2233)
|
||||
}, authKeyID, 2233)
|
||||
if err != nil {
|
||||
t.Fatalf("RecordReadStories: %v", err)
|
||||
}
|
||||
|
|
@ -350,7 +357,7 @@ func TestRecordStoryReadAndSentReactionExcludeCurrentSession(t *testing.T) {
|
|||
Reaction: reaction,
|
||||
Changed: true,
|
||||
Date: 1700000202,
|
||||
}, 2233)
|
||||
}, authKeyID, 2233)
|
||||
if err != nil {
|
||||
t.Fatalf("RecordSentStoryReaction: %v", err)
|
||||
}
|
||||
|
|
@ -384,7 +391,7 @@ func TestRecordNewStoryReactionDispatchesWithoutSavingDeviceState(t *testing.T)
|
|||
},
|
||||
Reaction: reaction,
|
||||
Date: 1700000101,
|
||||
}, 0)
|
||||
}, [8]byte{}, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("RecordNewStoryReaction: %v", err)
|
||||
}
|
||||
|
|
@ -493,7 +500,8 @@ func TestAcknowledgeCurrentStateAdvancesConfirmedWatermark(t *testing.T) {
|
|||
authKeyID[0] = 11
|
||||
userID := int64(1000000001)
|
||||
events := memory.NewUpdateEventStore()
|
||||
svc := NewService(memory.NewUpdateStateStore(), events)
|
||||
states := memory.NewUpdateStateStore()
|
||||
svc := NewService(states, events)
|
||||
if err := events.Append(ctx, userID, domain.UpdateEvent{
|
||||
UserID: userID, Type: domain.UpdateEventNewMessage, Pts: 1, PtsCount: 1,
|
||||
Date: 1700000001, Message: domain.Message{ID: 1, OwnerUserID: userID},
|
||||
|
|
@ -527,6 +535,138 @@ func TestAcknowledgeCurrentStateAdvancesConfirmedWatermark(t *testing.T) {
|
|||
if confirmed.Pts != 3 {
|
||||
t.Fatalf("confirmed watermark = %d, want advanced to 3", confirmed.Pts)
|
||||
}
|
||||
observed, ok := states.ObservedClientState(authKeyID, userID)
|
||||
if !ok || observed.Pts != 3 {
|
||||
t.Fatalf("getState observed watermark = %+v/%v, want pts=3", observed, ok)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetDifferenceRetainsOnlyClientObservedInputCursor(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
authKeyID := [8]byte{12}
|
||||
const userID int64 = 1000000012
|
||||
events := memory.NewUpdateEventStore()
|
||||
states := memory.NewUpdateStateStore()
|
||||
svc := NewService(states, events)
|
||||
for pts := 1; pts <= 2; pts++ {
|
||||
if err := events.Append(ctx, userID, domain.UpdateEvent{
|
||||
UserID: userID, Type: domain.UpdateEventNewMessage, Pts: pts, PtsCount: 1,
|
||||
Date: 1700000100 + pts, Message: domain.Message{ID: pts, OwnerUserID: userID},
|
||||
}); err != nil {
|
||||
t.Fatalf("append pts=%d: %v", pts, err)
|
||||
}
|
||||
}
|
||||
|
||||
// 服务端把 pts=1..2 放进 response,并不证明客户端收到了 response;observed 只能
|
||||
// 保持在本次 request 实际携带的 pts=0。
|
||||
diff, err := svc.GetDifference(ctx, authKeyID, userID, domain.UpdateState{Pts: 0, Date: 1700000100})
|
||||
if err != nil {
|
||||
t.Fatalf("first difference: %v", err)
|
||||
}
|
||||
if diff.State.Pts != 2 || len(diff.Events) != 2 {
|
||||
t.Fatalf("first difference = %+v, want response through pts=2", diff)
|
||||
}
|
||||
observed, ok := states.ObservedClientState(authKeyID, userID)
|
||||
if !ok || observed.Pts != 0 {
|
||||
t.Fatalf("observed after merely sending response = %+v/%v, want pts=0", observed, ok)
|
||||
}
|
||||
|
||||
// 客户端下一次明确带回 pts=2 后,才允许 retention 把共同安全水位推进到 2。
|
||||
if _, err := svc.GetDifference(ctx, authKeyID, userID, domain.UpdateState{Pts: 2, Date: 1700000102}); err != nil {
|
||||
t.Fatalf("confirming difference: %v", err)
|
||||
}
|
||||
observed, ok = states.ObservedClientState(authKeyID, userID)
|
||||
if !ok || observed.Pts != 2 {
|
||||
t.Fatalf("observed after client carried cursor = %+v/%v, want pts=2", observed, ok)
|
||||
}
|
||||
}
|
||||
|
||||
type retentionCheckpointEvents struct {
|
||||
*memory.UpdateEventStore
|
||||
pts int
|
||||
date int
|
||||
current int
|
||||
missFirst bool
|
||||
calls int
|
||||
}
|
||||
|
||||
func (s *retentionCheckpointEvents) UserUpdateRetentionCheckpoint(_ context.Context, _ [8]byte, _ int64) (int, int, bool, error) {
|
||||
s.calls++
|
||||
if s.missFirst && s.calls == 1 {
|
||||
return 0, 0, false, nil
|
||||
}
|
||||
return s.pts, s.date, s.pts > 0, nil
|
||||
}
|
||||
|
||||
func (s *retentionCheckpointEvents) MaxContiguousPts(_ context.Context, _ int64) (int, error) {
|
||||
return s.current, nil
|
||||
}
|
||||
|
||||
func TestGetDifferenceBelowRetainedFloorUsesEmptySliceCheckpoint(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
authKeyID := [8]byte{13}
|
||||
const userID int64 = 1000000013
|
||||
base := memory.NewUpdateEventStore()
|
||||
events := &retentionCheckpointEvents{UpdateEventStore: base, pts: 2, date: 1700000202, current: 3}
|
||||
states := memory.NewUpdateStateStore()
|
||||
svc := NewService(states, events)
|
||||
// Retention already removed pts 1..2; only the live tail remains.
|
||||
if err := base.Append(ctx, userID, domain.UpdateEvent{
|
||||
UserID: userID, Type: domain.UpdateEventNoop, Pts: 3, PtsCount: 1, Date: 1700000203,
|
||||
}); err != nil {
|
||||
t.Fatalf("append live tail: %v", err)
|
||||
}
|
||||
|
||||
checkpoint, err := svc.GetDifference(ctx, authKeyID, userID, domain.UpdateState{Pts: 0, Date: 1700000200})
|
||||
if err != nil {
|
||||
t.Fatalf("difference below retained floor: %v", err)
|
||||
}
|
||||
if !checkpoint.Partial || len(checkpoint.Events) != 0 || checkpoint.State.Pts != 2 || checkpoint.State.Date != 1700000202 {
|
||||
t.Fatalf("checkpoint difference = %+v, want empty differenceSlice at pts/date 2/1700000202", checkpoint)
|
||||
}
|
||||
|
||||
tail, err := svc.GetDifference(ctx, authKeyID, userID, checkpoint.State)
|
||||
if err != nil {
|
||||
t.Fatalf("difference from retained floor: %v", err)
|
||||
}
|
||||
if tail.Partial || len(tail.Events) != 1 || tail.Events[0].Pts != 3 || tail.State.Pts != 3 {
|
||||
t.Fatalf("tail difference = %+v, want normal event pts=3", tail)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetDifferenceRechecksCheckpointWhenRetentionRacesEventRead(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
authKeyID := [8]byte{14}
|
||||
const userID int64 = 1000000014
|
||||
base := memory.NewUpdateEventStore()
|
||||
events := &retentionCheckpointEvents{
|
||||
UpdateEventStore: base,
|
||||
pts: 2,
|
||||
date: 1700000302,
|
||||
current: 3,
|
||||
missFirst: true,
|
||||
}
|
||||
if err := base.Append(ctx, userID, domain.UpdateEvent{
|
||||
UserID: userID, Type: domain.UpdateEventNoop, Pts: 3, PtsCount: 1, Date: 1700000303,
|
||||
}); err != nil {
|
||||
t.Fatalf("append live tail: %v", err)
|
||||
}
|
||||
|
||||
diff, err := NewService(memory.NewUpdateStateStore(), events).GetDifference(
|
||||
ctx,
|
||||
authKeyID,
|
||||
userID,
|
||||
domain.UpdateState{Pts: 0, Date: 1700000300},
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("difference across retention race: %v", err)
|
||||
}
|
||||
if events.calls != 2 {
|
||||
t.Fatalf("checkpoint probes = %d, want pre-read plus post-gap recheck", events.calls)
|
||||
}
|
||||
if !diff.Partial || len(diff.Events) != 0 || diff.State.Pts != 2 || diff.State.Date != 1700000302 {
|
||||
t.Fatalf("race checkpoint difference = %+v, want empty differenceSlice at retained floor", diff)
|
||||
}
|
||||
}
|
||||
|
||||
type captureDispatchAppender struct {
|
||||
|
|
@ -559,8 +699,9 @@ func (s *failingDispatchAppender) AppendAllocatedWithDispatch(context.Context, i
|
|||
}
|
||||
|
||||
type captureStateStore struct {
|
||||
saveCount int
|
||||
states map[[16]byte]domain.UpdateState
|
||||
saveCount int
|
||||
lastSaveAuthKeyID [8]byte
|
||||
states map[[16]byte]domain.UpdateState
|
||||
}
|
||||
|
||||
func (s *captureStateStore) Get(_ context.Context, authKeyID [8]byte, userID int64) (domain.UpdateState, bool, error) {
|
||||
|
|
@ -576,10 +717,15 @@ func (s *captureStateStore) Save(_ context.Context, authKeyID [8]byte, userID in
|
|||
s.states = make(map[[16]byte]domain.UpdateState)
|
||||
}
|
||||
s.saveCount++
|
||||
s.lastSaveAuthKeyID = authKeyID
|
||||
s.states[captureStateKey(authKeyID, userID)] = state
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *captureStateStore) ObserveClientState(_ context.Context, _ [8]byte, _ int64, _ domain.UpdateState) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *captureStateStore) Delete(_ context.Context, authKeyID [8]byte, userID int64) error {
|
||||
if s.states != nil {
|
||||
delete(s.states, captureStateKey(authKeyID, userID))
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue