feat: sync login email verification support

This commit is contained in:
A 2026-07-08 17:09:44 +08:00
parent e0cabb4930
commit 9a501f900a
39 changed files with 2198 additions and 117 deletions

View file

@ -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

View file

@ -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)
}
}

View file

@ -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 再生成链路。
//

View file

@ -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)
}
}