feat: sync login email verification support
This commit is contained in:
parent
e0cabb4930
commit
9a501f900a
39 changed files with 2198 additions and 117 deletions
17
.env.example
17
.env.example
|
|
@ -5,6 +5,23 @@ TELESRV_LISTEN=0.0.0.0:2398
|
||||||
TELESRV_ADVERTISE_IP=127.0.0.1
|
TELESRV_ADVERTISE_IP=127.0.0.1
|
||||||
TELESRV_DC=2
|
TELESRV_DC=2
|
||||||
TELESRV_DEV_AUTH_CODE=12345
|
TELESRV_DEV_AUTH_CODE=12345
|
||||||
|
TELESRV_AUTH_CODE_TTL=5m
|
||||||
|
TELESRV_AUTH_CODE_MAX_ATTEMPTS=5
|
||||||
|
|
||||||
|
# Optional login-email verification. When enabled, accounts with a confirmed
|
||||||
|
# login email receive login codes by email; REQUIRE_SETUP also forces new/legacy
|
||||||
|
# accounts without a login email to set one during the phone login flow.
|
||||||
|
TELESRV_LOGIN_EMAIL_ENABLE=false
|
||||||
|
TELESRV_LOGIN_EMAIL_REQUIRE_SETUP=false
|
||||||
|
TELESRV_LOGIN_EMAIL_CODE_LENGTH=6
|
||||||
|
TELESRV_SMTP_HOST=
|
||||||
|
TELESRV_SMTP_PORT=587
|
||||||
|
TELESRV_SMTP_USERNAME=
|
||||||
|
TELESRV_SMTP_PASSWORD=
|
||||||
|
TELESRV_SMTP_FROM=
|
||||||
|
TELESRV_SMTP_FROM_NAME=telesrv
|
||||||
|
TELESRV_SMTP_TLS=starttls
|
||||||
|
TELESRV_SMTP_TIMEOUT=10s
|
||||||
|
|
||||||
# Client-visible telesrv links. Production uses https://telesrv.net.
|
# Client-visible telesrv links. Production uses https://telesrv.net.
|
||||||
# For local link/deeplink smoke tests use http://127.0.0.1:2401.
|
# For local link/deeplink smoke tests use http://127.0.0.1:2401.
|
||||||
|
|
|
||||||
|
|
@ -104,6 +104,10 @@ Useful local environment variables:
|
||||||
| `TELESRV_ADVERTISE_IP` | `127.0.0.1` | IP advertised to compatible clients |
|
| `TELESRV_ADVERTISE_IP` | `127.0.0.1` | IP advertised to compatible clients |
|
||||||
| `TELESRV_DC` | `2` | self-hosted DC id |
|
| `TELESRV_DC` | `2` | self-hosted DC id |
|
||||||
| `TELESRV_DEV_AUTH_CODE` | `12345` | fixed login code for local development |
|
| `TELESRV_DEV_AUTH_CODE` | `12345` | fixed login code for local development |
|
||||||
|
| `TELESRV_AUTH_CODE_MAX_ATTEMPTS` | `5` | wrong-code attempts before the code hash is deleted |
|
||||||
|
| `TELESRV_LOGIN_EMAIL_ENABLE` | `false` | send login codes to confirmed login email addresses through SMTP |
|
||||||
|
| `TELESRV_LOGIN_EMAIL_REQUIRE_SETUP` | `false` | force phone login/registration to set a login email first |
|
||||||
|
| `TELESRV_SMTP_HOST` | empty | SMTP host used when login email verification is enabled |
|
||||||
| `TELESRV_PUBLIC_BASE_URL` | `https://telesrv.net` | canonical base URL for public sticker/chatlist links |
|
| `TELESRV_PUBLIC_BASE_URL` | `https://telesrv.net` | canonical base URL for public sticker/chatlist links |
|
||||||
| `TELESRV_POSTGRES_DSN` | local Compose DSN | PostgreSQL connection string |
|
| `TELESRV_POSTGRES_DSN` | local Compose DSN | PostgreSQL connection string |
|
||||||
| `TELESRV_REDIS_ADDR` | `127.0.0.1:6399` | Redis address |
|
| `TELESRV_REDIS_ADDR` | `127.0.0.1:6399` | Redis address |
|
||||||
|
|
|
||||||
|
|
@ -96,6 +96,10 @@ go build -o bin/gramsrv ./cmd/telesrv
|
||||||
| `TELESRV_ADVERTISE_IP` | `127.0.0.1` | 下发给兼容客户端的连接 IP |
|
| `TELESRV_ADVERTISE_IP` | `127.0.0.1` | 下发给兼容客户端的连接 IP |
|
||||||
| `TELESRV_DC` | `2` | 自建 DC id |
|
| `TELESRV_DC` | `2` | 自建 DC id |
|
||||||
| `TELESRV_DEV_AUTH_CODE` | `12345` | 本地开发固定登录验证码 |
|
| `TELESRV_DEV_AUTH_CODE` | `12345` | 本地开发固定登录验证码 |
|
||||||
|
| `TELESRV_AUTH_CODE_MAX_ATTEMPTS` | `5` | 同一验证码 hash 允许的错误次数,达到后删除并要求重发 |
|
||||||
|
| `TELESRV_LOGIN_EMAIL_ENABLE` | `false` | 已绑定登录邮箱的账号通过 SMTP 接收登录验证码 |
|
||||||
|
| `TELESRV_LOGIN_EMAIL_REQUIRE_SETUP` | `false` | 登录/注册时强制先设置登录邮箱 |
|
||||||
|
| `TELESRV_SMTP_HOST` | 空 | 开启登录邮箱验证时使用的 SMTP host |
|
||||||
| `TELESRV_PUBLIC_BASE_URL` | `https://telesrv.net` | sticker/chatlist 公开链接的 canonical base URL |
|
| `TELESRV_PUBLIC_BASE_URL` | `https://telesrv.net` | sticker/chatlist 公开链接的 canonical base URL |
|
||||||
| `TELESRV_POSTGRES_DSN` | local Compose DSN | PostgreSQL 连接串 |
|
| `TELESRV_POSTGRES_DSN` | local Compose DSN | PostgreSQL 连接串 |
|
||||||
| `TELESRV_REDIS_ADDR` | `127.0.0.1:6399` | Redis 地址 |
|
| `TELESRV_REDIS_ADDR` | `127.0.0.1:6399` | Redis 地址 |
|
||||||
|
|
|
||||||
|
|
@ -54,6 +54,7 @@ import (
|
||||||
"telesrv/internal/botapi"
|
"telesrv/internal/botapi"
|
||||||
"telesrv/internal/config"
|
"telesrv/internal/config"
|
||||||
"telesrv/internal/domain"
|
"telesrv/internal/domain"
|
||||||
|
mailpkg "telesrv/internal/mail"
|
||||||
"telesrv/internal/mtprotoedge"
|
"telesrv/internal/mtprotoedge"
|
||||||
"telesrv/internal/rpc"
|
"telesrv/internal/rpc"
|
||||||
"telesrv/internal/seed/catalog"
|
"telesrv/internal/seed/catalog"
|
||||||
|
|
@ -464,7 +465,7 @@ func run(logger *zap.Logger) error {
|
||||||
// userCache 与 users 服务共享同一实例:bot 元数据写入(version bump)后必须
|
// userCache 与 users 服务共享同一实例:bot 元数据写入(version bump)后必须
|
||||||
// 失效缓存,否则 TTL 内 getUsers 回旧 first_name/旧 bot_info_version。
|
// 失效缓存,否则 TTL 内 getUsers 回旧 first_name/旧 bot_info_version。
|
||||||
userCache := redisstore.NewUserCache(rdb, redisstore.DefaultUserCacheTTL)
|
userCache := redisstore.NewUserCache(rdb, redisstore.DefaultUserCacheTTL)
|
||||||
accountService := account.NewService(passwordStore,
|
accountOptions := []account.ServiceOption{
|
||||||
account.WithReactionSettings(passwordStore),
|
account.WithReactionSettings(passwordStore),
|
||||||
account.WithAccountSettings(passwordStore),
|
account.WithAccountSettings(passwordStore),
|
||||||
account.WithNotifySettings(passwordStore),
|
account.WithNotifySettings(passwordStore),
|
||||||
|
|
@ -473,7 +474,24 @@ func run(logger *zap.Logger) error {
|
||||||
account.WithSavedMusic(passwordStore),
|
account.WithSavedMusic(passwordStore),
|
||||||
account.WithBusinessAutomation(passwordStore),
|
account.WithBusinessAutomation(passwordStore),
|
||||||
account.WithUsers(userStore),
|
account.WithUsers(userStore),
|
||||||
account.WithPublicBaseURL(cfg.PublicBaseURL))
|
account.WithPublicBaseURL(cfg.PublicBaseURL),
|
||||||
|
}
|
||||||
|
var loginEmailSender mailpkg.Sender
|
||||||
|
if cfg.LoginEmailEnable {
|
||||||
|
loginEmailSender = mailpkg.NewSMTP(mailpkg.Config{
|
||||||
|
Host: cfg.SMTPHost,
|
||||||
|
Port: cfg.SMTPPort,
|
||||||
|
Username: cfg.SMTPUsername,
|
||||||
|
Password: cfg.SMTPPassword,
|
||||||
|
From: cfg.SMTPFrom,
|
||||||
|
FromName: cfg.SMTPFromName,
|
||||||
|
TLSMode: cfg.SMTPTLSMode,
|
||||||
|
Timeout: cfg.SMTPTimeout,
|
||||||
|
})
|
||||||
|
accountOptions = append(accountOptions,
|
||||||
|
account.WithLoginEmailVerification(codeStore, loginEmailSender, cfg.AuthCodeTTL, cfg.AuthCodeMaxAttempts, cfg.LoginEmailCodeLength))
|
||||||
|
}
|
||||||
|
accountService := account.NewService(passwordStore, accountOptions...)
|
||||||
botsService := botsapp.NewService(userStore, botStore, messageStore,
|
botsService := botsapp.NewService(userStore, botStore, messageStore,
|
||||||
botsapp.WithLogger(logger.Named("bots")),
|
botsapp.WithLogger(logger.Named("bots")),
|
||||||
botsapp.WithBlockChecker(contactStore),
|
botsapp.WithBlockChecker(contactStore),
|
||||||
|
|
@ -611,7 +629,20 @@ func run(logger *zap.Logger) error {
|
||||||
messageapp.WithSendPermissionChecker(adminService),
|
messageapp.WithSendPermissionChecker(adminService),
|
||||||
messageapp.WithBusinessAutomation(passwordStore, businessAutomationOptions...),
|
messageapp.WithBusinessAutomation(passwordStore, businessAutomationOptions...),
|
||||||
)
|
)
|
||||||
authService := auth.NewService(userStore, authzStore, codeStore, authKeyStore, tempAuthKeyStore, cfg.DevAuthCode, auth.WithLoginMessages(messageStore, dialogStore), auth.WithPasswords(passwordStore), auth.WithBotLogin(botStore), auth.WithPremiumGrant(cfg.PremiumGrantMonths))
|
authService := auth.NewService(userStore, authzStore, codeStore, authKeyStore, tempAuthKeyStore, cfg.DevAuthCode,
|
||||||
|
auth.WithLoginMessages(messageStore, dialogStore),
|
||||||
|
auth.WithPasswords(passwordStore),
|
||||||
|
auth.WithBotLogin(botStore),
|
||||||
|
auth.WithPremiumGrant(cfg.PremiumGrantMonths),
|
||||||
|
auth.WithCodeTTL(cfg.AuthCodeTTL),
|
||||||
|
auth.WithCodeMaxAttempts(cfg.AuthCodeMaxAttempts),
|
||||||
|
auth.WithLoginEmail(auth.LoginEmailOptions{
|
||||||
|
Enabled: cfg.LoginEmailEnable,
|
||||||
|
RequireSetup: cfg.LoginEmailRequireSetup,
|
||||||
|
CodeLength: cfg.LoginEmailCodeLength,
|
||||||
|
Store: accountService,
|
||||||
|
Sender: loginEmailSender,
|
||||||
|
}))
|
||||||
updatesService := updates.NewService(updateStateStore, updateEventStore, updates.WithLogger(logger.Named("app").Named("updates")))
|
updatesService := updates.NewService(updateStateStore, updateEventStore, updates.WithLogger(logger.Named("app").Named("updates")))
|
||||||
router := rpc.New(rpc.Config{
|
router := rpc.New(rpc.Config{
|
||||||
DC: cfg.DC,
|
DC: cfg.DC,
|
||||||
|
|
|
||||||
1
deploy/migrations/0058_login_email_unique.down.sql
Normal file
1
deploy/migrations/0058_login_email_unique.down.sql
Normal file
|
|
@ -0,0 +1 @@
|
||||||
|
DROP INDEX IF EXISTS public.account_passwords_login_email_lower_unique_idx;
|
||||||
3
deploy/migrations/0058_login_email_unique.up.sql
Normal file
3
deploy/migrations/0058_login_email_unique.up.sql
Normal file
|
|
@ -0,0 +1,3 @@
|
||||||
|
CREATE UNIQUE INDEX account_passwords_login_email_lower_unique_idx
|
||||||
|
ON public.account_passwords (lower((login_email)::text))
|
||||||
|
WHERE ((login_email)::text <> ''::text);
|
||||||
7
deploy/migrations/0059_auth_key_client_info.down.sql
Normal file
7
deploy/migrations/0059_auth_key_client_info.down.sql
Normal file
|
|
@ -0,0 +1,7 @@
|
||||||
|
ALTER TABLE auth_keys
|
||||||
|
DROP COLUMN IF EXISTS app_version,
|
||||||
|
DROP COLUMN IF EXISTS api_id,
|
||||||
|
DROP COLUMN IF EXISTS system_version,
|
||||||
|
DROP COLUMN IF EXISTS platform,
|
||||||
|
DROP COLUMN IF EXISTS device_model,
|
||||||
|
DROP COLUMN IF EXISTS layer;
|
||||||
7
deploy/migrations/0059_auth_key_client_info.up.sql
Normal file
7
deploy/migrations/0059_auth_key_client_info.up.sql
Normal file
|
|
@ -0,0 +1,7 @@
|
||||||
|
ALTER TABLE auth_keys
|
||||||
|
ADD COLUMN layer integer NOT NULL DEFAULT 0,
|
||||||
|
ADD COLUMN device_model varchar(128) NOT NULL DEFAULT '',
|
||||||
|
ADD COLUMN platform varchar(64) NOT NULL DEFAULT '',
|
||||||
|
ADD COLUMN system_version varchar(64) NOT NULL DEFAULT '',
|
||||||
|
ADD COLUMN api_id integer NOT NULL DEFAULT 0,
|
||||||
|
ADD COLUMN app_version varchar(64) NOT NULL DEFAULT '';
|
||||||
|
|
@ -4,8 +4,10 @@ import (
|
||||||
"context"
|
"context"
|
||||||
"errors"
|
"errors"
|
||||||
"testing"
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
"telesrv/internal/domain"
|
"telesrv/internal/domain"
|
||||||
|
"telesrv/internal/store"
|
||||||
"telesrv/internal/store/memory"
|
"telesrv/internal/store/memory"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
@ -25,6 +27,17 @@ func createUser(t *testing.T, users *memory.UserStore, phone string) domain.User
|
||||||
return u
|
return u
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type captureMailSender struct {
|
||||||
|
to string
|
||||||
|
code string
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *captureMailSender) SendLoginCode(_ context.Context, to, code string, _ time.Duration) error {
|
||||||
|
s.to = to
|
||||||
|
s.code = code
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
// TestSetLoginEmailPersistsAndMasks 设置登录邮箱后,GetPassword 下发掩码 pattern,原始
|
// TestSetLoginEmailPersistsAndMasks 设置登录邮箱后,GetPassword 下发掩码 pattern,原始
|
||||||
// 地址只在 LoginEmail 读路径可见。
|
// 地址只在 LoginEmail 读路径可见。
|
||||||
func TestSetLoginEmailPersistsAndMasks(t *testing.T) {
|
func TestSetLoginEmailPersistsAndMasks(t *testing.T) {
|
||||||
|
|
@ -88,6 +101,24 @@ func TestSetLoginEmailRejectsInvalid(t *testing.T) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestSetLoginEmailRejectsDuplicateCaseInsensitive(t *testing.T) {
|
||||||
|
ctx := context.Background()
|
||||||
|
svc, users := newLoginEmailService(t)
|
||||||
|
u1 := createUser(t, users, "15550010103")
|
||||||
|
u2 := createUser(t, users, "15550010104")
|
||||||
|
|
||||||
|
if err := svc.SetLoginEmail(ctx, u1.ID, "Alice@Example.Test"); err != nil {
|
||||||
|
t.Fatalf("SetLoginEmail user1: %v", err)
|
||||||
|
}
|
||||||
|
if err := svc.SetLoginEmail(ctx, u2.ID, "alice@example.test"); !errors.Is(err, domain.ErrEmailOccupied) {
|
||||||
|
t.Fatalf("SetLoginEmail duplicate err = %v, want ErrEmailOccupied", err)
|
||||||
|
}
|
||||||
|
email, found, err := svc.LoginEmail(ctx, u1.ID)
|
||||||
|
if err != nil || !found || email != "alice@example.test" {
|
||||||
|
t.Fatalf("LoginEmail user1 = %q found=%v err=%v", email, found, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// TestRecoveryEmailDoesNotLeakIntoLoginEmailPattern 是核心解耦回归:设置 2FA 恢复邮箱
|
// TestRecoveryEmailDoesNotLeakIntoLoginEmailPattern 是核心解耦回归:设置 2FA 恢复邮箱
|
||||||
// 不得把恢复邮箱掩码写进 login_email_pattern(历史 bug)。
|
// 不得把恢复邮箱掩码写进 login_email_pattern(历史 bug)。
|
||||||
func TestRecoveryEmailDoesNotLeakIntoLoginEmailPattern(t *testing.T) {
|
func TestRecoveryEmailDoesNotLeakIntoLoginEmailPattern(t *testing.T) {
|
||||||
|
|
@ -114,3 +145,149 @@ func TestRecoveryEmailDoesNotLeakIntoLoginEmailPattern(t *testing.T) {
|
||||||
t.Fatal("HasRecovery = false, want true after setting recovery email")
|
t.Fatal("HasRecovery = false, want true after setting recovery email")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestSendLoginEmailCodeRejectsDuplicateBeforeSending(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))
|
||||||
|
u1 := createUser(t, users, "15550010105")
|
||||||
|
u2 := createUser(t, users, "15550010106")
|
||||||
|
if err := svc.SetLoginEmail(ctx, u1.ID, "taken@example.test"); err != nil {
|
||||||
|
t.Fatalf("SetLoginEmail user1: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if _, _, err := svc.SendLoginEmailCode(ctx, u2.ID, "", "", "TAKEN@example.test", false); !errors.Is(err, domain.ErrEmailOccupied) {
|
||||||
|
t.Fatalf("SendLoginEmailCode duplicate err = %v, want ErrEmailOccupied", err)
|
||||||
|
}
|
||||||
|
if sender.to != "" || sender.code != "" {
|
||||||
|
t.Fatalf("duplicate email sent to=%q code=%q, want no send", sender.to, sender.code)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLoginEmailSetupRejectsAlreadyOwnedEmailForNewPhone(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))
|
||||||
|
owner := createUser(t, users, "15550010107")
|
||||||
|
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 {
|
||||||
|
t.Fatalf("seed phone code: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if _, _, err := svc.SendLoginEmailCode(ctx, 0, "+1 555 001 0108", "new-phone-hash", "OWNER@example.test", true); !errors.Is(err, domain.ErrEmailOccupied) {
|
||||||
|
t.Fatalf("setup duplicate email err = %v, want ErrEmailOccupied", err)
|
||||||
|
}
|
||||||
|
if sender.to != "" || sender.code != "" {
|
||||||
|
t.Fatalf("duplicate setup email sent to=%q code=%q, want no send", sender.to, sender.code)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSendVerifyLoginEmailPersistsOnlyAfterVerify(t *testing.T) {
|
||||||
|
ctx := context.Background()
|
||||||
|
users := memory.NewUserStore()
|
||||||
|
codes := memory.NewCodeStore()
|
||||||
|
sender := &captureMailSender{}
|
||||||
|
svc := NewService(memory.NewPasswordStore(),
|
||||||
|
WithUsers(users),
|
||||||
|
WithLoginEmailVerification(codes, sender, time.Minute, 2, 6))
|
||||||
|
u := createUser(t, users, "15550010005")
|
||||||
|
|
||||||
|
pattern, length, err := svc.SendLoginEmailCode(ctx, u.ID, "", "", "alice@example.test", false)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("SendLoginEmailCode: %v", err)
|
||||||
|
}
|
||||||
|
if pattern != "a***e@example.test" || length != 6 || sender.to != "alice@example.test" || len(sender.code) != 6 {
|
||||||
|
t.Fatalf("send result pattern=%q length=%d to=%q code=%q", pattern, length, sender.to, sender.code)
|
||||||
|
}
|
||||||
|
if _, found, err := svc.LoginEmail(ctx, u.ID); err != nil || found {
|
||||||
|
t.Fatalf("LoginEmail before verify found=%v err=%v, want not found", found, err)
|
||||||
|
}
|
||||||
|
email, err := svc.VerifyLoginEmail(ctx, u.ID, "", "", sender.code, false)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("VerifyLoginEmail: %v", err)
|
||||||
|
}
|
||||||
|
if email != "alice@example.test" {
|
||||||
|
t.Fatalf("verified email = %q", email)
|
||||||
|
}
|
||||||
|
got, found, err := svc.LoginEmail(ctx, u.ID)
|
||||||
|
if err != nil || !found || got != "alice@example.test" {
|
||||||
|
t.Fatalf("LoginEmail after verify = %q found=%v err=%v", got, found, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLoginEmailSetupStoresPendingEmailOnPhoneCodeHash(t *testing.T) {
|
||||||
|
ctx := context.Background()
|
||||||
|
codes := memory.NewCodeStore()
|
||||||
|
sender := &captureMailSender{}
|
||||||
|
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 {
|
||||||
|
t.Fatalf("seed phone code: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if _, _, err := svc.SendLoginEmailCode(ctx, 0, "+1 555 001 0006", "phone-hash", "new@example.test", true); err != nil {
|
||||||
|
t.Fatalf("SendLoginEmailCode setup: %v", err)
|
||||||
|
}
|
||||||
|
email, err := svc.VerifyLoginEmail(ctx, 0, "+1 555 001 0006", "phone-hash", sender.code, true)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("VerifyLoginEmail setup: %v", err)
|
||||||
|
}
|
||||||
|
if email != "new@example.test" {
|
||||||
|
t.Fatalf("verified setup email = %q", email)
|
||||||
|
}
|
||||||
|
rec, found, err := codes.Get(ctx, "phone-hash")
|
||||||
|
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" {
|
||||||
|
t.Fatalf("phone code after verify = %+v", rec)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestVerifyLoginEmailDeletesCodeAfterMaxAttempts(t *testing.T) {
|
||||||
|
ctx := context.Background()
|
||||||
|
users := memory.NewUserStore()
|
||||||
|
codes := memory.NewCodeStore()
|
||||||
|
sender := &captureMailSender{}
|
||||||
|
svc := NewService(memory.NewPasswordStore(),
|
||||||
|
WithUsers(users),
|
||||||
|
WithLoginEmailVerification(codes, sender, time.Minute, 2, 6))
|
||||||
|
u := createUser(t, users, "15550010007")
|
||||||
|
|
||||||
|
if _, _, err := svc.SendLoginEmailCode(ctx, u.ID, "", "", "limit@example.test", false); err != nil {
|
||||||
|
t.Fatalf("SendLoginEmailCode: %v", err)
|
||||||
|
}
|
||||||
|
bad1 := "000000"
|
||||||
|
if bad1 == sender.code {
|
||||||
|
bad1 = "111111"
|
||||||
|
}
|
||||||
|
bad2 := "222222"
|
||||||
|
if bad2 == sender.code {
|
||||||
|
bad2 = "333333"
|
||||||
|
}
|
||||||
|
if _, err := svc.VerifyLoginEmail(ctx, u.ID, "", "", bad1, false); !errors.Is(err, domain.ErrEmailCodeInvalid) {
|
||||||
|
t.Fatalf("first bad VerifyLoginEmail err = %v, want ErrEmailCodeInvalid", err)
|
||||||
|
}
|
||||||
|
if _, err := svc.VerifyLoginEmail(ctx, u.ID, "", "", bad2, false); !errors.Is(err, domain.ErrEmailCodeInvalid) {
|
||||||
|
t.Fatalf("second bad VerifyLoginEmail err = %v, want ErrEmailCodeInvalid", err)
|
||||||
|
}
|
||||||
|
if _, err := svc.VerifyLoginEmail(ctx, u.ID, "", "", sender.code, false); !errors.Is(err, domain.ErrEmailCodeInvalid) {
|
||||||
|
t.Fatalf("VerifyLoginEmail after max attempts err = %v, want ErrEmailCodeInvalid", err)
|
||||||
|
}
|
||||||
|
if _, found, _ := svc.LoginEmail(ctx, u.ID); found {
|
||||||
|
t.Fatal("login email was set after exhausted verification code")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -4,19 +4,26 @@ import (
|
||||||
"context"
|
"context"
|
||||||
"crypto/rand"
|
"crypto/rand"
|
||||||
"crypto/subtle"
|
"crypto/subtle"
|
||||||
|
"fmt"
|
||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"telesrv/internal/domain"
|
"telesrv/internal/domain"
|
||||||
"telesrv/internal/links"
|
"telesrv/internal/links"
|
||||||
|
"telesrv/internal/mail"
|
||||||
"telesrv/internal/store"
|
"telesrv/internal/store"
|
||||||
)
|
)
|
||||||
|
|
||||||
var defaultSecureRandom = []byte("telesrv-tdesktop-dev-secure-rand")
|
var defaultSecureRandom = []byte("telesrv-tdesktop-dev-secure-rand")
|
||||||
|
|
||||||
const (
|
const (
|
||||||
passwordResetWait = 7 * 24 * time.Hour
|
passwordResetWait = 7 * 24 * time.Hour
|
||||||
passwordResetRetry = 24 * time.Hour
|
passwordResetRetry = 24 * time.Hour
|
||||||
|
loginEmailVerifyChangePrefix = "login-email-change:"
|
||||||
|
loginEmailVerifySetupPrefix = "login-email-setup:"
|
||||||
|
codeChannelEmailSetup = "email_setup"
|
||||||
|
codeChannelEmailChange = "email_change"
|
||||||
|
codeChannelEmailLogin = "email_login"
|
||||||
)
|
)
|
||||||
|
|
||||||
// Service 提供账号安全配置查询。
|
// Service 提供账号安全配置查询。
|
||||||
|
|
@ -30,8 +37,13 @@ type Service struct {
|
||||||
savedMusic store.SavedMusicStore
|
savedMusic store.SavedMusicStore
|
||||||
business store.BusinessAutomationStore
|
business store.BusinessAutomationStore
|
||||||
// users 仅用于登录邮箱的 phone→user 解析(sendCode 检测 / login-setup / reset 走 phone)。
|
// users 仅用于登录邮箱的 phone→user 解析(sendCode 检测 / login-setup / reset 走 phone)。
|
||||||
users store.UserStore
|
users store.UserStore
|
||||||
publicBaseURL string
|
publicBaseURL string
|
||||||
|
codes store.CodeStore
|
||||||
|
loginEmailSender mail.Sender
|
||||||
|
loginEmailCodeTTL time.Duration
|
||||||
|
loginEmailCodeMaxAttempts int
|
||||||
|
loginEmailCodeLength int
|
||||||
}
|
}
|
||||||
|
|
||||||
// ServiceOption 调整 account 服务依赖。
|
// ServiceOption 调整 account 服务依赖。
|
||||||
|
|
@ -99,9 +111,25 @@ func WithPublicBaseURL(baseURL string) ServiceOption {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func WithLoginEmailVerification(codes store.CodeStore, sender mail.Sender, ttl time.Duration, maxAttempts, length int) ServiceOption {
|
||||||
|
return func(s *Service) {
|
||||||
|
s.codes = codes
|
||||||
|
s.loginEmailSender = sender
|
||||||
|
if ttl > 0 {
|
||||||
|
s.loginEmailCodeTTL = ttl
|
||||||
|
}
|
||||||
|
if maxAttempts > 0 {
|
||||||
|
s.loginEmailCodeMaxAttempts = maxAttempts
|
||||||
|
}
|
||||||
|
if length > 0 {
|
||||||
|
s.loginEmailCodeLength = length
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// NewService 创建 account 服务。
|
// NewService 创建 account 服务。
|
||||||
func NewService(passwords store.PasswordStore, opts ...ServiceOption) *Service {
|
func NewService(passwords store.PasswordStore, opts ...ServiceOption) *Service {
|
||||||
s := &Service{passwords: passwords, publicBaseURL: links.DefaultPublicBaseURL}
|
s := &Service{passwords: passwords, publicBaseURL: links.DefaultPublicBaseURL, loginEmailCodeTTL: 5 * time.Minute, loginEmailCodeMaxAttempts: 5, loginEmailCodeLength: 6}
|
||||||
for _, opt := range opts {
|
for _, opt := range opts {
|
||||||
opt(s)
|
opt(s)
|
||||||
}
|
}
|
||||||
|
|
@ -200,6 +228,7 @@ func normalizePasswordSettings(settings domain.PasswordSettings) domain.Password
|
||||||
}
|
}
|
||||||
// login_email_pattern 始终从已确认的登录邮箱派生,与 2FA 恢复邮箱 RecoveryEmail
|
// login_email_pattern 始终从已确认的登录邮箱派生,与 2FA 恢复邮箱 RecoveryEmail
|
||||||
// 解耦(历史实现曾把恢复邮箱掩码误写进此字段,导致客户端把恢复邮箱当成登录邮箱显示)。
|
// 解耦(历史实现曾把恢复邮箱掩码误写进此字段,导致客户端把恢复邮箱当成登录邮箱显示)。
|
||||||
|
settings.LoginEmail = normalizeLoginEmail(settings.LoginEmail)
|
||||||
settings.LoginEmailPattern = emailPattern(settings.LoginEmail)
|
settings.LoginEmailPattern = emailPattern(settings.LoginEmail)
|
||||||
return settings
|
return settings
|
||||||
}
|
}
|
||||||
|
|
@ -467,25 +496,195 @@ func randomInt64() (int64, error) {
|
||||||
return out, nil
|
return out, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func randomDigits(n int) (string, error) {
|
||||||
|
if n <= 0 {
|
||||||
|
n = 6
|
||||||
|
}
|
||||||
|
b := make([]byte, n)
|
||||||
|
if _, err := rand.Read(b); err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
var out strings.Builder
|
||||||
|
out.Grow(n)
|
||||||
|
for _, v := range b {
|
||||||
|
out.WriteByte(byte('0') + v%10)
|
||||||
|
}
|
||||||
|
return out.String(), nil
|
||||||
|
}
|
||||||
|
|
||||||
func emailPattern(email string) string {
|
func emailPattern(email string) string {
|
||||||
return domain.MaskEmail(email)
|
return domain.MaskEmail(email)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func normalizeLoginEmail(email string) string {
|
||||||
|
return strings.ToLower(strings.TrimSpace(email))
|
||||||
|
}
|
||||||
|
|
||||||
// validLoginEmail 是登录邮箱的最小校验:非空且含 '@'。开发环境不做更严格的 RFC 校验。
|
// validLoginEmail 是登录邮箱的最小校验:非空且含 '@'。开发环境不做更严格的 RFC 校验。
|
||||||
func validLoginEmail(email string) bool {
|
func validLoginEmail(email string) bool {
|
||||||
|
email = normalizeLoginEmail(email)
|
||||||
return email != "" && strings.Contains(email, "@")
|
return email != "" && strings.Contains(email, "@")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (s *Service) SendLoginEmailCode(ctx context.Context, userID int64, phone, phoneCodeHash, email string, setup bool) (string, int, error) {
|
||||||
|
email = normalizeLoginEmail(email)
|
||||||
|
if !validLoginEmail(email) {
|
||||||
|
return "", 0, domain.ErrEmailInvalid
|
||||||
|
}
|
||||||
|
if s == nil || s.codes == nil || s.loginEmailSender == nil {
|
||||||
|
return "", 0, domain.ErrEmailNotAllowed
|
||||||
|
}
|
||||||
|
key := loginEmailVerifyChangePrefix + fmt.Sprint(userID)
|
||||||
|
rec := store.PhoneCode{
|
||||||
|
Code: "",
|
||||||
|
Channel: codeChannelEmailChange,
|
||||||
|
PendingEmail: email,
|
||||||
|
MaxAttempts: s.loginEmailCodeMaxAttempts,
|
||||||
|
}
|
||||||
|
if setup {
|
||||||
|
phone = domain.NormalizePhone(phone)
|
||||||
|
phoneRec, found, err := s.codes.Get(ctx, phoneCodeHash)
|
||||||
|
if err != nil {
|
||||||
|
return "", 0, err
|
||||||
|
}
|
||||||
|
if !found {
|
||||||
|
return "", 0, domain.ErrEmailCodeInvalid
|
||||||
|
}
|
||||||
|
if phoneRec.Phone != phone {
|
||||||
|
return "", 0, domain.ErrEmailInvalid
|
||||||
|
}
|
||||||
|
targetUserID := int64(0)
|
||||||
|
if existingUserID, found, err := s.userIDByPhone(ctx, phone); err != nil {
|
||||||
|
return "", 0, err
|
||||||
|
} else if found {
|
||||||
|
targetUserID = existingUserID
|
||||||
|
}
|
||||||
|
if err := s.ensureLoginEmailAvailable(ctx, targetUserID, email); err != nil {
|
||||||
|
return "", 0, err
|
||||||
|
}
|
||||||
|
key = loginEmailVerifySetupPrefix + phoneCodeHash
|
||||||
|
rec.Phone = phone
|
||||||
|
rec.Channel = codeChannelEmailSetup
|
||||||
|
} else if userID == 0 {
|
||||||
|
return "", 0, domain.ErrEmailInvalid
|
||||||
|
} else if err := s.ensureLoginEmailAvailable(ctx, userID, email); err != nil {
|
||||||
|
return "", 0, err
|
||||||
|
}
|
||||||
|
code, err := randomDigits(s.loginEmailCodeLength)
|
||||||
|
if err != nil {
|
||||||
|
return "", 0, err
|
||||||
|
}
|
||||||
|
rec.Code = code
|
||||||
|
if err := s.codes.Set(ctx, key, rec, s.loginEmailCodeTTL); err != nil {
|
||||||
|
return "", 0, err
|
||||||
|
}
|
||||||
|
if err := s.loginEmailSender.SendLoginCode(ctx, email, code, s.loginEmailCodeTTL); err != nil {
|
||||||
|
_ = s.codes.Del(ctx, key)
|
||||||
|
return "", 0, err
|
||||||
|
}
|
||||||
|
return emailPattern(email), len(code), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Service) VerifyLoginEmail(ctx context.Context, userID int64, phone, phoneCodeHash, code string, setup bool) (string, error) {
|
||||||
|
if s == nil || s.codes == nil {
|
||||||
|
return "", domain.ErrEmailNotAllowed
|
||||||
|
}
|
||||||
|
key := loginEmailVerifyChangePrefix + fmt.Sprint(userID)
|
||||||
|
if setup {
|
||||||
|
key = loginEmailVerifySetupPrefix + phoneCodeHash
|
||||||
|
}
|
||||||
|
rec, found, err := s.codes.Get(ctx, key)
|
||||||
|
if err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
if !found {
|
||||||
|
return "", domain.ErrEmailCodeInvalid
|
||||||
|
}
|
||||||
|
if strings.TrimSpace(code) == "" || subtle.ConstantTimeCompare([]byte(rec.Code), []byte(strings.TrimSpace(code))) != 1 {
|
||||||
|
return "", s.rejectEmailCode(ctx, key, rec)
|
||||||
|
}
|
||||||
|
email := normalizeLoginEmail(rec.PendingEmail)
|
||||||
|
if !validLoginEmail(email) {
|
||||||
|
_ = s.codes.Del(ctx, key)
|
||||||
|
return "", domain.ErrEmailInvalid
|
||||||
|
}
|
||||||
|
if setup {
|
||||||
|
phone = domain.NormalizePhone(phone)
|
||||||
|
phoneRec, found, err := s.codes.Get(ctx, phoneCodeHash)
|
||||||
|
if err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
if !found || phoneRec.Phone != phone {
|
||||||
|
return "", domain.ErrEmailCodeInvalid
|
||||||
|
}
|
||||||
|
targetUserID := int64(0)
|
||||||
|
if existingUserID, found, err := s.userIDByPhone(ctx, phone); err != nil {
|
||||||
|
return "", err
|
||||||
|
} else if found {
|
||||||
|
targetUserID = existingUserID
|
||||||
|
}
|
||||||
|
if err := s.ensureLoginEmailAvailable(ctx, targetUserID, email); err != nil {
|
||||||
|
_ = s.codes.Del(ctx, key)
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
_ = s.codes.Del(ctx, key)
|
||||||
|
phoneRec.Channel = codeChannelEmailLogin
|
||||||
|
phoneRec.Code = strings.TrimSpace(code)
|
||||||
|
phoneRec.Email = email
|
||||||
|
phoneRec.PendingEmail = email
|
||||||
|
phoneRec.VerifiedEmail = true
|
||||||
|
phoneRec.Attempts = 0
|
||||||
|
phoneRec.MaxAttempts = s.loginEmailCodeMaxAttempts
|
||||||
|
if err := s.codes.Update(ctx, phoneCodeHash, phoneRec); err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
if _, found, err := s.userIDByPhone(ctx, phone); err != nil {
|
||||||
|
return "", err
|
||||||
|
} else if found {
|
||||||
|
if err := s.SetLoginEmailByPhone(ctx, phone, email); err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return email, nil
|
||||||
|
}
|
||||||
|
if err := s.ensureLoginEmailAvailable(ctx, userID, email); err != nil {
|
||||||
|
_ = s.codes.Del(ctx, key)
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
_ = s.codes.Del(ctx, key)
|
||||||
|
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 {
|
||||||
|
rec.Attempts++
|
||||||
|
max := rec.MaxAttempts
|
||||||
|
if max <= 0 {
|
||||||
|
max = s.loginEmailCodeMaxAttempts
|
||||||
|
}
|
||||||
|
if max > 0 && rec.Attempts >= max {
|
||||||
|
_ = s.codes.Del(ctx, key)
|
||||||
|
return domain.ErrEmailCodeInvalid
|
||||||
|
}
|
||||||
|
_ = s.codes.Update(ctx, key, rec)
|
||||||
|
return domain.ErrEmailCodeInvalid
|
||||||
|
}
|
||||||
|
|
||||||
// SetLoginEmail 为已登录用户写入登录邮箱(authed 的 emailVerifyPurposeLoginChange)。
|
// SetLoginEmail 为已登录用户写入登录邮箱(authed 的 emailVerifyPurposeLoginChange)。
|
||||||
// 账号无 2FA 也可设置:account_passwords 行可在 has_password=false 下仅承载登录邮箱。
|
// 账号无 2FA 也可设置:account_passwords 行可在 has_password=false 下仅承载登录邮箱。
|
||||||
func (s *Service) SetLoginEmail(ctx context.Context, userID int64, email string) error {
|
func (s *Service) SetLoginEmail(ctx context.Context, userID int64, email string) error {
|
||||||
if s == nil || s.passwords == nil || userID == 0 {
|
if s == nil || s.passwords == nil || userID == 0 {
|
||||||
return domain.ErrEmailInvalid
|
return domain.ErrEmailInvalid
|
||||||
}
|
}
|
||||||
email = strings.TrimSpace(email)
|
email = normalizeLoginEmail(email)
|
||||||
if !validLoginEmail(email) {
|
if !validLoginEmail(email) {
|
||||||
return domain.ErrEmailInvalid
|
return domain.ErrEmailInvalid
|
||||||
}
|
}
|
||||||
|
if err := s.ensureLoginEmailAvailable(ctx, userID, email); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
settings, err := s.GetPasswordWithoutRefresh(ctx, userID)
|
settings, err := s.GetPasswordWithoutRefresh(ctx, userID)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
|
|
@ -520,7 +719,7 @@ func (s *Service) LoginEmail(ctx context.Context, userID int64) (string, bool, e
|
||||||
if !found || settings.LoginEmail == "" {
|
if !found || settings.LoginEmail == "" {
|
||||||
return "", false, nil
|
return "", false, nil
|
||||||
}
|
}
|
||||||
return settings.LoginEmail, true, nil
|
return normalizeLoginEmail(settings.LoginEmail), true, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// LoginEmailByPhone 按手机号返回登录邮箱原始地址(供 auth.sendCode 检测是否改投邮箱、
|
// LoginEmailByPhone 按手机号返回登录邮箱原始地址(供 auth.sendCode 检测是否改投邮箱、
|
||||||
|
|
@ -551,6 +750,24 @@ func (s *Service) ClearLoginEmailByPhone(ctx context.Context, phone string) erro
|
||||||
return s.passwords.Save(ctx, userID, settings)
|
return s.passwords.Save(ctx, userID, settings)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (s *Service) ensureLoginEmailAvailable(ctx context.Context, userID int64, email string) error {
|
||||||
|
if s == nil || s.passwords == nil {
|
||||||
|
return domain.ErrEmailInvalid
|
||||||
|
}
|
||||||
|
email = normalizeLoginEmail(email)
|
||||||
|
if !validLoginEmail(email) {
|
||||||
|
return domain.ErrEmailInvalid
|
||||||
|
}
|
||||||
|
ownerUserID, found, err := s.passwords.LoginEmailOwner(ctx, email)
|
||||||
|
if err != nil || !found {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if ownerUserID != userID {
|
||||||
|
return domain.ErrEmailOccupied
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
func (s *Service) userIDByPhone(ctx context.Context, phone string) (int64, bool, error) {
|
func (s *Service) userIDByPhone(ctx context.Context, phone string) (int64, bool, error) {
|
||||||
if s == nil || s.users == nil {
|
if s == nil || s.users == nil {
|
||||||
return 0, false, nil
|
return 0, false, nil
|
||||||
|
|
|
||||||
134
internal/app/auth/login_email_config_test.go
Normal file
134
internal/app/auth/login_email_config_test.go
Normal file
|
|
@ -0,0 +1,134 @@
|
||||||
|
package auth
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"telesrv/internal/domain"
|
||||||
|
"telesrv/internal/store/memory"
|
||||||
|
)
|
||||||
|
|
||||||
|
type testLoginEmailStore struct {
|
||||||
|
emails map[string]string
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *testLoginEmailStore) LoginEmailByPhone(_ context.Context, phone string) (string, bool, error) {
|
||||||
|
email, ok := s.emails[domain.NormalizePhone(phone)]
|
||||||
|
return email, ok, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *testLoginEmailStore) SetLoginEmailByPhone(_ context.Context, phone, email string) error {
|
||||||
|
s.emails[domain.NormalizePhone(phone)] = email
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
type testMailSender struct {
|
||||||
|
to string
|
||||||
|
code string
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *testMailSender) SendLoginCode(_ context.Context, to, code string, _ time.Duration) error {
|
||||||
|
s.to = to
|
||||||
|
s.code = code
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConfiguredEmailLoginSendsAndLimitsAttempts(t *testing.T) {
|
||||||
|
ctx := context.Background()
|
||||||
|
users := memory.NewUserStore()
|
||||||
|
authz := memory.NewAuthorizationStore()
|
||||||
|
if _, err := users.Create(ctx, domain.User{Phone: "15550009101", FirstName: "Email"}); err != nil {
|
||||||
|
t.Fatalf("create user: %v", err)
|
||||||
|
}
|
||||||
|
emails := &testLoginEmailStore{emails: map[string]string{"15550009101": "alice@example.test"}}
|
||||||
|
sender := &testMailSender{}
|
||||||
|
svc := NewService(users, authz, memory.NewCodeStore(), nil, nil, "12345",
|
||||||
|
WithLoginEmail(LoginEmailOptions{
|
||||||
|
Enabled: true,
|
||||||
|
CodeLength: 6,
|
||||||
|
Store: emails,
|
||||||
|
Sender: sender,
|
||||||
|
}),
|
||||||
|
WithCodeMaxAttempts(2))
|
||||||
|
|
||||||
|
hash, err := svc.SendCode(ctx, "+1 555 000 9101")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("SendCode: %v", err)
|
||||||
|
}
|
||||||
|
if sender.to != "alice@example.test" || len(sender.code) != 6 {
|
||||||
|
t.Fatalf("sent email to/code = %q/%q, want alice@example.test/6 digits", sender.to, sender.code)
|
||||||
|
}
|
||||||
|
delivery, found, err := svc.CodeDelivery(ctx, hash)
|
||||||
|
if err != nil || !found {
|
||||||
|
t.Fatalf("CodeDelivery found=%v err=%v", found, err)
|
||||||
|
}
|
||||||
|
if delivery.Kind != domain.AuthCodeDeliveryEmail || delivery.EmailPattern != "a***e@example.test" || delivery.Length != 6 {
|
||||||
|
t.Fatalf("delivery = %+v, want email masked length 6", delivery)
|
||||||
|
}
|
||||||
|
bad1 := wrongCode(sender.code, '0')
|
||||||
|
bad2 := wrongCode(sender.code, '1')
|
||||||
|
if bad2 == bad1 {
|
||||||
|
bad2 = wrongCode(sender.code, '2')
|
||||||
|
}
|
||||||
|
if _, _, _, err := svc.SignInWithEmail(ctx, domain.Authorization{}, "+15550009101", hash, bad1); !errors.Is(err, ErrCodeInvalid) {
|
||||||
|
t.Fatalf("first bad SignInWithEmail err = %v, want ErrCodeInvalid", err)
|
||||||
|
}
|
||||||
|
if _, _, _, err := svc.SignInWithEmail(ctx, domain.Authorization{}, "+15550009101", hash, bad2); !errors.Is(err, ErrCodeInvalid) {
|
||||||
|
t.Fatalf("second bad SignInWithEmail err = %v, want ErrCodeInvalid", err)
|
||||||
|
}
|
||||||
|
if _, _, _, err := svc.SignInWithEmail(ctx, domain.Authorization{}, "+15550009101", hash, sender.code); !errors.Is(err, ErrCodeExpired) {
|
||||||
|
t.Fatalf("SignInWithEmail after max attempts err = %v, want ErrCodeExpired", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func wrongCode(code string, digit byte) string {
|
||||||
|
if code == "" {
|
||||||
|
return string(digit)
|
||||||
|
}
|
||||||
|
out := make([]byte, len(code))
|
||||||
|
for i := range out {
|
||||||
|
out[i] = digit
|
||||||
|
}
|
||||||
|
if string(out) != code {
|
||||||
|
return string(out)
|
||||||
|
}
|
||||||
|
for i := range out {
|
||||||
|
out[i] = '9'
|
||||||
|
}
|
||||||
|
return string(out)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConfiguredEmailLoginAcceptsCorrectCode(t *testing.T) {
|
||||||
|
ctx := context.Background()
|
||||||
|
users := memory.NewUserStore()
|
||||||
|
authz := memory.NewAuthorizationStore()
|
||||||
|
u, err := users.Create(ctx, domain.User{Phone: "15550009102", FirstName: "Email"})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("create user: %v", err)
|
||||||
|
}
|
||||||
|
emails := &testLoginEmailStore{emails: map[string]string{"15550009102": "bob@example.test"}}
|
||||||
|
sender := &testMailSender{}
|
||||||
|
var key [8]byte
|
||||||
|
key[0] = 0x91
|
||||||
|
svc := NewService(users, authz, memory.NewCodeStore(), nil, nil, "12345",
|
||||||
|
WithLoginEmail(LoginEmailOptions{
|
||||||
|
Enabled: true,
|
||||||
|
CodeLength: 5,
|
||||||
|
Store: emails,
|
||||||
|
Sender: sender,
|
||||||
|
}))
|
||||||
|
|
||||||
|
hash, err := svc.SendCode(ctx, "+15550009102")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("SendCode: %v", err)
|
||||||
|
}
|
||||||
|
got, _, needSignUp, err := svc.SignInWithEmail(ctx, domain.Authorization{AuthKeyID: key}, "+15550009102", hash, sender.code)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("SignInWithEmail: %v", err)
|
||||||
|
}
|
||||||
|
if needSignUp || got.ID != u.ID {
|
||||||
|
t.Fatalf("SignInWithEmail got user=%d needSignUp=%v, want %d/false", got.ID, needSignUp, u.ID)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
@ -18,6 +18,7 @@ import (
|
||||||
mtcrypto "github.com/gotd/td/crypto"
|
mtcrypto "github.com/gotd/td/crypto"
|
||||||
|
|
||||||
"telesrv/internal/domain"
|
"telesrv/internal/domain"
|
||||||
|
"telesrv/internal/mail"
|
||||||
"telesrv/internal/store"
|
"telesrv/internal/store"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
@ -35,6 +36,12 @@ var (
|
||||||
ErrSystemUserLoginForbidden = errors.New("system user login forbidden")
|
ErrSystemUserLoginForbidden = errors.New("system user login forbidden")
|
||||||
)
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
codeChannelPhone = "phone"
|
||||||
|
codeChannelEmailLogin = "email_login"
|
||||||
|
codeChannelEmailSetupRequired = "email_setup_required"
|
||||||
|
)
|
||||||
|
|
||||||
// validPhone 校验规范化后的手机号:5-32 位纯数字(上限对齐 users.phone 列宽)。
|
// validPhone 校验规范化后的手机号:5-32 位纯数字(上限对齐 users.phone 列宽)。
|
||||||
// 核心目的是拒绝空/非数字 phone(防 0090 partial index 下无限铸造幽灵账号),
|
// 核心目的是拒绝空/非数字 phone(防 0090 partial index 下无限铸造幽灵账号),
|
||||||
// 长度上限从宽,不强求 E.164 精确位数(测试常用更长的唯一 phone)。
|
// 长度上限从宽,不强求 E.164 精确位数(测试常用更长的唯一 phone)。
|
||||||
|
|
@ -61,21 +68,40 @@ func systemLoginPhoneForbidden(phone string) bool {
|
||||||
|
|
||||||
// Service 实现登录/注册业务。第一阶段为开发固定验证码(不真实下发短信)。
|
// Service 实现登录/注册业务。第一阶段为开发固定验证码(不真实下发短信)。
|
||||||
type Service struct {
|
type Service struct {
|
||||||
users store.UserStore
|
users store.UserStore
|
||||||
auths store.AuthorizationStore
|
auths store.AuthorizationStore
|
||||||
codes store.CodeStore
|
codes store.CodeStore
|
||||||
authKeys store.AuthKeyStore
|
authKeys store.AuthKeyStore
|
||||||
tempKeys store.TempAuthKeyBindingStore
|
tempKeys store.TempAuthKeyBindingStore
|
||||||
passwords store.PasswordStore
|
passwords store.PasswordStore
|
||||||
messages store.MessageStore
|
messages store.MessageStore
|
||||||
dialogs store.DialogStore
|
dialogs store.DialogStore
|
||||||
bots store.BotStore
|
bots store.BotStore
|
||||||
fixedCode string
|
fixedCode string
|
||||||
codeTTL time.Duration
|
codeTTL time.Duration
|
||||||
|
codeMaxAttempts int
|
||||||
|
loginEmails loginEmailStore
|
||||||
|
loginEmailSender mail.Sender
|
||||||
|
loginEmailEnabled bool
|
||||||
|
loginEmailRequireSetup bool
|
||||||
|
loginEmailCodeLength int
|
||||||
// premiumGrantMonths 是新注册账号默认赠送的会员月数;0 表示关闭赠送。
|
// premiumGrantMonths 是新注册账号默认赠送的会员月数;0 表示关闭赠送。
|
||||||
premiumGrantMonths int
|
premiumGrantMonths int
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type loginEmailStore interface {
|
||||||
|
LoginEmailByPhone(ctx context.Context, phone string) (string, bool, error)
|
||||||
|
SetLoginEmailByPhone(ctx context.Context, phone, email string) error
|
||||||
|
}
|
||||||
|
|
||||||
|
type LoginEmailOptions struct {
|
||||||
|
Enabled bool
|
||||||
|
RequireSetup bool
|
||||||
|
CodeLength int
|
||||||
|
Store loginEmailStore
|
||||||
|
Sender mail.Sender
|
||||||
|
}
|
||||||
|
|
||||||
type authorizationRevoker interface {
|
type authorizationRevoker interface {
|
||||||
RevokeByHash(ctx context.Context, userID, hash int64) (domain.Authorization, bool, error)
|
RevokeByHash(ctx context.Context, userID, hash int64) (domain.Authorization, bool, error)
|
||||||
RevokeByUserExcept(ctx context.Context, userID int64, keepAuthKeyID [8]byte) ([]domain.Authorization, error)
|
RevokeByUserExcept(ctx context.Context, userID int64, keepAuthKeyID [8]byte) ([]domain.Authorization, error)
|
||||||
|
|
@ -114,9 +140,38 @@ func WithPremiumGrant(months int) Option {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func WithCodeTTL(ttl time.Duration) Option {
|
||||||
|
return func(s *Service) {
|
||||||
|
if ttl > 0 {
|
||||||
|
s.codeTTL = ttl
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func WithCodeMaxAttempts(max int) Option {
|
||||||
|
return func(s *Service) {
|
||||||
|
if max > 0 {
|
||||||
|
s.codeMaxAttempts = max
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func WithLoginEmail(opts LoginEmailOptions) Option {
|
||||||
|
return func(s *Service) {
|
||||||
|
s.loginEmailEnabled = opts.Enabled
|
||||||
|
s.loginEmailRequireSetup = opts.RequireSetup
|
||||||
|
s.loginEmailCodeLength = opts.CodeLength
|
||||||
|
if s.loginEmailCodeLength <= 0 {
|
||||||
|
s.loginEmailCodeLength = 6
|
||||||
|
}
|
||||||
|
s.loginEmails = opts.Store
|
||||||
|
s.loginEmailSender = opts.Sender
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// NewService 创建登录服务。fixedCode 为开发固定验证码。
|
// NewService 创建登录服务。fixedCode 为开发固定验证码。
|
||||||
func NewService(users store.UserStore, auths store.AuthorizationStore, codes store.CodeStore, authKeys store.AuthKeyStore, tempKeys store.TempAuthKeyBindingStore, fixedCode string, opts ...Option) *Service {
|
func NewService(users store.UserStore, auths store.AuthorizationStore, codes store.CodeStore, authKeys store.AuthKeyStore, tempKeys store.TempAuthKeyBindingStore, fixedCode string, opts ...Option) *Service {
|
||||||
s := &Service{users: users, auths: auths, codes: codes, authKeys: authKeys, tempKeys: tempKeys, fixedCode: fixedCode, codeTTL: 5 * time.Minute}
|
s := &Service{users: users, auths: auths, codes: codes, authKeys: authKeys, tempKeys: tempKeys, fixedCode: fixedCode, codeTTL: 5 * time.Minute, codeMaxAttempts: 5, loginEmailCodeLength: 6}
|
||||||
for _, opt := range opts {
|
for _, opt := range opts {
|
||||||
opt(s)
|
opt(s)
|
||||||
}
|
}
|
||||||
|
|
@ -214,7 +269,8 @@ func (s *Service) CompletePasswordSignIn(ctx context.Context, authKeyID [8]byte)
|
||||||
return s.auths.MarkPasswordPassed(ctx, authKeyID)
|
return s.auths.MarkPasswordPassed(ctx, authKeyID)
|
||||||
}
|
}
|
||||||
|
|
||||||
// SendCode 为 phone 生成 phone_code_hash,暂存(开发)固定验证码,返回 hash。
|
// SendCode 为 phone 生成 phone_code_hash,按配置选择开发 app code、登录邮箱 code
|
||||||
|
// 或登录邮箱 setup-required 状态,返回 hash。
|
||||||
func (s *Service) SendCode(ctx context.Context, phone string) (string, error) {
|
func (s *Service) SendCode(ctx context.Context, phone string) (string, error) {
|
||||||
phone = normalizePhone(phone)
|
phone = normalizePhone(phone)
|
||||||
if !validPhone(phone) {
|
if !validPhone(phone) {
|
||||||
|
|
@ -223,16 +279,105 @@ func (s *Service) SendCode(ctx context.Context, phone string) (string, error) {
|
||||||
if systemLoginPhoneForbidden(phone) {
|
if systemLoginPhoneForbidden(phone) {
|
||||||
return "", ErrSystemUserLoginForbidden
|
return "", ErrSystemUserLoginForbidden
|
||||||
}
|
}
|
||||||
|
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)
|
||||||
|
}
|
||||||
|
if s.loginEmailRequireSetup {
|
||||||
|
return s.createSetupRequiredCode(ctx, phone)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return s.createPhoneCode(ctx, phone)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Service) createPhoneCode(ctx context.Context, phone string) (string, error) {
|
||||||
hash, err := randomHex(8)
|
hash, err := randomHex(8)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return "", err
|
return "", err
|
||||||
}
|
}
|
||||||
if err := s.codes.Set(ctx, hash, store.PhoneCode{Phone: phone, Code: s.fixedCode}, s.codeTTL); err != nil {
|
if err := s.codes.Set(ctx, hash, store.PhoneCode{
|
||||||
|
Phone: phone,
|
||||||
|
Code: s.fixedCode,
|
||||||
|
Channel: codeChannelPhone,
|
||||||
|
MaxAttempts: s.codeMaxAttempts,
|
||||||
|
}, s.codeTTL); err != nil {
|
||||||
return "", fmt.Errorf("store code: %w", err)
|
return "", fmt.Errorf("store code: %w", err)
|
||||||
}
|
}
|
||||||
return hash, nil
|
return hash, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (s *Service) createSetupRequiredCode(ctx context.Context, phone string) (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,
|
||||||
|
}, s.codeTTL); err != nil {
|
||||||
|
return "", fmt.Errorf("store code: %w", err)
|
||||||
|
}
|
||||||
|
return hash, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Service) createEmailLoginCode(ctx context.Context, phone, email string) (string, error) {
|
||||||
|
hash, err := randomHex(8)
|
||||||
|
if err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
code, err := randomDigits(s.loginEmailCodeLength)
|
||||||
|
if err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
rec := store.PhoneCode{
|
||||||
|
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 s.loginEmailSender == nil {
|
||||||
|
_ = s.codes.Del(ctx, hash)
|
||||||
|
return "", 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 hash, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Service) CodeDelivery(ctx context.Context, phoneCodeHash string) (domain.AuthCodeDelivery, bool, error) {
|
||||||
|
rec, found, err := s.codes.Get(ctx, phoneCodeHash)
|
||||||
|
if err != nil || !found {
|
||||||
|
return domain.AuthCodeDelivery{}, found, err
|
||||||
|
}
|
||||||
|
return codeDelivery(rec), true, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func codeDelivery(rec store.PhoneCode) domain.AuthCodeDelivery {
|
||||||
|
switch rec.Channel {
|
||||||
|
case codeChannelEmailLogin:
|
||||||
|
return domain.AuthCodeDelivery{
|
||||||
|
Kind: domain.AuthCodeDeliveryEmail,
|
||||||
|
EmailPattern: domain.MaskEmail(rec.Email),
|
||||||
|
Length: len(rec.Code),
|
||||||
|
}
|
||||||
|
case codeChannelEmailSetupRequired:
|
||||||
|
return domain.AuthCodeDelivery{Kind: domain.AuthCodeDeliveryEmailSetupRequired}
|
||||||
|
default:
|
||||||
|
return domain.AuthCodeDelivery{Kind: domain.AuthCodeDeliveryPhone, Length: len(rec.Code)}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// ResendCode invalidates an existing code hash and sends a fresh code to the same phone.
|
// ResendCode invalidates an existing code hash and sends a fresh code to the same phone.
|
||||||
func (s *Service) ResendCode(ctx context.Context, phone, phoneCodeHash string) (string, error) {
|
func (s *Service) ResendCode(ctx context.Context, phone, phoneCodeHash string) (string, error) {
|
||||||
phone = normalizePhone(phone)
|
phone = normalizePhone(phone)
|
||||||
|
|
@ -247,6 +392,12 @@ func (s *Service) ResendCode(ctx context.Context, phone, phoneCodeHash string) (
|
||||||
return "", ErrCodeInvalid
|
return "", ErrCodeInvalid
|
||||||
}
|
}
|
||||||
_ = s.codes.Del(ctx, phoneCodeHash)
|
_ = s.codes.Del(ctx, phoneCodeHash)
|
||||||
|
if rec.Channel == codeChannelEmailLogin && strings.TrimSpace(rec.Email) != "" {
|
||||||
|
return s.createEmailLoginCode(ctx, phone, rec.Email)
|
||||||
|
}
|
||||||
|
if rec.Channel == codeChannelEmailSetupRequired {
|
||||||
|
return s.createSetupRequiredCode(ctx, phone)
|
||||||
|
}
|
||||||
return s.SendCode(ctx, phone)
|
return s.SendCode(ctx, phone)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -280,9 +431,15 @@ func (s *Service) SignIn(ctx context.Context, auth domain.Authorization, phone,
|
||||||
if !found {
|
if !found {
|
||||||
return domain.User{}, domain.Message{}, false, ErrCodeExpired
|
return domain.User{}, domain.Message{}, false, ErrCodeExpired
|
||||||
}
|
}
|
||||||
if rec.Phone != phone || rec.Code != code {
|
if rec.Phone != phone || rec.Channel == codeChannelEmailSetupRequired {
|
||||||
return domain.User{}, domain.Message{}, false, ErrCodeInvalid
|
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)
|
existing, found, err := s.users.ByPhone(ctx, phone)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|
@ -295,9 +452,10 @@ func (s *Service) SignIn(ctx context.Context, auth domain.Authorization, phone,
|
||||||
}
|
}
|
||||||
|
|
||||||
// SignInWithEmail 处理带 email_verification 的 auth.signIn:账号设置了登录邮箱后,新设备
|
// SignInWithEmail 处理带 email_verification 的 auth.signIn:账号设置了登录邮箱后,新设备
|
||||||
// 的验证码改投递到邮箱,客户端凭邮箱码(而非短信码)登录。开发环境接受任意非空邮箱码
|
// 的验证码改投递到邮箱,客户端凭邮箱码(而非短信码)登录。开启真实登录邮箱后必须匹配
|
||||||
// (与短信固定码同口径,"随意输入");仍校验 phone_code_hash 有效、手机号匹配,并与短信
|
// 随机邮箱码;未开启该特性时仅保留旧开发路径的任意非空兼容。仍校验 phone_code_hash
|
||||||
// 登录共用 2FA 门控——即便走邮箱验证,开启了两步验证的账号同样会停在 SESSION_PASSWORD_NEEDED。
|
// 有效、手机号匹配,并与短信登录共用 2FA 门控——即便走邮箱验证,开启了两步验证的账号
|
||||||
|
// 同样会停在 SESSION_PASSWORD_NEEDED。
|
||||||
func (s *Service) SignInWithEmail(ctx context.Context, auth domain.Authorization, phone, phoneCodeHash, code string) (domain.User, domain.Message, bool, error) {
|
func (s *Service) SignInWithEmail(ctx context.Context, auth domain.Authorization, phone, phoneCodeHash, code string) (domain.User, domain.Message, bool, error) {
|
||||||
phone = normalizePhone(phone)
|
phone = normalizePhone(phone)
|
||||||
if systemLoginPhoneForbidden(phone) {
|
if systemLoginPhoneForbidden(phone) {
|
||||||
|
|
@ -313,8 +471,15 @@ func (s *Service) SignInWithEmail(ctx context.Context, auth domain.Authorization
|
||||||
if rec.Phone != phone {
|
if rec.Phone != phone {
|
||||||
return domain.User{}, domain.Message{}, false, ErrCodeInvalid
|
return domain.User{}, domain.Message{}, false, ErrCodeInvalid
|
||||||
}
|
}
|
||||||
if strings.TrimSpace(code) == "" {
|
if rec.Channel != codeChannelEmailLogin {
|
||||||
return domain.User{}, domain.Message{}, false, ErrCodeInvalid
|
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.users.ByPhone(ctx, phone)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|
@ -378,6 +543,12 @@ func (s *Service) SignUp(ctx context.Context, auth domain.Authorization, phone,
|
||||||
if rec.Phone != phone {
|
if rec.Phone != phone {
|
||||||
return domain.User{}, domain.Message{}, ErrCodeInvalid
|
return domain.User{}, domain.Message{}, ErrCodeInvalid
|
||||||
}
|
}
|
||||||
|
if rec.Channel == codeChannelEmailSetupRequired {
|
||||||
|
return domain.User{}, domain.Message{}, ErrCodeInvalid
|
||||||
|
}
|
||||||
|
if s.loginEmailRequireSetup && !rec.VerifiedEmail && strings.TrimSpace(rec.PendingEmail) == "" {
|
||||||
|
return domain.User{}, domain.Message{}, ErrCodeInvalid
|
||||||
|
}
|
||||||
|
|
||||||
accessHash, err := randomInt64()
|
accessHash, err := randomInt64()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|
@ -398,6 +569,11 @@ func (s *Service) SignUp(ctx context.Context, auth domain.Authorization, phone,
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return domain.User{}, domain.Message{}, err
|
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 {
|
||||||
|
return domain.User{}, domain.Message{}, err
|
||||||
|
}
|
||||||
|
}
|
||||||
if err := s.bind(ctx, auth, u.ID); err != nil {
|
if err := s.bind(ctx, auth, u.ID); err != nil {
|
||||||
return domain.User{}, domain.Message{}, err
|
return domain.User{}, domain.Message{}, err
|
||||||
}
|
}
|
||||||
|
|
@ -525,6 +701,43 @@ func (s *Service) UpdateAuthorizationLayer(ctx context.Context, authKeyID [8]byt
|
||||||
return s.auths.UpdateLayer(ctx, authKeyID, layer)
|
return s.auths.UpdateLayer(ctx, authKeyID, layer)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (s *Service) AuthKeyClientInfo(ctx context.Context, authKeyID [8]byte) (domain.AuthKeyClientInfo, bool, error) {
|
||||||
|
if s == nil || s.authKeys == nil || authKeyID == ([8]byte{}) {
|
||||||
|
return domain.AuthKeyClientInfo{}, false, nil
|
||||||
|
}
|
||||||
|
key, found, err := s.authKeys.Get(ctx, authKeyID)
|
||||||
|
if err != nil || !found {
|
||||||
|
return domain.AuthKeyClientInfo{}, found, err
|
||||||
|
}
|
||||||
|
info := domain.AuthKeyClientInfo{
|
||||||
|
Layer: key.Layer,
|
||||||
|
DeviceModel: key.DeviceModel,
|
||||||
|
Platform: key.Platform,
|
||||||
|
SystemVersion: key.SystemVersion,
|
||||||
|
APIID: key.APIID,
|
||||||
|
AppVersion: key.AppVersion,
|
||||||
|
}
|
||||||
|
if info.Layer == 0 && info.DeviceModel == "" && info.Platform == "" &&
|
||||||
|
info.SystemVersion == "" && info.APIID == 0 && info.AppVersion == "" {
|
||||||
|
return domain.AuthKeyClientInfo{}, false, nil
|
||||||
|
}
|
||||||
|
return info, true, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Service) UpdateAuthKeyClientInfo(ctx context.Context, authKeyID [8]byte, info domain.AuthKeyClientInfo) error {
|
||||||
|
if s == nil || s.authKeys == nil || authKeyID == ([8]byte{}) {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return s.authKeys.UpdateClientInfo(ctx, authKeyID, store.AuthKeyClientInfo{
|
||||||
|
Layer: info.Layer,
|
||||||
|
DeviceModel: info.DeviceModel,
|
||||||
|
Platform: info.Platform,
|
||||||
|
SystemVersion: info.SystemVersion,
|
||||||
|
APIID: info.APIID,
|
||||||
|
AppVersion: info.AppVersion,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
func (s *Service) ListAuthorizations(ctx context.Context, userID int64) ([]domain.Authorization, error) {
|
func (s *Service) ListAuthorizations(ctx context.Context, userID int64) ([]domain.Authorization, error) {
|
||||||
if s == nil || s.auths == nil || userID == 0 {
|
if s == nil || s.auths == nil || userID == 0 {
|
||||||
return nil, nil
|
return nil, nil
|
||||||
|
|
@ -752,6 +965,20 @@ func authKeyIDInt64(id [8]byte) int64 {
|
||||||
return int64(binary.LittleEndian.Uint64(id[:]))
|
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 {
|
func normalizePhone(phone string) string {
|
||||||
return domain.NormalizePhone(phone)
|
return domain.NormalizePhone(phone)
|
||||||
}
|
}
|
||||||
|
|
@ -764,6 +991,22 @@ func randomHex(n int) (string, error) {
|
||||||
return hex.EncodeToString(b), nil
|
return hex.EncodeToString(b), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func randomDigits(n int) (string, error) {
|
||||||
|
if n <= 0 {
|
||||||
|
n = 6
|
||||||
|
}
|
||||||
|
b := make([]byte, n)
|
||||||
|
if _, err := rand.Read(b); err != nil {
|
||||||
|
return "", fmt.Errorf("rand: %w", err)
|
||||||
|
}
|
||||||
|
var out strings.Builder
|
||||||
|
out.Grow(n)
|
||||||
|
for _, v := range b {
|
||||||
|
out.WriteByte(byte('0') + v%10)
|
||||||
|
}
|
||||||
|
return out.String(), nil
|
||||||
|
}
|
||||||
|
|
||||||
func randomInt64() (int64, error) {
|
func randomInt64() (int64, error) {
|
||||||
var b [8]byte
|
var b [8]byte
|
||||||
if _, err := rand.Read(b[:]); err != nil {
|
if _, err := rand.Read(b[:]); err != nil {
|
||||||
|
|
|
||||||
|
|
@ -71,6 +71,26 @@ type Config struct {
|
||||||
|
|
||||||
// DevAuthCode 是开发固定验证码;生产短信/风控不在当前范围内。
|
// DevAuthCode 是开发固定验证码;生产短信/风控不在当前范围内。
|
||||||
DevAuthCode string
|
DevAuthCode string
|
||||||
|
// AuthCodeTTL 是登录/注册/邮箱验证 code 的有效期。
|
||||||
|
AuthCodeTTL time.Duration
|
||||||
|
// AuthCodeMaxAttempts 是同一 phone_code_hash / email verification code 的最大错误次数。
|
||||||
|
// 达到上限后验证码立即失效,用户必须重发。
|
||||||
|
AuthCodeMaxAttempts int
|
||||||
|
// LoginEmailEnable 启用手机号登录流程中的邮箱验证码投递。
|
||||||
|
LoginEmailEnable bool
|
||||||
|
// LoginEmailRequireSetup 为 true 时,没有登录邮箱的账号/新手机号会要求先设置邮箱。
|
||||||
|
LoginEmailRequireSetup bool
|
||||||
|
// LoginEmailCodeLength 是邮箱验证码长度。
|
||||||
|
LoginEmailCodeLength int
|
||||||
|
// SMTP* 是登录邮箱验证码的出站邮件配置。LoginEmailEnable=true 时必须可用。
|
||||||
|
SMTPHost string
|
||||||
|
SMTPPort int
|
||||||
|
SMTPUsername string
|
||||||
|
SMTPPassword string
|
||||||
|
SMTPFrom string
|
||||||
|
SMTPFromName string
|
||||||
|
SMTPTLSMode string
|
||||||
|
SMTPTimeout time.Duration
|
||||||
// MapboxToken 是服务端代理地图缩略图(upload.getWebFile)请求 Mapbox Static Images API
|
// MapboxToken 是服务端代理地图缩略图(upload.getWebFile)请求 Mapbox Static Images API
|
||||||
// 的 access token;为空则关闭代理、回退确定性占位图。客户端选点器 token 经 appConfig
|
// 的 access token;为空则关闭代理、回退确定性占位图。客户端选点器 token 经 appConfig
|
||||||
// `tdesktop_config_map` 下发(同源运行时配置)。
|
// `tdesktop_config_map` 下发(同源运行时配置)。
|
||||||
|
|
@ -322,6 +342,19 @@ func Load() (Config, error) {
|
||||||
RedisDB: envIntOr("TELESRV_REDIS_DB", 0),
|
RedisDB: envIntOr("TELESRV_REDIS_DB", 0),
|
||||||
|
|
||||||
DevAuthCode: envOr("TELESRV_DEV_AUTH_CODE", "12345"),
|
DevAuthCode: envOr("TELESRV_DEV_AUTH_CODE", "12345"),
|
||||||
|
AuthCodeTTL: envDurationOr("TELESRV_AUTH_CODE_TTL", 5*time.Minute),
|
||||||
|
AuthCodeMaxAttempts: envIntOr("TELESRV_AUTH_CODE_MAX_ATTEMPTS", 5),
|
||||||
|
LoginEmailEnable: envBoolOr("TELESRV_LOGIN_EMAIL_ENABLE", false),
|
||||||
|
LoginEmailRequireSetup: envBoolOr("TELESRV_LOGIN_EMAIL_REQUIRE_SETUP", false),
|
||||||
|
LoginEmailCodeLength: envIntOr("TELESRV_LOGIN_EMAIL_CODE_LENGTH", 6),
|
||||||
|
SMTPHost: envOr("TELESRV_SMTP_HOST", ""),
|
||||||
|
SMTPPort: envIntOr("TELESRV_SMTP_PORT", 587),
|
||||||
|
SMTPUsername: envOr("TELESRV_SMTP_USERNAME", ""),
|
||||||
|
SMTPPassword: envOr("TELESRV_SMTP_PASSWORD", ""),
|
||||||
|
SMTPFrom: envOr("TELESRV_SMTP_FROM", ""),
|
||||||
|
SMTPFromName: envOr("TELESRV_SMTP_FROM_NAME", "telesrv"),
|
||||||
|
SMTPTLSMode: strings.ToLower(strings.TrimSpace(envOr("TELESRV_SMTP_TLS", "starttls"))),
|
||||||
|
SMTPTimeout: envDurationOr("TELESRV_SMTP_TIMEOUT", 10*time.Second),
|
||||||
LangPackSeedDir: envOr("TELESRV_LANGPACK_SEED_DIR", "data/langpack"),
|
LangPackSeedDir: envOr("TELESRV_LANGPACK_SEED_DIR", "data/langpack"),
|
||||||
BlobDir: envOr("TELESRV_BLOB_DIR", "data/blobs"),
|
BlobDir: envOr("TELESRV_BLOB_DIR", "data/blobs"),
|
||||||
StickerSeedDir: envOr("TELESRV_STICKER_SEED_DIR", "data/sticker-seed"),
|
StickerSeedDir: envOr("TELESRV_STICKER_SEED_DIR", "data/sticker-seed"),
|
||||||
|
|
@ -406,9 +439,48 @@ func Load() (Config, error) {
|
||||||
LiveStreamWorkDir: envOr("TELESRV_LIVESTREAM_WORK_DIR", ""),
|
LiveStreamWorkDir: envOr("TELESRV_LIVESTREAM_WORK_DIR", ""),
|
||||||
LiveStreamSegmentKeep: envIntOr("TELESRV_LIVESTREAM_SEGMENT_KEEP", 32),
|
LiveStreamSegmentKeep: envIntOr("TELESRV_LIVESTREAM_SEGMENT_KEEP", 32),
|
||||||
}
|
}
|
||||||
|
if err := validateLoginEmailConfig(cfg); err != nil {
|
||||||
|
return Config{}, err
|
||||||
|
}
|
||||||
return cfg, nil
|
return cfg, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func validateLoginEmailConfig(cfg Config) error {
|
||||||
|
if cfg.LoginEmailRequireSetup && !cfg.LoginEmailEnable {
|
||||||
|
return fmt.Errorf("TELESRV_LOGIN_EMAIL_REQUIRE_SETUP requires TELESRV_LOGIN_EMAIL_ENABLE=true")
|
||||||
|
}
|
||||||
|
if cfg.AuthCodeTTL <= 0 {
|
||||||
|
return fmt.Errorf("TELESRV_AUTH_CODE_TTL must be positive")
|
||||||
|
}
|
||||||
|
if cfg.AuthCodeMaxAttempts <= 0 {
|
||||||
|
return fmt.Errorf("TELESRV_AUTH_CODE_MAX_ATTEMPTS must be positive")
|
||||||
|
}
|
||||||
|
if cfg.LoginEmailCodeLength < 4 || cfg.LoginEmailCodeLength > 10 {
|
||||||
|
return fmt.Errorf("TELESRV_LOGIN_EMAIL_CODE_LENGTH must be between 4 and 10")
|
||||||
|
}
|
||||||
|
switch cfg.SMTPTLSMode {
|
||||||
|
case "", "starttls", "tls", "none":
|
||||||
|
default:
|
||||||
|
return fmt.Errorf("TELESRV_SMTP_TLS must be starttls, tls, or none")
|
||||||
|
}
|
||||||
|
if !cfg.LoginEmailEnable {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if strings.TrimSpace(cfg.SMTPHost) == "" {
|
||||||
|
return fmt.Errorf("TELESRV_SMTP_HOST is required when TELESRV_LOGIN_EMAIL_ENABLE=true")
|
||||||
|
}
|
||||||
|
if cfg.SMTPPort <= 0 || cfg.SMTPPort > 65535 {
|
||||||
|
return fmt.Errorf("TELESRV_SMTP_PORT must be between 1 and 65535")
|
||||||
|
}
|
||||||
|
if strings.TrimSpace(cfg.SMTPFrom) == "" && strings.TrimSpace(cfg.SMTPUsername) == "" {
|
||||||
|
return fmt.Errorf("TELESRV_SMTP_FROM or TELESRV_SMTP_USERNAME is required when TELESRV_LOGIN_EMAIL_ENABLE=true")
|
||||||
|
}
|
||||||
|
if cfg.SMTPTimeout <= 0 {
|
||||||
|
return fmt.Errorf("TELESRV_SMTP_TIMEOUT must be positive")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
func loadAIProviders(env envSource) []AIProviderConfig {
|
func loadAIProviders(env envSource) []AIProviderConfig {
|
||||||
names := env.envListOr("TELESRV_AI_PROVIDERS", []string{"local"})
|
names := env.envListOr("TELESRV_AI_PROVIDERS", []string{"local"})
|
||||||
out := make([]AIProviderConfig, 0, len(names))
|
out := make([]AIProviderConfig, 0, len(names))
|
||||||
|
|
|
||||||
|
|
@ -63,6 +63,63 @@ func TestLoadBusinessAIProviderDefaultsToEcho(t *testing.T) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestLoadLoginEmailDefaultsDisabled(t *testing.T) {
|
||||||
|
disableDefaultConfigFile(t)
|
||||||
|
|
||||||
|
cfg, err := Load()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Load: %v", err)
|
||||||
|
}
|
||||||
|
if cfg.LoginEmailEnable {
|
||||||
|
t.Fatal("LoginEmailEnable = true, want false")
|
||||||
|
}
|
||||||
|
if cfg.LoginEmailRequireSetup {
|
||||||
|
t.Fatal("LoginEmailRequireSetup = true, want false")
|
||||||
|
}
|
||||||
|
if cfg.AuthCodeTTL != 5*time.Minute || cfg.AuthCodeMaxAttempts != 5 || cfg.LoginEmailCodeLength != 6 {
|
||||||
|
t.Fatalf("auth/login email defaults = %v/%d/%d", cfg.AuthCodeTTL, cfg.AuthCodeMaxAttempts, cfg.LoginEmailCodeLength)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLoadLoginEmailSMTPConfig(t *testing.T) {
|
||||||
|
disableDefaultConfigFile(t)
|
||||||
|
t.Setenv("TELESRV_LOGIN_EMAIL_ENABLE", "true")
|
||||||
|
t.Setenv("TELESRV_LOGIN_EMAIL_REQUIRE_SETUP", "true")
|
||||||
|
t.Setenv("TELESRV_AUTH_CODE_TTL", "3m")
|
||||||
|
t.Setenv("TELESRV_AUTH_CODE_MAX_ATTEMPTS", "4")
|
||||||
|
t.Setenv("TELESRV_LOGIN_EMAIL_CODE_LENGTH", "7")
|
||||||
|
t.Setenv("TELESRV_SMTP_HOST", "smtp.example.test")
|
||||||
|
t.Setenv("TELESRV_SMTP_PORT", "2525")
|
||||||
|
t.Setenv("TELESRV_SMTP_USERNAME", "smtp-user")
|
||||||
|
t.Setenv("TELESRV_SMTP_PASSWORD", "smtp-pass")
|
||||||
|
t.Setenv("TELESRV_SMTP_FROM", "noreply@example.test")
|
||||||
|
t.Setenv("TELESRV_SMTP_TLS", "none")
|
||||||
|
t.Setenv("TELESRV_SMTP_TIMEOUT", "2s")
|
||||||
|
|
||||||
|
cfg, err := Load()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Load: %v", err)
|
||||||
|
}
|
||||||
|
if !cfg.LoginEmailEnable || !cfg.LoginEmailRequireSetup {
|
||||||
|
t.Fatalf("login email flags = %v/%v, want true/true", cfg.LoginEmailEnable, cfg.LoginEmailRequireSetup)
|
||||||
|
}
|
||||||
|
if cfg.AuthCodeTTL != 3*time.Minute || cfg.AuthCodeMaxAttempts != 4 || cfg.LoginEmailCodeLength != 7 {
|
||||||
|
t.Fatalf("auth/login email config = %v/%d/%d", cfg.AuthCodeTTL, cfg.AuthCodeMaxAttempts, cfg.LoginEmailCodeLength)
|
||||||
|
}
|
||||||
|
if cfg.SMTPHost != "smtp.example.test" || cfg.SMTPPort != 2525 || cfg.SMTPUsername != "smtp-user" || cfg.SMTPPassword != "smtp-pass" || cfg.SMTPFrom != "noreply@example.test" || cfg.SMTPTLSMode != "none" || cfg.SMTPTimeout != 2*time.Second {
|
||||||
|
t.Fatalf("smtp config = %#v", cfg)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLoadLoginEmailRequiresSMTPWhenEnabled(t *testing.T) {
|
||||||
|
disableDefaultConfigFile(t)
|
||||||
|
t.Setenv("TELESRV_LOGIN_EMAIL_ENABLE", "true")
|
||||||
|
|
||||||
|
if _, err := Load(); err == nil {
|
||||||
|
t.Fatal("Load succeeded with login email enabled but no SMTP host")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestLoadKeepsAdminAndRtmpDefaultPortsSeparate(t *testing.T) {
|
func TestLoadKeepsAdminAndRtmpDefaultPortsSeparate(t *testing.T) {
|
||||||
disableDefaultConfigFile(t)
|
disableDefaultConfigFile(t)
|
||||||
t.Setenv("TELESRV_ADMIN_UI_ADDR", "")
|
t.Setenv("TELESRV_ADMIN_UI_ADDR", "")
|
||||||
|
|
|
||||||
|
|
@ -14,9 +14,25 @@ var (
|
||||||
ErrPasswordRecoveryNA = errors.New("password recovery not available")
|
ErrPasswordRecoveryNA = errors.New("password recovery not available")
|
||||||
ErrEmailCodeInvalid = errors.New("email code invalid")
|
ErrEmailCodeInvalid = errors.New("email code invalid")
|
||||||
ErrEmailInvalid = errors.New("email invalid")
|
ErrEmailInvalid = errors.New("email invalid")
|
||||||
|
ErrEmailNotAllowed = errors.New("email not allowed")
|
||||||
|
ErrEmailOccupied = errors.New("email occupied")
|
||||||
ErrSessionPasswordNeeded = errors.New("session password needed")
|
ErrSessionPasswordNeeded = errors.New("session password needed")
|
||||||
)
|
)
|
||||||
|
|
||||||
|
type AuthCodeDeliveryKind string
|
||||||
|
|
||||||
|
const (
|
||||||
|
AuthCodeDeliveryPhone AuthCodeDeliveryKind = "phone"
|
||||||
|
AuthCodeDeliveryEmail AuthCodeDeliveryKind = "email"
|
||||||
|
AuthCodeDeliveryEmailSetupRequired AuthCodeDeliveryKind = "email_setup_required"
|
||||||
|
)
|
||||||
|
|
||||||
|
type AuthCodeDelivery struct {
|
||||||
|
Kind AuthCodeDeliveryKind
|
||||||
|
EmailPattern string
|
||||||
|
Length int
|
||||||
|
}
|
||||||
|
|
||||||
// PasswordKDFAlgo 是业务层的 SRP KDF 算法描述,不依赖 tg.*。
|
// PasswordKDFAlgo 是业务层的 SRP KDF 算法描述,不依赖 tg.*。
|
||||||
type PasswordKDFAlgo struct {
|
type PasswordKDFAlgo struct {
|
||||||
Salt1 []byte
|
Salt1 []byte
|
||||||
|
|
@ -83,10 +99,10 @@ type PasswordSettings struct {
|
||||||
// LoginEmailPattern)。它独立于 2FA 恢复邮箱 RecoveryEmail:账号可只设登录邮箱而无 2FA。
|
// LoginEmailPattern)。它独立于 2FA 恢复邮箱 RecoveryEmail:账号可只设登录邮箱而无 2FA。
|
||||||
LoginEmail string
|
LoginEmail string
|
||||||
LoginEmailPattern string
|
LoginEmailPattern string
|
||||||
NewAlgo PasswordKDFAlgo
|
NewAlgo PasswordKDFAlgo
|
||||||
NewSecureAlgo SecurePasswordKDFAlgo
|
NewSecureAlgo SecurePasswordKDFAlgo
|
||||||
SecureRandom []byte
|
SecureRandom []byte
|
||||||
PendingResetDate int
|
PendingResetDate int
|
||||||
|
|
||||||
// Server-only SRP fields. They are persisted but never exposed to rpc/tg conversion.
|
// Server-only SRP fields. They are persisted but never exposed to rpc/tg conversion.
|
||||||
SRPVerifier []byte
|
SRPVerifier []byte
|
||||||
|
|
|
||||||
|
|
@ -21,3 +21,15 @@ type Authorization struct {
|
||||||
CreatedAt time.Time
|
CreatedAt time.Time
|
||||||
ActiveAt time.Time
|
ActiveAt time.Time
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// AuthKeyClientInfo 是未登录 auth_key 也需要保留的客户端协商元数据。
|
||||||
|
// 登录后的设备授权仍由 Authorization 表达;这里仅用于服务端重启后恢复
|
||||||
|
// pre-auth / setup 流程的 client type 与 layer。
|
||||||
|
type AuthKeyClientInfo struct {
|
||||||
|
Layer int
|
||||||
|
DeviceModel string
|
||||||
|
Platform string
|
||||||
|
SystemVersion string
|
||||||
|
APIID int
|
||||||
|
AppVersion string
|
||||||
|
}
|
||||||
|
|
|
||||||
154
internal/mail/sender.go
Normal file
154
internal/mail/sender.go
Normal file
|
|
@ -0,0 +1,154 @@
|
||||||
|
package mail
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"crypto/tls"
|
||||||
|
"fmt"
|
||||||
|
"mime"
|
||||||
|
"net"
|
||||||
|
stdmail "net/mail"
|
||||||
|
"net/smtp"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
type Config struct {
|
||||||
|
Host string
|
||||||
|
Port int
|
||||||
|
Username string
|
||||||
|
Password string
|
||||||
|
From string
|
||||||
|
FromName string
|
||||||
|
TLSMode string
|
||||||
|
Timeout time.Duration
|
||||||
|
}
|
||||||
|
|
||||||
|
type Sender interface {
|
||||||
|
SendLoginCode(ctx context.Context, to, code string, ttl time.Duration) error
|
||||||
|
}
|
||||||
|
|
||||||
|
type SMTP struct {
|
||||||
|
cfg Config
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewSMTP(cfg Config) *SMTP {
|
||||||
|
if cfg.Timeout <= 0 {
|
||||||
|
cfg.Timeout = 10 * time.Second
|
||||||
|
}
|
||||||
|
cfg.TLSMode = strings.ToLower(strings.TrimSpace(cfg.TLSMode))
|
||||||
|
if cfg.TLSMode == "" {
|
||||||
|
cfg.TLSMode = "starttls"
|
||||||
|
}
|
||||||
|
if strings.TrimSpace(cfg.From) == "" {
|
||||||
|
cfg.From = cfg.Username
|
||||||
|
}
|
||||||
|
return &SMTP{cfg: cfg}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *SMTP) SendLoginCode(ctx context.Context, to, code string, ttl time.Duration) error {
|
||||||
|
subject := "Your telesrv login code"
|
||||||
|
body := fmt.Sprintf("Your telesrv login code is %s.\n\nThis code expires in %s. If you did not request it, ignore this email.\n", code, humanTTL(ttl))
|
||||||
|
return s.send(ctx, to, subject, body)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *SMTP) send(ctx context.Context, to, subject, body string) error {
|
||||||
|
if strings.TrimSpace(s.cfg.Host) == "" {
|
||||||
|
return fmt.Errorf("smtp host is empty")
|
||||||
|
}
|
||||||
|
from := strings.TrimSpace(s.cfg.From)
|
||||||
|
if from == "" {
|
||||||
|
return fmt.Errorf("smtp from is empty")
|
||||||
|
}
|
||||||
|
if _, err := stdmail.ParseAddress(to); err != nil {
|
||||||
|
return fmt.Errorf("parse recipient: %w", err)
|
||||||
|
}
|
||||||
|
fromAddr := from
|
||||||
|
if s.cfg.FromName != "" {
|
||||||
|
fromAddr = (&stdmail.Address{Name: s.cfg.FromName, Address: from}).String()
|
||||||
|
}
|
||||||
|
addr := fmt.Sprintf("%s:%d", s.cfg.Host, s.cfg.Port)
|
||||||
|
var d net.Dialer
|
||||||
|
d.Timeout = s.cfg.Timeout
|
||||||
|
conn, err := d.DialContext(ctx, "tcp", addr)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("dial smtp: %w", err)
|
||||||
|
}
|
||||||
|
defer conn.Close()
|
||||||
|
|
||||||
|
mode := strings.ToLower(strings.TrimSpace(s.cfg.TLSMode))
|
||||||
|
var c *smtp.Client
|
||||||
|
if mode == "tls" {
|
||||||
|
tlsConn := tls.Client(conn, &tls.Config{ServerName: s.cfg.Host, MinVersion: tls.VersionTLS12})
|
||||||
|
if err := tlsConn.HandshakeContext(ctx); err != nil {
|
||||||
|
return fmt.Errorf("smtp tls handshake: %w", err)
|
||||||
|
}
|
||||||
|
c, err = smtp.NewClient(tlsConn, s.cfg.Host)
|
||||||
|
} else {
|
||||||
|
c, err = smtp.NewClient(conn, s.cfg.Host)
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("new smtp client: %w", err)
|
||||||
|
}
|
||||||
|
defer c.Close()
|
||||||
|
if mode == "starttls" {
|
||||||
|
if ok, _ := c.Extension("STARTTLS"); ok {
|
||||||
|
if err := c.StartTLS(&tls.Config{ServerName: s.cfg.Host, MinVersion: tls.VersionTLS12}); err != nil {
|
||||||
|
return fmt.Errorf("smtp starttls: %w", err)
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
return fmt.Errorf("smtp server does not support STARTTLS")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if s.cfg.Username != "" {
|
||||||
|
if err := c.Auth(smtp.PlainAuth("", s.cfg.Username, s.cfg.Password, s.cfg.Host)); err != nil {
|
||||||
|
return fmt.Errorf("smtp auth: %w", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if err := c.Mail(from); err != nil {
|
||||||
|
return fmt.Errorf("smtp mail from: %w", err)
|
||||||
|
}
|
||||||
|
if err := c.Rcpt(to); err != nil {
|
||||||
|
return fmt.Errorf("smtp rcpt: %w", err)
|
||||||
|
}
|
||||||
|
w, err := c.Data()
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("smtp data: %w", err)
|
||||||
|
}
|
||||||
|
msg := buildMessage(fromAddr, to, subject, body)
|
||||||
|
if _, err := w.Write(msg); err != nil {
|
||||||
|
_ = w.Close()
|
||||||
|
return fmt.Errorf("smtp write: %w", err)
|
||||||
|
}
|
||||||
|
if err := w.Close(); err != nil {
|
||||||
|
return fmt.Errorf("smtp close data: %w", err)
|
||||||
|
}
|
||||||
|
return c.Quit()
|
||||||
|
}
|
||||||
|
|
||||||
|
func buildMessage(from, to, subject, body string) []byte {
|
||||||
|
var b bytes.Buffer
|
||||||
|
b.WriteString("From: " + from + "\r\n")
|
||||||
|
b.WriteString("To: " + to + "\r\n")
|
||||||
|
b.WriteString("Subject: " + mime.QEncoding.Encode("utf-8", subject) + "\r\n")
|
||||||
|
b.WriteString("MIME-Version: 1.0\r\n")
|
||||||
|
b.WriteString("Content-Type: text/plain; charset=utf-8\r\n")
|
||||||
|
b.WriteString("Content-Transfer-Encoding: 8bit\r\n")
|
||||||
|
b.WriteString("\r\n")
|
||||||
|
b.WriteString(body)
|
||||||
|
return b.Bytes()
|
||||||
|
}
|
||||||
|
|
||||||
|
func humanTTL(ttl time.Duration) string {
|
||||||
|
if ttl <= 0 {
|
||||||
|
return "a short time"
|
||||||
|
}
|
||||||
|
if ttl%time.Minute == 0 {
|
||||||
|
minutes := int(ttl / time.Minute)
|
||||||
|
if minutes == 1 {
|
||||||
|
return "1 minute"
|
||||||
|
}
|
||||||
|
return fmt.Sprintf("%d minutes", minutes)
|
||||||
|
}
|
||||||
|
return ttl.String()
|
||||||
|
}
|
||||||
|
|
@ -31,8 +31,19 @@ import (
|
||||||
"telesrv/internal/store/memory"
|
"telesrv/internal/store/memory"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
type loginEmailTestSender struct {
|
||||||
|
to string
|
||||||
|
code string
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *loginEmailTestSender) SendLoginCode(_ context.Context, to, code string, _ time.Duration) error {
|
||||||
|
s.to = to
|
||||||
|
s.code = code
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
// TestLoginEmailEndToEnd 端到端验证登录邮箱:设备 A 注册并设置登录邮箱(loginChange),
|
// TestLoginEmailEndToEnd 端到端验证登录邮箱:设备 A 注册并设置登录邮箱(loginChange),
|
||||||
// 一个全新设备 B 调 sendCode 收到 sentCodeTypeEmailCode,凭任意邮箱验证码经 signIn
|
// 一个全新设备 B 调 sendCode 收到 sentCodeTypeEmailCode,凭真实邮箱验证码经 signIn
|
||||||
// (email_verification) 完成登录。
|
// (email_verification) 完成登录。
|
||||||
func TestLoginEmailEndToEnd(t *testing.T) {
|
func TestLoginEmailEndToEnd(t *testing.T) {
|
||||||
const (
|
const (
|
||||||
|
|
@ -59,10 +70,23 @@ func TestLoginEmailEndToEnd(t *testing.T) {
|
||||||
authKeyStore := memory.NewAuthKeyStore()
|
authKeyStore := memory.NewAuthKeyStore()
|
||||||
passwordStore := memory.NewPasswordStore()
|
passwordStore := memory.NewPasswordStore()
|
||||||
helpStore := memory.NewHelpStore()
|
helpStore := memory.NewHelpStore()
|
||||||
|
codeStore := memory.NewCodeStore()
|
||||||
|
emailSender := &loginEmailTestSender{}
|
||||||
|
accountService := account.NewService(passwordStore,
|
||||||
|
account.WithUsers(userStore),
|
||||||
|
account.WithLoginEmailVerification(codeStore, emailSender, 5*time.Minute, 5, 6))
|
||||||
|
authService := auth.NewService(userStore, authzStore, codeStore, authKeyStore, memory.NewTempAuthKeyBindingStore(), code,
|
||||||
|
auth.WithPasswords(passwordStore),
|
||||||
|
auth.WithLoginEmail(auth.LoginEmailOptions{
|
||||||
|
Enabled: true,
|
||||||
|
CodeLength: 6,
|
||||||
|
Store: accountService,
|
||||||
|
Sender: emailSender,
|
||||||
|
}))
|
||||||
|
|
||||||
deps := rpc.Deps{
|
deps := rpc.Deps{
|
||||||
Auth: auth.NewService(userStore, authzStore, memory.NewCodeStore(), authKeyStore, memory.NewTempAuthKeyBindingStore(), code, auth.WithPasswords(passwordStore)),
|
Auth: authService,
|
||||||
Account: account.NewService(passwordStore, account.WithUsers(userStore)),
|
Account: accountService,
|
||||||
Help: help.NewService(helpStore, helpStore),
|
Help: help.NewService(helpStore, helpStore),
|
||||||
Users: users.NewService(userStore),
|
Users: users.NewService(userStore),
|
||||||
Updates: updates.NewService(memory.NewUpdateStateStore(), memory.NewUpdateEventStore()),
|
Updates: updates.NewService(memory.NewUpdateStateStore(), memory.NewUpdateEventStore()),
|
||||||
|
|
@ -125,7 +149,7 @@ func TestLoginEmailEndToEnd(t *testing.T) {
|
||||||
}
|
}
|
||||||
verified, err := raw.AccountVerifyEmail(ctx, &tg.AccountVerifyEmailRequest{
|
verified, err := raw.AccountVerifyEmail(ctx, &tg.AccountVerifyEmailRequest{
|
||||||
Purpose: &tg.EmailVerifyPurposeLoginChange{},
|
Purpose: &tg.EmailVerifyPurposeLoginChange{},
|
||||||
Verification: &tg.EmailVerificationCode{Code: "whatever"},
|
Verification: &tg.EmailVerificationCode{Code: emailSender.code},
|
||||||
})
|
})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
|
|
@ -176,7 +200,7 @@ func TestLoginEmailEndToEnd(t *testing.T) {
|
||||||
signInRes, err := raw.AuthSignIn(ctx, &tg.AuthSignInRequest{
|
signInRes, err := raw.AuthSignIn(ctx, &tg.AuthSignInRequest{
|
||||||
PhoneNumber: phone,
|
PhoneNumber: phone,
|
||||||
PhoneCodeHash: sentCode.PhoneCodeHash,
|
PhoneCodeHash: sentCode.PhoneCodeHash,
|
||||||
EmailVerification: &tg.EmailVerificationCode{Code: "any-email-code"},
|
EmailVerification: &tg.EmailVerificationCode{Code: emailSender.code},
|
||||||
})
|
})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
|
|
|
||||||
|
|
@ -416,8 +416,7 @@ func (r *Router) onAccountCancelPasswordEmail(ctx context.Context) (bool, error)
|
||||||
}
|
}
|
||||||
|
|
||||||
// onAccountSendVerifyEmailCode 处理 account.sendVerifyEmailCode:为登录邮箱的设置/变更
|
// onAccountSendVerifyEmailCode 处理 account.sendVerifyEmailCode:为登录邮箱的设置/变更
|
||||||
// 发送验证码。开发环境不真正发邮件、验证码任意,故此处直接持久化待确认的登录邮箱地址,
|
// 发送邮箱验证码。loginChange 走已登录用户,loginSetup 走登录流程中的手机号 + phone_code_hash。
|
||||||
// 由后续 verifyEmail 做确认回显。loginChange 走已登录用户,loginSetup 走登录流程中的手机号。
|
|
||||||
func (r *Router) onAccountSendVerifyEmailCode(ctx context.Context, req *tg.AccountSendVerifyEmailCodeRequest) (*tg.AccountSentEmailCode, error) {
|
func (r *Router) onAccountSendVerifyEmailCode(ctx context.Context, req *tg.AccountSendVerifyEmailCodeRequest) (*tg.AccountSentEmailCode, error) {
|
||||||
if r.deps.Account == nil {
|
if r.deps.Account == nil {
|
||||||
return nil, internalErr()
|
return nil, internalErr()
|
||||||
|
|
@ -435,29 +434,32 @@ func (r *Router) onAccountSendVerifyEmailCode(ctx context.Context, req *tg.Accou
|
||||||
if userID == 0 {
|
if userID == 0 {
|
||||||
return nil, authKeyUnregisteredErr()
|
return nil, authKeyUnregisteredErr()
|
||||||
}
|
}
|
||||||
if err := r.deps.Account.SetLoginEmail(ctx, userID, email); err != nil {
|
pattern, length, err := r.deps.Account.SendLoginEmailCode(ctx, userID, "", "", email, false)
|
||||||
|
if err != nil {
|
||||||
return nil, passwordErr(err)
|
return nil, passwordErr(err)
|
||||||
}
|
}
|
||||||
|
return &tg.AccountSentEmailCode{EmailPattern: pattern, Length: length}, nil
|
||||||
case *tg.EmailVerifyPurposeLoginSetup:
|
case *tg.EmailVerifyPurposeLoginSetup:
|
||||||
if err := r.deps.Account.SetLoginEmailByPhone(ctx, p.PhoneNumber, email); err != nil {
|
pattern, length, err := r.deps.Account.SendLoginEmailCode(ctx, 0, p.PhoneNumber, p.PhoneCodeHash, email, true)
|
||||||
|
if err != nil {
|
||||||
return nil, passwordErr(err)
|
return nil, passwordErr(err)
|
||||||
}
|
}
|
||||||
|
return &tg.AccountSentEmailCode{EmailPattern: pattern, Length: length}, nil
|
||||||
default:
|
default:
|
||||||
return nil, emailInvalidErr()
|
return nil, emailInvalidErr()
|
||||||
}
|
}
|
||||||
return &tg.AccountSentEmailCode{EmailPattern: domain.MaskEmail(email), Length: devCodeLength}, nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// onAccountVerifyEmail 处理 account.verifyEmail:确认登录邮箱(验证码任意非空即通过)。
|
// onAccountVerifyEmail 处理 account.verifyEmail:确认登录邮箱验证码。
|
||||||
// loginChange(已登录)返回 emailVerified{email};loginSetup(登录流程中)返回
|
// loginChange(已登录)返回 emailVerified{email};loginSetup(登录流程中)返回
|
||||||
// emailVerifiedLogin{email, sent_code},其中 sent_code 是供客户端继续手机登录的新验证码。
|
// emailVerifiedLogin{email, sent_code}。TDesktop 能消费嵌套 auth.sentCodeSuccess,
|
||||||
|
// 直接进入注册/登录完成;DrKLO Android 12.8.1 该路径漏处理 sentCodeSuccess,
|
||||||
|
// 临时降级为普通 emailCode sentCode,待 Android 补齐后移除。
|
||||||
func (r *Router) onAccountVerifyEmail(ctx context.Context, req *tg.AccountVerifyEmailRequest) (tg.AccountEmailVerifiedClass, error) {
|
func (r *Router) onAccountVerifyEmail(ctx context.Context, req *tg.AccountVerifyEmailRequest) (tg.AccountEmailVerifiedClass, error) {
|
||||||
if r.deps.Account == nil {
|
if r.deps.Account == nil {
|
||||||
return nil, internalErr()
|
return nil, internalErr()
|
||||||
}
|
}
|
||||||
if strings.TrimSpace(emailVerificationCode(req.Verification)) == "" {
|
code := emailVerificationCode(req.Verification)
|
||||||
return nil, emailCodeInvalidErr()
|
|
||||||
}
|
|
||||||
switch p := req.Purpose.(type) {
|
switch p := req.Purpose.(type) {
|
||||||
case *tg.EmailVerifyPurposeLoginChange:
|
case *tg.EmailVerifyPurposeLoginChange:
|
||||||
userID, _, err := r.currentUserID(ctx)
|
userID, _, err := r.currentUserID(ctx)
|
||||||
|
|
@ -467,30 +469,36 @@ func (r *Router) onAccountVerifyEmail(ctx context.Context, req *tg.AccountVerify
|
||||||
if userID == 0 {
|
if userID == 0 {
|
||||||
return nil, authKeyUnregisteredErr()
|
return nil, authKeyUnregisteredErr()
|
||||||
}
|
}
|
||||||
email, found, err := r.deps.Account.LoginEmail(ctx, userID)
|
email, err := r.deps.Account.VerifyLoginEmail(ctx, userID, "", "", code, false)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, internalErr()
|
return nil, passwordErr(err)
|
||||||
}
|
|
||||||
if !found {
|
|
||||||
return nil, emailCodeInvalidErr()
|
|
||||||
}
|
}
|
||||||
return &tg.AccountEmailVerified{Email: email}, nil
|
return &tg.AccountEmailVerified{Email: email}, nil
|
||||||
case *tg.EmailVerifyPurposeLoginSetup:
|
case *tg.EmailVerifyPurposeLoginSetup:
|
||||||
email, found, err := r.deps.Account.LoginEmailByPhone(ctx, p.PhoneNumber)
|
|
||||||
if err != nil {
|
|
||||||
return nil, internalErr()
|
|
||||||
}
|
|
||||||
if !found {
|
|
||||||
return nil, emailCodeInvalidErr()
|
|
||||||
}
|
|
||||||
if r.deps.Auth == nil {
|
if r.deps.Auth == nil {
|
||||||
return nil, internalErr()
|
return nil, internalErr()
|
||||||
}
|
}
|
||||||
hash, err := r.deps.Auth.SendCode(ctx, p.PhoneNumber)
|
email, err := r.deps.Account.VerifyLoginEmail(ctx, 0, p.PhoneNumber, p.PhoneCodeHash, code, true)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, internalErr()
|
return nil, passwordErr(err)
|
||||||
}
|
}
|
||||||
return &tg.AccountEmailVerifiedLogin{Email: email, SentCode: tgSentCode(hash)}, nil
|
if ClientTypeFrom(ctx) == ClientTypeAndroid {
|
||||||
|
return &tg.AccountEmailVerifiedLogin{
|
||||||
|
Email: email,
|
||||||
|
SentCode: tgEmailSentCode(p.PhoneCodeHash, domain.MaskEmail(email), len(strings.TrimSpace(code))),
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
u, loginMessage, needSignUp, signInErr := r.deps.Auth.SignInWithEmail(ctx, r.authzFromCtx(ctx), p.PhoneNumber, p.PhoneCodeHash, code)
|
||||||
|
authorization, err := r.finishAuthSignIn(ctx, u, loginMessage, needSignUp, signInErr)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return &tg.AccountEmailVerifiedLogin{
|
||||||
|
Email: email,
|
||||||
|
SentCode: &tg.AuthSentCodeSuccess{
|
||||||
|
Authorization: authorization,
|
||||||
|
},
|
||||||
|
}, nil
|
||||||
default:
|
default:
|
||||||
return nil, emailInvalidErr()
|
return nil, emailInvalidErr()
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -241,6 +241,7 @@ func (r *Router) pushLoginTokenAccepted(ctx context.Context, target loginTokenTa
|
||||||
// 若该手机号账号设置了登录邮箱,验证码改投递到邮箱,返回 sentCodeTypeEmailCode
|
// 若该手机号账号设置了登录邮箱,验证码改投递到邮箱,返回 sentCodeTypeEmailCode
|
||||||
// (客户端据此进入"输入邮箱验证码"界面,随后用 auth.signIn 的 email_verification 完成登录)。
|
// (客户端据此进入"输入邮箱验证码"界面,随后用 auth.signIn 的 email_verification 完成登录)。
|
||||||
func (r *Router) onAuthSendCode(ctx context.Context, req *tg.AuthSendCodeRequest) (tg.AuthSentCodeClass, error) {
|
func (r *Router) onAuthSendCode(ctx context.Context, req *tg.AuthSendCodeRequest) (tg.AuthSentCodeClass, error) {
|
||||||
|
r.rememberClientAPIID(ctx, req.APIID)
|
||||||
hash, err := r.deps.Auth.SendCode(ctx, req.PhoneNumber)
|
hash, err := r.deps.Auth.SendCode(ctx, req.PhoneNumber)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
if errors.Is(err, auth.ErrPhoneNumberInvalid) ||
|
if errors.Is(err, auth.ErrPhoneNumberInvalid) ||
|
||||||
|
|
@ -249,35 +250,30 @@ func (r *Router) onAuthSendCode(ctx context.Context, req *tg.AuthSendCodeRequest
|
||||||
}
|
}
|
||||||
return nil, internalErr()
|
return nil, internalErr()
|
||||||
}
|
}
|
||||||
if pattern, ok := r.loginEmailPattern(ctx, req.PhoneNumber); ok {
|
return r.tgSentCodeForHash(ctx, hash)
|
||||||
return tgEmailSentCode(hash, pattern), nil
|
|
||||||
}
|
|
||||||
return tgSentCode(hash), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// loginEmailPattern 返回该手机号账号已确认登录邮箱的掩码,不存在则 ok=false。
|
|
||||||
func (r *Router) loginEmailPattern(ctx context.Context, phone string) (string, bool) {
|
|
||||||
if r.deps.Account == nil {
|
|
||||||
return "", false
|
|
||||||
}
|
|
||||||
email, found, err := r.deps.Account.LoginEmailByPhone(ctx, phone)
|
|
||||||
if err != nil || !found || email == "" {
|
|
||||||
return "", false
|
|
||||||
}
|
|
||||||
return domain.MaskEmail(email), true
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func tgSentCode(hash string) tg.AuthSentCodeClass {
|
func tgSentCode(hash string) tg.AuthSentCodeClass {
|
||||||
|
return tgSentCodeWithLength(hash, devCodeLength)
|
||||||
|
}
|
||||||
|
|
||||||
|
func tgSentCodeWithLength(hash string, length int) tg.AuthSentCodeClass {
|
||||||
|
if length <= 0 {
|
||||||
|
length = devCodeLength
|
||||||
|
}
|
||||||
return &tg.AuthSentCode{
|
return &tg.AuthSentCode{
|
||||||
Type: &tg.AuthSentCodeTypeApp{Length: devCodeLength},
|
Type: &tg.AuthSentCodeTypeApp{Length: length},
|
||||||
PhoneCodeHash: hash,
|
PhoneCodeHash: hash,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func tgEmailSentCode(hash, emailPattern string) tg.AuthSentCodeClass {
|
func tgEmailSentCode(hash, emailPattern string, length int) tg.AuthSentCodeClass {
|
||||||
|
if length <= 0 {
|
||||||
|
length = devCodeLength
|
||||||
|
}
|
||||||
codeType := &tg.AuthSentCodeTypeEmailCode{
|
codeType := &tg.AuthSentCodeTypeEmailCode{
|
||||||
EmailPattern: emailPattern,
|
EmailPattern: emailPattern,
|
||||||
Length: devCodeLength,
|
Length: length,
|
||||||
}
|
}
|
||||||
// reset_available_period=0 表示可立即调用 auth.resetLoginEmail(开发环境无等待期),
|
// reset_available_period=0 表示可立即调用 auth.resetLoginEmail(开发环境无等待期),
|
||||||
// 让客户端的"无法访问邮箱?"逃生入口可用。
|
// 让客户端的"无法访问邮箱?"逃生入口可用。
|
||||||
|
|
@ -288,6 +284,34 @@ func tgEmailSentCode(hash, emailPattern string) tg.AuthSentCodeClass {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func tgEmailSetupRequiredSentCode(hash string) tg.AuthSentCodeClass {
|
||||||
|
return &tg.AuthSentCode{
|
||||||
|
Type: &tg.AuthSentCodeTypeSetUpEmailRequired{},
|
||||||
|
PhoneCodeHash: hash,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *Router) tgSentCodeForHash(ctx context.Context, hash string) (tg.AuthSentCodeClass, error) {
|
||||||
|
if r.deps.Auth == nil {
|
||||||
|
return tgSentCode(hash), nil
|
||||||
|
}
|
||||||
|
delivery, found, err := r.deps.Auth.CodeDelivery(ctx, hash)
|
||||||
|
if err != nil {
|
||||||
|
return nil, internalErr()
|
||||||
|
}
|
||||||
|
if !found {
|
||||||
|
return nil, signInErr(auth.ErrCodeExpired)
|
||||||
|
}
|
||||||
|
switch delivery.Kind {
|
||||||
|
case domain.AuthCodeDeliveryEmail:
|
||||||
|
return tgEmailSentCode(hash, delivery.EmailPattern, delivery.Length), nil
|
||||||
|
case domain.AuthCodeDeliveryEmailSetupRequired:
|
||||||
|
return tgEmailSetupRequiredSentCode(hash), nil
|
||||||
|
default:
|
||||||
|
return tgSentCodeWithLength(hash, delivery.Length), nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// onAuthSignIn 处理 auth.signIn:校验验证码;用户不存在时返回 SignUpRequired。
|
// onAuthSignIn 处理 auth.signIn:校验验证码;用户不存在时返回 SignUpRequired。
|
||||||
// 带 email_verification 时走登录邮箱路径(验证码来自邮箱而非短信)。
|
// 带 email_verification 时走登录邮箱路径(验证码来自邮箱而非短信)。
|
||||||
func (r *Router) onAuthSignIn(ctx context.Context, req *tg.AuthSignInRequest) (tg.AuthAuthorizationClass, error) {
|
func (r *Router) onAuthSignIn(ctx context.Context, req *tg.AuthSignInRequest) (tg.AuthAuthorizationClass, error) {
|
||||||
|
|
@ -302,6 +326,10 @@ func (r *Router) onAuthSignIn(ctx context.Context, req *tg.AuthSignInRequest) (t
|
||||||
} else {
|
} else {
|
||||||
u, loginMessage, needSignUp, err = r.deps.Auth.SignIn(ctx, r.authzFromCtx(ctx), req.PhoneNumber, req.PhoneCodeHash, req.PhoneCode)
|
u, loginMessage, needSignUp, err = r.deps.Auth.SignIn(ctx, r.authzFromCtx(ctx), req.PhoneNumber, req.PhoneCodeHash, req.PhoneCode)
|
||||||
}
|
}
|
||||||
|
return r.finishAuthSignIn(ctx, u, loginMessage, needSignUp, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *Router) finishAuthSignIn(ctx context.Context, u domain.User, loginMessage domain.Message, needSignUp bool, err error) (tg.AuthAuthorizationClass, error) {
|
||||||
if err != nil {
|
if err != nil {
|
||||||
if errors.Is(err, domain.ErrSessionPasswordNeeded) && u.ID != 0 {
|
if errors.Is(err, domain.ErrSessionPasswordNeeded) && u.ID != 0 {
|
||||||
if err := r.clearAuthKeyStateOnUserChange(ctx, u.ID); err != nil {
|
if err := r.clearAuthKeyStateOnUserChange(ctx, u.ID); err != nil {
|
||||||
|
|
@ -337,7 +365,7 @@ func (r *Router) onAuthResendCode(ctx context.Context, req *tg.AuthResendCodeReq
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, signInErr(err)
|
return nil, signInErr(err)
|
||||||
}
|
}
|
||||||
return tgSentCode(hash), nil
|
return r.tgSentCodeForHash(ctx, hash)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *Router) onAuthCancelCode(ctx context.Context, req *tg.AuthCancelCodeRequest) (bool, error) {
|
func (r *Router) onAuthCancelCode(ctx context.Context, req *tg.AuthCancelCodeRequest) (bool, error) {
|
||||||
|
|
@ -502,7 +530,7 @@ func (r *Router) onAuthResetLoginEmail(ctx context.Context, req *tg.AuthResetLog
|
||||||
}
|
}
|
||||||
return nil, internalErr()
|
return nil, internalErr()
|
||||||
}
|
}
|
||||||
return tgSentCode(hash), nil
|
return r.tgSentCodeForHash(ctx, hash)
|
||||||
}
|
}
|
||||||
|
|
||||||
// emailVerificationCode 从 emailVerification 取出可校验的字符串(验证码 / Google·Apple
|
// emailVerificationCode 从 emailVerification 取出可校验的字符串(验证码 / Google·Apple
|
||||||
|
|
|
||||||
145
internal/rpc/auth_login_email_rpc_test.go
Normal file
145
internal/rpc/auth_login_email_rpc_test.go
Normal file
|
|
@ -0,0 +1,145 @@
|
||||||
|
package rpc
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/gotd/td/tg"
|
||||||
|
"go.uber.org/zap/zaptest"
|
||||||
|
|
||||||
|
"telesrv/internal/domain"
|
||||||
|
)
|
||||||
|
|
||||||
|
type loginEmailAccountService struct {
|
||||||
|
AccountService
|
||||||
|
verifiedEmail string
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s loginEmailAccountService) VerifyLoginEmail(context.Context, int64, string, string, string, bool) (string, error) {
|
||||||
|
return s.verifiedEmail, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestEmailSentCodeUsesDeliveryLength(t *testing.T) {
|
||||||
|
authSvc := &captureAuthService{
|
||||||
|
codeDelivery: domain.AuthCodeDelivery{
|
||||||
|
Kind: domain.AuthCodeDeliveryEmail,
|
||||||
|
EmailPattern: "a***e@example.test",
|
||||||
|
Length: 6,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
r := New(Config{}, Deps{Auth: authSvc}, zaptest.NewLogger(t), fixedClock{now: time.Unix(1700000000, 0)})
|
||||||
|
|
||||||
|
sent, err := r.tgSentCodeForHash(context.Background(), "hash-email")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("tgSentCodeForHash: %v", err)
|
||||||
|
}
|
||||||
|
code, ok := sent.(*tg.AuthSentCode)
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("sent = %T, want *tg.AuthSentCode", sent)
|
||||||
|
}
|
||||||
|
emailType, ok := code.Type.(*tg.AuthSentCodeTypeEmailCode)
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("sent type = %T, want *tg.AuthSentCodeTypeEmailCode", code.Type)
|
||||||
|
}
|
||||||
|
if emailType.Length != 6 {
|
||||||
|
t.Fatalf("email sent code length = %d, want 6", emailType.Length)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAccountVerifyEmailLoginSetupReturnsSentCodeSuccess(t *testing.T) {
|
||||||
|
user := domain.User{
|
||||||
|
ID: 100200300,
|
||||||
|
AccessHash: 900100200,
|
||||||
|
Phone: "8618800000020",
|
||||||
|
FirstName: "Alice",
|
||||||
|
}
|
||||||
|
authSvc := &captureAuthService{signInUser: user}
|
||||||
|
r := New(Config{}, Deps{
|
||||||
|
Auth: authSvc,
|
||||||
|
Account: loginEmailAccountService{verifiedEmail: "alice@example.test"},
|
||||||
|
}, zaptest.NewLogger(t), fixedClock{now: time.Unix(1700000000, 0)})
|
||||||
|
|
||||||
|
got, err := r.onAccountVerifyEmail(context.Background(), &tg.AccountVerifyEmailRequest{
|
||||||
|
Purpose: &tg.EmailVerifyPurposeLoginSetup{
|
||||||
|
PhoneNumber: "+86 188 0000 0020",
|
||||||
|
PhoneCodeHash: "hash-email-setup",
|
||||||
|
},
|
||||||
|
Verification: &tg.EmailVerificationCode{Code: "654321"},
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("onAccountVerifyEmail: %v", err)
|
||||||
|
}
|
||||||
|
verified, ok := got.(*tg.AccountEmailVerifiedLogin)
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("verified = %T, want *tg.AccountEmailVerifiedLogin", got)
|
||||||
|
}
|
||||||
|
if verified.Email != "alice@example.test" {
|
||||||
|
t.Fatalf("verified email = %q", verified.Email)
|
||||||
|
}
|
||||||
|
success, ok := verified.SentCode.(*tg.AuthSentCodeSuccess)
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("sent code = %T, want *tg.AuthSentCodeSuccess", verified.SentCode)
|
||||||
|
}
|
||||||
|
authorization, ok := success.Authorization.(*tg.AuthAuthorization)
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("authorization = %T, want *tg.AuthAuthorization", success.Authorization)
|
||||||
|
}
|
||||||
|
self, ok := authorization.User.(*tg.User)
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("authorization user = %T, want *tg.User", authorization.User)
|
||||||
|
}
|
||||||
|
if self.ID != user.ID || !self.Self {
|
||||||
|
t.Fatalf("authorization user = %+v, want self user %d", self, user.ID)
|
||||||
|
}
|
||||||
|
if authSvc.signInWithEmailCount != 1 {
|
||||||
|
t.Fatalf("SignInWithEmail calls = %d, want 1", authSvc.signInWithEmailCount)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAccountVerifyEmailLoginSetupAndroidReturnsEmailSentCode(t *testing.T) {
|
||||||
|
authSvc := &captureAuthService{}
|
||||||
|
r := New(Config{}, Deps{
|
||||||
|
Auth: authSvc,
|
||||||
|
Account: loginEmailAccountService{verifiedEmail: "alice@example.test"},
|
||||||
|
}, zaptest.NewLogger(t), fixedClock{now: time.Unix(1700000000, 0)})
|
||||||
|
|
||||||
|
ctx := WithClientInfo(context.Background(), ClientInfo{
|
||||||
|
Type: ClientTypeAndroid,
|
||||||
|
AppVersion: "12.8.1 (69169) pbeta",
|
||||||
|
})
|
||||||
|
got, err := r.onAccountVerifyEmail(ctx, &tg.AccountVerifyEmailRequest{
|
||||||
|
Purpose: &tg.EmailVerifyPurposeLoginSetup{
|
||||||
|
PhoneNumber: "+86 188 0000 0020",
|
||||||
|
PhoneCodeHash: "hash-email-setup",
|
||||||
|
},
|
||||||
|
Verification: &tg.EmailVerificationCode{Code: "654321"},
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("onAccountVerifyEmail: %v", err)
|
||||||
|
}
|
||||||
|
verified, ok := got.(*tg.AccountEmailVerifiedLogin)
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("verified = %T, want *tg.AccountEmailVerifiedLogin", got)
|
||||||
|
}
|
||||||
|
sent, ok := verified.SentCode.(*tg.AuthSentCode)
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("sent code = %T, want *tg.AuthSentCode", verified.SentCode)
|
||||||
|
}
|
||||||
|
if sent.PhoneCodeHash != "hash-email-setup" {
|
||||||
|
t.Fatalf("phone_code_hash = %q", sent.PhoneCodeHash)
|
||||||
|
}
|
||||||
|
emailType, ok := sent.Type.(*tg.AuthSentCodeTypeEmailCode)
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("sent type = %T, want *tg.AuthSentCodeTypeEmailCode", sent.Type)
|
||||||
|
}
|
||||||
|
if emailType.EmailPattern != "a***e@example.test" {
|
||||||
|
t.Fatalf("email pattern = %q", emailType.EmailPattern)
|
||||||
|
}
|
||||||
|
if emailType.Length != 6 {
|
||||||
|
t.Fatalf("email code length = %d, want 6", emailType.Length)
|
||||||
|
}
|
||||||
|
if authSvc.signInWithEmailCount != 0 {
|
||||||
|
t.Fatalf("SignInWithEmail calls = %d, want 0 for Android compat downgrade", authSvc.signInWithEmailCount)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
@ -2,12 +2,13 @@ package rpc
|
||||||
|
|
||||||
import "context"
|
import "context"
|
||||||
|
|
||||||
// withAndroidCompatMetadata 为「客户端构造器漂移」请求仅兜底 client 类型。
|
// withClientDriftMetadata 只在调用方已经用 constructor drift 证明客户端来源时
|
||||||
// DrKLO/OwpenGram Android 可能在不同版本使用不同 TL layer,client-private 构造器
|
// 补最小 client 类型。它不是 unknown fallback;不能在普通裸 RPC 上调用。
|
||||||
// 只能证明这是 Android 兼容路径,不能替代 invokeWithLayer 里的真实 layer。
|
// DrKLO Android 的 client-private constructor 只能证明 Android 兼容路径,
|
||||||
func (r *Router) withAndroidCompatMetadata(ctx context.Context) context.Context {
|
// 不能替代 invokeWithLayer/auth_keys/authorizations 里的真实 layer。
|
||||||
if ClientTypeFrom(ctx) == ClientTypeUnknown {
|
func (r *Router) withClientDriftMetadata(ctx context.Context, typ ClientType) context.Context {
|
||||||
ctx = WithClientInfo(ctx, ClientInfo{LangPack: string(ClientTypeAndroid), Type: ClientTypeAndroid})
|
if typ == ClientTypeUnknown || ClientTypeFrom(ctx) != ClientTypeUnknown {
|
||||||
|
return ctx
|
||||||
}
|
}
|
||||||
return ctx
|
return WithClientInfo(ctx, ClientInfo{Type: typ})
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -97,7 +97,23 @@ func knownClientType(t ClientType) bool {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func clientTypeFromAPIID(apiID int) ClientType {
|
||||||
|
switch apiID {
|
||||||
|
// DrKLO local BuildVars.APP_ID uses 4; TDesktop's active session
|
||||||
|
// classifier also recognizes the official Android ids below.
|
||||||
|
case 4, 5, 6, 24, 1026, 1083, 2458, 2521, 21724:
|
||||||
|
return ClientTypeAndroid
|
||||||
|
case 2040, 17349, 611335:
|
||||||
|
return ClientTypeTDesktop
|
||||||
|
default:
|
||||||
|
return ClientTypeUnknown
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func detectClientType(info ClientInfo) ClientType {
|
func detectClientType(info ClientInfo) ClientType {
|
||||||
|
if t := clientTypeFromAPIID(info.APIID); t != ClientTypeUnknown {
|
||||||
|
return t
|
||||||
|
}
|
||||||
if strings.EqualFold(info.LangPack, string(ClientTypeAndroid)) {
|
if strings.EqualFold(info.LangPack, string(ClientTypeAndroid)) {
|
||||||
return ClientTypeAndroid
|
return ClientTypeAndroid
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -25,6 +25,7 @@ type AuthService interface {
|
||||||
PendingPasswordUserID(ctx context.Context, authKeyID [8]byte) (int64, bool, error)
|
PendingPasswordUserID(ctx context.Context, authKeyID [8]byte) (int64, bool, error)
|
||||||
CompletePasswordSignIn(ctx context.Context, authKeyID [8]byte) error
|
CompletePasswordSignIn(ctx context.Context, authKeyID [8]byte) error
|
||||||
SendCode(ctx context.Context, phone string) (string, error)
|
SendCode(ctx context.Context, phone string) (string, error)
|
||||||
|
CodeDelivery(ctx context.Context, phoneCodeHash string) (domain.AuthCodeDelivery, bool, error)
|
||||||
ResendCode(ctx context.Context, phone, phoneCodeHash string) (string, error)
|
ResendCode(ctx context.Context, phone, phoneCodeHash string) (string, error)
|
||||||
CancelCode(ctx context.Context, phone, phoneCodeHash string) error
|
CancelCode(ctx context.Context, phone, phoneCodeHash string) error
|
||||||
SignIn(ctx context.Context, a domain.Authorization, phone, phoneCodeHash, code string) (domain.User, domain.Message, bool, error)
|
SignIn(ctx context.Context, a domain.Authorization, phone, phoneCodeHash, code string) (domain.User, domain.Message, bool, error)
|
||||||
|
|
@ -40,6 +41,8 @@ type AuthService interface {
|
||||||
LogOut(ctx context.Context, authKeyID [8]byte) error
|
LogOut(ctx context.Context, authKeyID [8]byte) error
|
||||||
Authorization(ctx context.Context, authKeyID [8]byte) (domain.Authorization, bool, error)
|
Authorization(ctx context.Context, authKeyID [8]byte) (domain.Authorization, bool, error)
|
||||||
UpdateAuthorizationLayer(ctx context.Context, authKeyID [8]byte, layer int) error
|
UpdateAuthorizationLayer(ctx context.Context, authKeyID [8]byte, layer int) error
|
||||||
|
AuthKeyClientInfo(ctx context.Context, authKeyID [8]byte) (domain.AuthKeyClientInfo, bool, error)
|
||||||
|
UpdateAuthKeyClientInfo(ctx context.Context, authKeyID [8]byte, info domain.AuthKeyClientInfo) error
|
||||||
ListAuthorizations(ctx context.Context, userID int64) ([]domain.Authorization, error)
|
ListAuthorizations(ctx context.Context, userID int64) ([]domain.Authorization, error)
|
||||||
ResetAuthorization(ctx context.Context, userID, hash int64) (domain.Authorization, bool, error)
|
ResetAuthorization(ctx context.Context, userID, hash int64) (domain.Authorization, bool, error)
|
||||||
ResetAuthorizations(ctx context.Context, userID int64, keepAuthKeyID [8]byte) ([]domain.Authorization, error)
|
ResetAuthorizations(ctx context.Context, userID int64, keepAuthKeyID [8]byte) ([]domain.Authorization, error)
|
||||||
|
|
@ -262,6 +265,8 @@ type AccountService interface {
|
||||||
ResendPasswordEmail(ctx context.Context, userID int64) error
|
ResendPasswordEmail(ctx context.Context, userID int64) error
|
||||||
CancelPasswordEmail(ctx context.Context, userID int64) error
|
CancelPasswordEmail(ctx context.Context, userID int64) error
|
||||||
// 登录邮箱(独立于 2FA 恢复邮箱):authed 走 userID,登录流程/重置走 phone。
|
// 登录邮箱(独立于 2FA 恢复邮箱):authed 走 userID,登录流程/重置走 phone。
|
||||||
|
SendLoginEmailCode(ctx context.Context, userID int64, phone, phoneCodeHash, email string, setup bool) (string, int, error)
|
||||||
|
VerifyLoginEmail(ctx context.Context, userID int64, phone, phoneCodeHash, code string, setup bool) (string, error)
|
||||||
SetLoginEmail(ctx context.Context, userID int64, email string) error
|
SetLoginEmail(ctx context.Context, userID int64, email string) error
|
||||||
SetLoginEmailByPhone(ctx context.Context, phone, email string) error
|
SetLoginEmailByPhone(ctx context.Context, phone, email string) error
|
||||||
LoginEmail(ctx context.Context, userID int64) (string, bool, error)
|
LoginEmail(ctx context.Context, userID int64) (string, bool, error)
|
||||||
|
|
|
||||||
|
|
@ -261,6 +261,7 @@ func srpPasswordChangedErr() error { return tgerr.New(400, "SRP_PASSWORD_CHAN
|
||||||
func newSettingsInvalidErr() error { return tgerr.New(400, "NEW_SETTINGS_INVALID") }
|
func newSettingsInvalidErr() error { return tgerr.New(400, "NEW_SETTINGS_INVALID") }
|
||||||
func newSaltInvalidErr() error { return tgerr.New(400, "NEW_SALT_INVALID") }
|
func newSaltInvalidErr() error { return tgerr.New(400, "NEW_SALT_INVALID") }
|
||||||
func emailInvalidErr() error { return tgerr.New(400, "EMAIL_INVALID") }
|
func emailInvalidErr() error { return tgerr.New(400, "EMAIL_INVALID") }
|
||||||
|
func emailNotAllowedErr() error { return tgerr.New(400, "EMAIL_NOT_ALLOWED") }
|
||||||
func emailCodeInvalidErr() error { return tgerr.New(400, "CODE_INVALID") }
|
func emailCodeInvalidErr() error { return tgerr.New(400, "CODE_INVALID") }
|
||||||
func passwordRecoveryNAErr() error { return tgerr.New(400, "PASSWORD_RECOVERY_NA") }
|
func passwordRecoveryNAErr() error { return tgerr.New(400, "PASSWORD_RECOVERY_NA") }
|
||||||
|
|
||||||
|
|
@ -393,6 +394,10 @@ func passwordErr(err error) error {
|
||||||
return newSaltInvalidErr()
|
return newSaltInvalidErr()
|
||||||
case errors.Is(err, domain.ErrEmailInvalid):
|
case errors.Is(err, domain.ErrEmailInvalid):
|
||||||
return emailInvalidErr()
|
return emailInvalidErr()
|
||||||
|
case errors.Is(err, domain.ErrEmailOccupied):
|
||||||
|
return emailNotAllowedErr()
|
||||||
|
case errors.Is(err, domain.ErrEmailNotAllowed):
|
||||||
|
return emailNotAllowedErr()
|
||||||
case errors.Is(err, domain.ErrEmailCodeInvalid):
|
case errors.Is(err, domain.ErrEmailCodeInvalid):
|
||||||
return emailCodeInvalidErr()
|
return emailCodeInvalidErr()
|
||||||
case errors.Is(err, domain.ErrPasswordRecoveryNA):
|
case errors.Is(err, domain.ErrPasswordRecoveryNA):
|
||||||
|
|
|
||||||
15
internal/rpc/errors_test.go
Normal file
15
internal/rpc/errors_test.go
Normal file
|
|
@ -0,0 +1,15 @@
|
||||||
|
package rpc
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/gotd/td/tgerr"
|
||||||
|
|
||||||
|
"telesrv/internal/domain"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestPasswordErrMapsOccupiedLoginEmailToNotAllowed(t *testing.T) {
|
||||||
|
if err := passwordErr(domain.ErrEmailOccupied); !tgerr.Is(err, "EMAIL_NOT_ALLOWED") {
|
||||||
|
t.Fatalf("passwordErr(ErrEmailOccupied) = %v, want EMAIL_NOT_ALLOWED", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
@ -34,8 +34,9 @@ type tempResolveResult struct {
|
||||||
}
|
}
|
||||||
|
|
||||||
const (
|
const (
|
||||||
authKeyResolveSingleflightPrefix = "resolve:"
|
authKeyResolveSingleflightPrefix = "resolve:"
|
||||||
authClientInfoSingleflightPrefix = "authinfo:"
|
authClientInfoSingleflightPrefix = "authinfo:"
|
||||||
|
authKeyClientInfoSingleflightPrefix = "authkeyinfo:"
|
||||||
)
|
)
|
||||||
|
|
||||||
var (
|
var (
|
||||||
|
|
@ -154,6 +155,7 @@ type clientSessionInfo struct {
|
||||||
layer int
|
layer int
|
||||||
clientInfo ClientInfo
|
clientInfo ClientInfo
|
||||||
hasClientInfo bool
|
hasClientInfo bool
|
||||||
|
authKeyInfoChecked bool
|
||||||
authorizationChecked bool
|
authorizationChecked bool
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -226,6 +228,19 @@ func (r *Router) Dispatch(ctx context.Context, authKeyID [8]byte, sessionID int6
|
||||||
}
|
}
|
||||||
tUser := r.clock.Now()
|
tUser := r.clock.Now()
|
||||||
info, hasClientMetadata, clientMetadataStored := r.clientSessionInfo(ctx)
|
info, hasClientMetadata, clientMetadataStored := r.clientSessionInfo(ctx)
|
||||||
|
if authInfo, ok := r.clientSessionInfoFromAuthKey(ctx, effectiveAuthKeyID, info); ok {
|
||||||
|
info = mergeClientSessionInfo(info, authInfo)
|
||||||
|
hasClientMetadata = true
|
||||||
|
r.rememberClientSessionInfo(ctx, info)
|
||||||
|
clientMetadataStored = true
|
||||||
|
if info.layer != 0 {
|
||||||
|
if binder, okBinder := r.deps.Sessions.(ClientLayerBinder); okBinder {
|
||||||
|
if rawAuthKeyID, okRaw := RawAuthKeyIDFrom(ctx); okRaw {
|
||||||
|
binder.SetClientLayerForAuthKey(rawAuthKeyID, sessionID, info.layer)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
if hasUserID {
|
if hasUserID {
|
||||||
if authInfo, ok := r.clientSessionInfoFromAuthorization(ctx, userID, effectiveAuthKeyID, info); ok {
|
if authInfo, ok := r.clientSessionInfoFromAuthorization(ctx, userID, effectiveAuthKeyID, info); ok {
|
||||||
info = mergeClientSessionInfo(info, authInfo)
|
info = mergeClientSessionInfo(info, authInfo)
|
||||||
|
|
@ -501,6 +516,7 @@ func (r *Router) invalidateAuthUserCache(authKeyID [8]byte) {
|
||||||
r.authUserSF.Forget(key)
|
r.authUserSF.Forget(key)
|
||||||
r.authUserSF.Forget(authKeyResolveSingleflightPrefix + key)
|
r.authUserSF.Forget(authKeyResolveSingleflightPrefix + key)
|
||||||
r.authUserSF.Forget(authClientInfoSingleflightPrefix + key)
|
r.authUserSF.Forget(authClientInfoSingleflightPrefix + key)
|
||||||
|
r.authUserSF.Forget(authKeyClientInfoSingleflightPrefix + key)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *Router) scopedSessions() (ScopedSessionBinder, bool) {
|
func (r *Router) scopedSessions() (ScopedSessionBinder, bool) {
|
||||||
|
|
@ -619,7 +635,7 @@ func (r *Router) dispatch(ctx context.Context, b *bin.Buffer, depth int) (bin.En
|
||||||
if clientDrift {
|
if clientDrift {
|
||||||
// 客户端漂移只能证明这是 Android 兼容路径;layer 仍以
|
// 客户端漂移只能证明这是 Android 兼容路径;layer 仍以
|
||||||
// invokeWithLayer 或授权记录里的真实观测值为准。
|
// invokeWithLayer 或授权记录里的真实观测值为准。
|
||||||
ctx = r.withAndroidCompatMetadata(ctx)
|
ctx = r.withClientDriftMetadata(ctx, ClientTypeAndroid)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
@ -699,6 +715,17 @@ func (r *Router) rememberClientInfo(ctx context.Context, info ClientInfo) {
|
||||||
sessionInfo.layer = layer
|
sessionInfo.layer = layer
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
|
r.persistAuthKeyClientInfo(ctx, clientSessionInfo{layer: layer, clientInfo: info, hasClientInfo: true})
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *Router) rememberClientAPIID(ctx context.Context, apiID int) {
|
||||||
|
if apiID == 0 {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
info := ClientInfo{APIID: apiID, Type: clientTypeFromAPIID(apiID)}
|
||||||
|
sessionInfo := clientSessionInfo{clientInfo: info, hasClientInfo: true}
|
||||||
|
r.rememberClientSessionInfo(ctx, sessionInfo)
|
||||||
|
r.persistAuthKeyClientInfo(ctx, sessionInfo)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *Router) rememberClientLayer(ctx context.Context, layer int) {
|
func (r *Router) rememberClientLayer(ctx context.Context, layer int) {
|
||||||
|
|
@ -747,6 +774,7 @@ func (r *Router) rememberClientLayer(ctx context.Context, layer int) {
|
||||||
binder.SetClientLayerForAuthKey(rawAuthKeyID, sessionID, layer)
|
binder.SetClientLayerForAuthKey(rawAuthKeyID, sessionID, layer)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
r.persistAuthKeyClientInfo(ctx, clientSessionInfo{layer: layer})
|
||||||
if persistAuthLayer && r.deps.Auth != nil {
|
if persistAuthLayer && r.deps.Auth != nil {
|
||||||
if err := r.deps.Auth.UpdateAuthorizationLayer(ctx, authKeyID, layer); err != nil {
|
if err := r.deps.Auth.UpdateAuthorizationLayer(ctx, authKeyID, layer); err != nil {
|
||||||
r.log.Warn("update authorization layer failed",
|
r.log.Warn("update authorization layer failed",
|
||||||
|
|
@ -757,6 +785,36 @@ func (r *Router) rememberClientLayer(ctx context.Context, layer int) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (r *Router) persistAuthKeyClientInfo(ctx context.Context, info clientSessionInfo) {
|
||||||
|
if r.deps.Auth == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
domainInfo := domainAuthKeyClientInfo(info)
|
||||||
|
if domainInfo.Layer == 0 && domainInfo.DeviceModel == "" && domainInfo.Platform == "" &&
|
||||||
|
domainInfo.SystemVersion == "" && domainInfo.APIID == 0 && domainInfo.AppVersion == "" {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
seen := make(map[[8]byte]struct{}, 2)
|
||||||
|
if rawAuthKeyID, ok := RawAuthKeyIDFrom(ctx); ok && rawAuthKeyID != ([8]byte{}) {
|
||||||
|
seen[rawAuthKeyID] = struct{}{}
|
||||||
|
if err := r.deps.Auth.UpdateAuthKeyClientInfo(ctx, rawAuthKeyID, domainInfo); err != nil {
|
||||||
|
r.log.Warn("update auth key client info failed",
|
||||||
|
zap.String("auth_key_id", fmt.Sprintf("%x", rawAuthKeyID[:])),
|
||||||
|
zap.Error(err))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if authKeyID, ok := AuthKeyIDFrom(ctx); ok && authKeyID != ([8]byte{}) {
|
||||||
|
if _, done := seen[authKeyID]; done {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if err := r.deps.Auth.UpdateAuthKeyClientInfo(ctx, authKeyID, domainInfo); err != nil {
|
||||||
|
r.log.Warn("update auth key client info failed",
|
||||||
|
zap.String("auth_key_id", fmt.Sprintf("%x", authKeyID[:])),
|
||||||
|
zap.Error(err))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// NegotiatedLayer returns the TL layer the given session negotiated via
|
// NegotiatedLayer returns the TL layer the given session negotiated via
|
||||||
// invokeWithLayer/initConnection. It is keyed first by (auth_key, session) then
|
// invokeWithLayer/initConnection. It is keyed first by (auth_key, session) then
|
||||||
// falls back to the stable auth_key — so a reconnect with a new session_id still
|
// falls back to the stable auth_key — so a reconnect with a new session_id still
|
||||||
|
|
@ -892,6 +950,9 @@ func clientSessionInfoContains(current, required clientSessionInfo) bool {
|
||||||
if required.authorizationChecked && !current.authorizationChecked {
|
if required.authorizationChecked && !current.authorizationChecked {
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
if required.authKeyInfoChecked && !current.authKeyInfoChecked {
|
||||||
|
return false
|
||||||
|
}
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -982,6 +1043,43 @@ func (r *Router) cachedResolvedAuthClientInfo(authKeyID [8]byte) (clientSessionI
|
||||||
return info, true
|
return info, true
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (r *Router) cachedResolvedAuthKeyClientInfo(authKeyID [8]byte) (clientSessionInfo, bool) {
|
||||||
|
r.clientInfoMu.RLock()
|
||||||
|
defer r.clientInfoMu.RUnlock()
|
||||||
|
info, ok := r.authInfo[authKeyID]
|
||||||
|
if !ok || clientSessionInfoNeedsAuthKeyInfo(info) {
|
||||||
|
return clientSessionInfo{}, false
|
||||||
|
}
|
||||||
|
return info, true
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *Router) clientSessionInfoFromAuthKey(ctx context.Context, authKeyID [8]byte, current clientSessionInfo) (clientSessionInfo, bool) {
|
||||||
|
if !clientSessionInfoNeedsAuthKeyInfo(current) || r.deps.Auth == nil || authKeyID == ([8]byte{}) {
|
||||||
|
return clientSessionInfo{}, false
|
||||||
|
}
|
||||||
|
v, err, _ := r.authUserSF.Do(authKeyClientInfoSingleflightPrefix+string(authKeyID[:]), func() (any, error) {
|
||||||
|
if cached, ok := r.cachedResolvedAuthKeyClientInfo(authKeyID); ok {
|
||||||
|
return cached, nil
|
||||||
|
}
|
||||||
|
info, found, err := r.deps.Auth.AuthKeyClientInfo(ctx, authKeyID)
|
||||||
|
if err != nil {
|
||||||
|
return clientSessionInfo{}, err
|
||||||
|
}
|
||||||
|
if !found {
|
||||||
|
return clientSessionInfo{authKeyInfoChecked: true}, nil
|
||||||
|
}
|
||||||
|
return clientSessionInfoFromAuthKeyClientInfo(info, current), nil
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return clientSessionInfo{}, false
|
||||||
|
}
|
||||||
|
info := v.(clientSessionInfo)
|
||||||
|
if info.layer == 0 && !info.hasClientInfo && !info.authKeyInfoChecked {
|
||||||
|
return clientSessionInfo{}, false
|
||||||
|
}
|
||||||
|
return info, true
|
||||||
|
}
|
||||||
|
|
||||||
func mergeClientSessionInfo(base, fallback clientSessionInfo) clientSessionInfo {
|
func mergeClientSessionInfo(base, fallback clientSessionInfo) clientSessionInfo {
|
||||||
if base.layer == 0 {
|
if base.layer == 0 {
|
||||||
base.layer = fallback.layer
|
base.layer = fallback.layer
|
||||||
|
|
@ -993,9 +1091,48 @@ func mergeClientSessionInfo(base, fallback clientSessionInfo) clientSessionInfo
|
||||||
if fallback.authorizationChecked {
|
if fallback.authorizationChecked {
|
||||||
base.authorizationChecked = true
|
base.authorizationChecked = true
|
||||||
}
|
}
|
||||||
|
if fallback.authKeyInfoChecked {
|
||||||
|
base.authKeyInfoChecked = true
|
||||||
|
}
|
||||||
return base
|
return base
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func clientSessionInfoFromAuthKeyClientInfo(item domain.AuthKeyClientInfo, current clientSessionInfo) clientSessionInfo {
|
||||||
|
info := clientSessionInfo{
|
||||||
|
layer: item.Layer,
|
||||||
|
authKeyInfoChecked: true,
|
||||||
|
clientInfo: ClientInfo{
|
||||||
|
APIID: item.APIID,
|
||||||
|
DeviceModel: item.DeviceModel,
|
||||||
|
SystemVersion: item.SystemVersion,
|
||||||
|
AppVersion: item.AppVersion,
|
||||||
|
Type: ClientType(item.Platform),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
info.clientInfo = normalizeClientInfo(info.clientInfo)
|
||||||
|
info.hasClientInfo = info.clientInfo.ClientType() != ClientTypeUnknown ||
|
||||||
|
info.clientInfo.DeviceModel != "" ||
|
||||||
|
info.clientInfo.SystemVersion != "" ||
|
||||||
|
info.clientInfo.AppVersion != "" ||
|
||||||
|
info.clientInfo.APIID != 0
|
||||||
|
if info.layer == 0 && current.layer != 0 {
|
||||||
|
info.layer = current.layer
|
||||||
|
}
|
||||||
|
return info
|
||||||
|
}
|
||||||
|
|
||||||
|
func domainAuthKeyClientInfo(info clientSessionInfo) domain.AuthKeyClientInfo {
|
||||||
|
out := domain.AuthKeyClientInfo{Layer: info.layer}
|
||||||
|
if info.hasClientInfo {
|
||||||
|
out.APIID = info.clientInfo.APIID
|
||||||
|
out.DeviceModel = info.clientInfo.DeviceModel
|
||||||
|
out.SystemVersion = info.clientInfo.SystemVersion
|
||||||
|
out.AppVersion = info.clientInfo.AppVersion
|
||||||
|
out.Platform = string(info.clientInfo.ClientType())
|
||||||
|
}
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
func (r *Router) clientSessionInfoFromAuthorization(ctx context.Context, userID int64, authKeyID [8]byte, current clientSessionInfo) (clientSessionInfo, bool) {
|
func (r *Router) clientSessionInfoFromAuthorization(ctx context.Context, userID int64, authKeyID [8]byte, current clientSessionInfo) (clientSessionInfo, bool) {
|
||||||
if !clientSessionInfoNeedsAuthorization(current) || r.deps.Auth == nil || userID == 0 {
|
if !clientSessionInfoNeedsAuthorization(current) || r.deps.Auth == nil || userID == 0 {
|
||||||
return clientSessionInfo{}, false
|
return clientSessionInfo{}, false
|
||||||
|
|
@ -1052,6 +1189,13 @@ func clientSessionInfoNeedsAuthorization(info clientSessionInfo) bool {
|
||||||
return info.layer == 0 || !info.hasClientInfo || info.clientInfo.ClientType() == ClientTypeUnknown
|
return info.layer == 0 || !info.hasClientInfo || info.clientInfo.ClientType() == ClientTypeUnknown
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func clientSessionInfoNeedsAuthKeyInfo(info clientSessionInfo) bool {
|
||||||
|
if info.authKeyInfoChecked {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return info.layer == 0 || !info.hasClientInfo || info.clientInfo.ClientType() == ClientTypeUnknown
|
||||||
|
}
|
||||||
|
|
||||||
// fallback 处理未注册的 RPC:记录到 compatibility trace(落兼容矩阵),
|
// fallback 处理未注册的 RPC:记录到 compatibility trace(落兼容矩阵),
|
||||||
// 返回 NOT_IMPLEMENTED rpc_error 让客户端继续运行而非断连。
|
// 返回 NOT_IMPLEMENTED rpc_error 让客户端继续运行而非断连。
|
||||||
func (r *Router) fallback(ctx context.Context, b *bin.Buffer) (bin.Encoder, error) {
|
func (r *Router) fallback(ctx context.Context, b *bin.Buffer) (bin.Encoder, error) {
|
||||||
|
|
|
||||||
|
|
@ -188,6 +188,146 @@ func TestDispatchRemembersLayerAndClientTypeForSession(t *testing.T) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestDispatchPersistsPreLoginClientMetadataOnInitConnection(t *testing.T) {
|
||||||
|
auth := &captureAuthService{}
|
||||||
|
r := New(Config{DC: 2, IP: "127.0.0.1", Port: 2398}, Deps{
|
||||||
|
Auth: auth,
|
||||||
|
}, zaptest.NewLogger(t), clock.System)
|
||||||
|
rawAuthKeyID := [8]byte{0x22, 0xdb, 0xcf, 0xc8, 0x0d, 0x4c, 0x77, 0x97}
|
||||||
|
sessionID := int64(8103956954238395544)
|
||||||
|
|
||||||
|
req := &tg.InvokeWithLayerRequest{
|
||||||
|
Layer: currentClientLayer,
|
||||||
|
Query: &tg.InitConnectionRequest{
|
||||||
|
APIID: 4,
|
||||||
|
DeviceModel: "GooglePixel 9a",
|
||||||
|
SystemVersion: "SDK 36",
|
||||||
|
AppVersion: "12.8.1 (69169) pbeta",
|
||||||
|
SystemLangCode: "en",
|
||||||
|
LangPack: "android",
|
||||||
|
LangCode: "en",
|
||||||
|
Query: &tg.HelpGetConfigRequest{},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
var in bin.Buffer
|
||||||
|
if err := req.Encode(&in); err != nil {
|
||||||
|
t.Fatalf("encode init request: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := r.Dispatch(context.Background(), rawAuthKeyID, sessionID, &in); err != nil {
|
||||||
|
t.Fatalf("dispatch init request: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
got, ok := auth.authKeyClientInfos[rawAuthKeyID]
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("auth key client metadata was not persisted")
|
||||||
|
}
|
||||||
|
if got.Layer != currentClientLayer {
|
||||||
|
t.Fatalf("persisted layer = %d, want %d", got.Layer, currentClientLayer)
|
||||||
|
}
|
||||||
|
if got.Platform != string(ClientTypeAndroid) {
|
||||||
|
t.Fatalf("persisted platform = %q, want android", got.Platform)
|
||||||
|
}
|
||||||
|
if got.DeviceModel != "GooglePixel 9a" || got.SystemVersion != "SDK 36" || got.APIID != 4 || got.AppVersion != "12.8.1 (69169) pbeta" {
|
||||||
|
t.Fatalf("persisted client metadata = %+v", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDispatchPersistsPreLoginClientMetadataFromSendCodeAPIID(t *testing.T) {
|
||||||
|
auth := &captureAuthService{}
|
||||||
|
rawAuthKeyID := [8]byte{0x33, 0xdb, 0xcf, 0xc8, 0x0d, 0x4c, 0x77, 0x97}
|
||||||
|
const sessionID = int64(8103956954238395544)
|
||||||
|
r := New(Config{DC: 2, IP: "127.0.0.1", Port: 2398}, Deps{
|
||||||
|
Auth: auth,
|
||||||
|
}, zaptest.NewLogger(t), clock.System)
|
||||||
|
|
||||||
|
var sendCode bin.Buffer
|
||||||
|
if err := (&tg.AuthSendCodeRequest{
|
||||||
|
PhoneNumber: "+8618800000020",
|
||||||
|
APIID: 4,
|
||||||
|
APIHash: "android",
|
||||||
|
Settings: tg.CodeSettings{},
|
||||||
|
}).Encode(&sendCode); err != nil {
|
||||||
|
t.Fatalf("encode auth.sendCode: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := r.Dispatch(context.Background(), rawAuthKeyID, sessionID, &sendCode); err != nil {
|
||||||
|
t.Fatalf("dispatch auth.sendCode: %v", err)
|
||||||
|
}
|
||||||
|
persisted, ok := auth.authKeyClientInfos[rawAuthKeyID]
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("auth.sendCode did not persist auth key client metadata")
|
||||||
|
}
|
||||||
|
if persisted.APIID != 4 || persisted.Platform != string(ClientTypeAndroid) {
|
||||||
|
t.Fatalf("persisted client metadata = %+v, want android api_id=4", persisted)
|
||||||
|
}
|
||||||
|
|
||||||
|
core, logs := observer.New(zap.DebugLevel)
|
||||||
|
afterRestart := New(Config{DC: 2, IP: "127.0.0.1", Port: 2398}, Deps{
|
||||||
|
Auth: auth,
|
||||||
|
}, zap.New(core), clock.System)
|
||||||
|
var help bin.Buffer
|
||||||
|
if err := (&tg.HelpGetConfigRequest{}).Encode(&help); err != nil {
|
||||||
|
t.Fatalf("encode help.getConfig: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := afterRestart.Dispatch(context.Background(), rawAuthKeyID, sessionID+1, &help); err != nil {
|
||||||
|
t.Fatalf("dispatch help.getConfig after restart: %v", err)
|
||||||
|
}
|
||||||
|
entries := logs.FilterMessage("RPC inner handled").All()
|
||||||
|
if len(entries) == 0 {
|
||||||
|
t.Fatalf("RPC inner handled log missing")
|
||||||
|
}
|
||||||
|
fields := entries[len(entries)-1].ContextMap()
|
||||||
|
if got := fields["client_type"]; got != string(ClientTypeAndroid) {
|
||||||
|
t.Fatalf("logged client_type = %v, want %s", got, ClientTypeAndroid)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDispatchRestoresPreLoginAndroidMetadataFromAuthKey(t *testing.T) {
|
||||||
|
core, logs := observer.New(zap.DebugLevel)
|
||||||
|
authKeyID := [8]byte{0x22, 0xdb, 0xcf, 0xc8, 0x0d, 0x4c, 0x77, 0x97}
|
||||||
|
auth := &captureAuthService{
|
||||||
|
authKeyClientInfos: map[[8]byte]domain.AuthKeyClientInfo{
|
||||||
|
authKeyID: {
|
||||||
|
Layer: currentClientLayer,
|
||||||
|
DeviceModel: "GooglePixel 9a",
|
||||||
|
Platform: string(ClientTypeAndroid),
|
||||||
|
SystemVersion: "SDK 36",
|
||||||
|
APIID: 4,
|
||||||
|
AppVersion: "12.8.1 (69169) pbeta",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
r := New(Config{DC: 2, IP: "127.0.0.1", Port: 2398}, Deps{
|
||||||
|
Auth: auth,
|
||||||
|
}, zap.New(core), clock.System)
|
||||||
|
|
||||||
|
var in bin.Buffer
|
||||||
|
if err := (&tg.HelpGetConfigRequest{}).Encode(&in); err != nil {
|
||||||
|
t.Fatalf("encode request: %v", err)
|
||||||
|
}
|
||||||
|
const sessionID = int64(8103956954238395544)
|
||||||
|
if _, err := r.Dispatch(context.Background(), authKeyID, sessionID, &in); err != nil {
|
||||||
|
t.Fatalf("dispatch plain request: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
entries := logs.FilterMessage("RPC inner handled").All()
|
||||||
|
if len(entries) == 0 {
|
||||||
|
t.Fatalf("RPC inner handled log missing")
|
||||||
|
}
|
||||||
|
fields := entries[len(entries)-1].ContextMap()
|
||||||
|
if got := intLogField(fields["layer"]); got != currentClientLayer {
|
||||||
|
t.Fatalf("logged layer = %d fields=%v, want %d", got, fields, currentClientLayer)
|
||||||
|
}
|
||||||
|
if got := fields["client_type"]; got != string(ClientTypeAndroid) {
|
||||||
|
t.Fatalf("logged client_type = %v, want %s", got, ClientTypeAndroid)
|
||||||
|
}
|
||||||
|
if got := fields["app_version"]; got != "12.8.1 (69169) pbeta" {
|
||||||
|
t.Fatalf("logged app_version = %v, want 12.8.1 (69169) pbeta", got)
|
||||||
|
}
|
||||||
|
if got, ok := r.NegotiatedLayer(authKeyID, sessionID+1); !ok || got != currentClientLayer {
|
||||||
|
t.Fatalf("auth-key fallback layer = (%d,%v), want (%d,true)", got, ok, currentClientLayer)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestAndroidLegacyCompatLogsClientMetadataWithoutInit(t *testing.T) {
|
func TestAndroidLegacyCompatLogsClientMetadataWithoutInit(t *testing.T) {
|
||||||
core, logs := observer.New(zap.DebugLevel)
|
core, logs := observer.New(zap.DebugLevel)
|
||||||
r := New(Config{DC: 2, IP: "127.0.0.1", Port: 2398}, Deps{}, zap.New(core), clock.System)
|
r := New(Config{DC: 2, IP: "127.0.0.1", Port: 2398}, Deps{}, zap.New(core), clock.System)
|
||||||
|
|
@ -210,9 +350,10 @@ func TestAndroidLegacyCompatLogsClientMetadataWithoutInit(t *testing.T) {
|
||||||
t.Fatalf("dispatch legacy updates.getDifference: %v", err)
|
t.Fatalf("dispatch legacy updates.getDifference: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// The legacy android constructor is upgraded by layerwire and dispatched
|
// The legacy Android constructor is upgraded by layerwire and dispatched
|
||||||
// normally; client metadata is still applied (withAndroidCompatMetadata for
|
// normally; client metadata is still applied only because IsClientDrift
|
||||||
// client drift), now surfaced on the standard "RPC inner handled" log.
|
// positively identified a DrKLO constructor, now surfaced on the standard
|
||||||
|
// "RPC inner handled" log.
|
||||||
entries := logs.FilterMessage("RPC inner handled").All()
|
entries := logs.FilterMessage("RPC inner handled").All()
|
||||||
if len(entries) == 0 {
|
if len(entries) == 0 {
|
||||||
t.Fatalf("RPC inner handled log missing")
|
t.Fatalf("RPC inner handled log missing")
|
||||||
|
|
@ -337,6 +478,16 @@ func TestClientTypeDetectsAndroidSDKVersion(t *testing.T) {
|
||||||
if got := info.ClientType(); got != ClientTypeUnknown {
|
if got := info.ClientType(); got != ClientTypeUnknown {
|
||||||
t.Fatalf("gotd test client type = %s, want %s", got, ClientTypeUnknown)
|
t.Fatalf("gotd test client type = %s, want %s", got, ClientTypeUnknown)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
info = normalizeClientInfo(ClientInfo{APIID: 4})
|
||||||
|
if got := info.ClientType(); got != ClientTypeAndroid {
|
||||||
|
t.Fatalf("DrKLO api_id=4 client type = %s, want %s", got, ClientTypeAndroid)
|
||||||
|
}
|
||||||
|
|
||||||
|
info = normalizeClientInfo(ClientInfo{APIID: 2040})
|
||||||
|
if got := info.ClientType(); got != ClientTypeTDesktop {
|
||||||
|
t.Fatalf("TDesktop api_id=2040 client type = %s, want %s", got, ClientTypeTDesktop)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestDispatchRestoresClientMetadataFromAuthorization(t *testing.T) {
|
func TestDispatchRestoresClientMetadataFromAuthorization(t *testing.T) {
|
||||||
|
|
@ -538,6 +689,28 @@ func TestDispatchCachesMissingClientMetadataAuthorizationLookup(t *testing.T) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestDispatchCachesMissingAuthKeyClientMetadataLookup(t *testing.T) {
|
||||||
|
authKeyID := [8]byte{0x68, 0x25, 0xc2, 0xee, 0xf8, 0x82, 0xef, 0x72}
|
||||||
|
auth := &captureAuthService{}
|
||||||
|
r := New(Config{DC: 2, IP: "127.0.0.1", Port: 2398}, Deps{
|
||||||
|
Auth: auth,
|
||||||
|
}, zaptest.NewLogger(t), clock.System)
|
||||||
|
|
||||||
|
for _, sessionID := range []int64{101, 102, 103} {
|
||||||
|
var in bin.Buffer
|
||||||
|
if err := (&tg.HelpGetConfigRequest{}).Encode(&in); err != nil {
|
||||||
|
t.Fatalf("encode request: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := r.Dispatch(context.Background(), authKeyID, sessionID, &in); err != nil {
|
||||||
|
t.Fatalf("dispatch session %d: %v", sessionID, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if auth.authKeyInfoLookups != 1 {
|
||||||
|
t.Fatalf("auth key client info lookups = %d, want 1 cached miss", auth.authKeyInfoLookups)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestCurrentUserIDUsesAuthUserCache(t *testing.T) {
|
func TestCurrentUserIDUsesAuthUserCache(t *testing.T) {
|
||||||
authKeyID := [8]byte{0x42, 0x42, 0x42, 0x42, 0x42, 0x42, 0x42, 0x42}
|
authKeyID := [8]byte{0x42, 0x42, 0x42, 0x42, 0x42, 0x42, 0x42, 0x42}
|
||||||
auth := &captureAuthService{userID: 1000000001}
|
auth := &captureAuthService{userID: 1000000001}
|
||||||
|
|
|
||||||
|
|
@ -25,11 +25,15 @@ type captureAuthService struct {
|
||||||
authorizationLookups int
|
authorizationLookups int
|
||||||
authorizationLists int
|
authorizationLists int
|
||||||
layerUpdates int
|
layerUpdates int
|
||||||
|
authKeyClientInfos map[[8]byte]domain.AuthKeyClientInfo
|
||||||
|
authKeyInfoLookups int
|
||||||
loggedOutAuthKeyID [8]byte
|
loggedOutAuthKeyID [8]byte
|
||||||
pendingPasswordUserID int64
|
pendingPasswordUserID int64
|
||||||
pendingPassword bool
|
pendingPassword bool
|
||||||
completedPasswordKey [8]byte
|
completedPasswordKey [8]byte
|
||||||
completePasswordCount int
|
completePasswordCount int
|
||||||
|
codeDelivery domain.AuthCodeDelivery
|
||||||
|
signInWithEmailCount int
|
||||||
}
|
}
|
||||||
|
|
||||||
type blockingUserAuthService struct {
|
type blockingUserAuthService struct {
|
||||||
|
|
@ -72,6 +76,10 @@ func (s *blockingUserAuthService) SendCode(context.Context, string) (string, err
|
||||||
return "", nil
|
return "", nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (s *blockingUserAuthService) CodeDelivery(context.Context, string) (domain.AuthCodeDelivery, bool, error) {
|
||||||
|
return domain.AuthCodeDelivery{Kind: domain.AuthCodeDeliveryPhone, Length: devCodeLength}, true, nil
|
||||||
|
}
|
||||||
|
|
||||||
func (s *blockingUserAuthService) ResendCode(context.Context, string, string) (string, error) {
|
func (s *blockingUserAuthService) ResendCode(context.Context, string, string) (string, error) {
|
||||||
return "", nil
|
return "", nil
|
||||||
}
|
}
|
||||||
|
|
@ -116,6 +124,14 @@ func (s *blockingUserAuthService) UpdateAuthorizationLayer(context.Context, [8]b
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (s *blockingUserAuthService) AuthKeyClientInfo(context.Context, [8]byte) (domain.AuthKeyClientInfo, bool, error) {
|
||||||
|
return domain.AuthKeyClientInfo{}, false, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *blockingUserAuthService) UpdateAuthKeyClientInfo(context.Context, [8]byte, domain.AuthKeyClientInfo) error {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
func (s *blockingUserAuthService) ListAuthorizations(context.Context, int64) ([]domain.Authorization, error) {
|
func (s *blockingUserAuthService) ListAuthorizations(context.Context, int64) ([]domain.Authorization, error) {
|
||||||
return nil, nil
|
return nil, nil
|
||||||
}
|
}
|
||||||
|
|
@ -154,6 +170,13 @@ func (s *captureAuthService) SendCode(context.Context, string) (string, error) {
|
||||||
return "", nil
|
return "", nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (s *captureAuthService) CodeDelivery(context.Context, string) (domain.AuthCodeDelivery, bool, error) {
|
||||||
|
if s.codeDelivery.Kind != "" {
|
||||||
|
return s.codeDelivery, true, nil
|
||||||
|
}
|
||||||
|
return domain.AuthCodeDelivery{Kind: domain.AuthCodeDeliveryPhone, Length: devCodeLength}, true, nil
|
||||||
|
}
|
||||||
|
|
||||||
func (s *captureAuthService) ResendCode(context.Context, string, string) (string, error) {
|
func (s *captureAuthService) ResendCode(context.Context, string, string) (string, error) {
|
||||||
return "", nil
|
return "", nil
|
||||||
}
|
}
|
||||||
|
|
@ -170,6 +193,7 @@ func (s *captureAuthService) SignIn(context.Context, domain.Authorization, strin
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *captureAuthService) SignInWithEmail(context.Context, domain.Authorization, string, string, string) (domain.User, domain.Message, bool, error) {
|
func (s *captureAuthService) SignInWithEmail(context.Context, domain.Authorization, string, string, string) (domain.User, domain.Message, bool, error) {
|
||||||
|
s.signInWithEmailCount++
|
||||||
if s.signInUser.ID != 0 {
|
if s.signInUser.ID != 0 {
|
||||||
return s.signInUser, domain.Message{}, false, nil
|
return s.signInUser, domain.Message{}, false, nil
|
||||||
}
|
}
|
||||||
|
|
@ -238,6 +262,39 @@ func (s *captureAuthService) UpdateAuthorizationLayer(_ context.Context, authKey
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (s *captureAuthService) AuthKeyClientInfo(_ context.Context, authKeyID [8]byte) (domain.AuthKeyClientInfo, bool, error) {
|
||||||
|
s.authKeyInfoLookups++
|
||||||
|
info, ok := s.authKeyClientInfos[authKeyID]
|
||||||
|
return info, ok, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *captureAuthService) UpdateAuthKeyClientInfo(_ context.Context, authKeyID [8]byte, info domain.AuthKeyClientInfo) error {
|
||||||
|
if s.authKeyClientInfos == nil {
|
||||||
|
s.authKeyClientInfos = make(map[[8]byte]domain.AuthKeyClientInfo)
|
||||||
|
}
|
||||||
|
current := s.authKeyClientInfos[authKeyID]
|
||||||
|
if info.Layer > 0 {
|
||||||
|
current.Layer = info.Layer
|
||||||
|
}
|
||||||
|
if info.DeviceModel != "" {
|
||||||
|
current.DeviceModel = info.DeviceModel
|
||||||
|
}
|
||||||
|
if info.Platform != "" {
|
||||||
|
current.Platform = info.Platform
|
||||||
|
}
|
||||||
|
if info.SystemVersion != "" {
|
||||||
|
current.SystemVersion = info.SystemVersion
|
||||||
|
}
|
||||||
|
if info.APIID != 0 {
|
||||||
|
current.APIID = info.APIID
|
||||||
|
}
|
||||||
|
if info.AppVersion != "" {
|
||||||
|
current.AppVersion = info.AppVersion
|
||||||
|
}
|
||||||
|
s.authKeyClientInfos[authKeyID] = current
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
func (s *captureAuthService) ListAuthorizations(context.Context, int64) ([]domain.Authorization, error) {
|
func (s *captureAuthService) ListAuthorizations(context.Context, int64) ([]domain.Authorization, error) {
|
||||||
s.authorizationLists++
|
s.authorizationLists++
|
||||||
return append([]domain.Authorization(nil), s.authorizations...), nil
|
return append([]domain.Authorization(nil), s.authorizations...), nil
|
||||||
|
|
|
||||||
|
|
@ -9,6 +9,7 @@ import (
|
||||||
// PasswordStore 持久化账号 2FA/SRP 配置。
|
// PasswordStore 持久化账号 2FA/SRP 配置。
|
||||||
type PasswordStore interface {
|
type PasswordStore interface {
|
||||||
GetByUser(ctx context.Context, userID int64) (domain.PasswordSettings, bool, error)
|
GetByUser(ctx context.Context, userID int64) (domain.PasswordSettings, bool, error)
|
||||||
|
LoginEmailOwner(ctx context.Context, email string) (int64, bool, error)
|
||||||
Save(ctx context.Context, userID int64, settings domain.PasswordSettings) error
|
Save(ctx context.Context, userID int64, settings domain.PasswordSettings) error
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -6,19 +6,37 @@ import "context"
|
||||||
//
|
//
|
||||||
// 不依赖 td 协议类型:连接层在边界做 crypto.AuthKey ↔ AuthKeyData 转换。
|
// 不依赖 td 协议类型:连接层在边界做 crypto.AuthKey ↔ AuthKeyData 转换。
|
||||||
type AuthKeyData struct {
|
type AuthKeyData struct {
|
||||||
ID [8]byte // auth_key_id(key 的 SHA1 低 64 位)
|
ID [8]byte // auth_key_id(key 的 SHA1 低 64 位)
|
||||||
Value [256]byte // 2048-bit auth key
|
Value [256]byte // 2048-bit auth key
|
||||||
ServerSalt int64 // 密钥交换产出的初始 server salt
|
ServerSalt int64 // 密钥交换产出的初始 server salt
|
||||||
CreatedAt int64 // unix 秒
|
CreatedAt int64 // unix 秒
|
||||||
|
Layer int
|
||||||
|
DeviceModel string
|
||||||
|
Platform string
|
||||||
|
SystemVersion string
|
||||||
|
APIID int
|
||||||
|
AppVersion string
|
||||||
// 用户绑定不在此处:auth_key 是协议产物,授权(auth_key↔user + 设备信息)由 authorization 承载(P2)。
|
// 用户绑定不在此处:auth_key 是协议产物,授权(auth_key↔user + 设备信息)由 authorization 承载(P2)。
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type AuthKeyClientInfo struct {
|
||||||
|
Layer int
|
||||||
|
DeviceModel string
|
||||||
|
Platform string
|
||||||
|
SystemVersion string
|
||||||
|
APIID int
|
||||||
|
AppVersion string
|
||||||
|
}
|
||||||
|
|
||||||
// AuthKeyStore 持久化 auth key。实现见 store/memory(测试替身)、store/postgres。
|
// AuthKeyStore 持久化 auth key。实现见 store/memory(测试替身)、store/postgres。
|
||||||
type AuthKeyStore interface {
|
type AuthKeyStore interface {
|
||||||
// Save 保存或覆盖一条 auth key 记录。
|
// Save 保存或覆盖一条 auth key 记录。
|
||||||
Save(ctx context.Context, k AuthKeyData) error
|
Save(ctx context.Context, k AuthKeyData) error
|
||||||
// Get 按 auth_key_id 查询;不存在时 found=false。
|
// Get 按 auth_key_id 查询;不存在时 found=false。
|
||||||
Get(ctx context.Context, id [8]byte) (data AuthKeyData, found bool, err error)
|
Get(ctx context.Context, id [8]byte) (data AuthKeyData, found bool, err error)
|
||||||
|
// UpdateClientInfo 合并更新 auth key 的客户端协商元数据。
|
||||||
|
// 空字段不覆盖已有值,layer/api_id 为 0 时不覆盖。
|
||||||
|
UpdateClientInfo(ctx context.Context, id [8]byte, info AuthKeyClientInfo) error
|
||||||
// Delete 删除一条 auth key 记录(destroy_auth_key)。不存在时静默成功。
|
// Delete 删除一条 auth key 记录(destroy_auth_key)。不存在时静默成功。
|
||||||
// 连接层每帧按 auth_key_id 回查本接口,删除后该 key 的入站帧立即失效。
|
// 连接层每帧按 auth_key_id 回查本接口,删除后该 key 的入站帧立即失效。
|
||||||
Delete(ctx context.Context, id [8]byte) error
|
Delete(ctx context.Context, id [8]byte) error
|
||||||
|
|
|
||||||
|
|
@ -5,10 +5,18 @@ import (
|
||||||
"time"
|
"time"
|
||||||
)
|
)
|
||||||
|
|
||||||
// PhoneCode 是一条登录验证码记录(与某次 sendCode 的 phone_code_hash 关联)。
|
// PhoneCode 是一条登录验证码记录(与某次 sendCode 的 phone_code_hash 或邮箱验证键关联)。
|
||||||
type PhoneCode struct {
|
type PhoneCode struct {
|
||||||
Phone string
|
Phone string
|
||||||
Code string
|
Code string
|
||||||
|
Channel string
|
||||||
|
Email string
|
||||||
|
PendingEmail string
|
||||||
|
Attempts int
|
||||||
|
MaxAttempts int
|
||||||
|
VerifiedEmail bool
|
||||||
|
RequireSignUp bool
|
||||||
|
LoginEmailHash string
|
||||||
}
|
}
|
||||||
|
|
||||||
// CodeStore 暂存登录验证码:phone_code_hash → 手机号 + 验证码,带 TTL。
|
// CodeStore 暂存登录验证码:phone_code_hash → 手机号 + 验证码,带 TTL。
|
||||||
|
|
@ -16,5 +24,6 @@ type PhoneCode struct {
|
||||||
type CodeStore interface {
|
type CodeStore interface {
|
||||||
Set(ctx context.Context, phoneCodeHash string, code PhoneCode, ttl time.Duration) error
|
Set(ctx context.Context, phoneCodeHash string, code PhoneCode, ttl time.Duration) error
|
||||||
Get(ctx context.Context, phoneCodeHash string) (PhoneCode, bool, error)
|
Get(ctx context.Context, phoneCodeHash string) (PhoneCode, bool, error)
|
||||||
|
Update(ctx context.Context, phoneCodeHash string, code PhoneCode) error
|
||||||
Del(ctx context.Context, phoneCodeHash string) error
|
Del(ctx context.Context, phoneCodeHash string) error
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -34,6 +34,38 @@ func (s *AuthKeyStore) Get(_ context.Context, id [8]byte) (store.AuthKeyData, bo
|
||||||
return k, ok, nil
|
return k, ok, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (s *AuthKeyStore) UpdateClientInfo(_ context.Context, id [8]byte, info store.AuthKeyClientInfo) error {
|
||||||
|
s.mu.Lock()
|
||||||
|
k, ok := s.keys[id]
|
||||||
|
if ok {
|
||||||
|
mergeAuthKeyClientInfo(&k, info)
|
||||||
|
s.keys[id] = k
|
||||||
|
}
|
||||||
|
s.mu.Unlock()
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func mergeAuthKeyClientInfo(k *store.AuthKeyData, info store.AuthKeyClientInfo) {
|
||||||
|
if info.Layer > 0 {
|
||||||
|
k.Layer = info.Layer
|
||||||
|
}
|
||||||
|
if info.DeviceModel != "" {
|
||||||
|
k.DeviceModel = info.DeviceModel
|
||||||
|
}
|
||||||
|
if info.Platform != "" {
|
||||||
|
k.Platform = info.Platform
|
||||||
|
}
|
||||||
|
if info.SystemVersion != "" {
|
||||||
|
k.SystemVersion = info.SystemVersion
|
||||||
|
}
|
||||||
|
if info.APIID != 0 {
|
||||||
|
k.APIID = info.APIID
|
||||||
|
}
|
||||||
|
if info.AppVersion != "" {
|
||||||
|
k.AppVersion = info.AppVersion
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func (s *AuthKeyStore) Delete(_ context.Context, id [8]byte) error {
|
func (s *AuthKeyStore) Delete(_ context.Context, id [8]byte) error {
|
||||||
s.mu.Lock()
|
s.mu.Lock()
|
||||||
delete(s.keys, id)
|
delete(s.keys, id)
|
||||||
|
|
@ -256,6 +288,18 @@ func (s *CodeStore) Get(_ context.Context, hash string) (store.PhoneCode, bool,
|
||||||
return e.code, true, nil
|
return e.code, true, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (s *CodeStore) Update(_ context.Context, hash string, code store.PhoneCode) error {
|
||||||
|
s.mu.Lock()
|
||||||
|
defer s.mu.Unlock()
|
||||||
|
e, ok := s.m[hash]
|
||||||
|
if !ok || time.Now().After(e.expires) {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
e.code = code
|
||||||
|
s.m[hash] = e
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
func (s *CodeStore) Del(_ context.Context, hash string) error {
|
func (s *CodeStore) Del(_ context.Context, hash string) error {
|
||||||
s.mu.Lock()
|
s.mu.Lock()
|
||||||
delete(s.m, hash)
|
delete(s.m, hash)
|
||||||
|
|
|
||||||
|
|
@ -3,6 +3,7 @@ package memory
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"sort"
|
"sort"
|
||||||
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
"telesrv/internal/domain"
|
"telesrv/internal/domain"
|
||||||
)
|
)
|
||||||
|
|
@ -75,11 +76,40 @@ func (s *PasswordStore) GetByUser(_ context.Context, userID int64) (domain.Passw
|
||||||
|
|
||||||
func (s *PasswordStore) Save(_ context.Context, userID int64, settings domain.PasswordSettings) error {
|
func (s *PasswordStore) Save(_ context.Context, userID int64, settings domain.PasswordSettings) error {
|
||||||
s.mu.Lock()
|
s.mu.Lock()
|
||||||
|
settings.LoginEmail = normalizeLoginEmail(settings.LoginEmail)
|
||||||
|
settings.LoginEmailPattern = domain.MaskEmail(settings.LoginEmail)
|
||||||
|
if settings.LoginEmail != "" {
|
||||||
|
for ownerUserID, existing := range s.m {
|
||||||
|
if ownerUserID != userID && strings.EqualFold(existing.LoginEmail, settings.LoginEmail) {
|
||||||
|
s.mu.Unlock()
|
||||||
|
return domain.ErrEmailOccupied
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
s.m[userID] = clonePasswordSettings(settings)
|
s.m[userID] = clonePasswordSettings(settings)
|
||||||
s.mu.Unlock()
|
s.mu.Unlock()
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (s *PasswordStore) LoginEmailOwner(_ context.Context, email string) (int64, bool, error) {
|
||||||
|
email = normalizeLoginEmail(email)
|
||||||
|
if email == "" {
|
||||||
|
return 0, false, nil
|
||||||
|
}
|
||||||
|
s.mu.RLock()
|
||||||
|
defer s.mu.RUnlock()
|
||||||
|
for userID, settings := range s.m {
|
||||||
|
if strings.EqualFold(settings.LoginEmail, email) {
|
||||||
|
return userID, true, nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return 0, false, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func normalizeLoginEmail(email string) string {
|
||||||
|
return strings.ToLower(strings.TrimSpace(email))
|
||||||
|
}
|
||||||
|
|
||||||
func clonePasswordSettings(in domain.PasswordSettings) domain.PasswordSettings {
|
func clonePasswordSettings(in domain.PasswordSettings) domain.PasswordSettings {
|
||||||
out := in
|
out := in
|
||||||
if in.CurrentAlgo != nil {
|
if in.CurrentAlgo != nil {
|
||||||
|
|
|
||||||
|
|
@ -5,14 +5,19 @@ import (
|
||||||
"database/sql"
|
"database/sql"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"github.com/jackc/pgerrcode"
|
||||||
"github.com/jackc/pgx/v5"
|
"github.com/jackc/pgx/v5"
|
||||||
|
"github.com/jackc/pgx/v5/pgconn"
|
||||||
|
|
||||||
"telesrv/internal/domain"
|
"telesrv/internal/domain"
|
||||||
"telesrv/internal/store/postgres/sqlcgen"
|
"telesrv/internal/store/postgres/sqlcgen"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
const accountPasswordsLoginEmailUniqueIdx = "account_passwords_login_email_lower_unique_idx"
|
||||||
|
|
||||||
// PasswordStore 用 PostgreSQL 实现 store.PasswordStore。
|
// PasswordStore 用 PostgreSQL 实现 store.PasswordStore。
|
||||||
type PasswordStore struct {
|
type PasswordStore struct {
|
||||||
db sqlcgen.DBTX
|
db sqlcgen.DBTX
|
||||||
|
|
@ -70,7 +75,29 @@ WHERE user_id = $1`, userID)
|
||||||
return settings, true, nil
|
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 {
|
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
|
algo := settings.NewAlgo
|
||||||
if settings.CurrentAlgo != nil {
|
if settings.CurrentAlgo != nil {
|
||||||
algo = *settings.CurrentAlgo
|
algo = *settings.CurrentAlgo
|
||||||
|
|
@ -117,11 +144,25 @@ ON CONFLICT (user_id) DO UPDATE SET
|
||||||
settings.RecoveryEmail, settings.RecoveryCode, recoveryExpires, settings.LoginEmail,
|
settings.RecoveryEmail, settings.RecoveryCode, recoveryExpires, settings.LoginEmail,
|
||||||
)
|
)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
if isAccountPasswordLoginEmailUnique(err) {
|
||||||
|
return domain.ErrEmailOccupied
|
||||||
|
}
|
||||||
return fmt.Errorf("upsert account password: %w", err)
|
return fmt.Errorf("upsert account password: %w", err)
|
||||||
}
|
}
|
||||||
return nil
|
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 {
|
func nonNilBytea(in []byte) []byte {
|
||||||
if in != nil {
|
if in != nil {
|
||||||
return in
|
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"
|
"fmt"
|
||||||
|
|
||||||
"github.com/jackc/pgx/v5"
|
"github.com/jackc/pgx/v5"
|
||||||
|
"github.com/jackc/pgx/v5/pgtype"
|
||||||
|
|
||||||
"telesrv/internal/store"
|
"telesrv/internal/store"
|
||||||
"telesrv/internal/store/postgres/sqlcgen"
|
"telesrv/internal/store/postgres/sqlcgen"
|
||||||
|
|
@ -26,11 +27,12 @@ func NewAuthKeyStore(db sqlcgen.DBTX) *AuthKeyStore {
|
||||||
// Save 实现 store.AuthKeyStore。auth_key_id 以小端解释为 int64 存入 BIGINT;
|
// Save 实现 store.AuthKeyStore。auth_key_id 以小端解释为 int64 存入 BIGINT;
|
||||||
// created_at 交由 DB 默认值(now()),故传入的 CreatedAt 不落库。
|
// created_at 交由 DB 默认值(now()),故传入的 CreatedAt 不落库。
|
||||||
func (s *AuthKeyStore) Save(ctx context.Context, k store.AuthKeyData) error {
|
func (s *AuthKeyStore) Save(ctx context.Context, k store.AuthKeyData) error {
|
||||||
if err := s.q.UpsertAuthKey(ctx, sqlcgen.UpsertAuthKeyParams{
|
if _, err := s.db.Exec(ctx, `
|
||||||
AuthKeyID: authKeyIDToInt64(k.ID),
|
INSERT INTO auth_keys (auth_key_id, body, server_salt)
|
||||||
Body: k.Value[:],
|
VALUES ($1, $2, $3)
|
||||||
ServerSalt: k.ServerSalt,
|
ON CONFLICT (auth_key_id) DO UPDATE
|
||||||
}); err != nil {
|
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 fmt.Errorf("upsert auth key: %w", err)
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
|
|
@ -38,24 +40,65 @@ func (s *AuthKeyStore) Save(ctx context.Context, k store.AuthKeyData) error {
|
||||||
|
|
||||||
// Get 实现 store.AuthKeyStore。不存在时 found=false。
|
// Get 实现 store.AuthKeyStore。不存在时 found=false。
|
||||||
func (s *AuthKeyStore) Get(ctx context.Context, id [8]byte) (store.AuthKeyData, bool, error) {
|
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 err != nil {
|
||||||
if errors.Is(err, pgx.ErrNoRows) {
|
if errors.Is(err, pgx.ErrNoRows) {
|
||||||
return store.AuthKeyData{}, false, nil
|
return store.AuthKeyData{}, false, nil
|
||||||
}
|
}
|
||||||
return store.AuthKeyData{}, false, fmt.Errorf("get auth key: %w", err)
|
return store.AuthKeyData{}, false, fmt.Errorf("get auth key: %w", err)
|
||||||
}
|
}
|
||||||
if len(row.Body) != len(store.AuthKeyData{}.Value) {
|
if len(body) != len(store.AuthKeyData{}.Value) {
|
||||||
return store.AuthKeyData{}, false, fmt.Errorf("auth key body length = %d, want 256", len(row.Body))
|
return store.AuthKeyData{}, false, fmt.Errorf("auth key body length = %d, want 256", len(body))
|
||||||
}
|
}
|
||||||
data := store.AuthKeyData{ID: id, ServerSalt: row.ServerSalt}
|
data := store.AuthKeyData{
|
||||||
copy(data.Value[:], row.Body)
|
ID: id,
|
||||||
if row.CreatedAt.Valid {
|
ServerSalt: serverSalt,
|
||||||
data.CreatedAt = row.CreatedAt.Time.Unix()
|
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
|
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。不存在时静默成功。
|
// Delete 实现 store.AuthKeyStore。不存在时静默成功。
|
||||||
// 手写 SQL 而非 sqlc 生成:避免触碰 sqlcgen 再生成链路。
|
// 手写 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)
|
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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -50,6 +50,25 @@ func (s *CodeStore) Get(ctx context.Context, hash string) (store.PhoneCode, bool
|
||||||
return code, true, nil
|
return code, true, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (s *CodeStore) Update(ctx context.Context, hash string, code store.PhoneCode) error {
|
||||||
|
key := codeKey(hash)
|
||||||
|
ttl, err := s.c.TTL(ctx, key).Result()
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("redis ttl phone code: %w", err)
|
||||||
|
}
|
||||||
|
if ttl <= 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
v, err := json.Marshal(code)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("marshal phone code: %w", err)
|
||||||
|
}
|
||||||
|
if err := s.c.Set(ctx, key, v, ttl).Err(); err != nil {
|
||||||
|
return fmt.Errorf("redis update phone code: %w", err)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
func (s *CodeStore) Del(ctx context.Context, hash string) error {
|
func (s *CodeStore) Del(ctx context.Context, hash string) error {
|
||||||
return s.c.Del(ctx, codeKey(hash)).Err()
|
return s.c.Del(ctx, codeKey(hash)).Err()
|
||||||
}
|
}
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue