diff --git a/server/internal/auth/device_register_test.go b/server/internal/auth/device_register_test.go index af37572..6c3924c 100644 --- a/server/internal/auth/device_register_test.go +++ b/server/internal/auth/device_register_test.go @@ -43,7 +43,7 @@ func TestService_RegisterDevice_OnRegisterAndLogin(t *testing.T) { t.Fatalf("SendCode: %v", err) } 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) } 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) } 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) } out, _, e := svc.Login(ctx, email, pw, "", meta) @@ -104,7 +104,7 @@ func TestService_RegisterDevice_BestEffort(t *testing.T) { t.Fatalf("SendCode: %v", err) } 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) } } @@ -120,7 +120,7 @@ func TestService_RegisterDevice_NoMeta(t *testing.T) { t.Fatalf("SendCode: %v", err) } 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) } if len(reg.calls) != 0 { diff --git a/server/internal/auth/handler.go b/server/internal/auth/handler.go index 24389e6..6152ae7 100644 --- a/server/internal/auth/handler.go +++ b/server/internal/auth/handler.go @@ -61,10 +61,11 @@ func (d deviceBody) toMeta() DeviceMeta { } type registerRequest struct { - Email string `json:"email"` - Code string `json:"code"` - Password string `json:"password"` - Device deviceBody `json:"device"` + Email string `json:"email"` + Code string `json:"code"` + Password string `json:"password"` + Device deviceBody `json:"device"` + InviteCode string `json:"invite_code"` // 选填:邀请码,注册后 best-effort 归因 } type loginRequest struct { @@ -104,7 +105,7 @@ func (h *Handler) Register(w http.ResponseWriter, r *http.Request) { if !decodeJSON(w, r, &req) { 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 { writeAPIErr(w, apiErr, 0) return diff --git a/server/internal/auth/integration_test.go b/server/internal/auth/integration_test.go index e3fb5bf..015aec7 100644 --- a/server/internal/auth/integration_test.go +++ b/server/internal/auth/integration_test.go @@ -140,7 +140,7 @@ func TestIntegration_FullChain(t *testing.T) { } // 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 { t.Fatalf("Register: %v", apiErr) } @@ -170,7 +170,7 @@ func TestIntegration_FullChain(t *testing.T) { } // Force a fresh code regardless of rate limit. _ = 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) } diff --git a/server/internal/auth/register_invite_test.go b/server/internal/auth/register_invite_test.go new file mode 100644 index 0000000..0cd4c78 --- /dev/null +++ b/server/internal/auth/register_invite_test.go @@ -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) + } +} diff --git a/server/internal/auth/service.go b/server/internal/auth/service.go index e99b78a..0b77e14 100644 --- a/server/internal/auth/service.go +++ b/server/internal/auth/service.go @@ -132,12 +132,22 @@ type Service struct { now func() time.Time devReg DeviceRegistrar // nil until wired; registration 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 -// (main keeps auth, devices and sessions decoupled). Call once during startup. +// ReferralHook is invoked after a successful registration to attribute the new +// 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) 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 // 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 // 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) if !ValidEmail(email) || len(password) < 8 || len(code) != 6 { return nil, ErrInvalidRequest @@ -301,6 +311,9 @@ func (s *Service) Register(ctx context.Context, rawEmail, code, password, ip str return nil, ErrInternal } 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 } diff --git a/server/internal/auth/service_test.go b/server/internal/auth/service_test.go index 33edc61..9f83d64 100644 --- a/server/internal/auth/service_test.go +++ b/server/internal/auth/service_test.go @@ -40,7 +40,7 @@ func TestService_RegisterFullFlow(t *testing.T) { } 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 { t.Fatalf("Register: %v", apiErr) } @@ -84,7 +84,7 @@ func TestService_DuplicateEmailConflict(t *testing.T) { // First registration. _, _ = 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) } @@ -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 { 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 { 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" _, _ = 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 { t.Fatalf("want code_invalid, got %v", apiErr) } @@ -131,7 +131,7 @@ func TestService_CodeExpired(t *testing.T) { // Expire the code key. 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 { t.Fatalf("want code_invalid after expiry, got %v", apiErr) } @@ -144,11 +144,11 @@ func TestService_CodeReuseRejected(t *testing.T) { _, _ = svc.SendCode(ctx, 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) } // 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 { 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. 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") } } // 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) } } @@ -205,7 +205,7 @@ func TestService_LoginAndLockout(t *testing.T) { const pw = "rightpassword" _, _ = 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) } @@ -247,7 +247,7 @@ func TestService_BannedUserRejected(t *testing.T) { const pw = "password1" _, _ = 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) } store.setStatus(email, "banned") @@ -264,7 +264,7 @@ func TestService_RefreshRotation(t *testing.T) { const email = "refresh@example.com" _, _ = 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 { t.Fatalf("register: %v", e) }