feat(server/auth): Register 加 inviteCode + ReferralHook(注册后 best-effort 归因)
Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
This commit is contained in:
@@ -43,7 +43,7 @@ func TestService_RegisterDevice_OnRegisterAndLogin(t *testing.T) {
|
|||||||
t.Fatalf("SendCode: %v", err)
|
t.Fatalf("SendCode: %v", err)
|
||||||
}
|
}
|
||||||
code := codeInRedis(t, svc, email)
|
code := codeInRedis(t, svc, email)
|
||||||
if _, e := svc.Register(ctx, email, code, pw, "1.2.3.4", meta); e != nil {
|
if _, e := svc.Register(ctx, email, code, pw, "1.2.3.4", meta, ""); e != nil {
|
||||||
t.Fatalf("Register: %v", e)
|
t.Fatalf("Register: %v", e)
|
||||||
}
|
}
|
||||||
if len(reg.calls) != 1 || reg.calls[0].meta.DeviceID != "dev-uuid-1" || reg.calls[0].meta.Platform != "macos" {
|
if len(reg.calls) != 1 || reg.calls[0].meta.DeviceID != "dev-uuid-1" || reg.calls[0].meta.Platform != "macos" {
|
||||||
@@ -79,7 +79,7 @@ func TestService_Login_SurfacesDeviceLimit(t *testing.T) {
|
|||||||
t.Fatalf("SendCode: %v", err)
|
t.Fatalf("SendCode: %v", err)
|
||||||
}
|
}
|
||||||
code := codeInRedis(t, svc, email)
|
code := codeInRedis(t, svc, email)
|
||||||
if _, e := svc.Register(ctx, email, code, pw, "", meta); e != nil {
|
if _, e := svc.Register(ctx, email, code, pw, "", meta, ""); e != nil {
|
||||||
t.Fatalf("Register: %v", e)
|
t.Fatalf("Register: %v", e)
|
||||||
}
|
}
|
||||||
out, _, e := svc.Login(ctx, email, pw, "", meta)
|
out, _, e := svc.Login(ctx, email, pw, "", meta)
|
||||||
@@ -104,7 +104,7 @@ func TestService_RegisterDevice_BestEffort(t *testing.T) {
|
|||||||
t.Fatalf("SendCode: %v", err)
|
t.Fatalf("SendCode: %v", err)
|
||||||
}
|
}
|
||||||
code := codeInRedis(t, svc, email)
|
code := codeInRedis(t, svc, email)
|
||||||
if _, e := svc.Register(ctx, email, code, "supersecret", "", DeviceMeta{DeviceID: "x", Platform: "windows"}); e != nil {
|
if _, e := svc.Register(ctx, email, code, "supersecret", "", DeviceMeta{DeviceID: "x", Platform: "windows"}, ""); e != nil {
|
||||||
t.Fatalf("Register must succeed despite registrar error: %v", e)
|
t.Fatalf("Register must succeed despite registrar error: %v", e)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -120,7 +120,7 @@ func TestService_RegisterDevice_NoMeta(t *testing.T) {
|
|||||||
t.Fatalf("SendCode: %v", err)
|
t.Fatalf("SendCode: %v", err)
|
||||||
}
|
}
|
||||||
code := codeInRedis(t, svc, email)
|
code := codeInRedis(t, svc, email)
|
||||||
if _, e := svc.Register(ctx, email, code, "supersecret", "", DeviceMeta{}); e != nil {
|
if _, e := svc.Register(ctx, email, code, "supersecret", "", DeviceMeta{}, ""); e != nil {
|
||||||
t.Fatalf("Register: %v", e)
|
t.Fatalf("Register: %v", e)
|
||||||
}
|
}
|
||||||
if len(reg.calls) != 0 {
|
if len(reg.calls) != 0 {
|
||||||
|
|||||||
@@ -61,10 +61,11 @@ func (d deviceBody) toMeta() DeviceMeta {
|
|||||||
}
|
}
|
||||||
|
|
||||||
type registerRequest struct {
|
type registerRequest struct {
|
||||||
Email string `json:"email"`
|
Email string `json:"email"`
|
||||||
Code string `json:"code"`
|
Code string `json:"code"`
|
||||||
Password string `json:"password"`
|
Password string `json:"password"`
|
||||||
Device deviceBody `json:"device"`
|
Device deviceBody `json:"device"`
|
||||||
|
InviteCode string `json:"invite_code"` // 选填:邀请码,注册后 best-effort 归因
|
||||||
}
|
}
|
||||||
|
|
||||||
type loginRequest struct {
|
type loginRequest struct {
|
||||||
@@ -104,7 +105,7 @@ func (h *Handler) Register(w http.ResponseWriter, r *http.Request) {
|
|||||||
if !decodeJSON(w, r, &req) {
|
if !decodeJSON(w, r, &req) {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
pair, apiErr := h.svc.Register(r.Context(), req.Email, req.Code, req.Password, clientIP(r), req.Device.toMeta())
|
pair, apiErr := h.svc.Register(r.Context(), req.Email, req.Code, req.Password, clientIP(r), req.Device.toMeta(), req.InviteCode)
|
||||||
if apiErr != nil {
|
if apiErr != nil {
|
||||||
writeAPIErr(w, apiErr, 0)
|
writeAPIErr(w, apiErr, 0)
|
||||||
return
|
return
|
||||||
|
|||||||
@@ -140,7 +140,7 @@ func TestIntegration_FullChain(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// 2. Register → trial subscription must exist for 7 days.
|
// 2. Register → trial subscription must exist for 7 days.
|
||||||
pair, apiErr := svc.Register(ctx, email, code, pw, "203.0.113.10", DeviceMeta{})
|
pair, apiErr := svc.Register(ctx, email, code, pw, "203.0.113.10", DeviceMeta{}, "")
|
||||||
if apiErr != nil {
|
if apiErr != nil {
|
||||||
t.Fatalf("Register: %v", apiErr)
|
t.Fatalf("Register: %v", apiErr)
|
||||||
}
|
}
|
||||||
@@ -170,7 +170,7 @@ func TestIntegration_FullChain(t *testing.T) {
|
|||||||
}
|
}
|
||||||
// Force a fresh code regardless of rate limit.
|
// Force a fresh code regardless of rate limit.
|
||||||
_ = rdb.Set(ctx, codeKey(email), code, 10*time.Minute).Err()
|
_ = rdb.Set(ctx, codeKey(email), code, 10*time.Minute).Err()
|
||||||
if _, e := svc.Register(ctx, email, code, pw, "203.0.113.10", DeviceMeta{}); e == nil || e.Code != ErrCodeInvalid.Code {
|
if _, e := svc.Register(ctx, email, code, pw, "203.0.113.10", DeviceMeta{}, ""); e == nil || e.Code != ErrCodeInvalid.Code {
|
||||||
t.Fatalf("want code_invalid (anti-enumeration), got %v", e)
|
t.Fatalf("want code_invalid (anti-enumeration), got %v", e)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,43 @@
|
|||||||
|
package auth
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
// fakeHook captures the arguments passed to OnRegister so tests can assert on
|
||||||
|
// them without depending on the real reward.Service.
|
||||||
|
type fakeHook struct {
|
||||||
|
gotInvitee int64
|
||||||
|
gotCode string
|
||||||
|
gotDev string
|
||||||
|
}
|
||||||
|
|
||||||
|
func (f *fakeHook) OnRegister(_ context.Context, inviteeID int64, code, dev string) {
|
||||||
|
f.gotInvitee, f.gotCode, f.gotDev = inviteeID, code, dev
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestRegister_CallsReferralHookWithInviteCode verifies that a successful
|
||||||
|
// registration best-effort invokes the ReferralHook with the invite code and
|
||||||
|
// this registration's device UUID (not some stale/historical value).
|
||||||
|
func TestRegister_CallsReferralHookWithInviteCode(t *testing.T) {
|
||||||
|
svc, _, _ := newService(t, ServiceConfig{})
|
||||||
|
ctx := context.Background()
|
||||||
|
const email = "new@x.com"
|
||||||
|
|
||||||
|
if _, err := svc.SendCode(ctx, email, "1.2.3.4"); err != nil {
|
||||||
|
t.Fatalf("SendCode: %v", err)
|
||||||
|
}
|
||||||
|
code := codeInRedis(t, svc, email)
|
||||||
|
|
||||||
|
h := &fakeHook{}
|
||||||
|
svc.SetReferralHook(h)
|
||||||
|
|
||||||
|
_, apiErr := svc.Register(ctx, email, code, "password123", "1.2.3.4", DeviceMeta{DeviceID: "dev-1"}, "INVCODE9")
|
||||||
|
if apiErr != nil {
|
||||||
|
t.Fatalf("Register: %v", apiErr)
|
||||||
|
}
|
||||||
|
if h.gotCode != "INVCODE9" || h.gotDev != "dev-1" || h.gotInvitee == 0 {
|
||||||
|
t.Fatalf("hook got invitee=%d code=%q dev=%q", h.gotInvitee, h.gotCode, h.gotDev)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -132,12 +132,22 @@ type Service struct {
|
|||||||
now func() time.Time
|
now func() time.Time
|
||||||
devReg DeviceRegistrar // nil until wired; registration is best-effort
|
devReg DeviceRegistrar // nil until wired; registration is best-effort
|
||||||
sessions SessionStore // nil until wired; session recording is best-effort
|
sessions SessionStore // nil until wired; session recording is best-effort
|
||||||
|
referral ReferralHook // nil until wired; referral attribution is best-effort
|
||||||
}
|
}
|
||||||
|
|
||||||
// SetDeviceRegistrar / SetSessionStore wire collaborators after construction
|
// ReferralHook is invoked after a successful registration to attribute the new
|
||||||
// (main keeps auth, devices and sessions decoupled). Call once during startup.
|
// account to the invite code it registered with (best-effort; implemented by
|
||||||
|
// reward.Service). Defined consumer-side to avoid an import cycle.
|
||||||
|
type ReferralHook interface {
|
||||||
|
OnRegister(ctx context.Context, inviteeID int64, inviteCode, deviceUUID string)
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetDeviceRegistrar / SetSessionStore / SetReferralHook wire collaborators
|
||||||
|
// after construction (main keeps auth, devices, sessions and reward decoupled).
|
||||||
|
// Call once during startup.
|
||||||
func (s *Service) SetDeviceRegistrar(r DeviceRegistrar) { s.devReg = r }
|
func (s *Service) SetDeviceRegistrar(r DeviceRegistrar) { s.devReg = r }
|
||||||
func (s *Service) SetSessionStore(st SessionStore) { s.sessions = st }
|
func (s *Service) SetSessionStore(st SessionStore) { s.sessions = st }
|
||||||
|
func (s *Service) SetReferralHook(h ReferralHook) { s.referral = h }
|
||||||
|
|
||||||
// recordLogin registers the device and binds a session to the freshly-issued
|
// recordLogin registers the device and binds a session to the freshly-issued
|
||||||
// refresh JTI. Best-effort: a registrar/session error (device cap, transient DB)
|
// refresh JTI. Best-effort: a registrar/session error (device cap, transient DB)
|
||||||
@@ -267,7 +277,7 @@ func (s *Service) SendCode(ctx context.Context, rawEmail, ip string) (retryAfter
|
|||||||
|
|
||||||
// Register verifies the code (one-time), creates the account plus a 7-day PRO
|
// Register verifies the code (one-time), creates the account plus a 7-day PRO
|
||||||
// trial in a single transaction, and returns a fresh token pair.
|
// trial in a single transaction, and returns a fresh token pair.
|
||||||
func (s *Service) Register(ctx context.Context, rawEmail, code, password, ip string, device DeviceMeta) (*TokenPair, *apierr.Error) {
|
func (s *Service) Register(ctx context.Context, rawEmail, code, password, ip string, device DeviceMeta, inviteCode string) (*TokenPair, *apierr.Error) {
|
||||||
email := NormalizeEmail(rawEmail)
|
email := NormalizeEmail(rawEmail)
|
||||||
if !ValidEmail(email) || len(password) < 8 || len(code) != 6 {
|
if !ValidEmail(email) || len(password) < 8 || len(code) != 6 {
|
||||||
return nil, ErrInvalidRequest
|
return nil, ErrInvalidRequest
|
||||||
@@ -301,6 +311,9 @@ func (s *Service) Register(ctx context.Context, rawEmail, code, password, ip str
|
|||||||
return nil, ErrInternal
|
return nil, ErrInternal
|
||||||
}
|
}
|
||||||
s.recordLogin(ctx, user.ID, jti, ip, device)
|
s.recordLogin(ctx, user.ID, jti, ip, device)
|
||||||
|
if s.referral != nil {
|
||||||
|
s.referral.OnRegister(ctx, user.ID, inviteCode, device.DeviceID) // best-effort
|
||||||
|
}
|
||||||
return pair, nil
|
return pair, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -40,7 +40,7 @@ func TestService_RegisterFullFlow(t *testing.T) {
|
|||||||
}
|
}
|
||||||
code := codeInRedis(t, svc, email)
|
code := codeInRedis(t, svc, email)
|
||||||
|
|
||||||
pair, apiErr := svc.Register(ctx, email, code, "supersecret", "", DeviceMeta{})
|
pair, apiErr := svc.Register(ctx, email, code, "supersecret", "", DeviceMeta{}, "")
|
||||||
if apiErr != nil {
|
if apiErr != nil {
|
||||||
t.Fatalf("Register: %v", apiErr)
|
t.Fatalf("Register: %v", apiErr)
|
||||||
}
|
}
|
||||||
@@ -84,7 +84,7 @@ func TestService_DuplicateEmailConflict(t *testing.T) {
|
|||||||
|
|
||||||
// First registration.
|
// First registration.
|
||||||
_, _ = svc.SendCode(ctx, email, "")
|
_, _ = svc.SendCode(ctx, email, "")
|
||||||
if _, e := svc.Register(ctx, email, codeInRedis(t, svc, email), "password1", "", DeviceMeta{}); e != nil {
|
if _, e := svc.Register(ctx, email, codeInRedis(t, svc, email), "password1", "", DeviceMeta{}, ""); e != nil {
|
||||||
t.Fatalf("first register: %v", e)
|
t.Fatalf("first register: %v", e)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -103,7 +103,7 @@ func TestService_DuplicateEmailConflict(t *testing.T) {
|
|||||||
if err := svc.rdb.Set(ctx, codeKey(email), "654321", 10*time.Minute).Err(); err != nil {
|
if err := svc.rdb.Set(ctx, codeKey(email), "654321", 10*time.Minute).Err(); err != nil {
|
||||||
t.Fatalf("force code: %v", err)
|
t.Fatalf("force code: %v", err)
|
||||||
}
|
}
|
||||||
_, apiErr := svc.Register(ctx, email, "654321", "password2", "", DeviceMeta{})
|
_, apiErr := svc.Register(ctx, email, "654321", "password2", "", DeviceMeta{}, "")
|
||||||
if apiErr == nil || apiErr.Code != ErrCodeInvalid.Code {
|
if apiErr == nil || apiErr.Code != ErrCodeInvalid.Code {
|
||||||
t.Fatalf("want code_invalid (anti-enumeration), got %v", apiErr)
|
t.Fatalf("want code_invalid (anti-enumeration), got %v", apiErr)
|
||||||
}
|
}
|
||||||
@@ -115,7 +115,7 @@ func TestService_CodeWrong(t *testing.T) {
|
|||||||
const email = "wrong@example.com"
|
const email = "wrong@example.com"
|
||||||
_, _ = svc.SendCode(ctx, email, "")
|
_, _ = svc.SendCode(ctx, email, "")
|
||||||
|
|
||||||
_, apiErr := svc.Register(ctx, email, "000000", "password1", "", DeviceMeta{})
|
_, apiErr := svc.Register(ctx, email, "000000", "password1", "", DeviceMeta{}, "")
|
||||||
if apiErr == nil || apiErr.Code != ErrCodeInvalid.Code {
|
if apiErr == nil || apiErr.Code != ErrCodeInvalid.Code {
|
||||||
t.Fatalf("want code_invalid, got %v", apiErr)
|
t.Fatalf("want code_invalid, got %v", apiErr)
|
||||||
}
|
}
|
||||||
@@ -131,7 +131,7 @@ func TestService_CodeExpired(t *testing.T) {
|
|||||||
// Expire the code key.
|
// Expire the code key.
|
||||||
svc.rdb.Del(ctx, codeKey(email))
|
svc.rdb.Del(ctx, codeKey(email))
|
||||||
|
|
||||||
_, apiErr := svc.Register(ctx, email, code, "password1", "", DeviceMeta{})
|
_, apiErr := svc.Register(ctx, email, code, "password1", "", DeviceMeta{}, "")
|
||||||
if apiErr == nil || apiErr.Code != ErrCodeInvalid.Code {
|
if apiErr == nil || apiErr.Code != ErrCodeInvalid.Code {
|
||||||
t.Fatalf("want code_invalid after expiry, got %v", apiErr)
|
t.Fatalf("want code_invalid after expiry, got %v", apiErr)
|
||||||
}
|
}
|
||||||
@@ -144,11 +144,11 @@ func TestService_CodeReuseRejected(t *testing.T) {
|
|||||||
_, _ = svc.SendCode(ctx, email, "")
|
_, _ = svc.SendCode(ctx, email, "")
|
||||||
code := codeInRedis(t, svc, email)
|
code := codeInRedis(t, svc, email)
|
||||||
|
|
||||||
if _, e := svc.Register(ctx, email, code, "password1", "", DeviceMeta{}); e != nil {
|
if _, e := svc.Register(ctx, email, code, "password1", "", DeviceMeta{}, ""); e != nil {
|
||||||
t.Fatalf("first register: %v", e)
|
t.Fatalf("first register: %v", e)
|
||||||
}
|
}
|
||||||
// Re-using the consumed code must fail.
|
// Re-using the consumed code must fail.
|
||||||
_, apiErr := svc.Register(ctx, "other@example.com", code, "password1", "", DeviceMeta{})
|
_, apiErr := svc.Register(ctx, "other@example.com", code, "password1", "", DeviceMeta{}, "")
|
||||||
if apiErr == nil || apiErr.Code != ErrCodeInvalid.Code {
|
if apiErr == nil || apiErr.Code != ErrCodeInvalid.Code {
|
||||||
t.Fatalf("want code_invalid on reuse, got %v", apiErr)
|
t.Fatalf("want code_invalid on reuse, got %v", apiErr)
|
||||||
}
|
}
|
||||||
@@ -163,12 +163,12 @@ func TestService_CodeBruteForceBurned(t *testing.T) {
|
|||||||
|
|
||||||
// 3 wrong attempts burn the code.
|
// 3 wrong attempts burn the code.
|
||||||
for i := 0; i < 3; i++ {
|
for i := 0; i < 3; i++ {
|
||||||
if _, e := svc.Register(ctx, email, "999999", "password1", "", DeviceMeta{}); e == nil {
|
if _, e := svc.Register(ctx, email, "999999", "password1", "", DeviceMeta{}, ""); e == nil {
|
||||||
t.Fatal("wrong code should fail")
|
t.Fatal("wrong code should fail")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
// Even the correct code no longer works.
|
// Even the correct code no longer works.
|
||||||
if _, e := svc.Register(ctx, email, good, "password1", "", DeviceMeta{}); e == nil || e.Code != ErrCodeInvalid.Code {
|
if _, e := svc.Register(ctx, email, good, "password1", "", DeviceMeta{}, ""); e == nil || e.Code != ErrCodeInvalid.Code {
|
||||||
t.Fatalf("burned code should reject correct value, got %v", e)
|
t.Fatalf("burned code should reject correct value, got %v", e)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -205,7 +205,7 @@ func TestService_LoginAndLockout(t *testing.T) {
|
|||||||
const pw = "rightpassword"
|
const pw = "rightpassword"
|
||||||
|
|
||||||
_, _ = svc.SendCode(ctx, email, "")
|
_, _ = svc.SendCode(ctx, email, "")
|
||||||
if _, e := svc.Register(ctx, email, codeInRedis(t, svc, email), pw, "", DeviceMeta{}); e != nil {
|
if _, e := svc.Register(ctx, email, codeInRedis(t, svc, email), pw, "", DeviceMeta{}, ""); e != nil {
|
||||||
t.Fatalf("register: %v", e)
|
t.Fatalf("register: %v", e)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -247,7 +247,7 @@ func TestService_BannedUserRejected(t *testing.T) {
|
|||||||
const pw = "password1"
|
const pw = "password1"
|
||||||
|
|
||||||
_, _ = svc.SendCode(ctx, email, "")
|
_, _ = svc.SendCode(ctx, email, "")
|
||||||
if _, e := svc.Register(ctx, email, codeInRedis(t, svc, email), pw, "", DeviceMeta{}); e != nil {
|
if _, e := svc.Register(ctx, email, codeInRedis(t, svc, email), pw, "", DeviceMeta{}, ""); e != nil {
|
||||||
t.Fatalf("register: %v", e)
|
t.Fatalf("register: %v", e)
|
||||||
}
|
}
|
||||||
store.setStatus(email, "banned")
|
store.setStatus(email, "banned")
|
||||||
@@ -264,7 +264,7 @@ func TestService_RefreshRotation(t *testing.T) {
|
|||||||
const email = "refresh@example.com"
|
const email = "refresh@example.com"
|
||||||
|
|
||||||
_, _ = svc.SendCode(ctx, email, "")
|
_, _ = svc.SendCode(ctx, email, "")
|
||||||
pair, e := svc.Register(ctx, email, codeInRedis(t, svc, email), "password1", "", DeviceMeta{})
|
pair, e := svc.Register(ctx, email, codeInRedis(t, svc, email), "password1", "", DeviceMeta{}, "")
|
||||||
if e != nil {
|
if e != nil {
|
||||||
t.Fatalf("register: %v", e)
|
t.Fatalf("register: %v", e)
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user