diff --git a/internal/logic/server/getServerUserListLogic.go b/internal/logic/server/getServerUserListLogic.go index a15f660..80bc085 100644 --- a/internal/logic/server/getServerUserListLogic.go +++ b/internal/logic/server/getServerUserListLogic.go @@ -19,6 +19,7 @@ import ( "github.com/perfect-panel/server/pkg/tool" "github.com/perfect-panel/server/pkg/uuidx" "github.com/perfect-panel/server/pkg/xerr" + "github.com/redis/go-redis/v9" ) type GetServerUserListLogic struct { @@ -27,6 +28,46 @@ type GetServerUserListLogic struct { svcCtx *svc.ServiceContext } +type serverUserListPerfStats struct { + serverID int64 + protocol string + cacheHit bool + nodesCount int + subsCount int + usersCount int + speedLimitZeroCount int + speedLimitPositiveCount int + trafficCalcMS int64 + startedAt time.Time +} + +type serverUserSpeedLimitCandidate struct { + userID int64 + userSubscribeID int64 + baseSpeed int64 + trafficLimit string +} + +type serverUserTrafficWindow struct { + statType string + statValue int64 + start time.Time + end time.Time +} + +type serverUserTrafficUsageKey struct { + userID int64 + userSubscribeID int64 + statType string + statValue int64 +} + +type serverUserTrafficUsage struct { + userID int64 + userSubscribeID int64 + usedGB float64 +} + // NewGetServerUserListLogic Get user list func NewGetServerUserListLogic(ctx *gin.Context, svcCtx *svc.ServiceContext) *GetServerUserListLogic { return &GetServerUserListLogic{ @@ -37,15 +78,27 @@ func NewGetServerUserListLogic(ctx *gin.Context, svcCtx *svc.ServiceContext) *Ge } func (l *GetServerUserListLogic) GetServerUserList(req *types.GetServerUserListRequest) (resp *types.GetServerUserListResponse, err error) { + startedAt := time.Now() protocolRequest := normalizeServerUserListProtocol(req.Protocol) + stats := serverUserListPerfStats{ + serverID: req.ServerId, + protocol: protocolRequest, + startedAt: startedAt, + cacheHit: false, + } cacheKey := fmt.Sprintf("%s%d:%s", node.ServerUserListCacheKey, req.ServerId, protocolRequest) cache, err := l.svcCtx.Redis.Get(l.ctx, cacheKey).Result() + if err != nil && err != redis.Nil { + l.Errorw("[ServerUserListCacheKey] redis get error", logger.Field("error", err.Error())) + } if cache != "" { + stats.cacheHit = true etag := tool.GenerateETag([]byte(cache)) resp = &types.GetServerUserListResponse{} // Check If-None-Match header if match := l.ctx.GetHeader("If-None-Match"); match == etag { + l.logServerUserListPerf(stats) return nil, xerr.StatusNotModified } l.ctx.Header("ETag", etag) @@ -54,6 +107,9 @@ func (l *GetServerUserListLogic) GetServerUserList(req *types.GetServerUserListR l.Errorw("[ServerUserListCacheKey] json unmarshal error", logger.Field("error", err.Error())) return nil, err } + stats.usersCount = len(resp.Users) + stats.recordSpeedLimits(resp.Users) + l.logServerUserListPerf(stats) return resp, nil } server, err := l.svcCtx.NodeModel.FindOneServer(l.ctx, req.ServerId) @@ -72,12 +128,16 @@ func (l *GetServerUserListLogic) GetServerUserList(req *types.GetServerUserListR l.Errorw("FilterNodeList error", logger.Field("error", err.Error())) return nil, err } + stats.nodesCount = len(nodes) if len(nodes) == 0 { l.Errorw("[ServerUserList] fallback: no nodes matched server+protocol, returning placeholder without cache", logger.Field("server_id", req.ServerId), logger.Field("protocol", req.Protocol), ) + stats.usersCount = 1 + stats.speedLimitZeroCount = 1 + l.logServerUserListPerf(stats) return &types.GetServerUserListResponse{ Users: []types.ServerUser{ { @@ -149,6 +209,9 @@ func (l *GetServerUserListLogic) GetServerUserList(req *types.GetServerUserListR logger.Field("server_id", req.ServerId), logger.Field("protocol", req.Protocol), ) + stats.usersCount = 1 + stats.speedLimitZeroCount = 1 + l.logServerUserListPerf(stats) return &types.GetServerUserListResponse{ Users: []types.ServerUser{ { @@ -158,7 +221,9 @@ func (l *GetServerUserListLogic) GetServerUserList(req *types.GetServerUserListR }, }, nil } + stats.subsCount = len(subs) users := make([]types.ServerUser, 0) + speedCandidates := make(map[int64]serverUserSpeedLimitCandidate) for _, sub := range subs { data, err := l.svcCtx.UserModel.FindUsersSubscribeBySubscribeId(l.ctx, sub.Id) if err != nil { @@ -169,17 +234,29 @@ func (l *GetServerUserListLogic) GetServerUserList(req *types.GetServerUserListR continue } - // 计算该用户的实际限速值(考虑按量限速规则) - effectiveSpeedLimit := l.calculateEffectiveSpeedLimit(sub, datum) - + baseSpeed, trafficLimit := serverUserSpeedLimitInputs(sub, datum) + speedCandidates[datum.Id] = serverUserSpeedLimitCandidate{ + userID: datum.UserId, + userSubscribeID: datum.Id, + baseSpeed: baseSpeed, + trafficLimit: trafficLimit, + } users = append(users, types.ServerUser{ Id: datum.Id, UUID: datum.UUID, - SpeedLimit: effectiveSpeedLimit, + SpeedLimit: baseSpeed, DeviceLimit: sub.DeviceLimit, }) } } + trafficCalcStartedAt := time.Now() + speedLimits := l.calculateServerUserSpeedLimits(speedCandidates) + stats.trafficCalcMS = time.Since(trafficCalcStartedAt).Milliseconds() + for i := range users { + if speedLimit, ok := speedLimits[users[i].Id]; ok { + users[i].SpeedLimit = speedLimit + } + } // 处理过期订阅用户:如果当前节点属于过期节点组,添加符合条件的过期用户 if len(nodeGroupIds) > 0 { @@ -197,6 +274,9 @@ func (l *GetServerUserListLogic) GetServerUserList(req *types.GetServerUserListR logger.Field("server_id", req.ServerId), logger.Field("protocol", req.Protocol), ) + stats.usersCount = 1 + stats.speedLimitZeroCount = 1 + l.logServerUserListPerf(stats) return &types.GetServerUserListResponse{ Users: []types.ServerUser{ { @@ -209,6 +289,8 @@ func (l *GetServerUserListLogic) GetServerUserList(req *types.GetServerUserListR resp = &types.GetServerUserListResponse{ Users: users, } + stats.usersCount = len(users) + stats.recordSpeedLimits(users) val, _ := json.Marshal(resp) etag := tool.GenerateETag(val) l.ctx.Header("ETag", etag) @@ -218,8 +300,10 @@ func (l *GetServerUserListLogic) GetServerUserList(req *types.GetServerUserListR } // Check If-None-Match header if match := l.ctx.GetHeader("If-None-Match"); match == etag { + l.logServerUserListPerf(stats) return nil, xerr.StatusNotModified } + l.logServerUserListPerf(stats) return resp, nil } @@ -334,8 +418,7 @@ func (l *GetServerUserListLogic) canUseExpiredNodeGroup(userSub *user.Subscribe, return true } -// calculateEffectiveSpeedLimit 计算用户的实际限速值(考虑按量限速规则) -func (l *GetServerUserListLogic) calculateEffectiveSpeedLimit(sub *subscribe.Subscribe, userSub *user.Subscribe) int64 { +func serverUserSpeedLimitInputs(sub *subscribe.Subscribe, userSub *user.Subscribe) (int64, string) { baseSpeed := sub.SpeedLimit if userSub.SpeedLimit > 0 { baseSpeed = userSub.SpeedLimit @@ -346,15 +429,170 @@ func (l *GetServerUserListLogic) calculateEffectiveSpeedLimit(sub *subscribe.Sub trafficLimit = *userSub.TrafficLimit } - result := speedlimit.CalculateWithCache( - l.ctx.Request.Context(), - l.svcCtx.Redis, - l.svcCtx.DB, - userSub.UserId, - userSub.Id, - baseSpeed, - trafficLimit, - 30*time.Second, - ) - return result.EffectiveSpeed + return baseSpeed, trafficLimit +} + +func (l *GetServerUserListLogic) calculateServerUserSpeedLimits(candidates map[int64]serverUserSpeedLimitCandidate) map[int64]int64 { + speedLimits := make(map[int64]int64, len(candidates)) + rulesByUserSubscribeID := make(map[int64][]speedlimit.TrafficLimitRule) + windowByKey := make(map[string]serverUserTrafficWindow) + now := time.Now() + + for userSubscribeID, candidate := range candidates { + speedLimits[userSubscribeID] = candidate.baseSpeed + if candidate.trafficLimit == "" { + continue + } + + var rules []speedlimit.TrafficLimitRule + if err := json.Unmarshal([]byte(candidate.trafficLimit), &rules); err != nil || len(rules) == 0 { + continue + } + rulesByUserSubscribeID[userSubscribeID] = rules + for _, rule := range rules { + window, ok := serverUserTrafficRuleWindow(rule, now) + if !ok { + continue + } + windowByKey[serverUserTrafficWindowKey(rule.StatType, rule.StatValue)] = window + } + } + + if len(rulesByUserSubscribeID) == 0 || len(windowByKey) == 0 { + return speedLimits + } + + usageByKey := make(map[serverUserTrafficUsageKey]float64) + for _, window := range windowByKey { + trafficUsage, err := l.queryServerUserTrafficUsage(candidates, window) + if err != nil { + l.Errorw("[ServerUserList] batch traffic usage query failed", + logger.Field("error", err.Error()), + logger.Field("stat_type", window.statType), + logger.Field("stat_value", window.statValue), + ) + continue + } + for _, usage := range trafficUsage { + usageByKey[serverUserTrafficUsageKey{ + userID: usage.userID, + userSubscribeID: usage.userSubscribeID, + statType: window.statType, + statValue: window.statValue, + }] = usage.usedGB + } + } + + for userSubscribeID, rules := range rulesByUserSubscribeID { + candidate := candidates[userSubscribeID] + for _, rule := range rules { + if rule.SpeedLimit <= 0 { + continue + } + if _, ok := serverUserTrafficRuleWindow(rule, now); !ok { + continue + } + usedGB := usageByKey[serverUserTrafficUsageKey{ + userID: candidate.userID, + userSubscribeID: candidate.userSubscribeID, + statType: rule.StatType, + statValue: rule.StatValue, + }] + if usedGB < float64(rule.TrafficUsage) { + continue + } + current := speedLimits[userSubscribeID] + if current == 0 || rule.SpeedLimit < current { + speedLimits[userSubscribeID] = rule.SpeedLimit + } + } + } + + return speedLimits +} + +func (l *GetServerUserListLogic) queryServerUserTrafficUsage( + candidates map[int64]serverUserSpeedLimitCandidate, + window serverUserTrafficWindow, +) ([]serverUserTrafficUsage, error) { + userSubscribeIDs := make([]int64, 0, len(candidates)) + for userSubscribeID := range candidates { + userSubscribeIDs = append(userSubscribeIDs, userSubscribeID) + } + + var rows []struct { + UserID int64 + UserSubscribeID int64 + Upload int64 + Download int64 + } + err := l.svcCtx.DB.WithContext(l.ctx.Request.Context()). + Table("traffic_log"). + Select("user_id, subscribe_id AS user_subscribe_id, COALESCE(SUM(upload), 0) AS upload, COALESCE(SUM(download), 0) AS download"). + Where("subscribe_id IN ? AND timestamp >= ? AND timestamp < ?", userSubscribeIDs, window.start, window.end). + Group("user_id, subscribe_id"). + Scan(&rows).Error + if err != nil { + return nil, err + } + + trafficUsage := make([]serverUserTrafficUsage, 0, len(rows)) + for _, row := range rows { + trafficUsage = append(trafficUsage, serverUserTrafficUsage{ + userID: row.UserID, + userSubscribeID: row.UserSubscribeID, + usedGB: float64(row.Upload+row.Download) / (1024 * 1024 * 1024), + }) + } + return trafficUsage, nil +} + +func serverUserTrafficRuleWindow(rule speedlimit.TrafficLimitRule, now time.Time) (serverUserTrafficWindow, bool) { + if rule.StatValue <= 0 { + return serverUserTrafficWindow{}, false + } + + window := serverUserTrafficWindow{ + statType: rule.StatType, + statValue: rule.StatValue, + end: now, + } + switch rule.StatType { + case "hour": + window.start = now.Add(-time.Duration(rule.StatValue) * time.Hour) + case "day": + window.start = now.AddDate(0, 0, -int(rule.StatValue)) + default: + return serverUserTrafficWindow{}, false + } + return window, true +} + +func serverUserTrafficWindowKey(statType string, statValue int64) string { + return fmt.Sprintf("%s:%d", statType, statValue) +} + +func (s *serverUserListPerfStats) recordSpeedLimits(users []types.ServerUser) { + for _, user := range users { + if user.SpeedLimit > 0 { + s.speedLimitPositiveCount++ + continue + } + s.speedLimitZeroCount++ + } +} + +func (l *GetServerUserListLogic) logServerUserListPerf(stats serverUserListPerfStats) { + l.Infow("[ServerUserList] performance", + logger.Field("server_id", stats.serverID), + logger.Field("protocol", stats.protocol), + logger.Field("cache_hit", stats.cacheHit), + logger.Field("nodes_count", stats.nodesCount), + logger.Field("subs_count", stats.subsCount), + logger.Field("users_count", stats.usersCount), + logger.Field("speed_limit_0_count", stats.speedLimitZeroCount), + logger.Field("speed_limit_positive_count", stats.speedLimitPositiveCount), + logger.Field("traffic_calc_ms", stats.trafficCalcMS), + logger.Field("total_ms", time.Since(stats.startedAt).Milliseconds()), + ) } diff --git a/internal/logic/server/getServerUserListLogic_test.go b/internal/logic/server/getServerUserListLogic_test.go index def0509..2744b3f 100644 --- a/internal/logic/server/getServerUserListLogic_test.go +++ b/internal/logic/server/getServerUserListLogic_test.go @@ -1,11 +1,24 @@ package server import ( + "context" "fmt" + "net/http" + "regexp" "slices" "testing" + "time" + "github.com/DATA-DOG/go-sqlmock" + "github.com/gin-gonic/gin" "github.com/perfect-panel/server/internal/model/node" + "github.com/perfect-panel/server/internal/model/subscribe" + "github.com/perfect-panel/server/internal/model/user" + "github.com/perfect-panel/server/internal/svc" + "github.com/perfect-panel/server/pkg/logger" + "github.com/perfect-panel/server/pkg/speedlimit" + "gorm.io/driver/mysql" + "gorm.io/gorm" ) // TestNormalizeServerUserListProtocol 验证客户端 hysteria2 兼容字段被映射回 @@ -76,3 +89,252 @@ func TestServerUserListCacheKey_ShapeMatchesDelPath(t *testing.T) { t.Fatalf("write-path key %q not present in Del-path enumeration %v", writeShape, delKeys) } } + +func TestServerUserSpeedLimitInputs(t *testing.T) { + planTrafficLimit := `[{"stat_type":"day","stat_value":1,"traffic_usage":10,"speed_limit":5}]` + userTrafficLimit := `[{"stat_type":"day","stat_value":1,"traffic_usage":1,"speed_limit":3}]` + + cases := []struct { + name string + sub *subscribe.Subscribe + userSub *user.Subscribe + wantSpeedLimit int64 + wantTrafficLimit string + }{ + { + name: "未设置用户覆盖时返回套餐基础限速", + sub: &subscribe.Subscribe{ + SpeedLimit: 100, + TrafficLimit: planTrafficLimit, + }, + userSub: &user.Subscribe{}, + wantSpeedLimit: 100, + wantTrafficLimit: planTrafficLimit, + }, + { + name: "用户 speed_limit 覆盖套餐基础限速", + sub: &subscribe.Subscribe{ + SpeedLimit: 100, + TrafficLimit: planTrafficLimit, + }, + userSub: &user.Subscribe{ + SpeedLimit: 30, + }, + wantSpeedLimit: 30, + wantTrafficLimit: planTrafficLimit, + }, + { + name: "用户 traffic_limit 覆盖套餐规则", + sub: &subscribe.Subscribe{ + SpeedLimit: 100, + TrafficLimit: planTrafficLimit, + }, + userSub: &user.Subscribe{ + TrafficLimit: &userTrafficLimit, + }, + wantSpeedLimit: 100, + wantTrafficLimit: userTrafficLimit, + }, + { + name: "用户 traffic_limit 为空时保留套餐规则", + sub: &subscribe.Subscribe{ + SpeedLimit: 100, + TrafficLimit: planTrafficLimit, + }, + userSub: &user.Subscribe{ + TrafficLimit: ptrString(""), + }, + wantSpeedLimit: 100, + wantTrafficLimit: planTrafficLimit, + }, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + gotSpeedLimit, gotTrafficLimit := serverUserSpeedLimitInputs(tc.sub, tc.userSub) + if gotSpeedLimit != tc.wantSpeedLimit { + t.Fatalf("speedLimit = %d, want %d", gotSpeedLimit, tc.wantSpeedLimit) + } + if gotTrafficLimit != tc.wantTrafficLimit { + t.Fatalf("trafficLimit = %q, want %q", gotTrafficLimit, tc.wantTrafficLimit) + } + }) + } +} + +func TestCalculateServerUserSpeedLimits_BatchTrafficRules(t *testing.T) { + db, mock, cleanup := newServerUserListTestDB(t) + defer cleanup() + + ctx, _ := gin.CreateTestContext(nil) + ctx.Request = httptestNewRequest() + logic := &GetServerUserListLogic{ + Logger: logger.WithContext(ctx.Request.Context()), + ctx: ctx, + svcCtx: &svc.ServiceContext{DB: db}, + } + + trafficLimit := `[{"stat_type":"day","stat_value":1,"traffic_usage":10,"speed_limit":5}]` + candidates := map[int64]serverUserSpeedLimitCandidate{ + 101: {userID: 1001, userSubscribeID: 101, baseSpeed: 0, trafficLimit: trafficLimit}, + 102: {userID: 1002, userSubscribeID: 102, baseSpeed: 0, trafficLimit: trafficLimit}, + 103: {userID: 1003, userSubscribeID: 103, baseSpeed: 20, trafficLimit: trafficLimit}, + } + + mock.ExpectQuery("FROM `traffic_log`"). + WithArgs(sqlmock.AnyArg(), sqlmock.AnyArg(), sqlmock.AnyArg(), sqlmock.AnyArg(), sqlmock.AnyArg()). + WillReturnRows(sqlmock.NewRows([]string{"user_id", "user_subscribe_id", "upload", "download"}). + AddRow(1001, 101, gb(4), gb(5)). + AddRow(1002, 102, gb(4), gb(6)). + AddRow(1003, 103, gb(12), int64(0))) + + got := logic.calculateServerUserSpeedLimits(candidates) + assertSpeedLimit(t, got, 101, 0) + assertSpeedLimit(t, got, 102, 5) + assertSpeedLimit(t, got, 103, 5) + assertServerUserListExpectations(t, mock) +} + +func TestCalculateServerUserSpeedLimits_UserTrafficLimitOverride(t *testing.T) { + db, mock, cleanup := newServerUserListTestDB(t) + defer cleanup() + + ctx, _ := gin.CreateTestContext(nil) + ctx.Request = httptestNewRequest() + logic := &GetServerUserListLogic{ + Logger: logger.WithContext(ctx.Request.Context()), + ctx: ctx, + svcCtx: &svc.ServiceContext{DB: db}, + } + + candidates := map[int64]serverUserSpeedLimitCandidate{ + 201: { + userID: 2001, + userSubscribeID: 201, + baseSpeed: 50, + trafficLimit: `[{"stat_type":"day","stat_value":1,"traffic_usage":2,"speed_limit":3}]`, + }, + 202: { + userID: 2002, + userSubscribeID: 202, + baseSpeed: 50, + trafficLimit: "", + }, + } + + mock.ExpectQuery("FROM `traffic_log`"). + WithArgs(sqlmock.AnyArg(), sqlmock.AnyArg(), sqlmock.AnyArg(), sqlmock.AnyArg()). + WillReturnRows(sqlmock.NewRows([]string{"user_id", "user_subscribe_id", "upload", "download"}). + AddRow(2001, 201, gb(2), int64(0))) + + got := logic.calculateServerUserSpeedLimits(candidates) + assertSpeedLimit(t, got, 201, 3) + assertSpeedLimit(t, got, 202, 50) + assertServerUserListExpectations(t, mock) +} + +func TestServerUserTrafficRuleWindow(t *testing.T) { + now := time.Date(2026, 6, 8, 12, 0, 0, 0, time.UTC) + + cases := []struct { + name string + rule speedlimit.TrafficLimitRule + want time.Time + ok bool + }{ + { + name: "hour window", + rule: speedlimit.TrafficLimitRule{StatType: "hour", StatValue: 2}, + want: now.Add(-2 * time.Hour), + ok: true, + }, + { + name: "day window", + rule: speedlimit.TrafficLimitRule{StatType: "day", StatValue: 1}, + want: now.AddDate(0, 0, -1), + ok: true, + }, + { + name: "invalid stat type", + rule: speedlimit.TrafficLimitRule{StatType: "week", StatValue: 1}, + ok: false, + }, + { + name: "invalid stat value", + rule: speedlimit.TrafficLimitRule{StatType: "day", StatValue: 0}, + ok: false, + }, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + got, ok := serverUserTrafficRuleWindow(tc.rule, now) + if ok != tc.ok { + t.Fatalf("ok = %v, want %v", ok, tc.ok) + } + if !ok { + return + } + if !got.start.Equal(tc.want) || !got.end.Equal(now) { + t.Fatalf("window = [%s, %s), want [%s, %s)", got.start, got.end, tc.want, now) + } + }) + } +} + +func newServerUserListTestDB(t *testing.T) (*gorm.DB, sqlmock.Sqlmock, func()) { + t.Helper() + + sqlDB, mock, err := sqlmock.New(sqlmock.QueryMatcherOption(sqlmock.QueryMatcherFunc(func(expectedSQL, actualSQL string) error { + if matched, _ := regexp.MatchString(expectedSQL, actualSQL); matched { + return nil + } + t.Fatalf("unexpected SQL\nexpected pattern: %s\nactual: %s", expectedSQL, actualSQL) + return nil + }))) + if err != nil { + t.Fatalf("create sqlmock: %v", err) + } + + db, err := gorm.Open(mysql.New(mysql.Config{ + Conn: sqlDB, + SkipInitializeWithVersion: true, + }), &gorm.Config{}) + if err != nil { + _ = sqlDB.Close() + t.Fatalf("open gorm db: %v", err) + } + + return db, mock, func() { + _ = sqlDB.Close() + } +} + +func assertServerUserListExpectations(t *testing.T, mock sqlmock.Sqlmock) { + t.Helper() + + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatalf("unmet sql expectations: %v", err) + } +} + +func assertSpeedLimit(t *testing.T, got map[int64]int64, userSubscribeID int64, want int64) { + t.Helper() + + if got[userSubscribeID] != want { + t.Fatalf("speed limit for user_subscribe_id=%d = %d, want %d", userSubscribeID, got[userSubscribeID], want) + } +} + +func gb(value int64) int64 { + return value * 1024 * 1024 * 1024 +} + +func ptrString(value string) *string { + return &value +} + +func httptestNewRequest() *http.Request { + req, _ := http.NewRequestWithContext(context.Background(), http.MethodGet, "/v1/server/user", nil) + return req +}