feat: sync login email verification support
This commit is contained in:
parent
e0cabb4930
commit
9a501f900a
39 changed files with 2198 additions and 117 deletions
|
|
@ -5,14 +5,19 @@ import (
|
|||
"database/sql"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/jackc/pgerrcode"
|
||||
"github.com/jackc/pgx/v5"
|
||||
"github.com/jackc/pgx/v5/pgconn"
|
||||
|
||||
"telesrv/internal/domain"
|
||||
"telesrv/internal/store/postgres/sqlcgen"
|
||||
)
|
||||
|
||||
const accountPasswordsLoginEmailUniqueIdx = "account_passwords_login_email_lower_unique_idx"
|
||||
|
||||
// PasswordStore 用 PostgreSQL 实现 store.PasswordStore。
|
||||
type PasswordStore struct {
|
||||
db sqlcgen.DBTX
|
||||
|
|
@ -70,7 +75,29 @@ WHERE user_id = $1`, userID)
|
|||
return settings, true, nil
|
||||
}
|
||||
|
||||
func (s *PasswordStore) LoginEmailOwner(ctx context.Context, email string) (int64, bool, error) {
|
||||
email = normalizeStoredLoginEmail(email)
|
||||
if email == "" {
|
||||
return 0, false, nil
|
||||
}
|
||||
row := s.db.QueryRow(ctx, `
|
||||
SELECT user_id
|
||||
FROM account_passwords
|
||||
WHERE login_email <> '' AND lower(login_email) = $1
|
||||
LIMIT 1`, email)
|
||||
var userID int64
|
||||
if err := row.Scan(&userID); err != nil {
|
||||
if errors.Is(err, pgx.ErrNoRows) {
|
||||
return 0, false, nil
|
||||
}
|
||||
return 0, false, fmt.Errorf("get login email owner: %w", err)
|
||||
}
|
||||
return userID, true, nil
|
||||
}
|
||||
|
||||
func (s *PasswordStore) Save(ctx context.Context, userID int64, settings domain.PasswordSettings) error {
|
||||
settings.LoginEmail = normalizeStoredLoginEmail(settings.LoginEmail)
|
||||
settings.LoginEmailPattern = domain.MaskEmail(settings.LoginEmail)
|
||||
algo := settings.NewAlgo
|
||||
if settings.CurrentAlgo != nil {
|
||||
algo = *settings.CurrentAlgo
|
||||
|
|
@ -117,11 +144,25 @@ ON CONFLICT (user_id) DO UPDATE SET
|
|||
settings.RecoveryEmail, settings.RecoveryCode, recoveryExpires, settings.LoginEmail,
|
||||
)
|
||||
if err != nil {
|
||||
if isAccountPasswordLoginEmailUnique(err) {
|
||||
return domain.ErrEmailOccupied
|
||||
}
|
||||
return fmt.Errorf("upsert account password: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func normalizeStoredLoginEmail(email string) string {
|
||||
return strings.ToLower(strings.TrimSpace(email))
|
||||
}
|
||||
|
||||
func isAccountPasswordLoginEmailUnique(err error) bool {
|
||||
var pgErr *pgconn.PgError
|
||||
return errors.As(err, &pgErr) &&
|
||||
pgErr.Code == pgerrcode.UniqueViolation &&
|
||||
pgErr.ConstraintName == accountPasswordsLoginEmailUniqueIdx
|
||||
}
|
||||
|
||||
func nonNilBytea(in []byte) []byte {
|
||||
if in != nil {
|
||||
return in
|
||||
|
|
|
|||
|
|
@ -0,0 +1,40 @@
|
|||
package postgres
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"testing"
|
||||
|
||||
"telesrv/internal/domain"
|
||||
)
|
||||
|
||||
func TestPasswordStoreLoginEmailUniqueCaseInsensitivePostgres(t *testing.T) {
|
||||
pool := testPool(t)
|
||||
ctx := context.Background()
|
||||
passwords := NewPasswordStore(pool)
|
||||
users := NewUserStore(pool)
|
||||
suffix := randomSuffix(t)
|
||||
u1, err := users.Create(ctx, domain.User{AccessHash: 101, Phone: "+1665" + suffix + "01", FirstName: "EmailOne"})
|
||||
if err != nil {
|
||||
t.Fatalf("create user1: %v", err)
|
||||
}
|
||||
u2, err := users.Create(ctx, domain.User{AccessHash: 102, Phone: "+1665" + suffix + "02", FirstName: "EmailTwo"})
|
||||
if err != nil {
|
||||
t.Fatalf("create user2: %v", err)
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
_, _ = pool.Exec(ctx, "DELETE FROM account_passwords WHERE user_id = ANY($1::bigint[])", []int64{u1.ID, u2.ID})
|
||||
_, _ = pool.Exec(ctx, "DELETE FROM users WHERE id = ANY($1::bigint[])", []int64{u1.ID, u2.ID})
|
||||
})
|
||||
|
||||
if err := passwords.Save(ctx, u1.ID, domain.PasswordSettings{LoginEmail: "Owner@Example.Test"}); err != nil {
|
||||
t.Fatalf("save user1 email: %v", err)
|
||||
}
|
||||
ownerID, found, err := passwords.LoginEmailOwner(ctx, "owner@example.test")
|
||||
if err != nil || !found || ownerID != u1.ID {
|
||||
t.Fatalf("LoginEmailOwner = id %d found %v err %v, want user1", ownerID, found, err)
|
||||
}
|
||||
if err := passwords.Save(ctx, u2.ID, domain.PasswordSettings{LoginEmail: "owner@example.test"}); !errors.Is(err, domain.ErrEmailOccupied) {
|
||||
t.Fatalf("save duplicate email err = %v, want ErrEmailOccupied", err)
|
||||
}
|
||||
}
|
||||
|
|
@ -7,6 +7,7 @@ import (
|
|||
"fmt"
|
||||
|
||||
"github.com/jackc/pgx/v5"
|
||||
"github.com/jackc/pgx/v5/pgtype"
|
||||
|
||||
"telesrv/internal/store"
|
||||
"telesrv/internal/store/postgres/sqlcgen"
|
||||
|
|
@ -26,11 +27,12 @@ func NewAuthKeyStore(db sqlcgen.DBTX) *AuthKeyStore {
|
|||
// Save 实现 store.AuthKeyStore。auth_key_id 以小端解释为 int64 存入 BIGINT;
|
||||
// created_at 交由 DB 默认值(now()),故传入的 CreatedAt 不落库。
|
||||
func (s *AuthKeyStore) Save(ctx context.Context, k store.AuthKeyData) error {
|
||||
if err := s.q.UpsertAuthKey(ctx, sqlcgen.UpsertAuthKeyParams{
|
||||
AuthKeyID: authKeyIDToInt64(k.ID),
|
||||
Body: k.Value[:],
|
||||
ServerSalt: k.ServerSalt,
|
||||
}); err != nil {
|
||||
if _, err := s.db.Exec(ctx, `
|
||||
INSERT INTO auth_keys (auth_key_id, body, server_salt)
|
||||
VALUES ($1, $2, $3)
|
||||
ON CONFLICT (auth_key_id) DO UPDATE
|
||||
SET body = EXCLUDED.body, server_salt = EXCLUDED.server_salt
|
||||
`, authKeyIDToInt64(k.ID), k.Value[:], k.ServerSalt); err != nil {
|
||||
return fmt.Errorf("upsert auth key: %w", err)
|
||||
}
|
||||
return nil
|
||||
|
|
@ -38,24 +40,65 @@ func (s *AuthKeyStore) Save(ctx context.Context, k store.AuthKeyData) error {
|
|||
|
||||
// Get 实现 store.AuthKeyStore。不存在时 found=false。
|
||||
func (s *AuthKeyStore) Get(ctx context.Context, id [8]byte) (store.AuthKeyData, bool, error) {
|
||||
row, err := s.q.GetAuthKey(ctx, authKeyIDToInt64(id))
|
||||
var (
|
||||
body []byte
|
||||
serverSalt int64
|
||||
createdAt pgtype.Timestamptz
|
||||
layer int
|
||||
deviceModel string
|
||||
platform string
|
||||
systemVersion string
|
||||
apiID int
|
||||
appVersion string
|
||||
)
|
||||
err := s.db.QueryRow(ctx, `
|
||||
SELECT auth_key_id, body, server_salt, created_at,
|
||||
layer, device_model, platform, system_version, api_id, app_version
|
||||
FROM auth_keys
|
||||
WHERE auth_key_id = $1
|
||||
`, authKeyIDToInt64(id)).Scan(new(int64), &body, &serverSalt, &createdAt, &layer, &deviceModel, &platform, &systemVersion, &apiID, &appVersion)
|
||||
if err != nil {
|
||||
if errors.Is(err, pgx.ErrNoRows) {
|
||||
return store.AuthKeyData{}, false, nil
|
||||
}
|
||||
return store.AuthKeyData{}, false, fmt.Errorf("get auth key: %w", err)
|
||||
}
|
||||
if len(row.Body) != len(store.AuthKeyData{}.Value) {
|
||||
return store.AuthKeyData{}, false, fmt.Errorf("auth key body length = %d, want 256", len(row.Body))
|
||||
if len(body) != len(store.AuthKeyData{}.Value) {
|
||||
return store.AuthKeyData{}, false, fmt.Errorf("auth key body length = %d, want 256", len(body))
|
||||
}
|
||||
data := store.AuthKeyData{ID: id, ServerSalt: row.ServerSalt}
|
||||
copy(data.Value[:], row.Body)
|
||||
if row.CreatedAt.Valid {
|
||||
data.CreatedAt = row.CreatedAt.Time.Unix()
|
||||
data := store.AuthKeyData{
|
||||
ID: id,
|
||||
ServerSalt: serverSalt,
|
||||
Layer: layer,
|
||||
DeviceModel: deviceModel,
|
||||
Platform: platform,
|
||||
SystemVersion: systemVersion,
|
||||
APIID: apiID,
|
||||
AppVersion: appVersion,
|
||||
}
|
||||
copy(data.Value[:], body)
|
||||
if createdAt.Valid {
|
||||
data.CreatedAt = createdAt.Time.Unix()
|
||||
}
|
||||
return data, true, nil
|
||||
}
|
||||
|
||||
func (s *AuthKeyStore) UpdateClientInfo(ctx context.Context, id [8]byte, info store.AuthKeyClientInfo) error {
|
||||
if _, err := s.db.Exec(ctx, `
|
||||
UPDATE auth_keys
|
||||
SET layer = CASE WHEN $2::integer > 0 THEN $2 ELSE layer END,
|
||||
device_model = CASE WHEN $3::text <> '' THEN $3 ELSE device_model END,
|
||||
platform = CASE WHEN $4::text <> '' THEN $4 ELSE platform END,
|
||||
system_version = CASE WHEN $5::text <> '' THEN $5 ELSE system_version END,
|
||||
api_id = CASE WHEN $6::integer <> 0 THEN $6 ELSE api_id END,
|
||||
app_version = CASE WHEN $7::text <> '' THEN $7 ELSE app_version END
|
||||
WHERE auth_key_id = $1
|
||||
`, authKeyIDToInt64(id), info.Layer, info.DeviceModel, info.Platform, info.SystemVersion, info.APIID, info.AppVersion); err != nil {
|
||||
return fmt.Errorf("update auth key client info: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Delete 实现 store.AuthKeyStore。不存在时静默成功。
|
||||
// 手写 SQL 而非 sqlc 生成:避免触碰 sqlcgen 再生成链路。
|
||||
//
|
||||
|
|
|
|||
|
|
@ -70,3 +70,62 @@ func TestAuthKeyStoreRoundTrip(t *testing.T) {
|
|||
t.Fatalf("missing key: found=%v err=%v, want found=false err=nil", found, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthKeyStoreClientInfoRoundTrip(t *testing.T) {
|
||||
pool := testPool(t)
|
||||
ctx := context.Background()
|
||||
|
||||
var id [8]byte
|
||||
var val [256]byte
|
||||
if _, err := rand.Read(id[:]); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := rand.Read(val[:]); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
_, _ = pool.Exec(ctx, "DELETE FROM auth_keys WHERE auth_key_id = $1", authKeyIDToInt64(id))
|
||||
})
|
||||
|
||||
keys := NewAuthKeyStore(pool)
|
||||
if err := keys.Save(ctx, store.AuthKeyData{ID: id, Value: val, ServerSalt: 0x0badf00d}); err != nil {
|
||||
t.Fatalf("save: %v", err)
|
||||
}
|
||||
if err := keys.UpdateClientInfo(ctx, id, store.AuthKeyClientInfo{
|
||||
Layer: 227,
|
||||
DeviceModel: "GooglePixel 9a",
|
||||
Platform: "android",
|
||||
SystemVersion: "SDK 36",
|
||||
APIID: 6,
|
||||
AppVersion: "12.8.1 (69169) pbeta",
|
||||
}); err != nil {
|
||||
t.Fatalf("update client info: %v", err)
|
||||
}
|
||||
|
||||
got, found, err := NewAuthKeyStore(pool).Get(ctx, id)
|
||||
if err != nil {
|
||||
t.Fatalf("get: %v", err)
|
||||
}
|
||||
if !found {
|
||||
t.Fatal("auth key not found after client info update")
|
||||
}
|
||||
if got.Layer != 227 || got.DeviceModel != "GooglePixel 9a" || got.Platform != "android" ||
|
||||
got.SystemVersion != "SDK 36" || got.APIID != 6 || got.AppVersion != "12.8.1 (69169) pbeta" {
|
||||
t.Fatalf("client info mismatch: %+v", got)
|
||||
}
|
||||
|
||||
if err := keys.UpdateClientInfo(ctx, id, store.AuthKeyClientInfo{AppVersion: "12.8.2"}); err != nil {
|
||||
t.Fatalf("partial update client info: %v", err)
|
||||
}
|
||||
got, found, err = NewAuthKeyStore(pool).Get(ctx, id)
|
||||
if err != nil {
|
||||
t.Fatalf("get after partial update: %v", err)
|
||||
}
|
||||
if !found {
|
||||
t.Fatal("auth key not found after partial client info update")
|
||||
}
|
||||
if got.Layer != 227 || got.DeviceModel != "GooglePixel 9a" || got.Platform != "android" ||
|
||||
got.SystemVersion != "SDK 36" || got.APIID != 6 || got.AppVersion != "12.8.2" {
|
||||
t.Fatalf("partial client info merge mismatch: %+v", got)
|
||||
}
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue