Update: Save current progress
This commit is contained in:
@@ -73,7 +73,7 @@ func (l *EmailLoginLogic) EmailLogin(req *types.EmailLoginRequest) (resp *types.
|
||||
if err := json.Unmarshal([]byte(value), &payload); err != nil {
|
||||
continue
|
||||
}
|
||||
if payload.Code == req.Code && time.Now().Unix()-payload.LastAt <= 900 {
|
||||
if payload.Code == req.Code && time.Now().Unix()-payload.LastAt <= l.svcCtx.Config.VerifyCode.ExpireTime {
|
||||
verified = true
|
||||
cacheKeyUsed = cacheKey
|
||||
break
|
||||
|
||||
@@ -84,7 +84,7 @@ func (l *ResetPasswordLogic) ResetPassword(req *types.ResetPasswordRequest) (res
|
||||
return nil, errors.Wrapf(xerr.NewErrCode(xerr.VerifyCodeError), "Verification code error")
|
||||
}
|
||||
// 校验有效期(15分钟)
|
||||
if time.Now().Unix()-payload.LastAt > 900 {
|
||||
if time.Now().Unix()-payload.LastAt > l.svcCtx.Config.VerifyCode.ExpireTime {
|
||||
l.Errorw("Verification code expired", logger.Field("cacheKey", cacheKey), logger.Field("error", "Verification code expired"), logger.Field("reqCode", req.Code), logger.Field("payloadCode", payload.Code))
|
||||
return nil, errors.Wrapf(xerr.NewErrCode(xerr.VerifyCodeError), "code expired")
|
||||
}
|
||||
|
||||
@@ -78,7 +78,7 @@ func (l *UserRegisterLogic) UserRegister(req *types.UserRegisterRequest) (resp *
|
||||
return nil, errors.Wrapf(xerr.NewErrCode(xerr.VerifyCodeError), "code error")
|
||||
}
|
||||
// 校验有效期(15分钟)
|
||||
if time.Now().Unix()-payload.LastAt > 900 {
|
||||
if time.Now().Unix()-payload.LastAt > l.svcCtx.Config.VerifyCode.ExpireTime {
|
||||
return nil, errors.Wrapf(xerr.NewErrCode(xerr.VerifyCodeError), "code expired")
|
||||
}
|
||||
l.svcCtx.Redis.Del(l.ctx, cacheKey)
|
||||
|
||||
@@ -91,7 +91,7 @@ func (l *SendEmailCodeLogic) SendEmailCode(req *types.SendCodeRequest) (resp *ty
|
||||
"Type": req.Type,
|
||||
"SiteLogo": l.svcCtx.Config.Site.SiteLogo,
|
||||
"SiteName": l.svcCtx.Config.Site.SiteName,
|
||||
"Expire": 15,
|
||||
"Expire": l.svcCtx.Config.VerifyCode.ExpireTime / 60,
|
||||
"Code": code,
|
||||
}
|
||||
// Save to Redis
|
||||
@@ -101,7 +101,7 @@ func (l *SendEmailCodeLogic) SendEmailCode(req *types.SendCodeRequest) (resp *ty
|
||||
}
|
||||
// Marshal the payload
|
||||
val, _ := json.Marshal(payload)
|
||||
if err = l.svcCtx.Redis.Set(l.ctx, cacheKey, string(val), time.Minute*15).Err(); err != nil {
|
||||
if err = l.svcCtx.Redis.Set(l.ctx, cacheKey, string(val), time.Second*time.Duration(l.svcCtx.Config.VerifyCode.ExpireTime)).Err(); err != nil {
|
||||
l.Errorw("[SendEmailCode]: Redis Error", logger.Field("error", err.Error()), logger.Field("cacheKey", cacheKey))
|
||||
return nil, errors.Wrap(xerr.NewErrCode(xerr.ERROR), "Failed to set verification code")
|
||||
}
|
||||
|
||||
@@ -0,0 +1,174 @@
|
||||
package user
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"reflect"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
"unsafe"
|
||||
|
||||
"github.com/alicebob/miniredis/v2"
|
||||
"github.com/gorilla/websocket"
|
||||
"github.com/perfect-panel/server/internal/config"
|
||||
"github.com/perfect-panel/server/internal/model/user"
|
||||
"github.com/perfect-panel/server/internal/svc"
|
||||
"github.com/perfect-panel/server/internal/types"
|
||||
"github.com/perfect-panel/server/pkg/constant"
|
||||
"github.com/perfect-panel/server/pkg/device"
|
||||
"github.com/redis/go-redis/v9"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
type MockEmailModel struct {
|
||||
MockUserModel
|
||||
}
|
||||
|
||||
func (m *MockEmailModel) FindUserAuthMethods(ctx context.Context, userId int64) ([]*user.AuthMethods, error) {
|
||||
return []*user.AuthMethods{
|
||||
{UserId: userId, AuthType: "device", AuthIdentifier: "device-1", Verified: true},
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (m *MockEmailModel) FindUserAuthMethodByOpenID(ctx context.Context, method, openID string) (*user.AuthMethods, error) {
|
||||
if openID == "test@example.com" {
|
||||
// 返回已存在的用户(不同的UserId)
|
||||
return &user.AuthMethods{Id: 99, UserId: 2, AuthType: "email", AuthIdentifier: openID}, nil
|
||||
}
|
||||
return nil, gorm.ErrRecordNotFound
|
||||
}
|
||||
|
||||
func (m *MockEmailModel) QueryDeviceList(ctx context.Context, userId int64) ([]*user.Device, int64, error) {
|
||||
// 模拟当前用户(User 1)持有设备 device-1
|
||||
if userId == 1 {
|
||||
return []*user.Device{
|
||||
{Id: 10, UserId: 1, Identifier: "device-1", Enabled: true},
|
||||
}, 1, nil
|
||||
}
|
||||
return nil, 0, nil
|
||||
}
|
||||
|
||||
func (m *MockEmailModel) UpdateDevice(ctx context.Context, data *user.Device, tx ...*gorm.DB) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// 模拟 Transaction 失败,以便在 KickDevice 后停止
|
||||
func (m *MockEmailModel) Transaction(ctx context.Context, fn func(db *gorm.DB) error) error {
|
||||
return fmt.Errorf("stop testing here")
|
||||
}
|
||||
|
||||
func (m *MockEmailModel) FindOne(ctx context.Context, id int64) (*user.User, error) {
|
||||
return &user.User{Id: id}, nil
|
||||
}
|
||||
|
||||
func TestBindEmailWithVerification_KickDevice(t *testing.T) {
|
||||
// 1. Redis Mock
|
||||
mr, err := miniredis.Run()
|
||||
assert.NoError(t, err)
|
||||
defer mr.Close()
|
||||
|
||||
rdb := redis.NewClient(&redis.Options{Addr: mr.Addr()})
|
||||
|
||||
// 准备验证码数据
|
||||
email := "test@example.com"
|
||||
code := "123456"
|
||||
payload := map[string]interface{}{
|
||||
"code": code,
|
||||
"lastAt": time.Now().Unix(),
|
||||
}
|
||||
bytes, _ := json.Marshal(payload)
|
||||
cacheKey := fmt.Sprintf("%s:%s:%s", config.AuthCodeCacheKey, constant.Register.String(), email)
|
||||
rdb.Set(context.Background(), cacheKey, string(bytes), time.Minute)
|
||||
|
||||
// 2. DeviceManager Mock
|
||||
// 启动 WebSocket 服务器以获取真实连接
|
||||
var serverConn *websocket.Conn
|
||||
connDone := make(chan struct{})
|
||||
|
||||
s := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
upgrader := websocket.Upgrader{}
|
||||
c, _ := upgrader.Upgrade(w, r, nil)
|
||||
serverConn = c
|
||||
close(connDone)
|
||||
// 保持连接直到测试结束 (read loop)
|
||||
for {
|
||||
if _, _, err := c.ReadMessage(); err != nil {
|
||||
break
|
||||
}
|
||||
}
|
||||
}))
|
||||
defer s.Close()
|
||||
|
||||
// 客户端连接
|
||||
wsURL := "ws" + strings.TrimPrefix(s.URL, "http")
|
||||
clientConn, _, err := websocket.DefaultDialer.Dial(wsURL, nil)
|
||||
assert.NoError(t, err)
|
||||
defer clientConn.Close()
|
||||
|
||||
<-connDone // 等待服务端获取连接
|
||||
|
||||
dm := device.NewDeviceManager(10, 10)
|
||||
|
||||
// 注入设备 (UserId=1, DeviceId="device-1")
|
||||
dev := &device.Device{
|
||||
Session: "session-1",
|
||||
DeviceID: "device-1",
|
||||
Conn: serverConn,
|
||||
}
|
||||
|
||||
// 使用反射注入
|
||||
v := reflect.ValueOf(dm).Elem()
|
||||
f := v.FieldByName("userDevices")
|
||||
// 直接获取指针
|
||||
userDevicesMap := (*sync.Map)(unsafe.Pointer(f.UnsafeAddr()))
|
||||
userDevicesMap.Store(int64(1), []*device.Device{dev})
|
||||
|
||||
// 3. User Mock
|
||||
mockModel := &MockEmailModel{}
|
||||
// 初始化内部 map,虽然这里只用到 override 的方法
|
||||
mockModel.users = make(map[int64]*user.User)
|
||||
|
||||
svcCtx := &svc.ServiceContext{
|
||||
UserModel: mockModel,
|
||||
Redis: rdb,
|
||||
DeviceManager: dm,
|
||||
Config: config.Config{
|
||||
VerifyCode: config.VerifyCode{ExpireTime: 900}, // Correct type
|
||||
JwtAuth: config.JwtAuth{MaxSessionsPerUser: 10},
|
||||
},
|
||||
}
|
||||
|
||||
// 4. Run Logic
|
||||
currentUser := &user.User{Id: 1} // 当前用户
|
||||
ctx := context.WithValue(context.Background(), constant.CtxKeyUser, currentUser)
|
||||
l := NewBindEmailWithVerificationLogic(ctx, svcCtx)
|
||||
|
||||
req := &types.BindEmailWithVerificationRequest{
|
||||
Email: email,
|
||||
Code: code,
|
||||
}
|
||||
|
||||
// 执行
|
||||
_, err = l.BindEmailWithVerification(req)
|
||||
// 我们预期这里会返回错误 ("stop testing here")
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "stop testing here")
|
||||
|
||||
// 5. Verify
|
||||
// 验证设备是否被移除 (KickDevice 会从 userDevices 中移除被踢出的设备)
|
||||
val, ok := userDevicesMap.Load(int64(1))
|
||||
|
||||
if ok {
|
||||
// 如果 key 还在,检查列表是否为空
|
||||
devices := val.([]*device.Device)
|
||||
assert.Empty(t, devices, "设备列表应为空 (KickDevice 应该移除设备)")
|
||||
} else {
|
||||
// key 不存在,说明已移除,符合预期
|
||||
}
|
||||
}
|
||||
@@ -67,7 +67,7 @@ func (l *BindEmailWithVerificationLogic) BindEmailWithVerification(req *types.Bi
|
||||
continue
|
||||
}
|
||||
// 校验验证码及有效期(15分钟)
|
||||
if p.Code == req.Code && time.Now().Unix()-p.LastAt <= 900 {
|
||||
if p.Code == req.Code && time.Now().Unix()-p.LastAt <= l.svcCtx.Config.VerifyCode.ExpireTime {
|
||||
_ = l.svcCtx.Redis.Del(l.ctx, cacheKey).Err()
|
||||
verified = true
|
||||
break
|
||||
@@ -141,6 +141,8 @@ func (l *BindEmailWithVerificationLogic) BindEmailWithVerification(req *types.Bi
|
||||
l.Errorw("查询用户设备列表失败", logger.Field("error", err.Error()), logger.Field("email_user_id", emailUserId))
|
||||
return nil, errors.Wrapf(xerr.NewErrCode(xerr.DatabaseDeletedError), "查询用户设备列表失败")
|
||||
}
|
||||
// 保存原用户ID,用于踢出旧连接
|
||||
originalUserId := u.Id
|
||||
for _, device := range devices {
|
||||
// 删除原本的设备记录
|
||||
// err = l.svcCtx.UserModel.DeleteDevice(l.ctx, device.Id)
|
||||
@@ -155,6 +157,27 @@ func (l *BindEmailWithVerificationLogic) BindEmailWithVerification(req *types.Bi
|
||||
l.Errorw("更新邮箱用户设备记录失败", logger.Field("error", err.Error()), logger.Field("email_user_id", emailUserId))
|
||||
return nil, errors.Wrapf(xerr.NewErrCode(xerr.DatabaseUpdateError), "更新原本的设备记录失败")
|
||||
}
|
||||
|
||||
// 踢出设备的旧 WebSocket 连接(使用原用户ID)
|
||||
l.svcCtx.DeviceManager.KickDevice(originalUserId, device.Identifier)
|
||||
l.Infow("已踢出设备旧连接",
|
||||
logger.Field("device_identifier", device.Identifier),
|
||||
logger.Field("original_user_id", originalUserId),
|
||||
logger.Field("new_user_id", emailUserId))
|
||||
|
||||
// 清理设备相关的 Redis 缓存
|
||||
deviceCacheKey := fmt.Sprintf("%v:%v", config.DeviceCacheKeyKey, device.Identifier)
|
||||
if sessionId, rerr := l.svcCtx.Redis.Get(l.ctx, deviceCacheKey).Result(); rerr == nil && sessionId != "" {
|
||||
_ = l.svcCtx.Redis.Del(l.ctx, deviceCacheKey).Err()
|
||||
sessionIdCacheKey := fmt.Sprintf("%v:%v", config.SessionIdKey, sessionId)
|
||||
_ = l.svcCtx.Redis.Del(l.ctx, sessionIdCacheKey).Err()
|
||||
// 清理 user_sessions
|
||||
sessionsKey := fmt.Sprintf("%s%v", config.UserSessionsKeyPrefix, originalUserId)
|
||||
_ = l.svcCtx.Redis.ZRem(l.ctx, sessionsKey, sessionId).Err()
|
||||
l.Infow("已清理设备缓存",
|
||||
logger.Field("device_identifier", device.Identifier),
|
||||
logger.Field("session_id", sessionId))
|
||||
}
|
||||
}
|
||||
// 再次更新 user_auth_method : 因为之前 默认 设备登录的时候 创建了一个设备认证数据
|
||||
// 现在需要 更新 为 邮箱认证
|
||||
|
||||
@@ -10,6 +10,7 @@ import (
|
||||
"github.com/perfect-panel/server/pkg/logger"
|
||||
"github.com/perfect-panel/server/pkg/xerr"
|
||||
"github.com/pkg/errors"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
type BindInviteCodeLogic struct {
|
||||
@@ -43,6 +44,9 @@ func (l *BindInviteCodeLogic) BindInviteCode(req *types.BindInviteCodeRequest) e
|
||||
// 查找邀请人
|
||||
referrer, err := l.svcCtx.UserModel.FindOneByReferCode(l.ctx, req.InviteCode)
|
||||
if err != nil {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return errors.Wrapf(xerr.NewErrCodeMsg(xerr.InviteCodeError, "无邀请码"), "invite code not found")
|
||||
}
|
||||
logger.WithContext(l.ctx).Error(err)
|
||||
return errors.Wrapf(xerr.NewErrCode(xerr.DatabaseQueryError), "query referrer failed: %v", err.Error())
|
||||
}
|
||||
|
||||
@@ -0,0 +1,136 @@
|
||||
package user
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/perfect-panel/server/internal/model/user"
|
||||
"github.com/perfect-panel/server/internal/svc"
|
||||
"github.com/perfect-panel/server/internal/types"
|
||||
"github.com/perfect-panel/server/pkg/constant"
|
||||
"github.com/perfect-panel/server/pkg/xerr"
|
||||
"github.com/pkg/errors"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// MockUserModel 只实现 bindInviteCodeLogic 需要的方法
|
||||
type MockUserModel struct {
|
||||
user.Model // 为了满足接口定义,嵌入 user.Model,未实现的方法会 panic
|
||||
users map[int64]*user.User
|
||||
}
|
||||
|
||||
func (m *MockUserModel) FindOneByReferCode(ctx context.Context, referCode string) (*user.User, error) {
|
||||
for _, u := range m.users {
|
||||
if u.ReferCode == referCode {
|
||||
return u, nil
|
||||
}
|
||||
}
|
||||
return nil, gorm.ErrRecordNotFound
|
||||
}
|
||||
|
||||
func (m *MockUserModel) Update(ctx context.Context, data *user.User, tx ...*gorm.DB) error {
|
||||
m.users[data.Id] = data
|
||||
return nil
|
||||
}
|
||||
|
||||
func TestBindInviteCodeLogic_BindInviteCode(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
currentUser user.User // 使用值类型,在 Run 中取地址,避免共享
|
||||
initUsers map[int64]*user.User
|
||||
inviteCode string
|
||||
expectError bool
|
||||
expectedCode uint32
|
||||
expectedMsg string
|
||||
}{
|
||||
{
|
||||
name: "成功绑定邀请码",
|
||||
currentUser: user.User{Id: 1, ReferCode: "CODE1", RefererId: 0},
|
||||
initUsers: map[int64]*user.User{
|
||||
1: {Id: 1, ReferCode: "CODE1", RefererId: 0},
|
||||
2: {Id: 2, ReferCode: "CODE2", RefererId: 0},
|
||||
},
|
||||
inviteCode: "CODE2",
|
||||
expectError: false,
|
||||
},
|
||||
{
|
||||
name: "邀请码不存在",
|
||||
currentUser: user.User{Id: 1, ReferCode: "CODE1", RefererId: 0},
|
||||
initUsers: map[int64]*user.User{
|
||||
1: {Id: 1, ReferCode: "CODE1", RefererId: 0},
|
||||
},
|
||||
inviteCode: "INVALID",
|
||||
expectError: true,
|
||||
expectedCode: xerr.InviteCodeError,
|
||||
expectedMsg: "无邀请码",
|
||||
},
|
||||
{
|
||||
name: "不允许绑定自己",
|
||||
currentUser: user.User{Id: 1, ReferCode: "CODE1", RefererId: 0},
|
||||
initUsers: map[int64]*user.User{
|
||||
1: {Id: 1, ReferCode: "CODE1", RefererId: 0},
|
||||
},
|
||||
inviteCode: "CODE1",
|
||||
expectError: true,
|
||||
expectedCode: xerr.InviteCodeError,
|
||||
expectedMsg: "不允许绑定自己",
|
||||
},
|
||||
{
|
||||
name: "用户已经绑定过",
|
||||
currentUser: user.User{Id: 3, ReferCode: "CODE3", RefererId: 2},
|
||||
initUsers: map[int64]*user.User{
|
||||
3: {Id: 3, ReferCode: "CODE3", RefererId: 2},
|
||||
2: {Id: 2, ReferCode: "CODE2", RefererId: 0},
|
||||
},
|
||||
inviteCode: "CODE2",
|
||||
expectError: true,
|
||||
expectedCode: xerr.UserExist,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
// 初始化 Mock 数据
|
||||
mockModel := &MockUserModel{
|
||||
users: tt.initUsers,
|
||||
}
|
||||
svcCtx := &svc.ServiceContext{
|
||||
UserModel: mockModel,
|
||||
}
|
||||
|
||||
// 确保 User 对象在 Mock DB 中也存在(Update操作需要)
|
||||
// 其实 MockUserModel.Update 会更新 map,所以这里不需要额外操作,
|
||||
// 只要 initUsers 配置正确即可。
|
||||
|
||||
// 将当前用户注入 context (使用拷贝的指针)
|
||||
u := tt.currentUser
|
||||
ctx := context.WithValue(context.Background(), constant.CtxKeyUser, &u)
|
||||
l := NewBindInviteCodeLogic(ctx, svcCtx)
|
||||
|
||||
err := l.BindInviteCode(&types.BindInviteCodeRequest{InviteCode: tt.inviteCode})
|
||||
|
||||
if tt.expectError {
|
||||
assert.Error(t, err)
|
||||
cause := errors.Cause(err)
|
||||
codeErr, ok := cause.(*xerr.CodeError)
|
||||
if !ok {
|
||||
// handle error
|
||||
} else {
|
||||
assert.Equal(t, tt.expectedCode, codeErr.GetErrCode())
|
||||
if tt.expectedMsg != "" {
|
||||
assert.Contains(t, codeErr.GetErrMsg(), tt.expectedMsg)
|
||||
}
|
||||
}
|
||||
if tt.expectedMsg != "" {
|
||||
assert.Contains(t, err.Error(), tt.expectedMsg)
|
||||
}
|
||||
} else {
|
||||
assert.NoError(t, err)
|
||||
if tt.name == "成功绑定邀请码" {
|
||||
assert.Equal(t, int64(2), mockModel.users[1].RefererId)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -119,6 +119,9 @@ func (l *UnbindDeviceLogic) UnbindDevice(req *types.UnbindDeviceRequest) error {
|
||||
_ = l.svcCtx.Redis.Del(ctx, deviceCacheKey).Err()
|
||||
sessionIdCacheKey := fmt.Sprintf("%v:%v", config.SessionIdKey, sessionId)
|
||||
_ = l.svcCtx.Redis.Del(ctx, sessionIdCacheKey).Err()
|
||||
// remove session from user sessions
|
||||
sessionsKey := fmt.Sprintf("%s%v", config.UserSessionsKeyPrefix, u.Id)
|
||||
_ = l.svcCtx.Redis.ZRem(ctx, sessionsKey, sessionId).Err()
|
||||
}
|
||||
l.svcCtx.DeviceManager.KickDevice(u.Id, identifier)
|
||||
l.Infow("设备解绑完成",
|
||||
|
||||
@@ -53,7 +53,7 @@ func (l *VerifyEmailLogic) VerifyEmail(req *types.VerifyEmailRequest) error {
|
||||
if payload.Code != req.Code { // 校验有效期(15分钟)
|
||||
return errors.Wrapf(xerr.NewErrCode(xerr.VerifyCodeError), "code error")
|
||||
}
|
||||
if time.Now().Unix()-payload.LastAt > 900 {
|
||||
if time.Now().Unix()-payload.LastAt > l.svcCtx.Config.VerifyCode.ExpireTime {
|
||||
return errors.Wrapf(xerr.NewErrCode(xerr.VerifyCodeError), "code expired")
|
||||
}
|
||||
l.svcCtx.Redis.Del(l.ctx, cacheKey)
|
||||
|
||||
Reference in New Issue
Block a user